Save a checkpoint with model state, optimizer state, and metadata.
(model, optimizer, global_step, config, is_main, data_config)
| 147 | |
| 148 | |
| 149 | def save_checkpoint(model, optimizer, global_step, config, is_main, data_config): |
| 150 | """Save a checkpoint with model state, optimizer state, and metadata.""" |
| 151 | if not is_main: |
| 152 | return |
| 153 | |
| 154 | # Only save if it's time to save or if it's the final step |
| 155 | if (global_step % config.save_interval == 0 and global_step > 0) or global_step == config.num_train_steps - 1: |
| 156 | # Create temporary directory for atomic checkpoint saving |
| 157 | final_ckpt_dir = config.checkpoint_dir / f"{global_step}" |
| 158 | tmp_ckpt_dir = config.checkpoint_dir / f"tmp_{global_step}" |
| 159 | |
| 160 | # Remove any existing temp directory and create new one |
| 161 | if tmp_ckpt_dir.exists(): |
| 162 | shutil.rmtree(tmp_ckpt_dir) |
| 163 | tmp_ckpt_dir.mkdir(parents=True, exist_ok=True) |
| 164 | |
| 165 | # Save model state using safetensors (handle shared tensors) |
| 166 | model_to_save = model.module if isinstance(model, torch.nn.parallel.DistributedDataParallel) else model |
| 167 | safetensors.torch.save_model(model_to_save, tmp_ckpt_dir / "model.safetensors") |
| 168 | |
| 169 | # Save optimizer state using PyTorch format |
| 170 | torch.save(optimizer.state_dict(), tmp_ckpt_dir / "optimizer.pt") |
| 171 | |
| 172 | # Save training metadata (avoid saving full config to prevent JAX/Flax compatibility issues) |
| 173 | metadata = { |
| 174 | "global_step": global_step, |
| 175 | "config": dataclasses.asdict(config), |
| 176 | "timestamp": time.time(), |
| 177 | } |
| 178 | torch.save(metadata, tmp_ckpt_dir / "metadata.pt") |
| 179 | |
| 180 | # save norm stats |
| 181 | norm_stats = data_config.norm_stats |
| 182 | if norm_stats is not None and data_config.asset_id is not None: |
| 183 | _normalize.save(tmp_ckpt_dir / "assets" / data_config.asset_id, norm_stats) |
| 184 | |
| 185 | # Atomically move temp directory to final location |
| 186 | if final_ckpt_dir.exists(): |
| 187 | shutil.rmtree(final_ckpt_dir) |
| 188 | tmp_ckpt_dir.rename(final_ckpt_dir) |
| 189 | |
| 190 | logging.info(f"Saved checkpoint at step {global_step} -> {final_ckpt_dir}") |
| 191 | |
| 192 | # Log checkpoint to wandb |
| 193 | if config.wandb_enabled: |
| 194 | wandb.log({"checkpoint_step": global_step}, step=global_step) |
| 195 | |
| 196 | |
| 197 | def load_checkpoint(model, optimizer, checkpoint_dir, device): |