1. 程式人生 > 程式設計 >keras的load_model實現載入含有引數的自定義模型

keras的load_model實現載入含有引數的自定義模型

網上的教程大多數是教大家如何載入自定義模型和函式,如下圖

keras的load_model實現載入含有引數的自定義模型

這個SelfAttention層是在訓練過程自己定義的一個class,但如果要載入這個自定義層,需要在load_model裡新增custom_objects字典,這個自定義的類,不要用import ,最好是直接複製進再訓練的模型中,這些是基本教程。

------------------分割線講重點------------------

如果直接執行上面的程式碼,會出現一個init初始化錯誤,如下圖,

keras的load_model實現載入含有引數的自定義模型

再來看看 這個SelfAttention 自定義的類的初始化

keras的load_model實現載入含有引數的自定義模型

這就說明再呼叫這個類的時候,輸入的ch=256並不會初始化這個類,需要先自定義好初始化值,如下圖

keras的load_model實現載入含有引數的自定義模型

呼叫方式不變

keras的load_model實現載入含有引數的自定義模型

這樣問題就解決啦!

補充知識:keras load model的時候,報錯('Keyword argument not understood:',u'******')如何解決

由於keras不同版本的API有變化,因此在一個keras版本下訓練的模型在另一個keras版本下載入時,可能會出現諸如('Keyword argument not understood:',u'data_format')等報錯。

通過開啟*.h5檔案,檢視該模型訓練所用環境,再安裝該環境即可解決報錯。

檢視Keras Model所用的Keras環境的方法

import h5py

f = h5py.File('Model.h5','r')
print(f.attrs.get('keras_version'))

根據輸出的keras版本安裝對應版本的keras即可解決載入問題。

以上這篇keras的load_model實現載入含有引數的自定義模型就是小編分享給大家的全部內容了,希望能給大家一個參考,也希望大家多多支援我們。