Colab中Keras模型保存后加载报嵌套结构不匹配错误
Keras模型保存后加载触发输入结构不匹配ValueError修复
问题复现
- 运行环境:Colab平台,基于Keras完成模型训练
- 模型保存执行代码:
model.save('/content/model') - 保存阶段输出日志包含两类警告:
- 存在未追踪函数,加载后无法直接调用
- 自定义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官方保存加载最佳实践,测试所有已知方案均失败,无法定位问题原因。
问题根因
该报错和保存阶段的两个警告无直接关联,核心是模型保存时记录的输入签名和实际模型结构不匹配:
- 从错误日志可直接看出,保存的模型记录的输入签名是仅包含
input_ids(序列长度5)的字典结构,但实际构建的是包含inputs、token_type_ids、attention_mask三个字段(序列长度128)的多输入模型,两边输入格式完全不一致,加载时结构校验直接失败。 - 该问题通常出现在多输入模型构建场景:构建模型时混用函数式API、自定义层硬编码输入格式、call方法参数声明不规范,都会导致Keras保存模型追踪计算图时错误识别输入结构。
- 优先排查低级错误:错误日志中加载的模型路径是
/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,改用权重保存+结构重建的方式,这是稳定性最高的方案:
- 训练完成后仅保存模型权重:
model.save_weights('/content/model_weights.h5') - 推理前先运行和训练阶段完全一致的模型构建代码,实例化出结构完全相同的模型实例,再加载权重:
# 此处执行和训练阶段完全相同的模型构建逻辑,得到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
相关产品推荐
相关产品推荐

