mirror of
https://github.com/unanmed/ginka-generator.git
synced 2026-08-14 18:12:28 +08:00
feat: 最终权重不保存优化器
This commit is contained in:
parent
9aefa74714
commit
4e69733759
@ -174,9 +174,8 @@ class SeperatedModels:
|
||||
self.scheduler.load_state_dict(ckpt["scheduler"])
|
||||
return ckpt.get("epoch", 0)
|
||||
|
||||
def save(self, path: str, epoch: int):
|
||||
# 保存完整检查点(模型权重 + 优化器/调度器状态)
|
||||
torch.save({
|
||||
def save(self, path: str, epoch: int, save_optim: bool = False):
|
||||
data = {
|
||||
"epoch": epoch,
|
||||
"vq1": self.vq1.state_dict(),
|
||||
"vq2": self.vq2.state_dict(),
|
||||
@ -188,6 +187,8 @@ class SeperatedModels:
|
||||
"quantizer2": self.quantizer2.state_dict(),
|
||||
"quantizer3": self.quantizer3.state_dict(),
|
||||
"latent_mask_embedding": self.latent_mask_embedding.data,
|
||||
"optimizer": self.optimizer.state_dict(),
|
||||
"scheduler": self.scheduler.state_dict(),
|
||||
}, path)
|
||||
}
|
||||
if save_optim:
|
||||
data["optimizer"] = self.optimizer.state_dict()
|
||||
data["scheduler"] = self.scheduler.state_dict()
|
||||
torch.save(data, path)
|
||||
|
||||
@ -473,12 +473,12 @@ def train(device: torch.device):
|
||||
visualize_mask_growth(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"
|
||||
result.save(ckpt_path, epoch + 1)
|
||||
result.save(ckpt_path, epoch + 1, save_optim=True)
|
||||
tqdm.write(f"Saved checkpoint: {ckpt_path}")
|
||||
|
||||
# 训练结束后保存最终完整权重(含优化器状态,可用于后续续训或推理)
|
||||
# 训练结束后保存最终完整权重
|
||||
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}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Loading…
Reference in New Issue
Block a user