Detect and repair catastrophic forgetting in sequential LLM fine-tuning.
Project description
avr-cl
Your fine-tune silently broke your model. avr-cl checks if it broke, and fixes it.
The forgetting-prevention layer for LLM post-training. After each fine-tuning stage, avr-cl detects if the model forgot prior tasks and repairs the damage in weight space — no replay buffer, no old training data, one LoRA snapshot in memory.
Left: naive sequential SFT — prior tasks collapse (GSM8K 66% → 9%). Right: avr-cl — prior tasks survive (GSM8K 62% → 47%). Same model, same data, same LoRA.
The problem
Every continual learning method in LLMs is just absorption with weight updates. Absorb new data → update weights → try not to forget. EWC, replay, SLAO — all variations of the same thing: absorb → update → hope.
But that's not learning. When a human learns something new, they don't just absorb it and move on. They absorb it, then verify it against what they already know — does this break something I learned before? If it does, they repair the conflict. Then they check again. Only when the old knowledge still holds do they call it "learned."
Real learning is: absorb → verify → repair → call it learned.
Current CL does: absorb → call it learned.
The verify and repair steps are missing. That's why fine-tuning silently breaks models — nobody checked.
avr-cl builds those missing steps:
- Anchor — snapshot the model before learning
- Verify — check old tasks after absorption. Did anything break?
- Repair — fix what broke. Closed-form weight interpolation, no gradients, no old data
Results
Qwen3-1.7B, 4-task stream: GSM8K → MATH(algebra) → AQuA-RAT → SVAMP. LoRA r=128, 5000 examples/task.
| Naive SFT | avr-cl | |
|---|---|---|
| BWT | −0.453 | −0.078 |
| GSM8K after all 4 tasks | 9% | 47% |
| ACC | 0.220 | 0.529 |
| Repair steps | — | 29 |
5.8× less forgetting. The repair loop fired 29 times across 3 tasks — each time detecting PPL drift on prior tasks and pulling weights back until the drift resolved.
Results: results/qwen3_1.7b/
Install
pip install git+https://github.com/ARYAN2302/tiny-cl.git
Use it
Option 1: Full loop — avr-cl handles everything (model loading, LoRA, SFT, verify, repair, eval):
import avr
result = avr.run(
model="Qwen/Qwen3-1.7B",
tasks=[
("task_a", train_examples, eval_examples),
("task_b", train_examples, eval_examples),
],
lora_rank=128,
)
print(f"BWT: {result['bwt']:+.3f} Repairs: {result['repairs']}")
Option 2: As a layer — keep your existing training loop (TRL, Axolotl, Unsloth), add avr-cl between stages:
import avr
from trl import SFTTrainer # your existing training
# Train on task A (your existing code)
trainer = SFTTrainer(model, train_dataset=task_a)
trainer.train()
# After training: check if the model forgot prior tasks
snapshot = avr.get_lora_state(model) # snapshot before training task B
# ... train on task B ...
drift = avr.check_drift(
current_ppls=avr.eval_ppls(model, tokenizer, prior_tasks, ...),
best_ppls=best_ppls,
completed_tasks=prior_task_names,
threshold=1.15,
)
if drift:
avr.repair(model, snapshot, alpha=0.1) # 2 lines, plugs into anything
The layer API: avr.get_lora_state(), avr.check_drift(), avr.repair(). Use them with any training framework.
Each task is a (name, train_examples, eval_examples) tuple. Each example is a (question, answer, gold) triple:
question— the input promptanswer— the full training target (reasoning + answer)gold— the short answer for scoring
The full-loop API handles: model loading, LoRA, chat templates, SFT training, PPL drift detection, weight repair, batched evaluation, R-matrix, BWT/FF/ACC.
The framework
Each phase is a separate module. Use the defaults, or swap your own:
from avr.learn import train_sft # LEARN: your training function
from avr.verify import check_drift # VERIFY: your drift detector
from avr.repair import repair # REPAIR: your repair operator
from avr.eval import evaluate # evaluation + scoring
avr/
├── model.py — load_model, chat template handling
├── learn.py — train_sft, consolidate (two-stream)
├── verify.py — compute_ppl, check_drift
├── repair.py — get/set/reset LoRA state, weight interpolation
├── eval.py — batched generation, scoring, R-matrix
├── run.py — orchestrator: wires LEARN → VERIFY → REPAIR
└── cli.py — avr train config.yaml
How it works
For each task in a sequential stream:
LEARN → fine-tune on the new task (SFT, any LoRA config)
VERIFY → compute PPL on prior tasks' data
if PPL_now / PPL_best > 1.15 → drift detected
REPAIR → θ ← (1−α)·θ + α·θ_snapshot (closed-form, no gradients)
repeat until drift resolves or 10-step cap
SNAPSHOT → save current LoRA state for next task's repair target
No replay buffer. No old training data. One LoRA snapshot in memory.
Why not just...?
| Approach | Problem |
|---|---|
| Replay buffers | Need to store old training data. Privacy, plumbing, maintenance. |
| Retrain from scratch | Expensive. Days of compute for every new task. |
| LoRA adapter switching | Needs a router at inference. Multiple adapters in memory. |
| EWC | Fisher penalty slows forgetting but doesn't prevent it. Barely better than Naive (see LFM2.5 results above). |
| mergekit | Merges N separately-trained models after the fact. avr-cl prevents the damage during one training run. Complementary, not competing. |
| TRL / Axolotl / Unsloth | Great training frameworks — but none of them check if your model forgot between stages. avr-cl plugs into them as a layer. |
| Letta / memory layers | Handles the retrieval/memory layer for agents. avr-cl handles the weight layer. Use both. |
avr-cl needs zero old data, zero gradients at repair time, one LoRA snapshot, and it knows when the model forgot. It's not a replacement for your training framework — it's the layer that watches for forgetting between stages.
Reproduce the headline result
# On Kaggle T4 or any GPU with 16GB+ VRAM
python scripts/avr_cl_math_qwen3_1.7b.py
This script is standalone and reproduces the Qwen3-1.7B math stream results shown above. The avr.run() API implements the same logic in a pip-installable package.
Limitations
- Validated on 1.7B. Smaller and larger models are next.
- SFT only. DPO/GRPO on the roadmap.
License
MIT.
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file avr_cl-0.1.0.tar.gz.
File metadata
- Download URL: avr_cl-0.1.0.tar.gz
- Upload date:
- Size: 15.3 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.12.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
2fa131f4d89a0ddefa3368a79687e673391ac1bbee1bca399e19970f0e849874
|
|
| MD5 |
4e9ba5ff3a33b43216101cd41cb053b5
|
|
| BLAKE2b-256 |
c66b04e4a7ff6939aba47ccc0c465d4a66c801f7a8026f160976400bfdef682b
|
File details
Details for the file avr_cl-0.1.0-py3-none-any.whl.
File metadata
- Download URL: avr_cl-0.1.0-py3-none-any.whl
- Upload date:
- Size: 15.1 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.12.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
7c6b0915399d8f8ede541c6a6b466ceb745640f4b17688abc47bfb63cd7d5692
|
|
| MD5 |
7f6e845d68132aa48829c8c0b52c4d1a
|
|
| BLAKE2b-256 |
28b8f9c1cb3fdae0f7d93e181239281e0c485757ff188b9bc91182b9f4ef9f1f
|