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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 06:36:26