"""Do gated specialists beat putting the fact in the prompt?

Everything measured tonight compared compositional systems against monoliths
trained on the same 22k records. Both arms of that comparison are far below
usable, so it establishes the mechanism and says nothing about whether any of
it beats the thing a practitioner would actually do: take a competent
pretrained base and give it the retrieved context.

This runs that comparison on a real base, with the specialists as LoRA
adapters rather than blocks trained from scratch.

    base            the pretrained model, nothing added. The floor.
    base + RAG      the nearest training record's answer pasted into the
                    prompt as context. No training at all. THE baseline.
    base + 1 LoRA   one adapter trained on every category. The generalist.
    base + routed   one adapter per gated category, selected at inference by
                    embedding similarity to a category centroid. Ours.
    routed + RAG    both, to see whether they compose or overlap.

Routing is real rather than oracle: the arm picks its own specialist and eats
its own routing errors, because that is what deployment does.

Adapters target v_proj and o_proj only, following the ablation in the
anchor-token-masking work where V/O-only beat full LoRA at half the parameters.

Metric is mean NLL on answer tokens of held-out records, plus generated
samples so the numbers can be read against actual text.

Usage:
    python lora_vs_rag.py --base HuggingFaceTB/SmolLM2-360M-Instruct --specialists 6
"""
from __future__ import annotations

import argparse
import json
import random
import sys
import time
from pathlib import Path

try:
    sys.stdout.reconfigure(encoding="utf-8")
except Exception:
    pass

import numpy as np
import torch
import torch.nn.functional as F

HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE))

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"


def build_data(manifest: Path, emb_path: Path, n_specialists: int, val_frac: float,
               seed: int):
    import build_specialists as bs
    man = json.loads(manifest.read_text(encoding="utf-8"))
    raw = bs.load_corpus(Path(man["corpus"]), None)
    X = np.load(emb_path)
    assign = man["assignments"]
    # keep the n most coherent categories
    keep = [c["id"] for c in man["kept"][:n_specialists]]
    per = {c: [] for c in keep}
    for k, g in assign.items():
        if g in keep:
            per[g].append(int(k))
    rng = random.Random(seed)
    train, val = {}, {}
    for g, idxs in per.items():
        rng.shuffle(idxs)
        k = max(2, int(len(idxs) * val_frac))
        val[g] = idxs[:k]
        train[g] = idxs[k:]
    return raw, X, keep, train, val


def prompt_of(rec, context=None):
    if context:
        return ("Use this context to answer.\nContext: %s\n\nQ: %s\nA:"
                % (context, rec["instruction"]))
    return "Q: %s\nA:" % rec["instruction"]


@torch.no_grad()
def answer_nll(model, tok, prompt, answer, max_len=320):
    p = tok(prompt, return_tensors="pt", truncation=True, max_length=max_len - 64)
    a = tok(" " + answer, return_tensors="pt", add_special_tokens=False,
            truncation=True, max_length=64)
    ids = torch.cat([p["input_ids"], a["input_ids"]], 1).to(model.device)
    if ids.shape[1] < 2 or a["input_ids"].shape[1] == 0:
        return None
    logits = model(ids).logits[:, :-1]
    tgt = ids[:, 1:]
    start = p["input_ids"].shape[1] - 1
    lg, tg = logits[0, start:], tgt[0, start:]
    if tg.numel() == 0:
        return None
    return float(F.cross_entropy(lg.float(), tg, reduction="mean"))


def evaluate(model, tok, raw, val, keep, retriever=None, limit=40, seed=0):
    rng = random.Random(seed)
    per, allv = {}, []
    for g in keep:
        idxs = val[g][:limit] if len(val[g]) <= limit else rng.sample(val[g], limit)
        vals = []
        for i in idxs:
            rec = raw[i]
            ctx = retriever(i, g) if retriever else None
            n = answer_nll(model, tok, prompt_of(rec, ctx), rec["output"])
            if n is not None and np.isfinite(n):
                vals.append(n)
        if vals:
            per[g] = float(np.mean(vals)); allv.extend(vals)
    return (float(np.mean(allv)) if allv else float("nan")), per, len(allv)


def train_lora(base_id, tok, raw, idxs, steps, batch, lr, rank, seed, max_len=256):
    from peft import LoraConfig, get_peft_model
    from transformers import AutoModelForCausalLM
    torch.manual_seed(seed)
    m = AutoModelForCausalLM.from_pretrained(base_id, dtype=torch.bfloat16).to(DEVICE)
    cfg = LoraConfig(r=rank, lora_alpha=rank * 2, lora_dropout=0.0, bias="none",
                     target_modules=["v_proj", "o_proj"], task_type="CAUSAL_LM")
    m = get_peft_model(m, cfg)
    m.train()
    params = [p for p in m.parameters() if p.requires_grad]
    opt = torch.optim.AdamW(params, lr=lr)
    rng = random.Random(seed)
    t0 = time.time(); losses = []
    for step in range(steps):
        pick = [idxs[rng.randrange(len(idxs))] for _ in range(batch)]
        seqs, starts = [], []
        for i in pick:
            rec = raw[i]
            p = tok(prompt_of(rec), add_special_tokens=True)["input_ids"]
            a = tok(" " + rec["output"], add_special_tokens=False)["input_ids"]
            s = (p + a)[:max_len]
            if len(s) <= len(p):
                continue
            seqs.append(s); starts.append(min(len(p), len(s)))
        if not seqs:
            continue
        L = max(len(s) for s in seqs)
        ids = torch.full((len(seqs), L), tok.pad_token_id or 0, dtype=torch.long)
        lm = torch.zeros((len(seqs), L), dtype=torch.bool)
        for r, (s, st) in enumerate(zip(seqs, starts)):
            ids[r, :len(s)] = torch.tensor(s); lm[r, st:len(s)] = True
        ids, lm = ids.to(DEVICE), lm.to(DEVICE)
        out = m(ids).logits[:, :-1]
        tgt, mm = ids[:, 1:], lm[:, 1:]
        if mm.sum() == 0:
            continue
        loss = F.cross_entropy(out[mm].float(), tgt[mm])
        opt.zero_grad(set_to_none=True); loss.backward()
        torch.nn.utils.clip_grad_norm_(params, 1.0); opt.step()
        losses.append(loss.item())
    n_tr = sum(p.numel() for p in params)
    print("      %d steps, %.0fs, %.2fM trainable, loss %.3f -> %.3f"
          % (steps, time.time() - t0, n_tr / 1e6, losses[0], np.mean(losses[-30:])))
    m.eval()
    return m


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    ap.add_argument("--base", default="HuggingFaceTB/SmolLM2-360M-Instruct")
    ap.add_argument("--manifest", type=Path, default=HERE / "specialist_build" / "manifest.json")
    ap.add_argument("--emb", type=Path, default=HERE / "specialist_build" / "emb.npy")
    ap.add_argument("--specialists", type=int, default=6)
    ap.add_argument("--steps", type=int, default=400)
    ap.add_argument("--gen-steps", type=int, default=1200)
    ap.add_argument("--batch", type=int, default=4)
    ap.add_argument("--lr", type=float, default=2e-4)
    ap.add_argument("--rank", type=int, default=16)
    ap.add_argument("--eval-per", type=int, default=25)
    ap.add_argument("--seed", type=int, default=42)
    ap.add_argument("--out", type=Path, default=HERE / "lora_vs_rag.json")
    args = ap.parse_args()

    from transformers import AutoModelForCausalLM, AutoTokenizer
    tok = AutoTokenizer.from_pretrained(args.base)
    if tok.pad_token is None:
        tok.pad_token = tok.eos_token

    raw, X, keep, train, val = build_data(args.manifest, args.emb, args.specialists,
                                          0.08, args.seed)
    print("base %s" % args.base)
    print("%d specialists: %s" % (len(keep), ", ".join("%s(n=%d)" % (g, len(train[g])) for g in keep)))

    # centroids for real routing, and a retrieval index for the RAG arm
    cents = {g: X[train[g]].mean(0) for g in keep}
    for g in cents:
        cents[g] = cents[g] / (np.linalg.norm(cents[g]) or 1.0)
    C = np.stack([cents[g] for g in keep])

    def route(i):
        return keep[int(np.argmax(C @ X[i]))]

    def retrieve(i, g):
        """Nearest training record anywhere in the library; its answer is the context."""
        pool = [j for gg in keep for j in train[gg]]
        sims = X[pool] @ X[i]
        return raw[pool[int(np.argmax(sims))]]["output"][:400]

    results = {}
    print("\n--- arm 1: base alone ---")
    base = AutoModelForCausalLM.from_pretrained(args.base, dtype=torch.bfloat16).to(DEVICE).eval()
    m, per, n = evaluate(base, tok, raw, val, keep, None, args.eval_per, args.seed)
    print("    mean answer NLL %.4f over %d held-out records" % (m, n))
    results["base"] = {"nll": m, "per": per}

    print("\n--- arm 2: base + RAG (no training) ---")
    m2, per2, _ = evaluate(base, tok, raw, val, keep, retrieve, args.eval_per, args.seed)
    print("    mean answer NLL %.4f" % m2)
    results["base_rag"] = {"nll": m2, "per": per2}

    # The RAG arm changes the prompt format as well as adding content, and the
    # LoRA arms are scored in the format they trained on. Without this control
    # the comparison silently charges RAG for the format change. A random
    # context holds the format fixed and removes only the relevance.
    print("\n--- arm 2b: control, same format with an irrelevant context ---")
    pool_all = [j for gg in keep for j in train[gg]]
    rng_ctl = random.Random(args.seed + 1)

    def irrelevant(i, g):
        return raw[pool_all[rng_ctl.randrange(len(pool_all))]]["output"][:400]

    m2b, per2b, _ = evaluate(base, tok, raw, val, keep, irrelevant, args.eval_per, args.seed)
    print("    mean answer NLL %.4f   (retrieval is worth %+.4f of this)"
          % (m2b, m2 - m2b))
    results["base_randctx"] = {"nll": m2b, "per": per2b}
    del base; torch.cuda.empty_cache()

    print("\n--- arm 3: base + one generalist LoRA ---")
    all_idx = [j for g in keep for j in train[g]]
    gen = train_lora(args.base, tok, raw, all_idx, args.gen_steps, args.batch,
                     args.lr, args.rank, args.seed)
    m3, per3, _ = evaluate(gen, tok, raw, val, keep, None, args.eval_per, args.seed)
    print("    mean answer NLL %.4f" % m3)
    results["generalist_lora"] = {"nll": m3, "per": per3}
    del gen; torch.cuda.empty_cache()

    print("\n--- arm 4: base + routed specialist LoRAs (real routing) ---")
    spec_models = {}
    for g in keep:
        print("    training specialist %s" % g)
        spec_models[g] = train_lora(args.base, tok, raw, train[g], args.steps,
                                    args.batch, args.lr, args.rank, args.seed)
    vals, per4, routed_ok = [], {}, 0
    rng = random.Random(args.seed)
    for g in keep:
        idxs = val[g][:args.eval_per] if len(val[g]) <= args.eval_per else rng.sample(val[g], args.eval_per)
        got = []
        for i in idxs:
            pick = route(i)
            routed_ok += int(pick == g)
            n = answer_nll(spec_models[pick], tok, prompt_of(raw[i]), raw[i]["output"])
            if n is not None and np.isfinite(n):
                got.append(n)
        if got:
            per4[g] = float(np.mean(got)); vals.extend(got)
    m4 = float(np.mean(vals))
    print("    mean answer NLL %.4f   routing accuracy %.3f"
          % (m4, routed_ok / max(1, sum(min(len(val[g]), args.eval_per) for g in keep))))
    results["routed_specialists"] = {"nll": m4, "per": per4}

    print("\n" + "=" * 62)
    print("  %-42s %10s %12s" % ("arm", "NLL", "vs RAG"))
    print("  " + "-" * 66)
    for name, label in (("base", "base alone, plain format"),
                        ("base_randctx", "base + irrelevant context (format control)"),
                        ("base_rag", "base + RAG (baseline)"),
                        ("generalist_lora", "base + 1 LoRA"),
                        ("routed_specialists", "base + routed specialists")):
        d = results[name]["nll"] - results["base_rag"]["nll"]
        print("  %-42s %10.4f %+12.4f" % (label, results[name]["nll"], d))
    print()
    if results["routed_specialists"]["nll"] < results["base_rag"]["nll"]:
        print("  routed specialists beat putting the fact in the prompt")
    else:
        print("  routed specialists do NOT beat putting the fact in the prompt;")
        print("  on this data the architecture does not yet earn its complexity")

    args.out.write_text(json.dumps(results, indent=2, default=str), encoding="utf-8")
    print("\nwritten to %s" % args.out)


if __name__ == "__main__":
    main()
