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:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Executable
+26
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user