Keras 自定义模型(尤其是含子模型嵌套结构的)在使用 model.save() 和 keras.saving.load_model() 时易因变量作用域混乱、层命名冲突或 call 方法调用方式不当,导致加载时报“Shape mismatch”或“expected 0 variables but received 2”等错误。本文详解根本原因并提供可落地的修复方案。
keras 自定义模型(尤其是含子模型嵌套结构的)在使用 `model.save()` 和 `keras.saving.load_model()` 时易因变量作用域混乱、层命名冲突或 `call` 方法调用方式不当,导致加载时报“shape mismatch”或“expected 0 variables but received 2”等错误。本文详解根本原因并提供可落地的修复方案。
在 TensorFlow/Keras 中,保存与加载自定义 keras.Model 子类的核心前提,是确保所有可训练层(包括嵌套子模型)的权重和结构能被完整、无歧义地序列化与反序列化。你遇到的 ValueError: Cannot assign value to variable 'dense_4/kernel:0': Shape mismatch 并非偶然的维度错误,而是模型加载时试图将权重赋给一个尚未正确构建(unbuilt)或命名冲突的层实例所引发的底层变量绑定失败。
? 根本原因:call 中绕过主模型 call 链导致子模型未被正确注册
从你的最小复现示例可精准定位问题:
# ❌ 错误写法:在 train_step 中直接调用子模型方法 y_pred = self.net2(self.net1(x)) # ← 这里跳过了 CustomModel3.call() # ✅ 正确写法:统一经由主模型 call 流程 z, y_pred = self(x) # ← 触发 CustomModel3.call() → 内部再调用 self.net1() 和 self.net2()
为什么这会导致加载失败?
Keras 的 save() 机制依赖于 模型的完整计算图构建状态(built state)和层名唯一性。当你在 train_step 中直接调用 self.net1(x):
- self.net1 和 self.net2 的 call() 被独立执行,但它们的 build() 可能未被主模型显式触发(尤其在动态图模式下);
- 更关键的是:Keras 序列化器无法识别这些“游离调用”的层参与了主模型的前向传播图,因此在保存时可能遗漏其权重初始化上下文,或在加载时因层名(如 dense, dense_1)重复生成(如 dense_3)而混淆变量归属;
- 最终表现为:加载时发现某层(如 dense_4)期望形状为 (24, 512),但实际权重文件中存的是 (784, 256)——这正是 Encoder 中第一个 Dense(256) 层(输入 784 维)与 Generator 中 Dense(512) 层(输入 24 维)的权重被错误互换或覆盖的典型症状。
✅ 正确实践:强制统一构建 + 显式层命名 + 安全序列化
1. 确保所有子模型在 __init__ 后立即 build()(推荐)
@tf.keras.saving.register_keras_serializable()
class HighNet(keras.Model):
def __init__(self, encoder, decoder, args, **kwargs):
super().__init__(**kwargs)
self.encoder = encoder
self.decoder = decoder
self.args = args
self.generator = Generator(args).build() # ✅ 立即 build,避免 lazy-build 不确定性
# ... 其他初始化
2. 在 call() 中完整封装子模型调用(强制纳入计算图)
def call(self, inputs, **kwargs):
# ✅ 所有子模型调用必须出现在 call() 中,确保被序列化器捕获
z = self.encoder(inputs) # → 触发 Encoder.call()
recon = self.decoder(z) # → 触发 Decoder.call()
return recon # 或根据需要返回中间结果
3. 重写 get_config() 与 from_config()(关键!)
仅 get_config() 不够,必须提供 from_config() 工厂方法,否则 Keras 无法重建嵌套模型实例:
@tf.keras.saving.register_keras_serializable()
class HighNet(keras.Model):
def __init__(self, encoder, decoder, args, **kwargs):
super().__init__(**kwargs)
self.encoder = encoder
self.decoder = decoder
self.args = args
self.generator = Generator(args).build()
def get_config(self):
# ✅ 序列化时只保存 args 和必要参数,不保存模型实例(会出错!)
return {
"args": self.args,
# 注意:不要序列化 encoder/decoder 实例!改用 config + from_config 重建
}
@classmethod
def from_config(cls, config):
# ✅ 反序列化时,用 config 重建子模型
encoder = Encoder(config["args"]).build()
decoder = Decoder(config["args"]).build()
return cls(encoder=encoder, decoder=decoder, args=config["args"])
⚠️ 重要提醒:get_config() 中绝不可直接序列化 encoder 或 decoder 模型对象(如 "encoder": self.encoder),这会导致 Pickle 式引用,破坏 Keras 原生序列化协议,引发 ValueError: Layer 'dense_3' expected 0 variables...。
4. 保存/加载时使用标准 Keras API(无需 pickle/dill)
# ✅ 正确保存(.keras 格式)
model.save("high_model.keras", save_format="keras")
# ✅ 正确加载(自动调用 from_config)
loaded_model = keras.saving.load_model("high_model.keras")
? 验证是否修复:检查加载后模型结构
加载后立即验证关键层输出形状与原始模型一致:
# 加载后检查
print("Loaded model summary:")
loaded_model.summary()
# 验证 encoder 输入/输出形状
test_input = tf.random.normal((1, 28, 28, 1))
z = loaded_model.encoder(test_input)
print(f"Encoder output shape: {z.shape}") # 应为 (1, 24)
# 验证 generator 输入/输出形状
noise = tf.random.normal((1, 24))
z_gen = loaded_model.generator(noise)
print(f"Generator output shape: {z_gen.shape}") # 应为 (1, 24)
? 总结:三大黄金法则
| 原则 | 错误做法 | 正确做法 |
|---|---|---|
| 构建时机 | 依赖 call() 时 lazy-build | __init__ 中显式调用 .build() |
| 调用路径 | train_step 中直调 submodel(x) | 所有子模型调用收束于主模型 call() |
| 序列化逻辑 | get_config() 包含模型实例 | get_config() 仅含参数,from_config() 重建子模型 |
遵循以上原则,即可彻底规避 Shape mismatch 和 variable assignment 类加载错误,确保自定义嵌套模型在训练、保存、部署全流程中稳定可靠。










