mirror of
https://github.com/wassname/jsteer.git
synced 2026-10-10 00:30:22 +08:00
The verified run-524 vectors were fit on chat-templated prompts (u4_prompts.json), so fitting raw WikiText diverged from what worked. Now: - config.chat_corpus wraps jlens WikiText in the chat template (fit J where we steer) - jsteer.demo.show_steer generates through apply_chat_template(enable_thinking) with the model's own generation_config sampling, splits </think>, shows lens_topk j-space readout + reasoning + answer as Tufte small-multiples per C - word_steering.ipynb rewired to Qwen3.5-4B, dim_batch=4 (3090-safe 4B), show_steer - fit.py defaults to Qwen3.5-4B + chat_corpus Co-Authored-By: Claudypoo <288921227+claudypoo@users.noreply.github.com>
57 lines
2.5 KiB
Python
57 lines
2.5 KiB
Python
"""Fit and cache any HF causal LM's Jacobian for the notebooks and README. (Claude)
|
|
|
|
Pass `--model`; the cache lands at `config.cache_path(model)` (e.g.
|
|
`artifacts/qwen3.5-4b.jac`). Prompts are jlens's WikiText-103 corpus wrapped in
|
|
the model's chat template (config.chat_corpus): we fit J at the chat operating
|
|
point where steering is applied, matching the verified run-524. jlens guidance:
|
|
~100 prompts is usable, the paper uses 1000; 128 is a cheap default. Idempotent:
|
|
re-running loads the existing cache instead of refitting (Jacobian.fit_cached).
|
|
|
|
uv run python scripts/fit.py --model Qwen/Qwen3.5-4B
|
|
uv run python scripts/fit.py --model Qwen/Qwen3-0.6B --dim-batch 8
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import sys
|
|
import time
|
|
from pathlib import Path
|
|
|
|
from loguru import logger
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) # repo root for config
|
|
import config # noqa: E402
|
|
from jsteer import Jacobian # noqa: E402
|
|
|
|
|
|
def main() -> None:
|
|
p = argparse.ArgumentParser(description=__doc__)
|
|
p.add_argument("--model", default="Qwen/Qwen3.5-4B")
|
|
p.add_argument("--n-prompts", type=int, default=128)
|
|
p.add_argument("--dim-batch", type=int, default=4,
|
|
help="d_model dims per backward batch (memory knob; 4 fits a 4B on a 24GB 3090, 8+ for smaller)")
|
|
p.add_argument("--layers", type=float, nargs=2, default=(0.3, 0.9),
|
|
metavar=("LO", "HI"), help="fractional layer band to fit")
|
|
p.add_argument("--max-seq-len", type=int, default=128)
|
|
args = p.parse_args()
|
|
|
|
out = config.cache_path(args.model)
|
|
logger.info(f"loading {args.model} ({config.DTYPE}) on {config.DEVICE}")
|
|
tok = AutoTokenizer.from_pretrained(args.model)
|
|
model = AutoModelForCausalLM.from_pretrained(
|
|
args.model, dtype=config.DTYPE).to(config.DEVICE).eval()
|
|
|
|
logger.info(f"fit-or-load {out} (layers={tuple(args.layers)}, "
|
|
f"dim_batch={args.dim_batch}, n_prompts={args.n_prompts} chat-templated WikiText)")
|
|
t0 = time.monotonic()
|
|
jac = Jacobian.fit_cached(model, tok, lambda: config.chat_corpus(tok, args.n_prompts), out,
|
|
layers=tuple(args.layers), dim_batch=args.dim_batch,
|
|
max_seq_len=args.max_seq_len,
|
|
checkpoint_path=str(config.cache_path(args.model, "ckpt")))
|
|
logger.info(f"{jac!r} -> {out} ({(time.monotonic() - t0) / 60:.1f} min)")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|