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

在Keras中使用tf.train.Checkpoint保存GAN后恢复训练异常如何解决?

问题更新

为解决该问题,我保持checkpoint结构不变,编写了自定义train_step函数,该函数自行计算梯度并调用apply_weights,而非编译模型后使用train_on_batch,实现了GAN完整状态的正常恢复。但遗憾的是,该方法下dropout层无法正常工作,判别器在训练初期就达到极高准确率,导致模型无法正常训练,不过原问题已得到解决。

原始提问

我目前正在Keras中训练GAN,需要实现模型保存及后续恢复训练的能力。通常Keras中直接使用model.save()即可,但对于GAN而言,若将判别器和GAN组合模型(由生成器与判别器组合,判别器权重设为不可训练)分开保存加载,二者的关联会断裂,导致GAN无法正常运行。此前已有用户提出类似问题,得到的解决方案是使用tf.train.Checkpoint一次性将完整模型保存为检查点。

我按如下方式实现该方案:

def train(epochs, batch_size):
    checkpoint = tf.train.Checkpoint(g_optimizer=g_optimizer,
                                     d_optimizer=d_optimizer,
                                     generator=generator,
                                     discriminator=discriminator,
                                     gan=gan
                                     )
    ckpt_manager = tf.train.CheckpointManager(checkpoint, 'checkpoints', max_to_keep=3)

    if ckpt_manager.latest_checkpoint:
        checkpoint.restore(ckpt_manager.latest_checkpoint)
        discriminator.compile(loss='binary_crossentropy', optimizer=d_optimizer)

        i = Input(shape=(None, latent_dims))
        lcs = generator(i)

        discriminator.trainable = False

        valid = discriminator(lcs)

        gan = Model(i, valid)
        gan.compile(loss='binary_crossentropy', optimizer=g_optimizer)

    for epoch in epochs:
        #train discriminator...
        #train generator...
        ckpt_manager.save()

其中g_optimizer、d_optimizer为tf.keras.optimizers.Adam实例,generator、discriminator、gan均为tf.keras.Model实例。

使用该方案时,加载检查点后GAN模型与判别器的关联得到保留,初始训练运行正常,但中断训练后从检查点恢复训练时,判别器损失会大幅上涨,生成的数据毫无意义。

加载检查点后重新编译模型是我能想到的可以沿用优化器最新状态的唯一方式,但显然该方法存在问题,没有从断点处恢复训练,反而严重干扰了训练过程。

我是否错误使用了tf.train.Checkpoint?若需要更多信息来排查问题,请告知我。

编辑:补充完整代码

以下是模型初始化及训练的完整代码,该配置下模型首次创建时会进行编译,若从检查点恢复训练则会使用检查点中的最新优化器状态重新编译。我知道二次编译的方式并不合理,但我没有想到其他可以沿用检查点中优化器状态的方法,若有更优方案我非常愿意调整。需要说明的是,这里使用基于GRU的特殊GAN结构是因为我正在测试可变长时间序列生成能力,代码中包含很多数据相关的特定逻辑,但整体逻辑是通顺的,train_df是存储所有训练数据的pandas DataFrame。

def build_generator():
    input = Input(shape=(None, latent_dims))
    gru1 = GRU(100, activation='relu', return_sequences=True)(input)
    gru2 = GRU(100, activation='relu', return_sequences=True)(gru1)
    output = GRU(9, return_sequences=True, activation='sigmoid')(gru2)
    model = Model(input, output)
    return model

def build_discriminator():
    input = Input(shape=(None, 9))
    gru1 = GRU(100, return_sequences=True)(input)
    gru2 = GRU(100, return_sequences=True)(gru1)
    output = GRU(1, activation='sigmoid')(gru2)
    model = Model(input, output)
    return model

d_optimizer = opt.Adam(learning_rate=lr)
g_optimizer = opt.Adam(learning_rate=lr)

# Build discriminator
discriminator = build_discriminator()
discriminator.compile(loss='binary_crossentropy', optimizer=d_optimizer)

# Build generator
generator = build_generator()

# Build combined model
i = Input(shape=(None, latent_dims))
lcs = generator(i)
discriminator.trainable = False
valid = discriminator(lcs)

gan = Model(i, valid)
gan.compile(loss='binary_crossentropy', optimizer=g_optimizer)

def train(epochs, batch_size=1): #Only works with batch size of 1 currently
    sne = train_df.sn.unique()
    n_batches = int(len(sne) / batch_size)
    rng = np.random.default_rng(123)

    checkpoint = tf.train.Checkpoint(g_optimizer=g_optimizer,
                                     d_optimizer=d_optimizer,
                                     generator=generator,
                                     discriminator=discriminator,
                                     gan=gan
                                     )
    ckpt_manager = tf.train.CheckpointManager(checkpoint, 'checkpoints', max_to_keep=3)
    if ckpt_manager.latest_checkpoint:
        checkpoint.restore(ckpt_manager.latest_checkpoint)
        discriminator.compile(loss='binary_crossentropy', optimizer=d_optimizer)

        i = Input(shape=(None, latent_dims))
        lcs = generator(i)

        discriminator.trainable = False
        valid = discriminator(lcs)

        gan = Model(i, valid)
        gan.compile(loss='binary_crossentropy', optimizer=g_optimizer)

    for epoch in range(epochs):
        rng.shuffle(sne)
        g_losses, d_losses = [], []
        for batch in range(n_batches):
            real = np.random.uniform(0.0, 0.1, (batch_size, 1)) # Used instead of np.zeros to avoid zero gradients
            fake = np.random.uniform(0.9, 1.0, (batch_size, 1)) # Used instead of np.ones to avoid zero gradients

        # Select real data
        sn = sne[batch]
        sndf = train_df[train_df.sn == sn]
        X = sndf[['g_t', 'r_t', 'i_t', 'z_t', 'g', 'r', 'i', 'z', 'g_err', 'r_err', 'i_err', 'z_err']].values

        X = X.reshape((1, *X.shape))

        noise = rand.normal(size=(batch_size, latent_dims))
        noise = np.reshape(noise, (batch_size, 1, latent_dims))
        noise = np.repeat(noise, X.shape[1], 1)

        gen_lcs = generator.predict(noise)

        # Train discriminator
        d_loss_real = discriminator.train_on_batch(X, real)
        d_loss_fake = discriminator.train_on_batch(gen_lcs, fake)
        d_loss = 0.5 * np.add(d_loss_real, d_loss_fake)
        
        # Train generator
        noise = rand.normal(size=(2 * batch_size, latent_dims))
        noise = np.reshape(noise, (2 * batch_size, 1, latent_dims))
        noise = np.repeat(noise, X.shape[1], 1)

        gen_labels = np.zeros((2 * batch_size, 1))
        g_loss = gan.train_on_batch(noise, gen_labels)
        g_losses.append(g_loss)
        d_losses.append(d_loss)
    ckpt_manager.save()
    full_g_loss = np.mean(g_losses)
    full_d_loss = np.mean(d_losses)
    print(f'{epoch + 1}/{epochs} g_loss={full_g_loss}, d_loss={full_d_loss}')

train()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 07:54:07