解讀Tensorflow2.0訓練損失值降低,但測試正確率基本不變的情況
Tensorflow2.0訓練損失值降低,但測試正確率基本不變的情況
問題描述
對于一個架構(gòu),在識別mnist手寫數(shù)字集精度較高的情況下,更換其他數(shù)據(jù)集,卻無法得到較高的識別結(jié)果。假設有n個類別,修改輸入端、輸出端及幾個卷積核的大小,識別時雖然loss在減小,但正確率acc穩(wěn)定在1/n左右不變化。
解決方法
修改參數(shù)
主要考慮的參數(shù)有batch、學習率和keep_prob:
- batch ,降低該值,使得網(wǎng)絡充分學習數(shù)據(jù);
- 學習率,降低該值,使得模型梯度下降;
- keep_prob ,降低該值,使得模型具有學習能力。
檢查模型
檢查模型是否有問題,修改網(wǎng)絡的架構(gòu)。
loss計算方法
選擇loss計算的公式方法是否有問題。
數(shù)據(jù)標簽
檢查數(shù)據(jù)的標簽是否轉(zhuǎn)換正確。
權(quán)重初始值
修改權(quán)重的初始化方法。
Tensorflow2.0準確率和損失值的可視化
進行準確率和損失值的可視化,就是將acc和loss使用matplot畫出來。
我們在使用model.fit()函數(shù)進行訓練時,同步記錄了訓練集和測試集的損失和準確率。
可以使用history進行調(diào)用,如下:
# 使用history將訓練集和測試集的loss和acc調(diào)出來
acc = history.history['sparse_categorical_accuracy'] # 訓練集準確率
val_acc = history.history['val_sparse_categorical_accuracy'] # 測試集準確率
loss = history.history['loss'] # 訓練集損失
val_loss = history.history['val_loss'] # 測試集損失
# 打印acc和loss,采用一個圖進行顯示。
# 將acc打印出來。
plt.subplot(1, 2, 1) # 將圖像分為一行兩列,將其顯示在第一列
plt.plot(acc, label='Training Accuracy')
plt.plot(val_acc, label='Validation Accuracy')
plt.title('Training and Validation Accuracy')
plt.legend()
plt.subplot(1, 2, 2) # 將其顯示在第二列
plt.plot(loss, label='Training Loss')
plt.plot(val_loss, label='Validation Loss')
plt.title('Training and Validation Loss')
plt.legend()
plt.show()將本篇代碼放在上篇文章代碼后,運行即可。
輸出結(jié)果:

總結(jié)
以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關(guān)文章
python3利用smtplib通過qq郵箱發(fā)送郵件方法示例
python實現(xiàn)郵件發(fā)送較為簡單,主要用到smtplib這個模塊,所以下面這篇文章主要給大家介紹了關(guān)于python3利用smtplib通過qq郵箱發(fā)送郵件的相關(guān)資料,文中通過示例代碼介紹的非常詳細,需要的朋友可以參考借鑒,下面隨著小編來一起看看吧。2017-12-12
python根據(jù)開頭和結(jié)尾字符串獲取中間字符串的方法
這篇文章主要介紹了python根據(jù)開頭和結(jié)尾字符串獲取中間字符串的方法,涉及Python操作字符串截取的相關(guān)技巧,具有一定參考借鑒價值,需要的朋友可以參考下2015-03-03
Python3 多線程(連接池)操作MySQL插入數(shù)據(jù)
本文將結(jié)合實例代碼,介紹Python3 多線程(連接池)操作MySQL插入數(shù)據(jù),具有一定的參考價值,感興趣的小伙伴們可以參考一下2021-06-06
Pandas實現(xiàn)復制dataframe中的每一行
這篇文章主要介紹了Pandas實現(xiàn)復制dataframe中的每一行方式,2024-02-02

