Hoppa till innehållet
AI-grafen
E· Universitetdeep-learning· ca 60 min· utvecklande· verifierad 2026-09-20

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å:

DelVarför
Modellens viktersjälvklart
OptimerartillståndAdams m och v — utan dem startar momentum om från noll
Schemaläggarens tillståndannars hoppar lärhastigheten tillbaka
Epok och stegför att veta var man var
Slumptillståndför att dataordning och dropout ska fortsätta likadant
Scaler-tillstånd (AMP)loss scaling-faktorn
Konfigurationhyperparametrar, 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:

CheckpointExport
Syfteåteruppta träningköra i produktion
Innehållallt ovanbara vikter och graf
Format.pt med state_dictsafetensors, ONNX, TorchScript
Storlekstor (optimerartillstånd dominerar)liten
Bindningtill kod och versionfristå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:

RegelVarför
Skriv till en temporärfil och byt namnavbrott mitt i skrivningen ger annars en trasig fil
Spara senaste och bästa separatsenaste för återstart, bästa för användning
Behåll de N senastedisken tar slut snabbare än man tror
Spara på tid, inte bara per epoklånga epoker ger för glesa sparpunkter
Logga git-commit i checkpointenannars 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

Alla källor och licenser