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

Colab中Keras模型保存后加载报嵌套结构不匹配错误

Keras模型保存后加载触发输入结构不匹配ValueError修复

问题复现

  • 运行环境:Colab平台,基于Keras完成模型训练
  • 模型保存执行代码:
    model.save('/content/model')
  • 保存阶段输出日志包含两类警告:
    1. 存在未追踪函数,加载后无法直接调用
    2. 自定义LSTMCell与Keras内置对象重名
      完整保存日志如下:
WARNING:absl:Found untraced functions such as embeddings_layer_call_fn, embeddings_layer_call_and_return_conditional_losses, encoder_layer_call_fn, encoder_layer_call_and_return_conditional_losses, pooler_layer_call_fn while saving (showing 5 of 430). These functions will not be directly callable after loading.
INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).
INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).
INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).
INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).
INFO:tensorflow:Assets written to: /content/model/assets
INFO:tensorflow:Assets written to: /content/model/assets
WARNING:absl:<keras.layers.recurrent.LSTMCell object at 0x7f95000fac10> has the same name 'LSTMCell' as a built-in Keras object. Consider renaming <class 'keras.layers.recurrent.LSTMCell'> to avoid naming conflicts when loading with `tf.keras.models.load_model`. If renaming is not possible, pass the object in the `custom_objects` parameter of the load function.
WARNING:absl:<keras.layers.recurrent.LSTMCell object at 0x7f950018eb50> has the same name 'LSTMCell' as a built-in Keras object. Consider renaming <class 'keras.layers.recurrent.LSTMCell'> to avoid naming conflicts when loading with `tf.keras.models.load_model`. If renaming is not possible, pass the object in the `custom_objects` parameter of the load function.
  • 模型加载执行代码:
    model_CSP = keras.models.load_model('/content/model')
  • 加载阶段触发ValueError,提示嵌套结构不匹配,完整错误日志:
2022-06-22 01:57:09.702116: W tensorflow/core/common_runtime/gpu/gpu_bfc_allocator.cc:39] Overriding allow_growth setting because the TF_FORCE_GPU_ALLOW_GROWTH environment variable is set. Original config value was 0.
Traceback (most recent call last):
  File "Do_predection_Combined_model_.py", line 730, in <module>
    model_CSP = keras.models.load_model('/content/CSP_model')
  File "/usr/local/lib/python3.7/dist-packages/keras/utils/traceback_utils.py", line 67, in error_handler
    raise e.with_traceback(filtered_tb) from None
  File "/usr/local/lib/python3.7/dist-packages/tensorflow/python/util/nest.py", line 573, in assert_same_structure
    % (str(e), str1, str2))
ValueError: The two structures don't have the same nested structure.

First structure: type=tuple str=(({'input_ids': TensorSpec(shape=(None, 5), dtype=tf.int32, name=None)},), {'training': False})

Second structure: type=tuple str=((TensorSpec(shape=(None, 128), dtype=tf.int32, name='inputs'),), {'token_type_ids': TensorSpec(shape=(None, 128), dtype=tf.int32, name='token_type_ids'), 'training': False, 'attention_mask': TensorSpec(shape=(None, 128), dtype=tf.int32, name='attention_mask')})

More specifically: Substructure "type=dict str={'input_ids': TensorSpec(shape=(None, 5), dtype=tf.int32, name=None)}" is a sequence, while substructure "type=TensorSpec str=TensorSpec(shape=(None, 128), dtype=tf.int32, name='inputs')" is not
Entire first structure:
(({'input_ids': .},), {'training': .})
Entire second structure:
((.,), {'token_type_ids': ., 'training': ., 'attention_mask': .})
  • 已尝试操作:严格遵循Keras官方保存加载最佳实践,测试所有已知方案均失败,无法定位问题原因。

问题根因

该报错和保存阶段的两个警告无直接关联,核心是模型保存时记录的输入签名和实际模型结构不匹配:

  1. 从错误日志可直接看出,保存的模型记录的输入签名是仅包含input_ids(序列长度5)的字典结构,但实际构建的是包含inputs、token_type_ids、attention_mask三个字段(序列长度128)的多输入模型,两边输入格式完全不一致,加载时结构校验直接失败。
  2. 该问题通常出现在多输入模型构建场景:构建模型时混用函数式API、自定义层硬编码输入格式、call方法参数声明不规范,都会导致Keras保存模型追踪计算图时错误识别输入结构。
  3. 优先排查低级错误:错误日志中加载的模型路径是/content/CSP_model,但执行保存的路径是/content/model,如果路径写错加载了其他旧版本模型,也会直接触发该错误。

修复方案

按优先级依次尝试:

  • 第一步:确认路径一致性,检查加载时传入的路径和实际保存模型的路径完全匹配,排除加载到旧版本错误模型的情况。
  • 第二步:解决自定义层命名冲突:自定义的LSTMCell类不要和Keras内置LSTMCell同名,直接修改自定义类名为CustomLSTMCell即可,实例化层时也可显式指定name参数避免冲突:
    self.lstm_cell = CustomLSTMCell(hidden_dim, name="custom_lstm_cell_0")
    如果不想修改类名,加载模型时需要把自定义类传入custom_objects参数:
    model_CSP = keras.models.load_model('/content/model', custom_objects={'LSTMCell': CustomLSTMCell})
  • 第三步:绕开Keras全量保存的签名追踪bug,改用权重保存+结构重建的方式,这是稳定性最高的方案:
    1. 训练完成后仅保存模型权重:
      model.save_weights('/content/model_weights.h5')
    2. 推理前先运行和训练阶段完全一致的模型构建代码,实例化出结构完全相同的模型实例,再加载权重:
    # 此处执行和训练阶段完全相同的模型构建逻辑,得到model实例
    model.load_weights('/content/model_weights.h5')
    
  • 第四步:如果必须保存可直接加载的完整模型,在模型构建完成后、训练前显式指定输入签名,强制Keras按定义的结构追踪计算图:
    # 按实际输入的shape、dtype、name定义输入签名
    input_signature = [
        (
            tf.TensorSpec(shape=(None, 128), dtype=tf.int32, name='inputs'),
        ),
        {
            'token_type_ids': tf.TensorSpec(shape=(None, 128), dtype=tf.int32, name='token_type_ids'),
            'attention_mask': tf.TensorSpec(shape=(None, 128), dtype=tf.int32, name='attention_mask'),
            'training': False
        }
    ]
    # 给模型call方法绑定输入签名
    model.call = tf.function(input_signature=input_signature)(model.call)
    # 再执行保存操作
    model.save('/content/model')
    

说明:保存时弹出的未追踪函数警告不影响模型正常推理,只要加载后不需要单独调用对应层的内部方法,可直接忽略。


内容的提问来源于stack exchange,提问作者kaloon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 18:48:19