From 01ac15d3e14482439d8de8f05a4afa41aa20b49e Mon Sep 17 00:00:00 2001 From: Nils Date: Tue, 14 Jul 2026 10:05:30 +0200 Subject: [PATCH] node kit: alpha flag through train/eval, carry env+width fixes, bootstrap script Co-Authored-By: Claude Fable 5 --- scripts/eval_carry.py | 7 +++++-- scripts/eval_loop_code.py | 4 +++- scripts/node_setup.sh | 26 ++++++++++++++++++++++++++ scripts/train_carry.py | 7 +++++-- scripts/train_merge_code.py | 8 ++++++-- 5 files changed, 45 insertions(+), 7 deletions(-) create mode 100755 scripts/node_setup.sh diff --git a/scripts/eval_carry.py b/scripts/eval_carry.py index c1a4878..d8d3556 100644 --- a/scripts/eval_carry.py +++ b/scripts/eval_carry.py @@ -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() diff --git a/scripts/eval_loop_code.py b/scripts/eval_loop_code.py index ec5edfd..b2dd04d 100644 --- a/scripts/eval_loop_code.py +++ b/scripts/eval_loop_code.py @@ -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() diff --git a/scripts/node_setup.sh b/scripts/node_setup.sh new file mode 100755 index 0000000..bfe4bb7 --- /dev/null +++ b/scripts/node_setup.sh @@ -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 diff --git a/scripts/train_carry.py b/scripts/train_carry.py index 5115dce..718b404 100644 --- a/scripts/train_carry.py +++ b/scripts/train_carry.py @@ -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 diff --git a/scripts/train_merge_code.py b/scripts/train_merge_code.py index d3bebe3..a5f81c7 100644 --- a/scripts/train_merge_code.py +++ b/scripts/train_merge_code.py @@ -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 # @@ -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"