最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

Tensorflow2.1 完成權(quán)重或模型的保存和加載

 更新時間:2022年11月17日 16:30:52   作者:我是王大你是誰  
這篇文章主要為大家介紹了Tensorflow2.1 完成權(quán)重或模型的保存和加載,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進步,早日升職加薪

前言

本文主要使用 cpu 版本的 tensorflow-2.1 來完成深度學習權(quán)重參數(shù)/模型的保存和加載操作。

在我們進行項目期間,很多時候都要在模型訓練期間、訓練結(jié)束之后對模型或者模型權(quán)重進行保存,然后我們可以從之前停止的地方恢復原模型效果繼續(xù)進行訓練或者直接投入實際使用,另外為了節(jié)省存儲空間我們還可以自定義保存內(nèi)容和保存頻率。

實現(xiàn)方法

1. 讀取數(shù)據(jù)

(1)本文重點介紹模型或者模型權(quán)重的保存和讀取的相關(guān)操作,使用到的是 MNIST 數(shù)據(jù)集僅是為了演示效果,我們無需關(guān)心模型訓練的質(zhì)量好壞。

(2)這里是常規(guī)的讀取數(shù)據(jù)操作,我們?yōu)榱四茌^快介紹本文重點內(nèi)容,只使用了 MNIST 前 1000 條數(shù)據(jù),然后對數(shù)據(jù)進行歸一化操作,加快模型訓練收斂速度,并且將每張圖片的數(shù)據(jù)從二維壓縮成一維。

import os
import tensorflow as tf
from tensorflow import keras
(train_images, train_labels), (test_images, test_labels) = tf.keras.datasets.mnist.load_data()
train_labels = train_labels[:1000]
test_labels = test_labels[:1000]
train_images = train_images[:1000].reshape(-1, 28 * 28) / 255.0
test_images = test_images[:1000].reshape(-1, 28 * 28) / 255.0

2. 搭建深度學習模型

(1)這里主要是搭建一個最簡單的深度學習模型。

(2)第一層將圖片的長度為 784 的一維向量轉(zhuǎn)換成 256 維向量的全連接操作,并且用到了 relu 激活函數(shù)。

(3)第二層緊接著使用了防止過擬合的 Dropout 操作,神經(jīng)元丟棄率為 50% 。

(4)第三層為輸出層,也就是輸出每張圖片屬于對應 10 種類別的分布概率。

(5)優(yōu)化器我們選擇了最常見的 Adam 。

(6)損失函數(shù)選擇了 SparseCategoricalCrossentropy 。

(7)評估指標選用了 SparseCategoricalAccuracy 。

def create_model():
    model = tf.keras.Sequential([keras.layers.Dense(256, activation='relu', input_shape=(784,)),
                                 keras.layers.Dropout(0.5),
                                 keras.layers.Dense(10) ])
    model.compile(optimizer='adam',
                  loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
                  metrics=[tf.keras.metrics.SparseCategoricalAccuracy()])
    return model

3. 使用回調(diào)函數(shù)在每個 epoch 后自動保存模型權(quán)重

(1)這里介紹一種在模型訓練期間保存權(quán)重參數(shù)的方法,我們定義一個回調(diào)函數(shù) callback ,它可以在訓練過程中將權(quán)重保存在自定義目錄中 weights_path ,在訓練過程中一共執(zhí)行 5 次 epoch ,每次 epoch 結(jié)束之后就會保存一次模型的權(quán)重到指定的目錄。

(2)可以看到最后使用測試集進行評估的 loss 為 0.4952 ,分類準確率為 0.8500 。

weights_path = "training_weights/cp.ckpt"
weights_dir = os.path.dirname(weights_path)
callback = tf.keras.callbacks.ModelCheckpoint(filepath=weights_path, save_weights_only=True,  verbose=1)
model = create_model()
model.fit(train_images, 
          train_labels,  
          epochs=5,
          validation_data=(test_images, test_labels),
          callbacks=[callback]) 

輸出結(jié)果為:

 val_loss: 0.4952 - val_sparse_categorical_accuracy: 0.8500             

(3)我們?yōu)g覽目標文件夾里,只有三個文件,每個 epoch 后自動都會保存三個文件,在下一次 epoch 之后會自動更新這三個文件的內(nèi)容。

os.listdir(weights_dir)

結(jié)果為:

['checkpoint', 'cp.ckpt.data-00000-of-00001', 'cp.ckpt.index']

(4) 我們通過 create_model 定義了一個新的模型實例,然后讓其在沒有訓練的情況下使用測試數(shù)據(jù)進行評估,結(jié)果可想而知,準確率差的離譜。

NewModel = create_model()
loss, acc = NewModel.evaluate(test_images, test_labels, verbose=2)

結(jié)果為:

loss: 2.3694 - sparse_categorical_accuracy: 0.1330

(5) tensorflow 中只要兩個模型有相同的模型結(jié)構(gòu),就可以在它們之間共享權(quán)重,所以我們使用 NewModel 讀取了之前訓練好的模型權(quán)重,再使用測試集對其進行評估發(fā)現(xiàn),損失值和準確率和舊模型的結(jié)果完全一樣,說明權(quán)重被相同結(jié)構(gòu)的新模型成功加載并使用。

NewModel.load_weights(checkpoint_path)
loss, acc = NewModel.evaluate(test_images, test_labels, verbose=2)

輸出結(jié)果:

loss: 0.4952 - sparse_categorical_accuracy: 0.8500

4. 使用回調(diào)函數(shù)每經(jīng)過 5 個 epoch 對模型權(quán)重保存一次

(1)如果我們想保留多個中間 epoch 的模型訓練的權(quán)重,或者我們想每隔幾個 epoch 保存一次模型訓練的權(quán)重,這時候我們可以通過設置保存頻率 period 來完成,我這里讓新建的模型訓練 30 個 epoch ,在每經(jīng)過 10 epoch 后保存一次模型訓練好的權(quán)重。

(2)使用測試集對此次模型進行評估,損失值為 0.4047 ,準確率為 0.8680 。

weights_path = "training_weights2/cp-{epoch:04d}.ckpt"
weights_dir = os.path.dirname(weights_path)
batch_size = 64
cp_callback = tf.keras.callbacks.ModelCheckpoint( filepath=weights_path, 
                                                  verbose=1, 
                                                  save_weights_only=True,
                                                  period=10)
model = create_model()
model.save_weights(weights_path.format(epoch=1))
model.fit(train_images, 
          train_labels,
          epochs=30, 
          batch_size=batch_size, 
          callbacks=[cp_callback],
          validation_data=(test_images, test_labels),
          verbose=1)

結(jié)果輸出為:

val_loss: 0.4047 - val_sparse_categorical_accuracy: 0.8680   

(3)這里我們能看到指定目錄中的文件組成,這里的 0001 是因為訓練時指定了要保存的 epoch 的權(quán)重,其他都是每 10 個 epoch 保存的權(quán)重參數(shù)文件。目錄中有一個 checkpoint ,它是一個檢查點文本文件,文件保存了一個目錄下所有的模型文件列表,首行記錄的是最后(最近)一次保存的模型名稱。

(4)每個 epoch 保存下來的文件都包含:

  • 一個索引文件,指示哪些權(quán)重存儲在哪個分片中
  • 一個或多個包含模型權(quán)重的分片

瀏覽文件夾內(nèi)容

os.listdir(weights_dir)

結(jié)果如下:

['checkpoint', 'cp-0001.ckpt.data-00000-of-00001', 'cp-0001.ckpt.index', 'cp-0010.ckpt.data-00000-of-00001', 'cp-0010.ckpt.index', 'cp-0020.ckpt.data-00000-of-00001', 'cp-0020.ckpt.index', 'cp-0030.ckpt.data-00000-of-00001', 'cp-0030.ckpt.index']

(5)我們將最后一次保存的權(quán)重讀取出來,然后創(chuàng)建一個新的模型去讀取剛剛保存的最新的之前訓練好的模型權(quán)重,然后通過測試集對新模型進行評估,發(fā)現(xiàn)損失值準確率和之前完全一樣,說明權(quán)重被成功讀取并使用。

latest = tf.train.latest_checkpoint(weights_dir)
newModel = create_model()
newModel.load_weights(latest)
loss, acc = newModel.evaluate(test_images, test_labels, verbose=2)

結(jié)果如下:

loss: 0.4047 - sparse_categorical_accuracy: 0.8680

5. 手動保存模型權(quán)重到指定目錄

(1)有時候我們還想手動將模型訓練好的權(quán)重保存到指定的目錄下,我們可以使用 save_weights 函數(shù),通過我們新建了一個同樣的新模型,然后使用 load_weights 函數(shù)去讀取權(quán)重并使用測試集對其進行評估,發(fā)現(xiàn)損失值和準確率仍然和之前的兩種結(jié)果完全一樣。

model.save_weights('./training_weights3/my_cp')
newModel = create_model()
newModel.load_weights('./training_weights3/my_cp')
loss, acc = newModel.evaluate(test_images, test_labels, verbose=2)

結(jié)果如下:

loss: 0.4047 - sparse_categorical_accuracy: 0.8680

6. 手動保存整個模型結(jié)構(gòu)和權(quán)重

(1)有時候我們還需要保存整個模型的結(jié)構(gòu)和權(quán)重,這時候我們直接使用 save 函數(shù)即可將這些內(nèi)容保存到指定目錄,使用該方法要保證目錄是存在的否則會報錯,所以這里我們要創(chuàng)建文件夾。我們能看到損失值為 0.4821,準確率為 0.8460 。

model = create_model()
model.fit(train_images, train_labels, epochs=5, validation_data=(test_images, test_labels), verbose=1)
!mkdir my_model
modelPath = './my_model'
model.save(modelPath)

輸出結(jié)果:

val_loss: 0.4821 - val_sparse_categorical_accuracy: 0.8460

(2)然后我們通過函數(shù) load_model 即可生成出一個新的完全一樣結(jié)構(gòu)和權(quán)重的模型,我們使用測試集對其進行評估,發(fā)現(xiàn)準確率和損失值和之前完全一樣,說明模型結(jié)構(gòu)和權(quán)重被完全讀取恢復。

new_model = tf.keras.models.load_model(modelPath)
loss, acc = new_model.evaluate(test_images, test_labels, verbose=2)

輸出結(jié)果:

 loss: 0.4821 - sparse_categorical_accuracy: 0.8460

以上就是Tensorflow2.1 完成權(quán)重或模型的保存和加載的詳細內(nèi)容,更多關(guān)于Tensorflow完成權(quán)重模型保存加載的資料請關(guān)注腳本之家其它相關(guān)文章!

相關(guān)文章

  • Python裝飾器原理與基本用法分析

    Python裝飾器原理與基本用法分析

    這篇文章主要介紹了Python裝飾器原理與基本用法,結(jié)合實例形式分析了Python裝飾器的基本功能、原理、用法與操作注意事項,需要的朋友可以參考下
    2020-01-01
  • 如何利用python讀取micaps文件詳解

    如何利用python讀取micaps文件詳解

    這篇文章主要給大家介紹了關(guān)于如何利用python讀取micaps文件的相關(guān)資料,文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友們下面隨著小編來一起學習學習吧
    2020-10-10
  • 解決pytorch load huge dataset(大數(shù)據(jù)加載)

    解決pytorch load huge dataset(大數(shù)據(jù)加載)

    這篇文章主要介紹了解決pytorch load huge dataset(大數(shù)據(jù)加載)的問題,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2021-05-05
  • Flask模板繼承深入理解與應用

    Flask模板繼承深入理解與應用

    Flask中的模板可以繼承,通過繼承可以把模板中許多重復出現(xiàn)的元素抽取出來,放在父模板中,并且父模板通過定義block給子模板開一個口,子模板根據(jù)需要,再實現(xiàn)這個block
    2022-09-09
  • python判斷鏈表是否有環(huán)的實例代碼

    python判斷鏈表是否有環(huán)的實例代碼

    在本篇文章里小編給大家整理的是關(guān)于python判斷鏈表是否有環(huán)的知識點及實例代碼,需要的朋友們參考下。
    2020-01-01
  • python mysqldb連接數(shù)據(jù)庫

    python mysqldb連接數(shù)據(jù)庫

    今天無事想弄下python做個gui開發(fā),最近發(fā)布的是python 3k,用到了數(shù)據(jù)庫,通過搜索發(fā)現(xiàn)有一個mysqldb這樣的控件,可以使用,就去官方看了下結(jié)果,沒有2.6以上的版本
    2009-03-03
  • Python自定義元類的實例講解

    Python自定義元類的實例講解

    在本篇文章里小編給大家整理的是一篇關(guān)于Python自定義元類的實例講解內(nèi)容,有興趣的朋友們可以學習參考下。
    2021-03-03
  • python嵌套異常的兩種處理器

    python嵌套異常的兩種處理器

    在Python中,異常也可以嵌套,本文主要介紹了python嵌套異常的兩種處理器,文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友們下面隨著小編來一起學習學習吧
    2024-01-01
  • Python人臉識別第三方庫face_recognition接口說明文檔

    Python人臉識別第三方庫face_recognition接口說明文檔

    Python人臉識別第三方庫face_recognition接口簡單說明,及簡單使用方法
    2019-05-05
  • python安裝cx_Oracle和wxPython的方法

    python安裝cx_Oracle和wxPython的方法

    這篇文章主要介紹了python安裝cx_Oracle和wxPython的方法,本文給大家介紹的非常詳細,對大家的學習或工作具有一定的參考借鑒價值,需要的朋友可以參考下
    2020-09-09

最新評論

响水县| 青田县| 四会市| 武宣县| 浦县| 长葛市| 贵港市| 页游| 汕头市| 将乐县| 辽宁省| 兴宁市| 双辽市| 阳西县| 福鼎市| 凭祥市| 资兴市| 贡山| 溧水县| 十堰市| 平舆县| 石城县| 京山县| 三穗县| 开江县| 高碑店市| 红安县| 育儿| 灵宝市| 通辽市| 凉山| 宿迁市| 泽普县| 巴南区| 漯河市| 观塘区| 栖霞市| 黑山县| 普安县| 邹平县| 古浪县|