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 argparse
import json import json
import os
import sys import sys
import time import time
from pathlib import Path 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)) sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402 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() @torch.no_grad()
@@ -58,7 +60,8 @@ def main():
model, tok = load_model(dtype=torch.bfloat16) model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left" tok.padding_side = "left"
looper = BandLooper(model) looper = BandLooper(model)
adapter = MergeAdapter().cuda() adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
if args.adapter: if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval() adapter.eval()
+3 -1
View File
@@ -74,6 +74,7 @@ def main():
help="no-recurrence control arm: adapter(e,e) once") help="no-recurrence control arm: adapter(e,e) once")
ap.add_argument("--pause", type=int, default=0, ap.add_argument("--pause", type=int, default=0,
help="append p pause tokens to each prompt") help="append p pause tokens to each prompt")
ap.add_argument("--alpha", type=float, default=0.3)
args = ap.parse_args() args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")] ks = [int(x) for x in args.ks.split(",")]
@@ -81,7 +82,8 @@ def main():
tok.padding_side = "left" tok.padding_side = "left"
looper = BandLooper(model) looper = BandLooper(model)
adapter = MergeAdapter( 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: if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda")) adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval() 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 argparse
import json import json
import os
import math import math
import random import random
import sys 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)) sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402 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 STEPS = 600
BATCH = 4 BATCH = 4
LR = 1e-3 LR = 1e-3
@@ -93,7 +95,8 @@ def main():
for pp in model.parameters(): for pp in model.parameters():
pp.requires_grad_(False) pp.requires_grad_(False)
looper = BandLooper(model) 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) opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
keep = [it for it in data 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, ap.add_argument("--pause", type=int, default=0,
help="pause-token control: p inert tokens after prompt, " help="pause-token control: p inert tokens after prompt, "
"feedforward adapter, no recurrence") "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() ARGS = ap.parse_args()
SEED = ARGS.seed 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> PAUSE_ID = 6 # <unused0>
@@ -101,7 +105,7 @@ def main():
p.requires_grad_(False) p.requires_grad_(False)
looper = BandLooper(model) looper = BandLooper(model)
d = model.config.get_text_config().hidden_size 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) opt = torch.optim.AdamW(adapter.parameters(), lr=LR, weight_decay=0.01)
train = [it for it in data if it["split"] == "train" train = [it for it in data if it["split"] == "train"