Compare commits

..
2 Commits
Author SHA1 Message Date
NilsandClaude Fable 5 f48011c470 decouple --pause from feedforward: enable loop+pause combo arm
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-14 10:09:18 +02:00
NilsandClaude Fable 5 5c4b76ba14 bootstrap via bucket git-bundle (no ssh keys on rented nodes)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-14 10:06:41 +02:00
2 changed files with 6 additions and 3 deletions
+2 -1
View File
@@ -4,7 +4,8 @@ set -x
source /venv/main/bin/activate
uv pip install -q "transformers==5.13.*" datasets accelerate hf_transfer
export HF_HOME=/dev/shm/hf; mkdir -p /dev/shm/hf
cd /dev/shm && git clone git@git.draic.info:nils/jspace.git 2>/dev/null || (cd jspace && git pull)
cd /dev/shm && rclone copy jspace:jspace/git/jspace.bundle /dev/shm/
git clone /dev/shm/jspace.bundle jspace 2>/dev/null || (cd jspace && git pull /dev/shm/jspace.bundle main)
cd /dev/shm/jspace
rclone copy jspace:jspace/results-loop /dev/shm/jspace/results-loop/ --include "*.json" --include "adapter_code_e399.pt"
export HF_HUB_ENABLE_HF_TRANSFER=1
+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)