node kit: alpha flag through train/eval, carry env+width fixes, bootstrap script

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-14 10:05:30 +02:00
co-authored by Claude Fable 5
parent 674b814b65
commit 01ac15d3e1
5 changed files with 45 additions and 7 deletions
+5 -2
View File
@@ -7,6 +7,7 @@ Grid entries are k:p pairs; k=0,p=0 is the plain baseline (same harness).
import argparse
import json
import os
import sys
import time
from pathlib import Path
@@ -20,7 +21,8 @@ from loop_common import (DIRECT_SUFFIX, BandLooper, MergeAdapter, chat_prompt,
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
@torch.no_grad()
@@ -58,7 +60,8 @@ def main():
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter().cuda()
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
+3 -1
View File
@@ -74,6 +74,7 @@ def main():
help="no-recurrence control arm: adapter(e,e) once")
ap.add_argument("--pause", type=int, default=0,
help="append p pause tokens to each prompt")
ap.add_argument("--alpha", type=float, default=0.3)
args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")]
@@ -81,7 +82,8 @@ def main():
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
d=model.config.get_text_config().hidden_size,
alpha=args.alpha).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
+26
View File
@@ -0,0 +1,26 @@
#!/bin/bash
# vast.ai node bootstrap for jspace (see vast-ai-notes.md for pitfalls)
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/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
hf download google/gemma-4-12B-it
hf download google/gemma-4-E2B-it
# probe gate: one verified generation before anything batch
export HF_HUB_OFFLINE=1 JLENS_MODEL=google/gemma-4-12B-it
python - <<'PY'
import torch, sys
sys.path.insert(0, "/dev/shm/jspace/scripts"); sys.path.insert(0, "/dev/shm/jspace")
from jlens.core import load_model
m, tok = load_model(dtype=torch.bfloat16)
ids = tok.apply_chat_template([{"role":"user","content":"What is 2+2? Just the number."}],
return_tensors="pt", add_generation_prompt=True, return_dict=True)["input_ids"].cuda()
out = m.generate(ids, max_new_tokens=8, do_sample=False)
txt = tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True)
print("PROBE:", repr(txt)); assert "4" in txt, "PROBE FAILED"
PY
echo NODE-READY
+5 -2
View File
@@ -8,6 +8,7 @@ Curriculum: easy p=2, hard p=6 (hard needs the longer latent chain).
import argparse
import json
import os
import math
import random
import sys
@@ -23,7 +24,8 @@ from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
STEPS = 600
BATCH = 4
LR = 1e-3
@@ -93,7 +95,8 @@ def main():
for pp in model.parameters():
pp.requires_grad_(False)
looper = BandLooper(model)
adapter = MergeAdapter().cuda()
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
keep = [it for it in data
+6 -2
View File
@@ -40,9 +40,13 @@ 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("--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()
SEED = ARGS.seed
SUFFIX = (f"_s{SEED}" if SEED else "") + (f"_p{ARGS.pause}" if ARGS.pause else "")
SUFFIX = ((f"_s{SEED}" if SEED else "")
+ (f"_p{ARGS.pause}" if ARGS.pause else "")
+ (f"_a{ARGS.alpha}" if ARGS.alpha != 0.3 else ""))
PAUSE_ID = 6 # <unused0>
@@ -101,7 +105,7 @@ def main():
p.requires_grad_(False)
looper = BandLooper(model)
d = model.config.get_text_config().hidden_size
adapter = MergeAdapter(d=d).cuda()
adapter = MergeAdapter(d=d, alpha=ARGS.alpha).cuda()
opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
train = [it for it in data if it["split"] == "train"