(rate, params)
| 257 | |
| 258 | def save(self): |
| 259 | def save_checkpoint(rate, params): |
| 260 | state_dict = self.mp_trainer.master_params_to_state_dict(params) |
| 261 | # if dist.get_rank() == 0: |
| 262 | logger.log(f"saving model {rate}...") |
| 263 | if not rate: |
| 264 | filename = f"model{(self.step+self.resume_step):06d}.pt" |
| 265 | else: |
| 266 | filename = f"ema_{rate}_{(self.step+self.resume_step):06d}.pt" |
| 267 | with bf.BlobFile(bf.join(get_blob_logdir(), filename), "wb") as f: |
| 268 | th.save(state_dict, f) |
| 269 | |
| 270 | # save_checkpoint(0, self.mp_trainer.master_params) |
| 271 | for rate, params in zip(self.ema_rate, self.ema_params): |
nothing calls this directly
no test coverage detected