Skip to content

Latest commit

 

History

3 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 

Repository files navigation

J-space Early Stopping for Reasoning Models

Reading the model's mind, not its mouth. We use a Jacobian lens to decode what the model wants to say at Layer 24 — long before the output layer makes up its mind.


What this is

A set of experiments asking a simple question: does Qwen3-8B internally know when to stop thinking, well before it actually outputs </think>?

Answer: Yes — by ~99 tokens on average, at Layer 24, using a pre-fitted Jacobian lens.

This repo contains the full pipeline: trajectory generation, lens fitting, dense J-space scanning, threshold calibration, counterfactual verification, and cumulative Φ benchmarking against SAGE-RL.


Quick Results

Metric Value
Probe accuracy at L24 (stop vs. thinking) 99.99%
J-space Φ detects earlier than output Φ (50% threshold) +99 tokens
% of samples with J-space advantage 99.6%
Counterfactual: force-early-stop at z=1.5 96.6% accuracy preserved, 338 tokens saved
Counterfactual: force-early-stop at z=2.0 100% accuracy preserved, 377 tokens saved
Correct vs. incorrect discrimination (Cohen's d) 1.97–3.10 across layers

How it works (in 30 seconds)

  1. Trajectory generation (step1/generate_trajectories.py): Run Qwen3-8B on GSM8K, record hidden states at selected layers.
  2. Lens fitting (step1/fit_jlens.py): Fit Jacobian matrices J_l that transport hidden states from layer l directly to the vocabulary space (bypassing the model's own late-layer processing).
  3. Dense sweep (step2/scan_jspace_early.py): For every token position in the thinking span, compute J-score = unembed(J_l · h_l[pos]).
  4. Calibration + counterfactual (step2/level2_calibrate.py): Z-score normalization, threshold sweep, force-inject </think> and regenerate.
  5. Φ benchmark (step2/level3_jspace_phi.py): Cumulative Φ comparison against SAGE-RL's output-layer baseline.

The key insight: J-space at L24 reads the "stop decision" at the point of crystallization (67% model depth), before the Verification Gate (L27–L35) processes and gates it to the output. The output layer only sees a binary cliff; J-space sees a gradual semantic ramp.


Repo Structure

├── step1/                          # Data generation + lens fitting
│   ├── generate_trajectories.py    # Run model, save hidden states per layer
│   ├── fit_jlens.py                # Fit Jacobian lens matrices
│   ├── utils.py                    # Shared utilities (model loading, paths)
│   └── verify_setup.py             # Environment check
├── step2/                          # Analysis + experiments
│   ├── scan_jspace_early.py        # Level 1: Dense J-space sweep (every position)
│   ├── analyze_sweep.py            # Level 1.5: Per-sample z-score analysis
│   ├── level2_calibrate.py         # Level 2: Threshold calibration + counterfactual
│   ├── level3_jspace_phi.py        # Level 3: J-space Φ vs SAGE-RL Φ benchmark
│   ├── probe_v2.py                 # Linear probe on hidden states (auxiliary)
│   ├── probe_stop_signal.py        # Probe specifically for stop signal
│   ├── steer_v2.py                 # Steering experiments (auxiliary)
│   └── steer_stop_signal.py        # Steering targeting stop signal
└── README.md

Technical reports: Full experimental reports (probing, circuit ablation, dense sweep, Φ benchmark) are coming soon. They will be linked here when published.


Running the Experiments

Prerequisites

  • Qwen3-8B model (HuggingFace)
  • Python 3.10+, PyTorch 2.x, Transformers, jlens package

Level 1: Dense J-space Sweep

CUDA_VISIBLE_DEVICES=2,3,4,5,6,7 python step2/scan_jspace_early.py \
    --data /path/to/trajectories_8b_v3/ \
    --lens /path/to/lenses/qwen3_8b/lens.pt \
    --model /path/to/Qwen3-8B \
    --output /path/to/jspace_sweep_8b/ \
    --max-samples 100

Level 2: Calibration + Counterfactual

# Parts A+B only (CPU, fast):
python step2/level2_calibrate.py \
    --sweep /path/to/jspace_sweep_8b/ \
    --data /path/to/trajectories_8b_v3/ \
    --output /path/to/level2_calibration/ \
    --layer 24 --no-counterfactual

# Full pipeline (A+B+C) with GPUs:
CUDA_VISIBLE_DEVICES=2,3,4 python step2/level2_calibrate.py \
    --sweep /path/to/jspace_sweep_8b/ \
    --data /path/to/trajectories_8b_v3/ \
    --model /path/to/Qwen3-8B \
    --lens /path/to/lenses/qwen3_8b/lens.pt \
    --output /path/to/level2_calibration/ \
    --layer 24 --z-thresholds 1.5 2.0 2.5 3.0 --counterfactual-n 30

Level 3: Φ Benchmark

python step2/level3_jspace_phi.py \
    --sweep-dir /path/to/jspace_sweep_8b/ \
    --data /path/to/trajectories_8b_v3/ \
    --output /path/to/level3_phi/

Citation

If you use this code or build on these findings, please cite:

@misc{gao2026jspace,
  title={J-space Early Stopping: Reading the Stop Decision Before the Output Cliff},
  author={Xueqing Gao},
  year={2026},
  note={Independent research. Code at \url{https://github.com/tensor2023/jspace-stop}},
}

License

MIT. The Jacobian lens code (jlens package) is from Anthropic's jacobian-lens (Apache 2.0).

About

Investigating Early Stopping in Reasoning Models Using J-space

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages