升级TensorFlow/Keras后无法加载旧模型,求无需降级重训的解决方法
无需降级/重训加载旧模型的解决方案
针对你遇到的Conv2D不识别batch_input_shape参数的问题,以下是两种无需降级或重新训练的解决方法:
方法一:自定义Conv2D适配类加载
通过自定义兼容旧参数的Conv2D子类,自动将batch_input_shape转换为新版本支持的input_shape:
from tensorflow.keras.layers import Conv2D from tensorflow.keras.models import load_model class LegacyConv2D(Conv2D): def __init__(self, *args, **kwargs): # 移除旧参数并转换为input_shape(忽略batch维度) if 'batch_input_shape' in kwargs: kwargs['input_shape'] = kwargs.pop('batch_input_shape')[1:] super().__init__(*args, **kwargs) models_folder = '/savedModels/models_stack/' model1 = load_model(f'{models_folder}model_best1.keras', custom_objects={'Conv2D': LegacyConv2D})
旧版本Keras允许在Conv2D中通过batch_input_shape指定输入形状,而Keras 3已废弃该参数,改用input_shape(无需指定batch维度)。自定义类会自动完成参数转换,让模型正常加载。
方法二:手动修改模型配置文件
如果模型是以.keras格式保存的(实际为文件夹结构),可以直接修改配置文件后重建模型:
import json from tensorflow.keras.models import model_from_json models_folder = '/savedModels/models_stack/' # 读取模型配置 with open(f'{models_folder}model_best1.keras/config.json', 'r') as f: config = json.load(f) # 遍历所有层,修复Conv2D的参数问题 for layer in config['config']['layers']: if layer['class_name'] == 'Conv2D' and 'batch_input_shape' in layer['config']: # 将batch_input_shape转换为input_shape,去掉batch维度 layer['config']['input_shape'] = layer['config'].pop('batch_input_shape')[1:] # 从修改后的配置重建模型并加载权重 model = model_from_json(json.dumps(config)) model.load_weights(f'{models_folder}model_best1.keras/variables/variables')
.keras格式包含模型配置文件和权重文件,修改配置中Conv2D层的参数后,用新配置重建模型再加载权重,即可适配新版本Keras。
内容的提问来源于stack exchange,提问作者Romário Carvalho Neto
相关产品推荐
相关产品推荐

