(self, ckpt_path: Path)
| 660 | logger.info(f"Current step: {self.step}") |
| 661 | |
| 662 | def load_model(self, ckpt_path: Path): |
| 663 | state_dict = torch.load(ckpt_path) |
| 664 | self.vae.load_state_dict(state_dict['vae']) |
| 665 | if self.optimizer is not None: |
| 666 | self.optimizer.load_state_dict(state_dict['optimizer']) |
| 667 | self.step = state_dict['step'] |
| 668 | |
| 669 | # 加载EMA模型 |
| 670 | if self.use_ema and 'ema_models' in state_dict: |
| 671 | for name, ema_state in state_dict['ema_models'].items(): |
| 672 | if name in self.ema_models: |
| 673 | self.ema_models[name].load_state_dict(ema_state) |
| 674 | logger.info(f"Loaded EMA model: {name}") |
| 675 | |
| 676 | logger.info(f"Loaded CKPT model & optimizer from {ckpt_path}") |
| 677 | logger.info(f"CKPT step: {self.step}") |
| 678 | |
| 679 | def update_ema_models(self): |
| 680 | """更新EMA模型""" |
no test coverage detected