(self, step, text_encoder, save_samples=True, path=None)
| 486 | |
| 487 | @torch.no_grad() |
| 488 | def checkpoint(self, step, text_encoder, save_samples=True, path=None): |
| 489 | print("Saving checkpoint for step %d..." % step) |
| 490 | with torch.autocast("cuda"): |
| 491 | if path is None: |
| 492 | checkpoints_path = f"{self.output_dir}/checkpoints" |
| 493 | os.makedirs(checkpoints_path, exist_ok=True) |
| 494 | |
| 495 | unwrapped = self.accelerator.unwrap_model(text_encoder) |
| 496 | |
| 497 | # Save a checkpoint |
| 498 | learned_embeds = unwrapped.get_input_embeddings().weight[ |
| 499 | self.placeholder_token_id |
| 500 | ] |
| 501 | learned_embeds_dict = { |
| 502 | self.placeholder_token: learned_embeds.detach().cpu() |
| 503 | } |
| 504 | |
| 505 | filename = "%s_%d.bin" % (slugify(self.placeholder_token), step) |
| 506 | if path is not None: |
| 507 | torch.save(learned_embeds_dict, path) |
| 508 | else: |
| 509 | torch.save(learned_embeds_dict, f"{checkpoints_path}/{filename}") |
| 510 | torch.save(learned_embeds_dict, f"{checkpoints_path}/last.bin") |
| 511 | del unwrapped |
| 512 | del learned_embeds |
| 513 | |
| 514 | @torch.no_grad() |
| 515 | def save_samples( |
no test coverage detected