160 lines
6.4 KiB
Python
160 lines
6.4 KiB
Python
"""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()
|