MCPcopy Create free account
hub / github.com/Physical-Intelligence/openpi / save_checkpoint

Function save_checkpoint

scripts/train_pytorch.py:149–194  ·  view source on GitHub ↗

Save a checkpoint with model state, optimizer state, and metadata.

(model, optimizer, global_step, config, is_main, data_config)

Source from the content-addressed store, hash-verified

147
148
149def 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
197def load_checkpoint(model, optimizer, checkpoint_dir, device):

Callers 1

train_loopFunction · 0.85

Calls 1

saveMethod · 0.80

Tested by

no test coverage detected