Checkpoints, sparande och återstart
Kunna spara och ladda modeller och optimerartillstånd och återuppta träning.
Förkunskaper
Intuition
Att spara bara modellvikterna räcker inte för att kunna fortsätta träna. En checkpoint måste innehålla allt som träningen beror på:
| Del | Varför |
|---|---|
| Modellens vikter | självklart |
| Optimerartillstånd | Adams m och v — utan dem startar momentum om från noll |
| Schemaläggarens tillstånd | annars hoppar lärhastigheten tillbaka |
| Epok och steg | för att veta var man var |
| Slumptillstånd | för att dataordning och dropout ska fortsätta likadant |
| Scaler-tillstånd (AMP) | loss scaling-faktorn |
| Konfiguration | hyperparametrar, så att checkpointen är självbeskrivande |
Det vanligaste misstaget är att spara bara state_dict() för modellen och tro att träningen kan återupptas. Den kan fortsätta, men inte identiskt — och en omstartad Adam ger ofta en synlig hicka i förlustkurvan.
Formellt
Två helt olika syften, två format:
| Checkpoint | Export | |
|---|---|---|
| Syfte | återuppta träning | köra i produktion |
| Innehåll | allt ovan | bara vikter och graf |
| Format | .pt med state_dict | safetensors, ONNX, TorchScript |
| Storlek | stor (optimerartillstånd dominerar) | liten |
| Bindning | till kod och version | fristående |
torch.load kör pickle och kan därmed exekvera kod vid inläsning. Ladda aldrig en checkpoint från en okänd källa. Använd weights_only=True när du bara behöver vikterna, och safetensors för allt som delas — formatet kan per konstruktion inte innehålla kod.
En sparrutin som håller i praktiken:
| Regel | Varför |
|---|---|
| Skriv till en temporärfil och byt namn | avbrott mitt i skrivningen ger annars en trasig fil |
| Spara senaste och bästa separat | senaste för återstart, bästa för användning |
| Behåll de N senaste | disken tar slut snabbare än man tror |
| Spara på tid, inte bara per epok | långa epoker ger för glesa sparpunkter |
| Logga git-commit i checkpointen | annars vet ingen vilken kod som skapade den |
strict=False vid laddning är bekvämt och farligt: det tillåter att lager saknas eller är för många, vilket är precis vad man vill vid finjustering av en delmängd — men också vad som döljer att du laddat fel modell. Kontrollera alltid returvärdet:
saknade, oväntade = modell.load_state_dict(sd, strict=False)
Är listorna icke-tomma när de inte borde vara det har du just tyst laddat en modell med slumpmässiga vikter i några lager.
Kod
import os, random, subprocess, tempfile
from pathlib import Path
import numpy as np, torch
def git_commit():
try:
return subprocess.check_output(["git", "rev-parse", "--short", "HEAD"], text=True).strip()
except Exception:
return "okänd"
def spara(sokvag: Path, modell, opt, schema, scaler, epok, steg, basta, config):
ck = {
"modell": modell.state_dict(),
"optimerare": opt.state_dict(),
"schema": schema.state_dict() if schema else None,
"scaler": scaler.state_dict() if scaler else None,
"epok": epok, "steg": steg, "basta": basta,
"config": config, "git_commit": git_commit(),
"rng": {"python": random.getstate(),
"numpy": np.random.get_state(),
"torch": torch.get_rng_state(),
"cuda": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None},
}
sokvag.parent.mkdir(parents=True, exist_ok=True)
fd, tmp = tempfile.mkstemp(dir=sokvag.parent, suffix=".tmp")
os.close(fd)
torch.save(ck, tmp)
os.replace(tmp, sokvag) # atomiskt — aldrig en halvskriven checkpoint
def ladda(sokvag: Path, modell, opt=None, schema=None, scaler=None):
ck = torch.load(sokvag, map_location="cpu", weights_only=False)
saknade, ovantade = modell.load_state_dict(ck["modell"], strict=True)
if opt:
opt.load_state_dict(ck["optimerare"])
if schema and ck.get("schema"):
schema.load_state_dict(ck["schema"])
if scaler and ck.get("scaler"):
scaler.load_state_dict(ck["scaler"])
r = ck["rng"]
random.setstate(r["python"]); np.random.set_state(r["numpy"])
torch.set_rng_state(r["torch"])
if r["cuda"] and torch.cuda.is_available():
torch.cuda.set_rng_state_all(r["cuda"])
return ck["epok"], ck["steg"], ck["basta"], ck["git_commit"]
def stada(katalog: Path, behall=3):
filer = sorted(katalog.glob("steg-*.pt"), key=lambda p: p.stat().st_mtime)
for f in filer[:-behall]:
f.unlink()
# Export för produktion — safetensors kan inte innehålla kod
from safetensors.torch import save_file, load_file
save_file({k: v.contiguous() for k, v in modell.state_dict().items()}, "modell.safetensors")
modell.load_state_dict(load_file("modell.safetensors"))
Testa återstarten innan du behöver den. Spara vid steg 100, ladda, kör till steg 110, och jämför förlusten med en obruten körning. Skiljer de sig har du glömt något i checkpointen — och det märker du hellre nu än efter att ett dygns träning avbrutits.
Behärskning innebär
- Sparar och laddar modell och optimerartillstånd
- Återupptar träning exakt
- Vet skillnaden mellan checkpoint och exportformat
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- safetensors (Apache-2.0) — Apache-2.0
- Hugging Face — dokumentation (Apache-2.0) — Apache-2.0