You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

升级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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 04:12:42