Part 4: Methods
Full fine-tuning vs parameter-efficient tuning
Full fine-tuning updates every weight in the model. For TinyLM that is 1,071,872 numbers; for a 7-billion-parameter model it is 7 billion, plus the optimizer's bookkeeping. AdamW keeps two extra numbers per trained weight, so training memory is several times the model size, and every fine-tuned version is a full copy of the model on disk.
Parameter-efficient fine-tuning (PEFT) freezes the original weights and trains a small number of new ones. The most widely used method is LoRA (Low-Rank Adaptation, Hu et al., 2021). The idea: the change a fine-tune makes to a big weight matrix W (size out by in) can be approximated by the product of two thin matrices, B (out by r) times A (r by in), where the rank r is small (4, 8, 16). The layer computes
y = W x + (alpha / r) * B A x
with W frozen and only A and B trained. B starts at zero, so a fresh adapter changes nothing, and training moves it only as far as the data pushes it. alpha is a scaling knob: the update is multiplied by alpha / r, so if you change r you can keep alpha / r fixed and the learning rate still behaves the same. The trained A and B are the adapter, a small file you can load onto the base model, swap out, or fold into W (merge) after training. The LoRA paper reports cutting trainable parameters by 10,000 times and GPU memory by 3 times on GPT-3 175B, with quality on par with full fine-tuning on their benchmarks.
examples/m09_lora.py
"""LoRA from scratch for TinyGPT: a frozen linear layer plus a trainable low-rank update B @ A."""
from __future__ import annotations
import math
from pathlib import Path
import torch
import torch.nn as nn
from examples.m09_setup import MODELS
from supportdesk.tinylm import TinyGPT, load
DEFAULT_TARGETS = ("qkv", "proj", "mlp.0", "mlp.2") # every linear layer inside the transformer blocks
class LoRALinear(nn.Module):
"""y = W x + b + (alpha / r) * B A x, with W and b frozen and only A and B trained."""
def __init__(self, base: nn.Linear, r: int = 8, alpha: float = 16.0) -> None:
super().__init__()
self.base = base
self.r, self.alpha = r, alpha
self.scaling = alpha / r
self.lora_A = nn.Parameter(torch.empty(r, base.in_features))
self.lora_B = nn.Parameter(torch.zeros(base.out_features, r)) # zero: the adapter starts as a no-op
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
for p in self.base.parameters():
p.requires_grad = False
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.base(x) + (x @ self.lora_A.T @ self.lora_B.T) * self.scaling
def add_lora(model: TinyGPT, r: int = 8, alpha: float = 16.0, targets: tuple[str, ...] = DEFAULT_TARGETS) -> TinyGPT:
"""Freeze every base weight, then wrap the target linear layers of each block with LoRA (in place)."""
for p in model.parameters():
p.requires_grad = False
for block in model.blocks:
for name in targets:
parent, attr = (block.mlp, int(name.split(".")[1])) if name.startswith("mlp.") else (block, name)
layer = parent[attr] if isinstance(attr, int) else getattr(parent, attr)
wrapped = LoRALinear(layer, r, alpha)
if isinstance(attr, int):
parent[attr] = wrapped
else:
setattr(parent, attr, wrapped)
return model
def lora_state_dict(model: nn.Module) -> dict[str, torch.Tensor]:
"""Only the adapter tensors: this is everything you need to store per task."""
return {k: v.detach().clone() for k, v in model.state_dict().items() if "lora_" in k}
def save_adapter(model: nn.Module, path: Path | str, meta: dict | None = None) -> int:
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
torch.save({"lora": lora_state_dict(model), "meta": meta or {}}, path)
return path.stat().st_size
def load_adapter(model: nn.Module, path_or_state) -> dict:
"""Copy adapter tensors into a model that already has LoRA layers of the same shape."""
state = torch.load(path_or_state, weights_only=True) if isinstance(path_or_state, (str, Path)) else path_or_state
missing = model.load_state_dict(state["lora"], strict=False)
assert not missing.unexpected_keys, missing.unexpected_keys
return state.get("meta", {})
def with_adapter(name: str) -> TinyGPT:
"""A fresh base model with the saved adapter models/<name>/adapter.pt attached (rank read from its meta)."""
model, _ = load()
meta = torch.load(MODELS / name / "adapter.pt", weights_only=True)["meta"]
add_lora(model, r=meta["r"], alpha=meta["alpha"])
load_adapter(model, MODELS / name / "adapter.pt")
model.eval()
return model
@torch.no_grad()
def merge_lora(model: TinyGPT) -> TinyGPT:
"""Fold B A into W so inference costs exactly what the base model costs (in place)."""
for block in model.blocks:
for parent in (block, block.mlp):
for name, child in list(parent.named_children()):
if isinstance(child, LoRALinear):
child.base.weight += (child.lora_B @ child.lora_A) * child.scaling
if isinstance(parent, nn.Sequential):
parent[int(name)] = child.base
else:
setattr(parent, name, child.base)
return model
if __name__ == "__main__":
from examples.m09_setup import count_trainable
torch.manual_seed(0)
base, tok = load()
print(f"full fine-tuning trains {count_trainable(base):,} parameters")
for r in (1, 2, 4, 8, 16):
m, _ = load()
add_lora(m, r=r, alpha=2 * r)
n = count_trainable(m)
print(f"LoRA r={r:<2d} trains {n:>7,} parameters ({n / base.num_parameters():.2%} of the model)")
m, _ = load()
add_lora(m, r=8, alpha=16)
x = torch.tensor([tok.encode("Customer (Ana): Hi, how do I cancel my Team subscription?").ids])
with torch.no_grad():
same = torch.allclose(base(x), m(x), atol=1e-6)
print("fresh adapter changes nothing (B starts at zero):", same)
print("wrapped layer:", m.blocks[0].qkv.__class__.__name__, "base weight trainable:",
m.blocks[0].qkv.base.weight.requires_grad, "| lora_A", tuple(m.blocks[0].qkv.lora_A.shape),
"lora_B", tuple(m.blocks[0].qkv.lora_B.shape))
Code explained
- In simple words: wrap each frozen linear layer with a small trainable side path, and save only that side path.
- What happens:
LoRALinear: holds the original layer asbaseand freezes it.lora_A(r by in) gets a small random start;lora_B(out by r) starts at zero.forwardreturns the frozen output plus the scaled low-rank update.add_lora: freezes every parameter in the model (embeddings and layer norms included), then replaces the four target layers in every block withLoRALinearwrappers. After this, only A and B haverequires_grad=True, which is allsft_trainneeds to know.lora_state_dict,save_adapter,load_adapter: save and restore just the tensors whose names containlora_, plus the rank and alpha inmeta.with_adapter: load a fresh base model and attach a saved adapter by name. Later scripts use it.merge_lora: adds (alpha / r) * B A into W and removes the wrapper, so the merged model runs exactly as fast as the base model. The tests check that merged and unmerged outputs match.- The main block counts trainable parameters for full fine-tuning and ranks 1 to 16, checks that a fresh adapter leaves the model's outputs unchanged, and prints one wrapped layer's shapes. Run it with
python -m examples.m09_lora.
- Comes out:
full fine-tuning trains 1,071,872 parameters
LoRA r=1 trains 8,192 parameters (0.76% of the model)
LoRA r=2 trains 16,384 parameters (1.53% of the model)
LoRA r=4 trains 32,768 parameters (3.06% of the model)
LoRA r=8 trains 65,536 parameters (6.11% of the model)
LoRA r=16 trains 131,072 parameters (12.23% of the model)
fresh adapter changes nothing (B starts at zero): True
wrapped layer: LoRALinear base weight trainable: False | lora_A (8, 128) lora_B (384, 8)
Rank 4 trains 32,768 parameters, 3% of TinyLM. For each wrapped layer that is r times (in + out): for qkv, 4 times (128 + 384) = 2,048, and 16 layers add up to 32,768. The percentage looks large only because TinyLM is tiny. On a 7B model, rank 8 on the attention layers is typically well under 1%.
Hyperparameters that matter for small datasets
With 48 examples, three knobs dominate: the learning rate, the number of epochs (too many and the model memorizes the training tickets), and for LoRA the rank. Batch size, warmup, and weight decay matter much less at this scale. We run a small sweep.
The problem with sweeping on 48 tickets is noise. A single 13-ticket validation split cannot rank anything: one ticket is 8 points of accuracy. So the sweep uses 4-fold cross-validation: split the dev tickets into 4 groups, train on 3 and validate on the 4th, rotate, and add up. Every dev ticket is validated exactly once. The selection rule is fixed before running: highest cross-validated accuracy, ties broken by lower validation loss. The test set is never touched.
examples/m09_sweep.py
"""A small hyperparameter sweep scored by 4-fold cross-validation on the 48 dev tickets (never on test).
With so few tickets, one 13-ticket validation split is too noisy to rank configurations, so every
dev ticket takes a turn as validation. Selection rule, fixed before running: highest cross-validated
accuracy, ties broken by lower validation loss.
"""
import json
import random
from examples.m09_data import augment, dedupe
from examples.m09_lora import add_lora
from examples.m09_setup import HeldOut, MODELS, answer_loss, evaluate, load_rows, sft_train, to_pairs
from supportdesk.data import CATEGORIES
from supportdesk.tinylm import load
_, tok = load()
dev = load_rows("dev_real")
FOLDS = 4
fold_of = {}
for c in CATEGORIES: # stratified: each category spread across the folds
idx = [i for i, r in enumerate(dev) if r["category"] == c]
random.Random(0).shuffle(idx)
for k, i in enumerate(idx):
fold_of[i] = k % FOLDS
CHECK = (4, 8, 12)
def run(method, lr, rank=0, use_aug=True, synth=False):
correct = {e: 0 for e in CHECK}
losses = {e: 0.0 for e in CHECK}
seconds = 0.0
for f in range(FOLDS):
train = [r for i, r in enumerate(dev) if fold_of[i] != f]
val = [r for i, r in enumerate(dev) if fold_of[i] == f]
rows = train + (dedupe(train + [a for r in train for a in augment(r)])[len(train):] if use_aug else [])
rows += load_rows("train_synth") if synth else []
model, _ = load()
if method == "lora":
add_lora(model, r=rank, alpha=2 * rank)
val_t, val_p = [HeldOut(r) for r in val], to_pairs(tok, val)
def on_epoch(epoch):
if epoch in CHECK:
correct[epoch] += evaluate(model, tok, val_t)["correct"]
losses[epoch] += answer_loss(model, tok, val_p) * len(val) / len(dev)
seconds += sft_train(model, tok, to_pairs(tok, rows), epochs=max(CHECK), lr=lr, on_epoch=on_epoch)["seconds"]
best = max(CHECK, key=lambda e: (correct[e], -losses[e]))
data = "real+aug" + ("+synth" if synth else "") if use_aug else "real"
cells = " ".join(f"ep{e}: {correct[e]:>2}/48 loss {losses[e]:.2f}" for e in CHECK)
print(f"{method:<4} lr={lr:<6g} r={rank:<2d} {data:<14} {seconds:>4.0f}s {cells}", flush=True)
return {"method": method, "lr": lr, "rank": rank, "data": data, "best_epoch": best,
"cv_correct": correct[best], "cv_loss": round(losses[best], 3), "seconds": round(seconds, 1)}
if __name__ == "__main__":
runs = [run("lora", lr, rank) for lr in (1e-3, 3e-3) for rank in (4, 16)]
runs += [run("full", lr) for lr in (1e-4, 3e-4)]
pick = lambda rs: max(rs, key=lambda r: (r["cv_correct"], -r["cv_loss"])) # noqa: E731
best_lora = pick([r for r in runs if r["method"] == "lora"])
runs += [run("lora", best_lora["lr"], best_lora["rank"], use_aug=False),
run("lora", best_lora["lr"], best_lora["rank"], synth=True)]
chosen = {"lora": best_lora, "full": pick([r for r in runs if r["method"] == "full"])}
MODELS.mkdir(exist_ok=True)
(MODELS / "m09-sweep.json").write_text(json.dumps({"runs": runs, "chosen": chosen}, indent=1) + "\n")
print(f"total sweep training time {sum(r['seconds'] for r in runs):.0f}s; chosen:")
for k, r in chosen.items():
print(f" {k}: lr={r['lr']:g} rank={r['rank']} epochs={r['best_epoch']} cv accuracy {r['cv_correct']}/48")
Code explained
- In simple words: try 8 settings, score each on all 48 dev tickets by rotating which quarter is held out, and pick by a rule written in advance.
- What happens:
fold_ofassigns each dev ticket to one of 4 folds, stratified by category.runtrains a fresh model per fold (augmenting only that fold's training tickets, after the split) and uses theon_epochhook to score the held-out fold after epochs 4, 8, and 12, so one 12-epoch run yields three epoch counts. The first six runs cross LoRA learning rates 1e-3 and 3e-3 with ranks 4 and 16, plus full fine-tuning at 1e-4 and 3e-4 (full fine-tuning needs a much smaller learning rate because every weight moves). The last two rerun the best LoRA setting with real tickets only, and with the synthetic rows added. Results go tomodels/m09-sweep.json. Run it withpython -m examples.m09_sweep(about 10 minutes on one thread). - Comes out:
lora lr=0.001 r=4 real+aug 66s ep4: 9/48 loss 1.40 ep8: 10/48 loss 1.26 ep12: 8/48 loss 1.38
lora lr=0.001 r=16 real+aug 73s ep4: 13/48 loss 1.28 ep8: 12/48 loss 1.71 ep12: 12/48 loss 2.41
lora lr=0.003 r=4 real+aug 94s ep4: 11/48 loss 1.28 ep8: 12/48 loss 1.59 ep12: 16/48 loss 1.86
lora lr=0.003 r=16 real+aug 72s ep4: 13/48 loss 1.34 ep8: 13/48 loss 2.36 ep12: 13/48 loss 3.00
full lr=0.0001 r=0 real+aug 73s ep4: 14/48 loss 1.27 ep8: 15/48 loss 1.32 ep12: 14/48 loss 1.96
full lr=0.0003 r=0 real+aug 73s ep4: 14/48 loss 1.20 ep8: 11/48 loss 2.17 ep12: 13/48 loss 2.66
lora lr=0.003 r=4 real 18s ep4: 12/48 loss 1.64 ep8: 13/48 loss 1.29 ep12: 12/48 loss 1.26
lora lr=0.003 r=4 real+aug+synth 98s ep4: 11/48 loss 1.25 ep8: 17/48 loss 1.47 ep12: 19/48 loss 2.01
total sweep training time 566s; chosen:
lora: lr=0.003 rank=4 epochs=12 cv accuracy 16/48
full: lr=0.0001 rank=0 epochs=8 cv accuracy 15/48
Read this with the noise in mind. The best cell is 16/48 (33%), whose 95% interval is about 22% to 47%; almost every other cell is inside it. The sweep's honest message is "nothing here is clearly better than anything else", and it picks by the rule anyway: LoRA at lr 3e-3, rank 4, 12 epochs, and full fine-tuning at lr 1e-4, 8 epochs. Three patterns are worth noticing, even through the noise:
- Validation loss rises while accuracy holds or improves (LoRA 3e-3, rank 16: loss 1.34 to 3.00). The model grows more confident on the labels it gets wrong, a typical sign of overfitting on tiny data. Watch both numbers.
- Augmentation roughly quadruples training time (18s for real only against 66 to 94s with variants) for +3 tickets out of 48, which is within noise.
- Synthetic rows reach 19/48, the best number in the table, but +3 over the chosen setting is still within noise, and the plan fixed "real + augmented" as the data before the sweep. A paired comparison on more data would be needed before trusting it, and Part 3 explained what synthetic rows cost in diversity.
Training both methods and scoring the test set once
Now we train the final models on all 48 dev tickets plus their variants with the chosen settings, and score the sealed test set once.
examples/m09_train.py
"""Train the triage model twice, full fine-tuning and LoRA, with the settings the sweep chose; score on test once."""
import json
from collections import Counter
from examples.m09_lora import add_lora, save_adapter
from examples.m09_setup import (HeldOut, MODELS, count_trainable, evaluate, fmt, load_rows, seed_everything,
sft_train, to_pairs)
from supportdesk.data import load_tickets
from supportdesk.tinylm import load, save
chosen = json.loads((MODELS / "m09-sweep.json").read_text())["chosen"]
test = load_tickets("test")
_, tok = load()
dev_real = load_rows("dev_real")
pairs = to_pairs(tok, dev_real + load_rows("dev_aug")) # all 48 dev tickets plus their variants
print(f"training pairs: {len(pairs)} (48 real dev tickets + augmented variants); test tickets: {len(test)}")
results = {}
for method in ("full", "lora"):
cfg = chosen[method]
seed_everything(0)
model, _ = load()
if method == "lora":
add_lora(model, r=cfg["rank"], alpha=2 * cfg["rank"])
trainable = count_trainable(model)
info = sft_train(model, tok, pairs, epochs=cfg["best_epoch"], lr=cfg["lr"])
res = evaluate(model, tok, test)
seen = evaluate(model, tok, [HeldOut(r) for r in dev_real])
free = evaluate(model, tok, test, mode="free")
if method == "full":
save(model, tok, MODELS / "m09-full", meta={"task": "triage", **cfg})
size = (MODELS / "m09-full" / "model.pt").stat().st_size
else:
size = save_adapter(model, MODELS / "m09-lora-triage" / "adapter.pt",
{"task": "triage", "r": cfg["rank"], "alpha": 2 * cfg["rank"]})
results[method] = res["preds"]
print(f"{method:<4} lr={cfg['lr']:g} epochs={cfg['best_epoch']} trainable={trainable:>9,} "
f"train {info['seconds']:>4.1f}s file {size / 1024:>5.0f} KB")
print(f" accuracy on its own training tickets {seen['correct']}/48; on test {fmt(res)}")
print(f" free generation (no scoring): test {free['correct']}/24, valid label rate {free['valid_label_rate']:.0%}")
print(" predicted:", dict(Counter(res["preds"])))
agree = sum(a == b for a, b in zip(results["full"], results["lora"]))
print(f"full and LoRA agree on {agree}/{len(test)} test tickets")
(MODELS / "m09-test-preds.json").write_text(json.dumps(results) + "\n")
Code explained
- In simple words: train the full fine-tune and the LoRA adapter with the chosen settings, save both, and compare them on everything the brief asks for.
- What happens: for each method it reloads the base model, adds LoRA if needed, counts trainable parameters, trains, and scores three things: accuracy on its own 48 training tickets, accuracy on the 24 test tickets (log-probability scoring), and free generation on test (does the model now write a valid label on its own?). The full model is saved with the canonical
save(4 MB); the adapter withsave_adapter. Predictions go tomodels/m09-test-preds.jsonfor Part 7. Run it withpython -m examples.m09_train(about 45 seconds). - Comes out:
training pairs: 189 (48 real dev tickets + augmented variants); test tickets: 24
full lr=0.0001 epochs=8 trainable=1,071,872 train 17.4s file 4201 KB
accuracy on its own training tickets 47/48; on test 5/24 = 21% (95% CI 9% to 40%)
free generation (no scoring): test 5/24, valid label rate 100%
predicted: {'how_to': 10, 'bug': 5, 'account_access': 6, 'billing': 1, 'feature_request': 2}
lora lr=0.003 epochs=12 trainable= 32,768 train 25.0s file 137 KB
accuracy on its own training tickets 44/48; on test 6/24 = 25% (95% CI 12% to 45%)
free generation (no scoring): test 6/24, valid label rate 100%
predicted: {'cancellation': 2, 'how_to': 10, 'billing': 1, 'account_access': 5, 'bug': 3, 'feature_request': 3}
full and LoRA agree on 9/24 test tickets
| Full fine-tuning | LoRA (r = 4) | |
|---|---|---|
| Trainable parameters | 1,071,872 | 32,768 (3%) |
| Training time (one thread, this machine) | 17 s | 25 s |
| Saved file | 4,201 KB (whole model) | 137 KB (adapter only) |
| Accuracy on its own training tickets | 47/48 | 44/48 |
| Accuracy on test (n = 24) | 5/24 = 21% (9% to 40%) | 6/24 = 25% (12% to 45%) |
| Valid label rate, free generation | 100% (was 0%) | 100% (was 0%) |
Four things to read here.
- Format was learned perfectly. Before training, free generation produced a valid label 0% of the time; now 100%. That is the "fine-tuning fixes format" claim from Part 1, measured.
- The task was not learned. 47/48 on training tickets and 5/24 on test is memorization: the model learned these 48 tickets, not the skill of triage. Test accuracy for both methods is at the "always say
how_to" floor, and the two models agree on only 9 of 24 test tickets, so they are not even making the same guesses. TinyLM's pretraining was templated support dialogue, with no general knowledge of language to build on, so 48 examples have nothing to steer. A real pretrained model starts from representations in which "I can't log in" and "my password doesn't work" are already close, which is why SFT on a few hundred examples usually works there. - LoRA is slower here, not faster. LoRA saves memory (optimizer state for 32,768 numbers instead of 1 million, and a 137 KB file instead of 4 MB). It does not save compute: the forward and backward passes still go through the whole network, plus the extra low-rank matrix multiplications. On a 1M-parameter model the memory saving is irrelevant and the overhead shows. On a 7B model on a GPU, the memory saving is what makes training possible at all.
- The file size is the real LoRA win for Brightlane: 137 KB per task means one base model can carry dozens of task adapters (Part 6).
A failure diagnosed: why the labels are single common words
An earlier version of this module trained on the raw category names (cancellation, account_access, feature_request). LoRA's training loss stalled well above zero and the model kept misspelling labels. Before changing any hyperparameter, read the evidence. This script looks at the label tokens themselves:
examples/m09_label_tokens.py
"""Why did LoRA stall on the category names? Look at the label tokens the base model has (never) seen."""
import json
from collections import Counter
import torch
import torch.nn.functional as F
from examples.m09_lora import add_lora
from examples.m09_setup import (DATA_OUT, LABEL_WORDS, RAW_NAMES, answer_for, encode_example, make_batch,
seed_everything, ticket_prompt)
from supportdesk.tinylm import load, loss_on
model, tok = load()
counts = Counter(tok.encode((DATA_OUT.parents[0] / "corpus.txt").read_text(encoding="utf-8")).ids)
norms = model.tok_emb.weight.norm(dim=1)
print(f"median embedding norm of tokens seen in pretraining: {norms[[i for i in counts]].median():.2f}")
for name, words in (("category names", RAW_NAMES), ("one-word labels", LABEL_WORDS)):
print(name)
for c, w in words.items():
enc = tok.encode(" " + w)
print(f" {w:<16}", " ".join(f"{tok.decode([i])!r}x{counts[i]}(|e|={norms[i]:.2f})" for i in enc.ids))
def first_vs_rest(m, x, y, mask):
"""Loss on the first label token (the decision) vs the remaining label tokens (the spelling)."""
with torch.no_grad():
losses = F.cross_entropy(m(x).reshape(-1, m.cfg.vocab_size), y.reshape(-1), reduction="none").view(y.shape)
first = mask.argmax(dim=1)
first_loss = losses[torch.arange(len(x)), first]
rest = (losses * mask).sum() - first_loss.sum()
return round(first_loss.mean().item(), 2), round((rest / (mask.sum() - len(x))).item(), 2)
rows = [json.loads(line) for line in (DATA_OUT / "train_real.jsonl").open()][:16]
for name, words in (("category names", RAW_NAMES), ("one-word labels", LABEL_WORDS)):
seed_everything(0)
m, _ = load()
add_lora(m, r=8, alpha=16)
x, y, mask = make_batch([encode_example(tok, ticket_prompt(tok, r["subject"], r["body"]),
answer_for(r["category"], words)) for r in rows])
params = [p for p in m.parameters() if p.requires_grad]
opt = torch.optim.AdamW(params, lr=3e-3, weight_decay=0.0)
m.train()
for step in range(60):
loss = loss_on(m, x, y, mask)
opt.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(params, 1.0)
opt.step()
m.eval()
print(f"{name}: after 60 LoRA steps on 16 tickets, loss (first token, later tokens) =", first_vs_rest(m, x, y, mask))
Code explained
- In simple words: check how each label is tokenized, how often each token occurred in pretraining, and whether LoRA can learn to produce it.
- What happens: it counts every token in the pretraining corpus, and reports the length (norm) of each token's embedding vector: tokens that appeared in pretraining got trained embeddings, tokens that never appeared kept small random ones. Then it trains a LoRA adapter for 60 steps on 16 tickets with each label set and splits the answer loss into the first label token (the decision) and the remaining tokens (the spelling). Run it with
python -m examples.m09_label_tokens(about 20 seconds). - Comes out:
median embedding norm of tokens seen in pretraining: 1.08
category names
billing ' billing'x505(|e|=1.23)
cancellation ' cancell'x0(|e|=0.62) 'ation'x0(|e|=0.59)
account_access ' account'x492(|e|=1.26) '_'x0(|e|=0.57) 'ac'x0(|e|=0.64) 'cess'x0(|e|=0.61)
bug ' bug'x1(|e|=1.02)
how_to ' how'x764(|e|=0.94) '_'x0(|e|=0.57) 'to'x0(|e|=0.59)
feature_request ' feat'x0(|e|=0.60) 'ure'x0(|e|=0.61) '_'x0(|e|=0.57) 're'x0(|e|=0.62) 'qu'x0(|e|=0.59) 'est'x0(|e|=0.62)
one-word labels
billing ' billing'x505(|e|=1.23)
cancel ' cancel'x494(|e|=1.13)
account ' account'x492(|e|=1.26)
error ' error'x2(|e|=1.13)
how ' how'x764(|e|=0.94)
request ' request'x259(|e|=1.25)
category names: after 60 LoRA steps on 16 tickets, loss (first token, later tokens) = (1.1, 3.16)
one-word labels: after 60 LoRA steps on 16 tickets, loss (first token, later tokens) = (0.0, 0.0)
The diagnosis is in the counts. cancell + ation, _ + ac + cess, feat + ure + _ + re + qu + est: these are tokens the tokenizer knows but that occurred zero times in the pretraining data, with untrained, short embeddings (about 0.6 against a median of 1.08). Producing them requires the model to output tokens it has never produced. LoRA cannot fix that. It only adjusts the attention and feed-forward layers, while the embedding and output layers, where a token's identity lives, stay frozen. After 60 steps the "spelling" loss is still 3.16. With one common word per label the same run reaches 0.0. The fix was a data decision, not a training one: choose label strings the base model can already say. The same applies to real models. If a label is an odd string the model rarely saw, it costs extra tokens and training effort. (This is also why the tokenizer's 1,503 real tokens against a 2,048-row table matters: rows beyond 1,503 are never trained at all.)
LoRA and QLoRA: the practical knobs
| Knob | What it does | Sensible start | Watch for |
|---|---|---|---|
| Rank r | Capacity of the update | 8 to 16 | Rank 16 overfit faster than 4 in our sweep; more capacity only helps with more data |
| alpha | Scale of the update (alpha / r) | 2r (as here) or 16 to 32 | Changing r without adjusting alpha changes the effective learning rate |
| Target layers | Where adapters go | All linear layers in the blocks | Attention-only adapters are smaller but often weaker |
| Learning rate | Step size | About 1e-4 to 3e-4 on real models; higher for tiny ones (we used 3e-3) | Full fine-tuning needs roughly 10 times smaller rates than LoRA |
| Epochs | Passes over the data | 1 to 3 on real models with hundreds of examples | Training accuracy near 100% with flat validation is memorization |
QLoRA (Dettmers et al., 2023) adds one idea: store the frozen base weights in 4 bits instead of 16, and train LoRA adapters in higher precision on top. Gradients still flow through the quantized base, but only the adapters change. The paper introduced a 4-bit NormalFloat data type, double quantization of the scales, and paged optimizers, and fine-tuned a 65B model on a single 48 GB GPU. The practical effect: the GPU you need is set by the 4-bit base, not the 16-bit one. Here is the mechanism in miniature, with an int8 base and a 4-bit base:
examples/m09_qlora.py
"""QLoRA in miniature: quantize the frozen base weights to int8 or 4-bit, then train LoRA on top in float32."""
import json
import torch
import torch.nn as nn
import torch.nn.functional as F
from examples.m09_lora import add_lora
from examples.m09_setup import (MODELS, corpus_validation_text, evaluate, fmt, load_rows, seed_everything, sft_train,
to_pairs)
from supportdesk.data import load_tickets
from supportdesk.tinylm import load, perplexity
class QuantLinear(nn.Module):
"""A frozen linear layer stored as small integers plus one float scale per group of weights."""
def __init__(self, layer: nn.Linear, bits: int = 8, group: int = 32) -> None:
super().__init__()
self.in_features, self.out_features, self.bits, self.group = layer.in_features, layer.out_features, bits, group
w = layer.weight.detach().reshape(-1, group) # groups of `group` consecutive weights
qmax = 2 ** (bits - 1) - 1 # 127 for int8, 7 for 4-bit
scale = w.abs().amax(dim=1, keepdim=True).clamp(min=1e-8) / qmax
self.register_buffer("q", torch.round(w / scale).clamp(-qmax, qmax).to(torch.int8))
self.register_buffer("scale", scale.to(torch.float16))
self.register_buffer("bias", layer.bias.detach().clone())
def weight(self) -> torch.Tensor:
return (self.q.float() * self.scale.float()).reshape(self.out_features, self.in_features)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return F.linear(x, self.weight(), self.bias) # dequantize on the fly
def storage_bytes(self) -> int:
"""Bytes if packed: bits per weight plus a 2-byte scale per group (4-bit packs two per byte)."""
return self.q.numel() * self.bits // 8 + self.scale.numel() * 2 + self.bias.numel() * 4
def quantize_blocks(model, bits):
for block in model.blocks:
block.qkv, block.proj = QuantLinear(block.qkv, bits), QuantLinear(block.proj, bits)
block.mlp[0], block.mlp[2] = QuantLinear(block.mlp[0], bits), QuantLinear(block.mlp[2], bits)
return model
if __name__ == "__main__":
base, tok = load()
val_text = corpus_validation_text(tok)
fp32_bytes = sum(m.weight.numel() * 4 + m.bias.numel() * 4 for b in base.blocks
for m in (b.qkv, b.proj, b.mlp[0], b.mlp[2]))
cfg = json.loads((MODELS / "m09-sweep.json").read_text())["chosen"]["lora"]
pairs = to_pairs(tok, load_rows("dev_real") + load_rows("dev_aug"))
test = load_tickets("test")
print(f"{'base weights':<14} {'block bytes':>11} {'corpus ppl':>10} LoRA triage result on test")
for bits in (32, 8, 4):
seed_everything(0)
model, _ = load()
if bits < 32:
quantize_blocks(model, bits)
size = fp32_bytes if bits == 32 else sum(m.storage_bytes() for b in model.blocks
for m in (b.qkv, b.proj, b.mlp[0], b.mlp[2]))
ppl = perplexity(model, tok, val_text)
add_lora(model, r=cfg["rank"], alpha=2 * cfg["rank"])
info = sft_train(model, tok, pairs, epochs=cfg["best_epoch"], lr=cfg["lr"])
res = evaluate(model, tok, test)
name = "float32" if bits == 32 else f"int{bits}"
print(f"{name:<14} {size / 1024:>8.0f} KB {ppl:>10.3f} {fmt(res)} ({info['seconds']:.0f}s)")
Code explained
- In simple words: replace the frozen weights with small integers plus a scale per group of 32, then train the same LoRA adapter on top and compare.
- What happens:
QuantLinearsplits a layer's weights into groups of 32, stores each group as integers between -127 and 127 (int8) or -7 and 7 (4-bit) plus one float16 scale, and rebuilds approximate float weights on every forward pass (dequantization).storage_bytescounts what the packed format would take on disk (4-bit packs two weights per byte; we store them in int8 tensors for simplicity).quantize_blocksswaps all 16 block layers. For each precision the script measures block weight bytes, perplexity on the corpus validation text before adding LoRA (the cost of quantization itself), then trains the chosen LoRA recipe and scores the test set. This is a simplified symmetric quantizer, not QLoRA's NF4, and there is no GPU memory to save here; it shows the mechanism and its accuracy cost. Run it withpython -m examples.m09_qlora(about 70 seconds). - Comes out:
Continued pretraining for domain language
SFT teaches a behavior from labeled pairs. Continued pretraining (CPT) is different: you keep doing plain next-token training, the way the model was pretrained, on raw text from your domain (help-center articles, product docs, internal wikis). The goal is vocabulary and style, not a task. The standard risk is forgetting the original distribution, and the standard defense is replay: mix some original pretraining text into every batch.
examples/m09_cpt.py
"""Continued pretraining: keep training the base LM on new domain text, with and without replay."""
import time
import torch
from examples.m09_setup import MODELS, ROOT, corpus_validation_text, seed_everything
from supportdesk.data import load_articles, load_tickets
from supportdesk.tinylm import batches, load, loss_on, perplexity, save
HELD_OUT_ARTICLES = {"data-privacy", "mobile-app", "status-incidents"}
def domain_text(articles, tickets) -> str:
parts = [f"# {a.title}\n{a.body}\n" for a in articles]
parts += [f"Ticket: {t.subject}\n{t.body}\n" for t in tickets if t.language == "en"]
return "\n".join(parts)
def continue_pretraining(tok, text: str, replay_ids=None, steps=150, lr=3e-4, seed=0):
seed_everything(seed)
model, _ = load()
ids = torch.tensor(tok.encode(text).ids)
g = torch.Generator().manual_seed(seed)
new = batches(ids, 64, 8, g) # shorter windows: the new text is only a few thousand tokens
old = batches(replay_ids, 64, 8, g) if replay_ids is not None else None
opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.0)
model.train()
started = time.perf_counter()
for step in range(steps):
x, y = next(new)
if old is not None: # replay: half of every batch is original pretraining text
ox, oy = next(old)
x, y = torch.cat([x[:4], ox[:4]]), torch.cat([y[:4], oy[:4]])
loss = loss_on(model, x, y)
opt.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
model.eval()
return model, time.perf_counter() - started, len(ids)
if __name__ == "__main__":
base, tok = load()
articles = load_articles()
train_text = domain_text([a for a in articles if a.id not in HELD_OUT_ARTICLES], load_tickets("dev"))
held_kb = domain_text([a for a in articles if a.id in HELD_OUT_ARTICLES], [])
held_tickets = domain_text([], load_tickets("test"))
val_text = corpus_validation_text(tok)
corpus_ids = tok.encode((ROOT / "data" / "corpus.txt").read_text(encoding="utf-8")).ids
replay = torch.tensor(corpus_ids[: int(len(corpus_ids) * 0.95)]) # training part only
seen_kb = domain_text([a for a in articles if a.id not in HELD_OUT_ARTICLES], [])
def row(name, m):
print(f"{name:<19} seen KB {perplexity(m, tok, seen_kb):6.2f} held-out KB {perplexity(m, tok, held_kb):6.2f} held-out tickets "
f"{perplexity(m, tok, held_tickets):6.2f} corpus val {perplexity(m, tok, val_text):5.2f}")
row("base", base)
cpt, secs, n = continue_pretraining(tok, train_text)
print(f" (continued pretraining on {n:,} new tokens, 150 steps, {secs:.0f}s)")
row("CPT, new text only", cpt)
cpt_replay, secs, _ = continue_pretraining(tok, train_text, replay_ids=replay)
row("CPT, 50% replay", cpt_replay)
save(cpt_replay, tok, MODELS / "m09-cpt", meta={"steps": 150, "lr": 3e-4, "replay": 0.5})
Code explained
- In simple words: keep pretraining TinyLM on Brightlane's help center and English dev tickets, with and without mixing in old corpus text, and measure perplexity on text it did and did not train on.
- What happens: three KB articles (
data-privacy,mobile-app,status-incidents) are held out. The training text is the other 9 articles plus English dev tickets, about 2,600 tokens.continue_pretrainingruns 150 steps of next-token training on random 64-token windows (the text is too short for 128), optionally replacing half of each batch with windows from the original corpus's training portion. Perplexity (Module 1) is exp of the average loss: roughly how many tokens the model is choosing between at each step, so lower means the text looks more familiar. It is measured on four texts: the KB articles it trained on, the held-out articles, the English test tickets (never seen), and the corpus validation text (the "old" skill). The replay model is saved tomodels/m09-cpt. Run it withpython -m examples.m09_cpt(about 25 seconds). - Comes out:
- Seen KB perplexity falls from 6.42 to 2.39. The model absorbed the text it trained on.
- Held-out tickets improve enormously (5,415 to about 300). The base model had never seen the
Ticket:format or real customer phrasing, so even 2,600 tokens teach it a lot about how tickets look. This is the real benefit of CPT: domain language. - Held-out KB articles get worse (14.7 to 20.7). 2,600 tokens for 150 steps is enough to memorize 9 articles but not enough to learn "how Brightlane help articles are written" in general; the model has overfit to the specific articles. More text, fewer steps, or both, would be needed.
- Replay protects the old skill. Without replay, corpus validation perplexity worsens by 9% (1.65 to 1.80). With 50% replay it slightly improves (1.55), because the replay batches are extra training on the original distribution. Replay also slightly helps the held-out text. When you do CPT, always measure the original distribution and mix old data in.
| Situation | Use this | Why |
|---|---|---|
| You need a behavior (format, label, style) and have labeled pairs | SFT with LoRA | Small, cheap, swappable; the frozen base limits damage |
| LoRA underfits after a proper sweep and you have plenty of data | Full fine-tuning (or LoRA with higher rank on all layers) | More capacity; accept a full model copy per task |
| The base model barely knows your domain's language | Continued pretraining with replay, then SFT | Vocabulary comes from raw text, behavior from pairs |
| Your GPU cannot hold the 16-bit base model for LoRA | QLoRA (4-bit base plus LoRA) | Memory set by the 4-bit base, small accuracy cost |
| You need several tasks on one deployment | One base plus one LoRA adapter per task | Adapters are KBs to MBs and swap in milliseconds (Part 6) |