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

如何在VAE中使用随Epoch变化的自定义参数且无需启用run_eagerly?

问题描述

我正在构建一个VAE,参考Keras官方VAE示例实现,自定义了损失函数,希望给损失项乘以一个随训练轮次(Epoch)递增的系数。当前的实现方式如下:

  • 模型初始化时定义非训练型权重变量:
def __init__(self, encoder, decoder, **kwargs):
    self.eloss_weight = tf.Variable(initial_value=args.eloss_weight, trainable=False)
  • 编译模型时启用即时执行模式:
vae.compile(optimizer=tf.keras.optimizers.Adam(jit_compile=False), run_eagerly=True)
  • 训练阶段通过回调函数更新系数:
def eloss_weight_increase(epoch, logs):
    vae.eloss_weight = vae.eloss_weight + 1

increase_eloss_weight = tf.keras.callbacks.LambdaCallback(on_epoch_end=eloss_weight_increase)

vae.fit(
    X_train, V_train, batch_size=args.batch_size, epochs=args.epochs,
    callbacks=[increase_eloss_weight], verbose=1,)

但启用run_eagerly=True后训练速度骤降4倍,原本4.5小时的训练耗时拉长至近一天。想寻求无需启用即时执行的替代实现方案。

解决方案

方法1:使用assign_add更新变量值(最简洁方案)

当前回调中直接赋值的方式会替换变量对象,导致图模式下的计算图无法追踪到更新。改用TensorFlow原生的assign_add方法修改已有变量的内部值,就能让计算图正确识别变量更新,无需启用即时执行:

修改回调函数

def eloss_weight_increase(epoch, logs):
    vae.eloss_weight.assign_add(1.0)  # 用assign_add更新变量值,而非替换变量

increase_eloss_weight = tf.keras.callbacks.LambdaCallback(on_epoch_end=eloss_weight_increase)

编译模型时移除run_eagerly=True

vae.compile(optimizer=tf.keras.optimizers.Adam(jit_compile=True))

方法2:将训练轮次作为模型输入传入

把当前训练轮次作为额外输入传入模型,在损失计算逻辑中直接基于轮次计算系数,完全适配图模式:

自定义VAE模型

class VAE(tf.keras.Model):
    def __init__(self, encoder, decoder, initial_eloss_weight, **kwargs):
        super().__init__(**kwargs)
        self.encoder = encoder
        self.decoder = decoder
        self.initial_eloss_weight = initial_eloss_weight
        # 定义采样函数
        self.sampling = lambda args: args[0] + tf.exp(0.5 * args[1]) * tf.random.normal(tf.shape(args[0]))

    def call(self, inputs):
        x, epoch = inputs
        # 原有encoder逻辑,得到z_mean、z_log_var
        z_mean, z_log_var = self.encoder(x)
        z = self.sampling((z_mean, z_log_var))
        # 原有decoder逻辑
        reconstructed = self.decoder(z)
        return reconstructed, z_mean, z_log_var

    def train_step(self, data):
        x, y, epoch = data  # 训练数据包含输入、标签、当前轮次
        with tf.GradientTape() as tape:
            reconstructed, z_mean, z_log_var = self([x, epoch])
            # 计算重构损失
            reconstruction_loss = tf.reduce_mean(
                tf.keras.losses.mean_squared_error(y, reconstructed)
            )
            # 计算KL损失
            kl_loss = -0.5 * tf.reduce_mean(1 + z_log_var - tf.square(z_mean) - tf.exp(z_log_var))
            # 基于轮次计算动态系数
            eloss_weight = self.initial_eloss_weight + epoch
            total_loss = reconstruction_loss + eloss_weight * kl_loss
        
        # 更新权重
        grads = tape.gradient(total_loss, self.trainable_weights)
        self.optimizer.apply_gradients(zip(grads, self.trainable_weights))
        
        return {
            "total_loss": total_loss,
            "reconstruction_loss": reconstruction_loss,
            "kl_loss": kl_loss
        }

训练时构造含轮次的数据集

def generate_training_data(X_train, V_train, epochs):
    for epoch in range(epochs):
        # 为当前轮次生成对应维度的轮次张量
        epoch_tensor = tf.constant(epoch, dtype=tf.float32)
        epoch_batch = tf.repeat(epoch_tensor, repeats=X_train.shape[0])
        epoch_batch = tf.reshape(epoch_batch, (-1, 1))
        yield (X_train, V_train, epoch_batch)

# 初始化模型
vae = VAE(encoder, decoder, initial_eloss_weight=args.eloss_weight)
vae.compile(optimizer=tf.keras.optimizers.Adam(jit_compile=True))

# 启动训练
vae.fit(
    generate_training_data(X_train, V_train, args.epochs),
    batch_size=args.batch_size,
    epochs=args.epochs,
    verbose=1
)

方法3:自定义带权重调度的损失层

可以将动态系数的逻辑封装为自定义损失层,在层内部维护权重变量并通过回调更新,核心逻辑和方法1一致,适合需要模块化管理损失的场景。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 05:02:49