runwhere CLI
Checkpoint 最佳实践
在 RunWhere 训练任务中保存、恢复和控制 checkpoint 成本。
Checkpoint 决定了训练中断后能从哪里继续。RunWhere 可以为训练任务挂载 checkpoint 路径,并在失败重试或手动恢复时保留这部分数据;但平台不会替你的训练代码调用 save 或 load。
一句话原则:在 YAML 里声明 checkpoint 路径,在训练代码里定期保存,并在启动时主动寻找最近一次 checkpoint。
平台提供什么
| 能力 | 说明 |
|---|---|
| 持久路径 | 在 YAML 中声明 storage.checkpoint.path 后,训练容器内会出现对应 checkpoint 目录。 |
| 快速暂存路径 | 平台会在 checkpoint 目录下提供 fast 子目录,适合单次运行内高频写入。 |
| 持久恢复路径 | 平台会在 checkpoint 目录下提供 durable 子目录,适合保存关键 checkpoint。 |
| 失败重试 | Training 声明 checkpoint 后,平台会开启训练失败后的重试策略。 |
| 手动恢复 | 停止或失败后的训练可用 runwhere resume <name> 恢复。 |
平台不做这些事:
- 不自动调用你的框架保存 checkpoint。
- 不解析 checkpoint 文件内容。
- 不决定保存频率。
- 不自动清理历史 checkpoint 文件。
YAML 配置
kind: Training
version: v1
job:
name: llama3-lora
environment:
image: pytorch/pytorch:2.4.1-cuda12.1-cudnn9-runtime
command: python train.py
resources:
gpuNum: 8
gpuType: a100-80g
storage:
checkpoint:
path: /ckpts任务启动后,训练代码可以使用:
| 路径 | 用途 |
|---|---|
/ckpts/fast | 高频、临时 checkpoint。适合减少训练阻塞,但不应作为唯一恢复来源。 |
/ckpts/durable | 关键 checkpoint。适合故障恢复、手动恢复和长期保留。 |
如果你只想用一个路径,优先写 /ckpts/durable。
最小训练代码
保存:
import torch
CKPT_DIR = "/ckpts/durable"
for step in range(start_step, total_steps):
loss = train_step(model, batch)
if step % save_interval == 0:
torch.save(
{
"step": step,
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
},
f"{CKPT_DIR}/step_{step}.pt",
)启动时恢复:
import glob
import os
import torch
CKPT_DIR = "/ckpts/durable"
def step_of(path: str) -> int:
name = os.path.basename(path)
return int(name.removeprefix("step_").removesuffix(".pt"))
ckpts = sorted(glob.glob(f"{CKPT_DIR}/step_*.pt"), key=step_of)
if ckpts:
ckpt = torch.load(ckpts[-1], map_location="cpu")
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
start_step = ckpt["step"] + 1
print(f"Resumed from step {start_step}")
else:
start_step = 0做到这两件事后,平台重启任务时,训练代码就能从最近一次 checkpoint 继续。
两级保存策略
长时间训练建议把“高频暂存”和“关键持久化”分开:
FAST_DIR = "/ckpts/fast"
DURABLE_DIR = "/ckpts/durable"
for step in range(start_step, total_steps):
loss = train_step(model, batch)
if step % 100 == 0:
torch.save(state, f"{FAST_DIR}/latest.pt")
if step % 1000 == 0:
torch.save(state, f"{DURABLE_DIR}/step_{step}.pt")恢复时优先检查 fast,没有再检查 durable:
import glob
import os
def find_latest_ckpt():
fast = "/ckpts/fast/latest.pt"
if os.path.exists(fast):
return fast
durable = sorted(glob.glob("/ckpts/durable/step_*.pt"))
return durable[-1] if durable else Nonefast 适合减少频繁保存的开销,durable 才是跨重试、跨恢复时最可靠的落点。
框架接入
HuggingFace Trainer
from transformers import Trainer, TrainingArguments
args = TrainingArguments(
output_dir="/ckpts/durable",
save_strategy="steps",
save_steps=500,
save_total_limit=3,
resume_from_checkpoint=True,
)
trainer = Trainer(model=model, args=args, ...)
trainer.train()DeepSpeed
{
"checkpoint": {
"tag_latest": true,
"save_on_each_node": false
}
}model.save_checkpoint("/ckpts/durable", tag=f"step_{step}")
model.load_checkpoint("/ckpts/durable")PyTorch Distributed Checkpoint
FSDP 或 DDP 场景下,可以用 PyTorch Distributed Checkpoint 让多卡并行保存 shard:
import torch.distributed.checkpoint as dcp
dcp.save(
{"model": model, "optimizer": optimizer},
checkpoint_id=f"/ckpts/durable/step_{step}",
)
dcp.load(
{"model": model, "optimizer": optimizer},
checkpoint_id=f"/ckpts/durable/step_{step}",
)保存频率建议
| 模型规模 | 常见 checkpoint 大小 | 建议保存间隔 |
|---|---|---|
| 小于 1B | 小于 2 GB | 每 100 step 左右 |
| 1B 到 7B | 2 到 14 GB | 每 500 step 左右 |
| 7B 到 70B | 14 到 140 GB | 每 1000 step 左右 |
| 大于 70B | 大于 140 GB | 每 2000 step 左右,并考虑异步保存或分布式 checkpoint |
保存太频繁会影响训练吞吐;保存太稀疏会增加故障后的回退步数。实际间隔应按训练速度、模型大小和可接受的回退范围调整。
停止、恢复与删除
runwhere stop job <name>
runwhere resume <name>
runwhere resume <name> --from-step 12000
runwhere delete job <name>stop 会停止训练并释放 GPU,checkpoint 数据保留。resume 会重新提交训练,让你的启动逻辑从 checkpoint 继续。delete 删除的是作业记录,不应该当作正常停机方式使用。
常见问题
| 问题 | 回答 |
|---|---|
| 平台会自动帮我保存 checkpoint 吗? | 不会。平台提供路径和恢复条件,保存与加载由训练代码或训练框架负责。 |
| 不写 checkpoint 代码会怎样? | 任务失败后可以重启,但训练通常会从头开始。 |
| 用 HuggingFace Trainer 还需要手写保存吗? | 通常不需要,把 output_dir 指向 /ckpts/durable 并启用恢复即可。 |
| 多卡训练怎么保存? | DDP 常见做法是 rank 0 保存;FSDP 或大模型训练建议使用 PyTorch Distributed Checkpoint 或框架自带分片 checkpoint。 |
| checkpoint 会一直堆积吗? | 平台不会替你清理历史文件。建议用 save_total_limit、save_top_k 或自己的清理逻辑控制保留数量。 |