riscautious / evaluate.py
Sendy08's picture
RiscAutious: LoRA from scratch, 92.2% on banking77
a75ccfd verified
Raw
History Blame Contribute Delete
11.9 kB
#!/usr/bin/env python3
"""Compare the LoRA and full fine-tuned models side by side.
python train.py --mode lora
python train.py --mode full
python evaluate.py
Writes ``results/comparison.md`` (paste into the README) and
``results/comparison.csv``, and prints the table.
Everything is measured on the held-out **test** split — never validation, which
was used to pick the best epoch and would therefore give an optimistic number.
"""
from __future__ import annotations
import argparse
import json
import logging
import sys
import time
from collections import Counter
from pathlib import Path
from typing import Any, Sequence
import torch
from sklearn.metrics import classification_report, confusion_matrix, f1_score
from data.dataset import build_dataloaders, label_mappings
from models.classifier import TextClassifier
from train import evaluate_split, pick_device, synchronize
log = logging.getLogger("evaluate")
MODES = ("full", "lora")
def measure_latency(
model: TextClassifier,
loader: torch.utils.data.DataLoader,
device: torch.device,
n_examples: int = 200,
warmup: int = 20,
) -> float:
"""Mean single-example inference latency, in milliseconds.
Measured at **batch size 1** on purpose. That is what a deployed demo
actually does — one user, one text box, one request. Batched throughput is a
different (and much prettier) number that would not reflect the Space.
The warmup iterations are not optional: the first CUDA call pays one-off
context and kernel-compilation costs, and on CPU the first pass populates
caches. Including them inflates the mean by a wide margin.
"""
model.eval()
dataset = loader.dataset
def run_one(i: int) -> None:
item = dataset[i]
# unsqueeze(0): (L,) -> (1, L). The model always expects a batch
# dimension, even when the batch is a single example.
input_ids = item["input_ids"].unsqueeze(0).to(device)
attention_mask = item["attention_mask"].unsqueeze(0).to(device)
with torch.no_grad():
model(input_ids, attention_mask)
for i in range(min(warmup, len(dataset))):
run_one(i)
n = min(n_examples, len(dataset))
synchronize(device)
start = time.perf_counter()
for i in range(n):
run_one(i)
synchronize(device)
elapsed = time.perf_counter() - start
return 1000.0 * elapsed / n
def majority_baseline(
loader: torch.utils.data.DataLoader, labels: Sequence[str]
) -> dict[str, float]:
"""Score of a model that always predicts the most common training grade.
The number every other row has to beat. An accuracy figure without this next
to it is not interpretable — on this dataset, always saying "B" scores
around 1.7% on banking77.
"""
ids = list(loader.dataset.labels)
most_common, count = Counter(ids).most_common(1)[0]
predictions = [most_common] * len(ids)
_, id_to_label = label_mappings(labels)
return {
"accuracy": count / len(ids),
"macro_f1": f1_score(ids, predictions, average="macro", zero_division=0),
"class": id_to_label[most_common],
}
def evaluate_mode(
mode: str,
args: argparse.Namespace,
loader: torch.utils.data.DataLoader,
device: torch.device,
) -> dict[str, Any] | None:
"""Load one trained model, score it on the test split, and time it.
Returns None (with a warning) if that mode has not been trained yet, so the
script still works when you have only run one of them.
"""
checkpoint = Path(args.checkpoint_dir) / f"{mode}.pt"
metrics_path = Path(args.results_dir) / f"{mode}_metrics.json"
if not checkpoint.exists() or not metrics_path.exists():
log.warning("Skipping %s — run: python train.py --mode %s", mode, mode)
return None
train_metrics = json.loads(metrics_path.read_text())
model = TextClassifier.load(checkpoint, device=device)
_, accuracy, predictions, truths = evaluate_split(model, loader, device)
latency_ms = measure_latency(model, loader, device, args.latency_examples)
return {
"mode": mode,
"test_accuracy": accuracy,
"macro_f1": f1_score(truths, predictions, average="macro", zero_division=0),
"trainable_params": train_metrics["trainable_params"],
"total_params": train_metrics["total_params"],
"trainable_pct": train_metrics["trainable_pct"],
"train_seconds": train_metrics["train_seconds"],
"best_val_acc": train_metrics["best_val_acc"],
"checkpoint_kb": train_metrics["checkpoint_kb"],
"latency_ms": latency_ms,
"lora_rank": train_metrics.get("lora_rank"),
"device": str(device),
"predictions": predictions,
"truths": truths,
}
def build_markdown(
rows: list[dict], baseline: dict, device: torch.device, n_classes: int
) -> str:
"""Render the comparison as a markdown table, ready to paste into the README."""
lines = [
"# LoRA vs. full fine-tuning",
"",
f"DistilBERT, {n_classes}-class intent classification. Test split, device: `{device}`.",
"",
"| Model | Test accuracy | Macro F1 | Trainable params | % of total | "
"Train time | Latency / example | Checkpoint |",
"| --- | --- | --- | --- | --- | --- | --- | --- |",
f"| Majority baseline ({baseline['class']}) | {baseline['accuracy']:.1%} | "
f"{baseline['macro_f1']:.3f} | 0 | 0% | — | — | — |",
]
for row in rows:
name = "Full fine-tuning" if row["mode"] == "full" else f"LoRA (r={row['lora_rank']})"
size = (f"{row['checkpoint_kb'] / 1024:.0f} MB" if row["checkpoint_kb"] > 1024
else f"{row['checkpoint_kb']:.0f} KB")
lines.append(
f"| {name} | {row['test_accuracy']:.1%} | {row['macro_f1']:.3f} | "
f"{row['trainable_params']:,} | {row['trainable_pct']:.3f}% | "
f"{row['train_seconds']:.1f}s | {row['latency_ms']:.1f} ms | {size} |"
)
by_mode = {row["mode"]: row for row in rows}
if "lora" in by_mode and "full" in by_mode:
lora, full = by_mode["lora"], by_mode["full"]
param_ratio = full["trainable_params"] / lora["trainable_params"]
acc_delta = (lora["test_accuracy"] - full["test_accuracy"]) * 100
speedup = full["train_seconds"] / lora["train_seconds"]
size_ratio = full["checkpoint_kb"] / lora["checkpoint_kb"]
lines += [
"",
"## Takeaways",
"",
f"- LoRA trains **{param_ratio:,.0f}x fewer parameters** "
f"({lora['trainable_params']:,} vs {full['trainable_params']:,}).",
f"- Accuracy difference: **{acc_delta:+.1f} percentage points** "
f"({lora['test_accuracy']:.1%} vs {full['test_accuracy']:.1%}).",
f"- Training was **{speedup:.2f}x** "
f"{'faster' if speedup > 1 else 'slower'} with LoRA.",
f"- The LoRA checkpoint is **{size_ratio:,.0f}x smaller**, which is what "
f"makes free-tier deployment practical.",
"",
"Inference latency is essentially identical by construction: both models run "
"the same 66M-parameter forward pass. LoRA saves on *training* cost and "
"*storage*, not on inference.",
]
return "\n".join(lines) + "\n"
def build_csv(rows: list[dict], baseline: dict) -> str:
"""Render the same numbers as CSV."""
header = (
"model,test_accuracy,macro_f1,trainable_params,total_params,"
"trainable_pct,train_seconds,latency_ms,checkpoint_kb"
)
lines = [
header,
f"majority_baseline,{baseline['accuracy']:.6f},{baseline['macro_f1']:.6f},"
f"0,0,0,,,",
]
for row in rows:
lines.append(
f"{row['mode']},{row['test_accuracy']:.6f},{row['macro_f1']:.6f},"
f"{row['trainable_params']},{row['total_params']},"
f"{row['trainable_pct']:.6f},{row['train_seconds']:.3f},"
f"{row['latency_ms']:.3f},{row['checkpoint_kb']:.1f}"
)
return "\n".join(lines) + "\n"
def run(args: argparse.Namespace) -> int:
"""Evaluate every trained mode and write the comparison outputs."""
device = pick_device(args.device)
_, _, test_loader, _, labels = build_dataloaders(
data_path=args.data, batch_size=args.batch_size,
max_length=args.max_length, seed=args.seed,
)
log.info("Test split: %d examples on %s", len(test_loader.dataset), device)
baseline = majority_baseline(test_loader, labels)
rows = [r for mode in MODES if (r := evaluate_mode(mode, args, test_loader, device))]
if not rows:
log.error("No trained models found. Run train.py first.")
return 1
out_dir = Path(args.results_dir)
out_dir.mkdir(parents=True, exist_ok=True)
markdown = build_markdown(rows, baseline, device, len(labels))
(out_dir / "comparison.md").write_text(markdown)
(out_dir / "comparison.csv").write_text(build_csv(rows, baseline))
print("\n" + markdown)
if args.per_class:
_ = labels
for row in rows:
print(f"\n Per-class breakdown — {row['mode']}")
print(" " + "-" * 60)
print(classification_report(
row["truths"], row["predictions"],
labels=list(range(len(labels))), target_names=list(labels),
zero_division=0, digits=3,
))
# A 77x77 confusion matrix is unreadable in a terminal; list the
# most frequent confusions instead, which is what you actually act on.
matrix = confusion_matrix(
row["truths"], row["predictions"], labels=list(range(len(labels)))
)
confusions = [
(matrix[i][j], labels[i], labels[j])
for i in range(len(labels)) for j in range(len(labels))
if i != j and matrix[i][j] > 0
]
print(" Top confusions (true -> predicted)")
for count, true_name, pred_name in sorted(confusions, reverse=True)[:12]:
print(f" {count:>3}x {true_name} -> {pred_name}")
print(f"\n Wrote {out_dir / 'comparison.md'} and {out_dir / 'comparison.csv'}\n")
return 0
def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
"""Define and parse the command-line interface."""
p = argparse.ArgumentParser(
description="Compare trained LoRA and full fine-tuned models.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
p.add_argument("--data", type=Path, default=Path("data/processed/dataset.csv"))
p.add_argument("--results-dir", type=Path, default=Path("results"))
p.add_argument("--checkpoint-dir", type=Path, default=Path("checkpoints"))
p.add_argument("--batch-size", type=int, default=32)
p.add_argument("--max-length", type=int, default=128)
p.add_argument("--seed", type=int, default=42,
help="Must match training, or the test split differs.")
p.add_argument("--latency-examples", type=int, default=200)
p.add_argument("--per-class", action="store_true",
help="Also print per-grade precision/recall and a confusion matrix.")
p.add_argument("--device", default=None)
return p.parse_args(argv)
def main(argv: Sequence[str] | None = None) -> int:
"""Entry point. Returns a process exit code."""
logging.basicConfig(level=logging.INFO, format="%(levelname)-7s %(message)s")
for noisy in ("httpx", "urllib3", "filelock", "huggingface_hub"):
logging.getLogger(noisy).setLevel(logging.WARNING)
import transformers
transformers.logging.set_verbosity_error()
return run(parse_args(argv))
if __name__ == "__main__":
sys.exit(main())