Skip to content

RVQ-VAE 训练阶段权重保存问题 #9

Description

@cangnai

RVQ-VAE 训练阶段可能存在权重保存问题

您好,感谢开源 SemTalk 代码。

我在训练第一阶段 RVQ-VAE 模型时,发现 checkpoint 保存逻辑可能存在一个问题,例如训练 upperhandslower 等部位的 RVQ-VAE。

configs/cnn_vqvae_upper_30.yaml 为例,其中配置为:

trainer: ae

因此训练时会使用 ae_trainer.py 中的 trainer。

但是在 train.py 中,保存 best_<epoch>.bin 的逻辑似乎依赖于 trainer.test(epoch) 的返回值:

fid = trainer.test(epoch)
is_best = (fid is not None) and (fid < getattr(trainer, "best_fid", float("inf")))

if is_best:
    trainer.best_fid = fid
    save_checkpoints(... best_<epoch>.bin ...)

也就是说,只有当 trainer.test(epoch) 返回一个有效的 FID 数值时,才会触发 best checkpoint 的保存。

但在 ae_trainer.py 中,test() 函数主要是保存重建结果,例如 gt.npzres.npz,在正常执行路径下似乎并没有返回有效的 FID 值。因此 fid 通常会是 None,导致 is_best 一直为 False,从而不会保存 best_<epoch>.bin

另外,我还注意到 train.py 中似乎会多次调用 trainer.test(epoch)。由于 ae_trainer.test() 在发现当前 epoch 的结果目录已经存在时会直接 return 0,这可能会导致第二次调用时错误地返回 0,从而意外触发 best checkpoint 的保存逻辑。但这个 0 并不是真正计算得到的 FID,因此可能不是一个可靠的 best checkpoint。

想请问一下,这种行为是作者预期的吗?

我的理解是,对于第一阶段 RVQ-VAE 的训练,checkpoint 保存可能不应该依赖 FID,而更适合采用周期性保存,或者根据重建损失进行保存。例如每隔一定 epoch 保存 last_<epoch>.binrvq_<epoch>.bin,这样可能会更稳定。

谢谢!

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions