TensorFlow 2.7.0保存包含数据增强层的模型时出现KeyError报错
TensorFlow 2.7.0 保存含数据增强层模型KeyError报错解决方案
问题根因
报错来源于两个层面:
- 自定义
RandomColorDistortion层未实现标准序列化接口,SavedModel无法追踪层的配置参数 - TensorFlow 2.7版本的内置Keras预处理层存在已知序列化缺陷,即使从实验性API迁移为正式API仍会触发资源张量追踪失败
可行解决方案
- 修复自定义层序列化逻辑
给RandomColorDistortion层补充get_config和from_config方法,示例实现如下:class RandomColorDistortion(layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) # 原有初始化逻辑保留 def call(self, inputs, training=None): # 原有前向传播逻辑保留 def get_config(self): config = super().get_config() # 若有自定义初始化参数,一并更新到config字典中 return config @classmethod def from_config(cls, config): return cls(**config) - 换用H5格式保存模型
避开默认SavedModel格式的序列化限制,修改保存代码为:
加载时使用model.save("./model.h5", save_format="h5")tf.keras.models.load_model("./model.h5", custom_objects={"RandomColorDistortion": RandomColorDistortion})即可正常读取。 - 剥离数据增强层后保存推理模型
推理阶段不需要数据增强逻辑,可以重构推理模型复用训练权重,避免序列化异常:# 提取训练模型中除数据增强层外的所有层 inference_layers = model.layers[1:] inference_model = tf.keras.Sequential([ layers.Input(input_shape), *inference_layers ]) # 直接保存无数据增强层的推理模型 inference_model.save("./") - 升级TensorFlow版本
该内置预处理层序列化bug在TensorFlow 2.8及以上版本已被修复,可直接升级环境版本解决问题。
内容的提问来源于stack exchange,提问作者moumed
相关产品推荐
相关产品推荐

