TensorFlow 2.6自定义层使用已注册自定义激活函数触发未知报错如何解决
TensorFlow 2.6+ 自定义激活函数在自定义层中识别失败的解决方案
该问题是TensorFlow 2.6调整Keras自定义对象解析逻辑导致的:内置层实例化时会默认读取get_custom_objects()的全局注册内容,而自定义层的激活参数解析逻辑不再默认读取全局注册项,仅识别显式传入的自定义对象,因此出现2.5版本运行正常、2.6+版本报错的情况。
可选择以下任意一种方案解决:
- 方案1:直接传入激活函数对象,不使用字符串标识
这是兼容性最强的方案,无需注册即可跨版本生效:dense = 你的自定义层类(3, activation=my_act) - 方案2:用官方装饰器注册自定义激活
该方案支持字符串调用,同时适配内置层和自定义层,比手动更新全局自定义对象更稳定:import tensorflow as tf @tf.keras.utils.register_keras_serializable() def my_act(x): return x # 后续可直接用字符串调用 dense = 你的自定义层类(3, activation="my_act") - 方案3:实例化自定义层时显式传入custom_objects参数
如果你坚持用get_custom_objects().update的注册方式,可在实例化时补充参数:dense = 你的自定义层类(3, activation="my_act", custom_objects={"my_act": my_act}) - 方案4:加载已保存模型时报错的处理方式
如果是加载已有模型触发该错误,加载时传入自定义对象即可:model = tf.keras.models.load_model("你的模型路径", custom_objects={"my_act": my_act})
内容的提问来源于stack exchange,提问作者J Agustin Barrachina
相关产品推荐
相关产品推荐

