【发布时间】:2021-08-11 17:04:42
【问题描述】:
我有一个自定义 Keras 层,它从 pickle 文件中读取以初始化一些权重,我希望能够在其上使用 tf.keras.utils.register_keras_serializable()。问题是我的__init__ 函数采用了pickle 文件的路径,当层再次反序列化时,该路径可能不可用。 Keras Assets 理论上应该使层更便携,但我不知道如何让它与层的get_config() 一起工作。
我的代码的准系统版本:
@tf.keras.utils.register_keras_serializable()
class AssetLayer(tf.keras.layers.Layer):
def __init__(self, asset_path, **kwargs):
super().__init__(**kwargs)
self.asset_path = asset_path
self.asset = tf.saved_model.Asset(asset_path)
data = tf.io.read_file(self.asset)
# do something with data
def get_config(self):
return {
**super().get_config(),
"asset_path": self.asset_path,
}
def call(self, arg):
# arbitrary call function
return arg
如果使用该层的模型是使用tf.keras.models.load_model() 加载的,Keras 将调用get_config() 以使用保存的asset_path 重新初始化该层,该asset_path 在反序列化时可能未指向正确的位置。理想情况下,它会指向已保存资产的路径,但我不知道如何做到这一点。
例如,我试过这段代码
!echo abcd > file.txt
model = tf.keras.Sequential([AssetLayer("file.txt")])
model(tf.ones(3))
model.save("test")
# reloading
!rm file.txt
reloaded_model = tf.keras.models.load_model("test")
这给了我一个错误,提示找不到 file.txt。
我也尝试过完全删除 get_config() 函数。这使得图层可以成功重新加载,同时保留对asset 变量的访问权限,但图层中的其他属性(例如self.asset_path)不可访问。这对于调试目的来说并不理想,所以我想知道是否有更好的方法。
我目前正在使用 Tensorflow 2.5.0`
【问题讨论】:
标签: python tensorflow keras deep-learning