From f48011c470cdf3855aacd4077ee6f2ffbcbde002 Mon Sep 17 00:00:00 2001 From: Nils Date: Tue, 14 Jul 2026 10:09:18 +0200 Subject: [PATCH] decouple --pause from feedforward: enable loop+pause combo arm Co-Authored-By: Claude Fable 5 --- scripts/train_merge_code.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/scripts/train_merge_code.py b/scripts/train_merge_code.py index a5f81c7..7eafcc6 100644 --- a/scripts/train_merge_code.py +++ b/scripts/train_merge_code.py @@ -40,6 +40,8 @@ ap.add_argument("--seed", type=int, default=0) ap.add_argument("--pause", type=int, default=0, help="pause-token control: p inert tokens after prompt, " "feedforward adapter, no recurrence") +ap.add_argument("--feedforward", action="store_true", + help="apply adapter once, no recurrence (pause-FF control)") ap.add_argument("--alpha", type=float, default=0.3, help="merge weight (2B-tuned default 0.3; try 0.1-0.15 at 12B)") ARGS = ap.parse_args() @@ -87,7 +89,7 @@ def val_loss(looper, adapter, tok, items, k): for i in range(0, len(items), BATCH): ids, msk, lab, lmask = build_code_batch(tok, items[i : i + BATCH]) logits = looper.loop_logits(adapter, ids, k, attention_mask=msk, - loop_mask=lmask, feedforward=bool(ARGS.pause)) + loop_mask=lmask, feedforward=ARGS.feedforward) loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(), lab[:, 1:].flatten(), ignore_index=-100) tot += loss.item() * len(ids) @@ -133,7 +135,7 @@ def main(): g["lr"] = lr_at(step) logits = looper.loop_logits(adapter, ids, k, attention_mask=msk, use_checkpoint=True, loop_mask=lmask, - feedforward=bool(ARGS.pause)) + feedforward=ARGS.feedforward) loss = F.cross_entropy(logits[:, :-1].flatten(0, 1).float(), lab[:, 1:].flatten(), ignore_index=-100) opt.zero_grad(set_to_none=True)