"""Train the merge adapter on Blocksworld (prompt-only latent planning). Same recipe as train_merge_code.py; supervision = the model's own verified passing plan (direct for easy, CoT-derived for hard — the plan section only, CE on plan tokens).""" import argparse import json import math import os import random import re import sys import time from pathlib import Path import torch import torch.nn.functional as F from bw_prep import DIRECT_SUFFIX, chat from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter 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 LR = 1e-3 WARMUP = 20 K_BUCKETS = [(1, ("easy",)), (2, ("easy", "hard")), (4, ("hard",))] ap = argparse.ArgumentParser() ap.add_argument("--seed", type=int, default=0) ap.add_argument("--adaptive", action="store_true") ARGS = ap.parse_args() MOVELIST_RE = re.compile(r"(?:^|\n)\s*1\s*[.)]", re.M) def plan_only(text): """From a CoT output, keep from the numbered list onward.""" m = MOVELIST_RE.search(text) return text[m.start():].strip() if m else text.strip() 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 build_batch(tok, items, device="cuda"): seqs, labs = [], [] for it in items: p = tok(chat(tok, it, DIRECT_SUFFIX), add_special_tokens=False)["input_ids"] a = tok(plan_only(it["sol_plan"]) + "", add_special_tokens=False)["input_ids"] seqs.append(p + a) labs.append([-100] * len(p) + a) T = max(len(s) for s in seqs) pad = tok.pad_token_id or 0 ids = torch.full((len(seqs), T), pad, dtype=torch.long) lab = torch.full((len(seqs), T), -100, dtype=torch.long) msk = torch.zeros((len(seqs), T), dtype=torch.long) for i, (s, l) in enumerate(zip(seqs, labs)): ids[i, : len(s)] = torch.tensor(s) lab[i, : len(s)] = torch.tensor(l) msk[i, : len(s)] = 1 lmask = (lab == -100) & (msk == 1) return ids.to(device), msk.to(device), lab.to(device), lmask.to(device) def main(): rng = random.Random(ARGS.seed) torch.manual_seed(ARGS.seed) data = json.load(open(OUT / "bw_data.json")) model, tok = load_model(dtype=torch.bfloat16) for p in model.parameters(): p.requires_grad_(False) looper = BandLooper(model) cls = AdaptiveMergeAdapter if ARGS.adaptive else MergeAdapter adapter = cls( 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.get("sol_plan")] pool = {l: [it for it in train if it["label"] == l] for l in ("easy", "hard")} val = {l: pool[l][:12] for l in pool} pool = {l: pool[l][12:] for l in pool} print(f"pool: easy={len(pool['easy'])} hard={len(pool['hard'])}", flush=True) buckets = K_BUCKETS if len(pool["easy"]) < BATCH: # BW direct pass ~1%: easy bucket is empty by construction; train # hard-only rather than aborting the arm buckets = [(2, ("hard",)), (4, ("hard",))] print(f"easy pool too small ({len(pool['easy'])}) — " "hard-only curriculum k=2/4", flush=True) if len(pool["hard"]) < BATCH: print("INSUFFICIENT POOL — aborting", flush=True) return log = [] t0 = time.time() for step in range(STEPS): k, labels = buckets[step % len(buckets)] cand = [it for lbl in labels for it in pool[lbl]] batch = rng.sample(cand, BATCH) ids, msk, lab, lmask = build_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) opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0) opt.step() log.append({"step": step, "k": k, "loss": loss.item()}) 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 % 100 == 99 or step == STEPS - 1: torch.save(adapter.state_dict(), OUT / f"adapter_bw{"_ad" if ARGS.adaptive else ""}_e{step+1}.pt") json.dump(log, open(OUT / "train_bw_log.json", "w"), indent=1) print("done", flush=True) if __name__ == "__main__": main()