121 lines
4.4 KiB
Python
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()
|