Kunskapsdestillation
Kunna träna en liten modell på en stor modells utdata och mäta kvalitetsförlust.
Förkunskaper
Intuition
Destillation tränar en liten modell (eleven) att härma en stor (läraren).
Det förvånande är att eleven ofta blir bättre än om den tränats direkt på samma märkta data. Förklaringen ligger i vad lärarens utdata innehåller.
En hård etikett säger: «det här är en 7».
Lärarens mjuka fördelning säger: «det här är en 7 (0,82), men det liknar en 1 (0,11) och lite en 9 (0,04)».
Den andra bär information om hur klasserna förhåller sig till varandra — Hinton kallade det «dark knowledge». En bild som är svår att skilja från en 1 lär eleven något som den hårda etiketten inte kan.
| Hård etikett | Lärarens fördelning | |
|---|---|---|
| Information per exempel | log₂(10) ≈ 3,3 bitar | mycket mer |
| Säger något om osäkerhet | nej | ja |
| Kräver märkt data | ja | nej — omärkt räcker |
Den sista raden är ofta den praktiskt viktigaste: destillation fungerar på omärkt data, eftersom läraren producerar målen.
Formellt
Förlusten kombinerar två termer:
| Del | Betyder |
|---|---|
| temperatur, typiskt 2–10 — mjukar upp båda fördelningarna | |
| kompenserar att gradienterna skalas med | |
| vikt mellan lärare och hårda etiketter, ofta 0,5–0,9 |
Temperaturen är hela poängen. Vid är lärarens fördelning ofta nästan one-hot och bär lite mer information än etiketten. Vid framträder de små sannolikheterna, och det är där informationen finns.
-faktorn missas ofta. Gradienten av KL-termen skalas med , så utan kompensationen blir lärartermen försumbar vid hög temperatur.
Tre former av destillation:
| Form | Vad eleven matchar | Kommentar |
|---|---|---|
| Response-based | utdatafördelningen | enklast, fungerar bra |
| Feature-based | mellanliggande representationer | kräver projektion mellan olika dimensioner |
| Sekvensnivå | lärarens genererade sekvenser | vanligast för språkmodeller |
För språkmodeller är den tredje dominerande: låt läraren generera svar på många prompter och finjustera eleven på dem. Det är i praktiken vad de flesta små «instruct»-modeller tränats med.
Vad man realistiskt kan vänta sig:
| Kompression | Typisk kvalitetsförlust |
|---|---|
| 2× | nästan ingen |
| 4× | liten |
| 10× | märkbar, ofta acceptabel |
| 50×+ | stor |
Två juridiska och praktiska förbehåll:
- Villkoren. Många kommersiella API:er förbjuder uttryckligen att utdata används för att träna konkurrerande modeller. Läs dem.
- Fel ärvs. Eleven lär sig lärarens misstag, skevheter och hallucinationer — inklusive dem ingen upptäckt. Utvärdera eleven självständigt, inte bara mot läraren.
Kombinera med annat. Destillation, kvantisering och beskärning är ortogonala: destillera till en mindre arkitektur, kvantisera den, och beskär om det behövs. Tillsammans ger de betydligt mer än var för sig.
Kod
import torch, torch.nn as nn, torch.nn.functional as F
def destillationsforlust(elev_logits, larare_logits, y=None, T=4.0, alfa=0.7):
mjuk = F.kl_div(F.log_softmax(elev_logits / T, dim=-1),
F.softmax(larare_logits / T, dim=-1),
reduction="batchmean") * (T ** 2) # T² kompenserar gradientskalningen
if y is None:
return mjuk # ren destillation, omärkt data
hard = F.cross_entropy(elev_logits, y)
return alfa * mjuk + (1 - alfa) * hard
def trana_elev(elev, larare, dataloader, opt, T=4.0, alfa=0.7, epoker=3):
larare.eval()
for _ in range(epoker):
for x, y in dataloader:
with torch.no_grad():
lt = larare(x)
loss = destillationsforlust(elev(x), lt, y, T=T, alfa=alfa)
opt.zero_grad(); loss.backward(); opt.step()
# Varför temperaturen spelar roll — samma logits, olika informationsinnehåll
logits = torch.tensor([[8.0, 2.0, 1.5, 0.2, -1.0]])
for T in (1.0, 2.0, 4.0, 8.0):
p = F.softmax(logits / T, dim=-1)[0]
entropi = float(-(p * p.clamp_min(1e-12).log()).sum())
print(f"T={T}: {[round(float(v), 4) for v in p]} entropi {entropi:.3f}")
# T=1.0: [0.9955, 0.0025, 0.0015, 0.0004, 0.0001] entropi 0.033
# T=4.0: [0.5996, 0.1338, 0.1181, 0.0853, 0.0632] entropi 1.213
# ↑ vid T=1 är fördelningen nästan one-hot; vid T=4 syns klassrelationerna
# Sekvensnivå för språkmodeller: låt läraren generera träningsdatan
def generera_destillationsdata(larare, tokenizer, prompter, max_tokens=512):
par = []
for p in prompter:
ids = tokenizer(p, return_tensors="pt")
with torch.no_grad():
ut = larare.generate(**ids, max_new_tokens=max_tokens, do_sample=False)
svar = tokenizer.decode(ut[0, ids["input_ids"].shape[1]:], skip_special_tokens=True)
par.append({"instruktion": p, "svar": svar})
return par
# Mät kvalitetsförlust MOT kostnadsvinst — båda sidorna behövs
def utvardera_destillation(larare, elev, testset, kor):
resultat = {}
for namn, m in (("lärare", larare), ("elev", elev)):
ratt = sum(kor(m, f["fraga"]) == f["facit"] for f in testset) / len(testset)
params = sum(p.numel() for p in m.parameters())
resultat[namn] = {"traffsakerhet": round(ratt, 4), "parametrar_m": round(params / 1e6, 1)}
l, e = resultat["lärare"], resultat["elev"]
resultat["kompression"] = round(l["parametrar_m"] / e["parametrar_m"], 1)
resultat["kvalitetsförlust_pe"] = round(
(l["traffsakerhet"] - e["traffsakerhet"]) * 100, 2)
return resultat
# {'lärare': {...}, 'elev': {...}, 'kompression': 8.0, 'kvalitetsförlust_pe': 1.8}
Utvärdera eleven självständigt, inte bara mot läraren. En elev som perfekt härmar en lärare med ett systematiskt fel har lärt sig felet lika bra som allt annat.
Behärskning innebär
- Tränar en elevmodell på en lärares utdata
- Förklarar varför mjuka mål bär mer information
- Mäter kvalitetsförlusten mot kostnadsvinsten
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Distilling the Knowledge in a Neural Network — arXiv (öppen åtkomst; licens per artikel)
- arXiv — DistilBERT, a distilled version of BERT — arXiv (öppen åtkomst; licens per artikel)
- Hugging Face — dokumentation (Apache-2.0) — Apache-2.0