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

TensorFlow CycleGAN模型fit方法报错:无法保存模型求助

Keras CycleGAN模型保存报错的解决办法

问题根源

官方CycleGAN示例中的自定义CycleGan模型类仅实现了train_step方法,未明确定义call方法,导致Keras无法确定模型的输入输出结构,进而在使用ModelCheckpoint回调保存模型时触发报错。

解决办法

方法1:为CycleGan类添加call方法并显式构建模型

在CycleGan类中添加call方法定义前向传播逻辑,同时显式指定输入形状:

class CycleGan(keras.Model):
    # 保留原有的__init__、compile、train_step等方法
    def call(self, inputs):
        x, y = inputs
        fake_y = self.generator_g(x)
        fake_x = self.generator_f(y)
        return fake_y, fake_x

初始化模型后,显式构建输入形状(假设输入尺寸为256×256×3):

cycle_gan = CycleGan(generator_g, generator_f, discriminator_x, discriminator_y)
# None表示batch size可变
cycle_gan.build([(None, 256, 256, 3), (None, 256, 256, 3)])

之后正常使用原ModelCheckpoint回调即可:

model_checkpoint_callback = keras.callbacks.ModelCheckpoint(
    filepath="./checkpoints",
    save_weights_only=False,
    save_freq="epoch",
)
cycle_gan.fit(train_dataset, epochs=EPOCHS, callbacks=[model_checkpoint_callback])

方法2:自定义回调保存子模型

如果不想修改CycleGan类,可以自定义回调直接保存生成器和判别器等子模型,避免保存整个CycleGan实例:

import os
from tensorflow.keras.callbacks import Callback

class CycleGANCheckpoint(Callback):
    def __init__(self, save_dir):
        super().__init__()
        self.save_dir = save_dir
        os.makedirs(save_dir, exist_ok=True)
    
    def on_epoch_end(self, epoch, logs=None):
        self.model.generator_g.save(f"{self.save_dir}/generator_g_epoch_{epoch}.h5")
        self.model.generator_f.save(f"{self.save_dir}/generator_f_epoch_{epoch}.h5")
        self.model.discriminator_x.save(f"{self.save_dir}/discriminator_x_epoch_{epoch}.h5")
        self.model.discriminator_y.save(f"{self.save_dir}/discriminator_y_epoch_{epoch}.h5")

使用时替换原回调:

checkpoint_callback = CycleGANCheckpoint(save_dir="./cyclegan_checkpoints")
cycle_gan.fit(train_dataset, epochs=EPOCHS, callbacks=[checkpoint_callback])

说明

方法1让Keras能正确识别整个CycleGAN模型的结构,支持完整模型保存;方法2直接保存核心子模型,更轻量化,适合后续仅使用生成器进行图像转换的场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 16:25:19