RVQ-VAE 训练阶段可能存在权重保存问题
您好,感谢开源 SemTalk 代码。
我在训练第一阶段 RVQ-VAE 模型时,发现 checkpoint 保存逻辑可能存在一个问题,例如训练 upper、hands、lower 等部位的 RVQ-VAE。
以 configs/cnn_vqvae_upper_30.yaml 为例,其中配置为:
因此训练时会使用 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.npz 和 res.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>.bin 或 rvq_<epoch>.bin,这样可能会更稳定。
谢谢!
RVQ-VAE 训练阶段可能存在权重保存问题
您好,感谢开源 SemTalk 代码。
我在训练第一阶段 RVQ-VAE 模型时,发现 checkpoint 保存逻辑可能存在一个问题,例如训练
upper、hands、lower等部位的 RVQ-VAE。以
configs/cnn_vqvae_upper_30.yaml为例,其中配置为:因此训练时会使用
ae_trainer.py中的 trainer。但是在
train.py中,保存best_<epoch>.bin的逻辑似乎依赖于trainer.test(epoch)的返回值:也就是说,只有当
trainer.test(epoch)返回一个有效的 FID 数值时,才会触发 best checkpoint 的保存。但在
ae_trainer.py中,test()函数主要是保存重建结果,例如gt.npz和res.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>.bin或rvq_<epoch>.bin,这样可能会更稳定。谢谢!