From 4e69733759e1ecd38ae8a7806b28b411bb9dc587 Mon Sep 17 00:00:00 2001 From: unanmed <1319491857@qq.com> Date: Wed, 5 Aug 2026 14:08:01 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9C=80=E7=BB=88=E6=9D=83=E9=87=8D?= =?UTF-8?q?=E4=B8=8D=E4=BF=9D=E5=AD=98=E4=BC=98=E5=8C=96=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ginka/model.py | 13 +++++++------ ginka/train_seperated.py | 6 +++--- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/ginka/model.py b/ginka/model.py index 9db7a61..8019bc7 100644 --- a/ginka/model.py +++ b/ginka/model.py @@ -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) diff --git a/ginka/train_seperated.py b/ginka/train_seperated.py index 50e868b..cb7014e 100644 --- a/ginka/train_seperated.py +++ b/ginka/train_seperated.py @@ -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__":