decouple --pause from feedforward: enable loop+pause combo arm

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-14 10:09:18 +02:00
co-authored by Claude Fable 5
parent 5c4b76ba14
commit f48011c470
+4 -2
View File
@@ -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)