Unlearn — remove a capability¶
Remove a capability from a finished model using optimizer steps on forget/retain data.
flowchart LR
subgraph RECOVER[Recover — training-free]
A1[Finished model] --> B1[Weight-patching function<br/>e.g. RESTA · C-Θ · LoX]
B1 --> C1[Patched weights<br/>no training run]
end
subgraph UNLEARN[Unlearn — forget-set training]
A2[Finished model] --> B2[Training loop<br/>optimizer steps on forget + retain]
A3[Forget set] --> B2
A4[Retain set] --> B2
B2 --> C2[Unlearned model<br/>new checkpoint]
end
Input contract¶
A finished model + a forget set + a retain set. The forget set defines what to remove; the retain set preserves everything else. Unlike Recover, Unlearn trains: it runs optimizer steps.
Quick example¶
from safetune.runner import unlearn
trainer = unlearn.RMUTrainer(model)
trainer.unlearn(forget=forget_batches, retain=retain_batches)
Data format:
forgetandretainare iterables of tokenized batches — dicts withinput_ids,attention_mask, andlabels. For a quick start,unlearn.load_unlearn_data(model_id)returns a(forget, retain)pair. FLAT and SimDPO train on refusal/harmful preference pairs; pass rawforgetbatches and they build the pairs for you (see their pages).Sequence length:
load_unlearn_datatokenizes withmax_len=Noneby default: the length is sized from the chat template asmax(256, longest templated prompt among the first rows + 256), capped at 2048, so a long system preamble (Tiny Aya's template adds about 366 tokens) does not fill the sequence. An explicitmax_lenis used exactly; if it leaves no row with a supervised token the loader raises (naming the templated prompt length), and it warns with a count when only some rows are fully masked.
Precision¶
Every unlearn trainer takes upcast: bool = True: a model loaded in fp16 or
bf16 is cast to fp32 before unlearn() runs (with a warning). fp16 overflows,
and bf16 has 8 mantissa bits, so at lr=1e-5 most updates round away (about 80%
lost on Tiny Aya) and unlearning silently does little. upcast=False keeps the
low-precision weights, e.g. to save memory. upcast_fp16 is a deprecated alias
of upcast. Results therefore depend on the dtype you train in.
Catalog of alternatives¶
| Method | Mechanism | Guide |
|---|---|---|
RMU |
representation misdirection — steers harmful hidden states to random anchors | RMU |
NPO |
negative preference optimization — sigmoid-bounded NLL on forget set | NPO |
GradientAscent / GradDiff |
gradient ascent on forget set (+ KL preservation on retain) | Gradient Ascent |
FLATTrainer |
f-divergence variational bound, no reference model needed | FLAT |
SimDPOTrainer |
SimDPO-style unlearning, reference-free | SimDPO |