Files
jspace/scripts/train_rung2.py
T

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()