TensorFlow自定义train_step封装模型时损失无法正常最小化
你的自定义MyModel子类存在两个关键问题,导致训练时优化器根本没有更新内层模型的参数,损失自然无法正常下降:
- 第一,你没有为子类模型实现Keras要求的标准
call()方法,框架无法自动跟踪内层Functional模型的权重,导致self.trainable_weights为空列表。你在train_step里传入梯度计算的可训练参数列表是空的,优化器执行apply_gradients时没有更新任何参数,训练全程内层模型的权重都是初始随机值。你看到的损失波动只是因为训练时每个batch的样本被随机打乱,不同批次数据算出来的MSE自然有浮动,整体不会收敛。 - 第二,你将内层模型赋值给
self.model属性,和tf.keras.Model父类的内部保留属性重名,会进一步干扰Keras的权重跟踪、模型构建逻辑,加剧异常。
你可以在开始训练前加一行print(len(mm.trainable_weights))验证:输出会是0;而直接打印内层原生模型的len(model.trainable_weights)会输出10(5个Dense层,每个层对应kernel、bias两个参数张量),就能证明外层模型根本没有收集到内层的可训练参数。
你直接编译内层Functional模型训练时,走的是Functional模型自身的标准训练流程,它的trainable_weights是完整的,参数可以正常更新,所以损失会收敛,和封装后的外层模型逻辑完全不同。另外你用来做对比的Sequential模型只有3个10维的隐藏层,参数量远小于你写的4个30维隐藏层的Functional模型,正常训练下后者拟合能力更强,你之前得到的高损失和模型结构、SELU激活本身没有关系。
标准写法(推荐)
按照Keras子类模型的开发规范实现call()方法,同时修改内层模型的属性名避免冲突:
import tensorflow as tf # 内层模型定义保持不变 mod_input = tf.keras.layers.Input(shape=(4,)) mod = tf.keras.layers.Dense(30, "selu")(mod_input) mod = tf.keras.layers.Dense(30, "selu")(mod) mod = tf.keras.layers.Dense(30, "selu")(mod) pre_mu = tf.keras.layers.Dense(30, "selu")(mod) mu = tf.keras.layers.Dense(1, "linear")(pre_mu) base_model = tf.keras.Model(mod_input, mu, name="base_model") class MyModel(tf.keras.Model): def __init__(self, base_model, **kwargs): super().__init__(**kwargs) # 不要用self.model作为属性名,避免和父类内部属性冲突 self.base_model = base_model def call(self, inputs, training=None): # 实现标准前向传播逻辑,Keras会自动跟踪base_model的所有权重 return self.base_model(inputs, training=training) def train_step(self, data): x, y = data with tf.GradientTape() as tape: z_mean = self(x, training=True) # 通过self()调用走标准前向流程 total_loss = tf.math.reduce_mean(tf.math.square(z_mean - y)) # 此时self.trainable_weights已经包含内层模型的全部可训练参数 grads = tape.gradient(total_loss, self.trainable_weights) self.optimizer.apply_gradients(zip(grads, self.trainable_weights)) return {"loss": total_loss}
修复后你再打印len(mm.trainable_weights)就会得到10,和内层模型的参数数量一致,训练时参数会被正常更新,损失会正常收敛,且收敛后的损失会低于你写的小参数量Sequential模型。
临时修复(不推荐)
如果你暂时不想实现call()方法,可以直接在计算梯度的时候指定内层模型的可训练参数,同时修改属性名避免冲突:
class MyModel(tf.keras.Model): def __init__(self, base_model, **kwargs): super().__init__(**kwargs) self.base_model = base_model def train_step(self, data): x, y = data with tf.GradientTape() as tape: z_mean = self.base_model(x, training=True) total_loss = tf.math.reduce_mean(tf.math.square(z_mean - y)) # 直接指定内层模型的参数计算梯度、更新 grads = tape.gradient(total_loss, self.base_model.trainable_weights) self.optimizer.apply_gradients(zip(grads, self.base_model.trainable_weights)) return {"loss": total_loss}
这种写法虽然能跑通训练,但不符合Keras的模型开发规范,后续做模型保存、加载、推理、分布式训练的时候很容易触发其他兼容性问题。
内容的提问来源于stack exchange,提问作者Alberto

