feat: 最终权重不保存优化器

This commit is contained in:
unanmed 2026-08-05 14:08:01 +08:00
parent 9aefa74714
commit 4e69733759
2 changed files with 10 additions and 9 deletions

View File

@ -174,9 +174,8 @@ class SeperatedModels:
self.scheduler.load_state_dict(ckpt["scheduler"]) self.scheduler.load_state_dict(ckpt["scheduler"])
return ckpt.get("epoch", 0) return ckpt.get("epoch", 0)
def save(self, path: str, epoch: int): def save(self, path: str, epoch: int, save_optim: bool = False):
# 保存完整检查点(模型权重 + 优化器/调度器状态) data = {
torch.save({
"epoch": epoch, "epoch": epoch,
"vq1": self.vq1.state_dict(), "vq1": self.vq1.state_dict(),
"vq2": self.vq2.state_dict(), "vq2": self.vq2.state_dict(),
@ -188,6 +187,8 @@ class SeperatedModels:
"quantizer2": self.quantizer2.state_dict(), "quantizer2": self.quantizer2.state_dict(),
"quantizer3": self.quantizer3.state_dict(), "quantizer3": self.quantizer3.state_dict(),
"latent_mask_embedding": self.latent_mask_embedding.data, "latent_mask_embedding": self.latent_mask_embedding.data,
"optimizer": self.optimizer.state_dict(), }
"scheduler": self.scheduler.state_dict(), if save_optim:
}, path) data["optimizer"] = self.optimizer.state_dict()
data["scheduler"] = self.scheduler.state_dict()
torch.save(data, path)

View File

@ -473,12 +473,12 @@ def train(device: torch.device):
visualize_mask_growth(dataset, result, device, tile_dict, epoch + 1) visualize_mask_growth(dataset, result, device, tile_dict, epoch + 1)
visualize_mask_maskgit(dataset, result, device, tile_dict, epoch + 1) visualize_mask_maskgit(dataset, result, device, tile_dict, epoch + 1)
ckpt_path = f"result/seperated/sep-{epoch + 1}.pth" ckpt_path = f"result/seperated/sep-{epoch + 1}.pth"
result.save(ckpt_path, epoch + 1) result.save(ckpt_path, epoch + 1, save_optim=True)
tqdm.write(f"Saved checkpoint: {ckpt_path}") tqdm.write(f"Saved checkpoint: {ckpt_path}")
# 训练结束后保存最终完整权重(含优化器状态,可用于后续续训或推理) # 训练结束后保存最终完整权重
final_path = "result/seperated.pth" final_path = "result/seperated.pth"
result.save(final_path, EPOCHS) result.save(final_path, EPOCHS, save_optim=False)
tqdm.write(f"Training complete. Final model saved: {final_path}") tqdm.write(f"Training complete. Final model saved: {final_path}")
if __name__ == "__main__": if __name__ == "__main__":