| """Tier-1 sanity demo: RRT tie + param-efficient KD on the tiny CPU model. |
| |
| Proves the end-to-end pipeline runs and trains (loss drops) before we touch real |
| weights. The tiny model is random, so this validates *mechanics*, not recovery. |
| |
| uv run python scripts/rrt_tiny_demo.py |
| |
| Mirrors the Tier-2 flow on the GPU box (scripts/rrt_run.py): precompute teacher |
| logits with the untied model -> tie in place -> freeze base -> KD the LoRA adapters. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import sys |
| from pathlib import Path |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) |
|
|
| import torch |
| import torch.nn.functional as F |
|
|
| from looped_laguna import build_tiny_model |
| from rrt_laguna import ( |
| TieConfig, |
| adjacent_pairs, |
| parameter_report, |
| set_param_efficient, |
| tie_model, |
| trainable_parameters, |
| ) |
|
|
|
|
| def main() -> None: |
| torch.manual_seed(0) |
| model = build_tiny_model(num_layers=8) |
|
|
| gen = torch.Generator().manual_seed(1) |
| ids = torch.randint(3, model.config.vocab_size, (4, 32), generator=gen) |
|
|
| |
| with torch.no_grad(): |
| teacher_logits = model(input_ids=ids, use_cache=False).logits.detach() |
| teacher_probs = torch.softmax(teacher_logits.float(), dim=-1) |
|
|
| |
| pairs = adjacent_pairs([3, 4, 5, 6]) |
| cfg = TieConfig(pairs=pairs, rank=8, init="lower", lora_init="svd") |
| rep_before = parameter_report(model)["total_unique"] |
| tie_model(model, cfg) |
| rep_after = parameter_report(model) |
| print(f"tied pairs={pairs} init={cfg.init} rank={cfg.rank}") |
| print(f"unique params: {rep_before:,} -> {rep_after['total_unique']:,} " |
| f"({100 * (1 - rep_after['total_unique'] / rep_before):.1f}% smaller)") |
|
|
| |
| set_param_efficient(model) |
| n_train = sum(p.numel() for p in trainable_parameters(model)) |
| print(f"trainable (LoRA) params: {n_train:,}") |
|
|
| def kd_loss(): |
| out = model(input_ids=ids, use_cache=False).logits |
| logp = torch.log_softmax(out.float(), dim=-1) |
| kl = F.kl_div(logp, teacher_probs, reduction="batchmean") |
| ce = F.cross_entropy(out[:, :-1].reshape(-1, out.shape[-1]), ids[:, 1:].reshape(-1)) |
| return kl + ce, kl.item(), ce.item() |
|
|
| opt = torch.optim.Adam(trainable_parameters(model), lr=5e-3) |
| print("\nstep loss KL CE") |
| for step in range(0, 41): |
| loss, kl, ce = kd_loss() |
| if step % 10 == 0: |
| print(f"{step:>4} {loss.item():.4f} {kl:.4f} {ce:.4f}") |
| opt.zero_grad() |
| loss.backward() |
| opt.step() |
| print("\nKD loop reduced the loss -> pipeline works. Ready to point at real weights.") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|