"""Rung-2 pilot: warm-start trained adapter, unfreeze band via loop-only LoRA. Entrance-weighted LoRA (rank 8, q/v/down of L14-22, scale fading 1->0) active ONLY inside band re-runs — k=0 remains bit-exact base model. Joint training of adapter (warm) + LoRA on the standard MBPP curriculum; inline slow-path eval (LoRA applies to every band traversal, so the fast prefill trick would be inconsistent). """ import argparse import json import math import os import random import sys import time from concurrent.futures import ThreadPoolExecutor from pathlib import Path import torch import torch.nn.functional as F from lora_band import entrance_faded_scales, inject_band_lora from loop_common import BandLooper, MergeAdapter, BAND from prep_mbpp import DIRECT_SUFFIX, extract_code, mbpp_prompt, run_tests from train_merge_code import build_code_batch # noqa: F401 (shared format) sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from jlens.core import load_model, _text_model # noqa: E402 OUT = Path(os.environ.get("LOOP_OUT", Path(__file__).resolve().parent.parent / "results-loop")) STEPS = 600 BATCH = 4 WARMUP = 20 K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))] ap = argparse.ArgumentParser() ap.add_argument("--warm-adapter", default=None, help="path to trained MergeAdapter to warm-start from") ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--rank", type=int, default=8) ap.add_argument("--seed", type=int, default=0) ap.add_argument("--eval-n", type=int, default=250) ARGS = ap.parse_args() def lr_at(step): if step < WARMUP: return ARGS.lr * (step + 1) / WARMUP t = (step - WARMUP) / max(1, STEPS - WARMUP) return 0.1 * ARGS.lr + 0.45 * ARGS.lr * (1 + math.cos(math.pi * t)) def main(): rng = random.Random(ARGS.seed) torch.manual_seed(ARGS.seed) data = json.load(open(OUT / "mbpp_data.json")) model, tok = load_model(dtype=torch.bfloat16) for p in model.parameters(): p.requires_grad_(False) tm = _text_model(model) looper = BandLooper(model) d = model.config.get_text_config().hidden_size adapter = MergeAdapter(d=d).cuda() if ARGS.warm_adapter: adapter.load_state_dict(torch.load(ARGS.warm_adapter, map_location="cuda")) print("warm-started adapter from", ARGS.warm_adapter, flush=True) scales = entrance_faded_scales(BAND[0], BAND[0] + 9, full_until=BAND[0] + 3) lora_params = inject_band_lora(tm, BAND[0], scales, rank=ARGS.rank) model.cuda() print(f"LoRA on layers {sorted(scales)} " f"({sum(p.numel() for p in lora_params)/1e6:.1f}M params)", flush=True) opt = torch.optim.AdamW( [{"params": adapter.parameters(), "lr": ARGS.lr}, {"params": lora_params, "lr": ARGS.lr}], weight_decay=0.01) train = [it for it in data if it["split"] == "train" and it["label"] != "drop" and it["sol_code"]] train = [it for it in train if len(tok(mbpp_prompt(tok, it, DIRECT_SUFFIX))["input_ids"]) + len(tok(it["sol_code"])["input_ids"]) + 12 <= 512] pool = {"easy": [it for it in train if it["label"] == "easy"][16:], "hard": [it for it in train if it["label"] == "hard"][16:]} print(f"pool: easy={len(pool['easy'])} hard={len(pool['hard'])}", flush=True) t0 = time.time() for step in range(STEPS): k, labels = K_BUCKETS[step % len(K_BUCKETS)] cand = [it for lbl in labels for it in pool[lbl]] batch = rng.sample(cand, min(BATCH, len(cand))) ids, msk, lab, lmask = build_code_batch(tok, batch) for g in opt.param_groups: g["lr"] = lr_at(step) logits = looper.loop_logits(adapter, ids, k, attention_mask=msk, use_checkpoint=True, loop_mask=lmask) loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(), lab[:, 1:].flatten(), ignore_index=-100) if not torch.isfinite(loss): print(f"NON-FINITE LOSS step {step} — halting (McLeish warning)", flush=True) break opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_( [p for g in opt.param_groups for p in g["params"]], 1.0) opt.step() if step % 10 == 0: print(f"step {step:4d} k={k} loss={loss.item():.4f} " f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True) if step % 200 == 199 or step == STEPS - 1: torch.save({"adapter": adapter.state_dict(), "lora": [ (p.detach().cpu()) for p in lora_params ]}, OUT / f"rung2_e{step+1}.pt") # inline slow-path eval (LoRA active in every band traversal) items = [it for it in data if it["split"] == "test"][: ARGS.eval_n] res = {} for k in (0, 2, 4): codes = [] for i in range(0, len(items), 8): torch.cuda.empty_cache() chunk = items[i : i + 8] enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk], return_tensors="pt", padding=True, add_special_tokens=False).to("cuda") gen = looper.loop_generate(adapter, tok, enc["input_ids"], k, max_new_tokens=220, attention_mask=enc["attention_mask"], loop_prompt_only=True) for j in range(len(chunk)): codes.append(extract_code( tok.decode(gen[j, enc["input_ids"].shape[1]:], skip_special_tokens=True))) with ThreadPoolExecutor(8) as ex: oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]), zip(codes, items))) by = {} for it, ok in zip(items, oks): dd = by.setdefault(it["label"], [0, 0]) dd[0] += ok dd[1] += 1 res[k] = {"acc": sum(oks) / len(items), "by_label": {l: c / n for l, (c, n) in by.items()}} print(f"k={k}: pass@1={res[k]['acc']:.3f} " f"by_label={ {l: round(v,3) for l,v in res[k]['by_label'].items()} }", flush=True) json.dump(res, open(OUT / "eval_rung2.json", "w"), indent=1) print("wrote", OUT / "eval_rung2.json") if __name__ == "__main__": main()