From b461970c9a1570815f2ef72b44ffc247eb9c92f8 Mon Sep 17 00:00:00 2001 From: Nils Date: Tue, 14 Jul 2026 13:02:46 +0200 Subject: [PATCH] warm-start flag (stacking arm), adaptive flags for Blocksworld Co-Authored-By: Claude Fable 5 --- scripts/eval_bw.py | 6 ++++-- scripts/train_merge_bw.py | 8 +++++--- scripts/train_merge_code.py | 7 ++++++- 3 files changed, 15 insertions(+), 6 deletions(-) diff --git a/scripts/eval_bw.py b/scripts/eval_bw.py index 1cc9aec..9d0dd15 100644 --- a/scripts/eval_bw.py +++ b/scripts/eval_bw.py @@ -11,7 +11,7 @@ import torch from bw_common import verify_plan from bw_prep import DIRECT_SUFFIX, chat -from loop_common import BandLooper, MergeAdapter +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 @@ -52,13 +52,15 @@ def main(): ap.add_argument("--adapter", default=None) ap.add_argument("--tag", default="bw") ap.add_argument("--ks", default="0,2,4") + ap.add_argument("--adaptive", action="store_true") args = ap.parse_args() ks = [int(x) for x in args.ks.split(",")] model, tok = load_model(dtype=torch.bfloat16) tok.padding_side = "left" looper = BandLooper(model) - adapter = MergeAdapter( + cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter + adapter = cls( d=model.config.get_text_config().hidden_size).cuda() if args.adapter: adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) diff --git a/scripts/train_merge_bw.py b/scripts/train_merge_bw.py index 605aad6..fffd1ce 100644 --- a/scripts/train_merge_bw.py +++ b/scripts/train_merge_bw.py @@ -18,7 +18,7 @@ import torch import torch.nn.functional as F from bw_prep import DIRECT_SUFFIX, chat -from loop_common import BandLooper, MergeAdapter +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 @@ -33,6 +33,7 @@ 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) @@ -81,7 +82,8 @@ def main(): for p in model.parameters(): p.requires_grad_(False) looper = BandLooper(model) - adapter = MergeAdapter( + 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) @@ -127,7 +129,7 @@ def main(): 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_e{step+1}.pt") + 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) diff --git a/scripts/train_merge_code.py b/scripts/train_merge_code.py index 8f8acfe..6631310 100644 --- a/scripts/train_merge_code.py +++ b/scripts/train_merge_code.py @@ -51,6 +51,7 @@ ap.add_argument("--deepk", type=int, default=0, ap.add_argument("--bptt", type=int, default=0, help="truncated BPTT: grads only through last N iterations") ap.add_argument("--lr", type=float, default=1e-3) +ap.add_argument("--warm", default=None, help="warm-start adapter checkpoint") ARGS, _ = ap.parse_known_args() SEED = ARGS.seed LR = ARGS.lr @@ -59,7 +60,8 @@ SUFFIX = ((f"_s{SEED}" if SEED else "") + (f"_a{ARGS.alpha}" if ARGS.alpha != 0.3 else "") + ("_ad" if ARGS.adaptive else "") + (f"_dk{ARGS.deepk}" if ARGS.deepk else "") - + (f"_lr{ARGS.lr}" if ARGS.lr != 1e-3 else "")) + + (f"_lr{ARGS.lr}" if ARGS.lr != 1e-3 else "") + + ("_warm" if ARGS.warm else "")) PAUSE_ID = 6 # @@ -122,6 +124,9 @@ def main(): adapter = AdaptiveMergeAdapter(d=d, alpha0=ARGS.alpha).cuda() else: adapter = MergeAdapter(d=d, alpha=ARGS.alpha).cuda() + if ARGS.warm: + adapter.load_state_dict(torch.load(ARGS.warm, map_location="cuda")) + print("warm-started from", ARGS.warm, flush=True) global K_BUCKETS if ARGS.deepk: f = ARGS.deepk / 4