CourseLarge Language Models · Module 9: Adaptation: Fine-Tuning and Customization · part 47 of 80
Part 47 · Module 9: Adaptation: Fine-Tuning and Customization

Part 4: Methods

25 min read·22 Sept 2026

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

python
"""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 as base and freezes it. lora_A (r by in) gets a small random start; lora_B (out by r) starts at zero. forward returns 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 with LoRALinear wrappers. After this, only A and B have requires_grad=True, which is all sft_train needs to know.
    • lora_state_dict, save_adapter, load_adapter: save and restore just the tensors whose names contain lora_, plus the rank and alpha in meta.
    • 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:
text
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

python
"""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_of assigns each dev ticket to one of 4 folds, stratified by category. run trains a fresh model per fold (augmenting only that fold's training tickets, after the split) and uses the on_epoch hook 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 to models/m09-sweep.json. Run it with python -m examples.m09_sweep (about 10 minutes on one thread).
  • Comes out:
text
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

python
"""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 with save_adapter. Predictions go to models/m09-test-preds.json for Part 7. Run it with python -m examples.m09_train (about 45 seconds).
  • Comes out:
text
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-tuningLoRA (r = 4)
Trainable parameters1,071,87232,768 (3%)
Training time (one thread, this machine)17 s25 s
Saved file4,201 KB (whole model)137 KB (adapter only)
Accuracy on its own training tickets47/4844/48
Accuracy on test (n = 24)5/24 = 21% (9% to 40%)6/24 = 25% (12% to 45%)
Valid label rate, free generation100% (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

python
"""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:
text
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

KnobWhat it doesSensible startWatch for
Rank rCapacity of the update8 to 16Rank 16 overfit faster than 4 in our sweep; more capacity only helps with more data
alphaScale of the update (alpha / r)2r (as here) or 16 to 32Changing r without adjusting alpha changes the effective learning rate
Target layersWhere adapters goAll linear layers in the blocksAttention-only adapters are smaller but often weaker
Learning rateStep sizeAbout 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
EpochsPasses over the data1 to 3 on real models with hundreds of examplesTraining 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

python
"""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: QuantLinear splits 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_bytes counts what the packed format would take on disk (4-bit packs two weights per byte; we store them in int8 tensors for simplicity). quantize_blocks swaps 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 with python -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

python
"""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_pretraining runs 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 to models/m09-cpt. Run it with python -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.
SituationUse thisWhy
You need a behavior (format, label, style) and have labeled pairsSFT with LoRASmall, cheap, swappable; the frozen base limits damage
LoRA underfits after a proper sweep and you have plenty of dataFull 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 languageContinued pretraining with replay, then SFTVocabulary comes from raw text, behavior from pairs
Your GPU cannot hold the 16-bit base model for LoRAQLoRA (4-bit base plus LoRA)Memory set by the 4-bit base, small accuracy cost
You need several tasks on one deploymentOne base plus one LoRA adapter per taskAdapters are KBs to MBs and swap in milliseconds (Part 6)