Files
jspace/scripts/train_gate_code.py

121 lines
4.4 KiB
Python

"""E1 trainer (PLAN_SELFPACED): halting-gate merge on MBPP prompt-only loop.
Same data/recipe as train_merge_code.py (STaR-verified code supervision,
prompt-only loop_mask), but NO difficulty->depth curriculum: every batch
mixes difficulties and runs the soft-halting loop to k_max — the GATE must
learn the allocation the curriculum used to hand-code.
Loss: CE(code tokens | s_mix) + lam * mean(E[iters]).
"""
import argparse
import json
import os
import math
import random
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from halting_common import HaltingMergeAdapter, halting_loop_logits
from loop_common import BandLooper
from prep_mbpp import DIRECT_SUFFIX, mbpp_prompt
from train_merge_code import build_code_batch, MAX_TOK
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
STEPS = 600
BATCH = 4
WARMUP = 20
ap = argparse.ArgumentParser()
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--lam", type=float, default=1e-3,
help="compute penalty weight on E[iterations]")
ap.add_argument("--kmax", type=int, default=4)
ap.add_argument("--lr", type=float, default=1e-3)
ap.add_argument("--lam-warmup", type=int, default=100,
help="steps before the penalty ramps in (anti-collapse)")
ARGS = ap.parse_args()
LR = ARGS.lr
TAG = (f"gate_l{ARGS.lam:g}" + (f"_s{ARGS.seed}" if ARGS.seed else ""))
def lr_at(step):
if step < WARMUP:
return LR * (step + 1) / WARMUP
t = (step - WARMUP) / max(1, STEPS - WARMUP)
return 1e-4 + 0.5 * (LR - 1e-4) * (1 + math.cos(math.pi * t))
def lam_at(step):
return ARGS.lam * min(1.0, max(0.0, (step - ARGS.lam_warmup) / 100))
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)
looper = BandLooper(model)
adapter = HaltingMergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=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 <= MAX_TOK]
val = train[:24]
pool = train[24:]
print(f"pool={len(pool)} val={len(val)} lam={ARGS.lam} kmax={ARGS.kmax}",
flush=True)
log = []
t0 = time.time()
for step in range(STEPS):
batch = rng.sample(pool, BATCH)
ids, msk, lab, lmask = build_code_batch(tok, batch)
# last prompt position = last position where label == -100 & mask==1
pl = (lmask.long().cumsum(-1).argmax(-1))
for g in opt.param_groups:
g["lr"] = lr_at(step)
logits, q, ei = halting_loop_logits(
looper, adapter, ids, ARGS.kmax, msk, pl, loop_mask=lmask)
ce = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(),
lab[:, 1:].flatten(), ignore_index=-100)
loss = ce + lam_at(step) * ei.mean()
opt.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
opt.step()
hard_b = [it["label"] == "hard" for it in batch]
log.append({"step": step, "ce": ce.item(), "ei": ei.mean().item(),
"ei_hard": (ei[torch.tensor(hard_b)].mean().item()
if any(hard_b) else None),
"ei_easy": (ei[~torch.tensor(hard_b)].mean().item()
if not all(hard_b) else None)})
if step % 10 == 0:
r = log[-1]
print(f"step {step:4d} ce={r['ce']:.4f} E[k]={r['ei']:.2f} "
f"(hard {r['ei_hard']} / easy {r['ei_easy']}) "
f"({(time.time()-t0)/(step+1):.1f}s/step)", flush=True)
if step % 100 == 99 or step == STEPS - 1:
torch.save(adapter.state_dict(),
OUT / f"adapter_{TAG}_e{step+1}.pt")
json.dump(log, open(OUT / f"train_{TAG}_log.json", "w"))
print("done", flush=True)
if __name__ == "__main__":
main()