Skip to content

Repository files navigation

RepDist

DDPM modeling of hidden-state representations.

Given last-prompt-token hidden states (\mathbf{h}) at transformer layers 1, 8, 16, 24, and 32, this repo asks:

  1. Does (p_H) have nontrivial low-dimensional geometry (anisotropy, spectral concentration, effective rank)?
  2. Can a diffusion model (p_\theta) recover that distribution on held-out hidden states, beyond a Gaussian baseline?

The implementation follows the experimental memo in the slides: cosine DDPM, full-rank skip noise predictor, covariance spectra, effective rank, PCA, and sliced Wasserstein-2. The residual branch is zero-initialized, and the training normalizer is fitted once and frozen.

Setup

Python 3.13+, a CUDA GPU for the 7B extraction (RTX 4090 is enough in bf16), and uv.

uv sync --extra dev
uv run pytest

Data

Item Value
Model allenai/Olmo-3-7B-Think
Dataset allenai/Dolci-Think-SFT-7B (streaming)
Layers 1, 8, 16, 24, 32 (hidden_states[k] after block k), aligned in one forward
Token last prompt token after chat template (add_generation_prompt=True, left padding)
Splits 5k val, 5k test (frozen), 10k train initially
Streaming if val diffusion loss plateaus, append more train hidden states

User/assistant traces are truncated to the user prompt. Assistant <think> traces are never encoded.

Method

Hidden states are centered with the training mean and divided by a global RMS scale (s_{\mathrm{train}}).

Forward process (cosine schedule, (T=1000), (s=0.008), (\beta_t \le 0.999)):

[ \mathbf{x}_t = \sqrt{\bar\alpha_t},\mathbf{x}_0 + \sqrt{1-\bar\alpha_t},\boldsymbol{\varepsilon} ]

The noise predictor is a two-layer MLP with an 8196-dimensional SiLU hidden layer and a sinusoidal timestep embedding. Training minimizes (|\boldsymbol{\varepsilon}-\boldsymbol{\varepsilon}_\theta(\mathbf{x}_t,t)|_2^2).

Reverse sampling uses the standard DDPM mean and posterior variance (\tilde\beta_t), with no extra noise at (t=1). Samples are mapped back by (\tilde{\mathbf{h}} = s_{\mathrm{train}}\tilde{\mathbf{x}}0 + \boldsymbol{\mu}{\mathrm{train}}). The corrected checkpoint format is v2; old checkpoints are incompatible and must not be resumed.

Commands

# CPU smoke (synthetic hidden states)
uv run repdist --config configs/smoke.yaml all

# Full experiment (extract all five layers, then train/eval each)
uv run repdist --config configs/default.yaml extract
uv run repdist --config configs/default.yaml train
uv run repdist --config configs/default.yaml eval

# One layer only
uv run repdist --config configs/default.yaml --layer 16 train
uv run repdist --config configs/default.yaml --layer 16 eval

# Isolated corrected run (pilot, then full)
uv run repdist --config configs/rankfix-pilot.yaml train --no-resume
uv run repdist --config configs/rankfix-pilot.yaml eval
uv run repdist --config configs/rankfix.yaml train --no-resume
uv run repdist --config configs/rankfix.yaml eval

extract writes aligned per-layer stores under data/hidden_states/layer_XX/. train / eval loop extract.layers unless --layer is set. Resume reads each layer's checkpoints/layer_XX/latest.pt by default. Use --no-resume for a fresh corrected run; never resume the legacy 50k run with a rankfix config.

Artifacts

Directory Contents
data/hidden_states/layer_XX/ val.pt, test.pt, train_*.pt, manifest.json
checkpoints/layer_XX/ latest.pt, best.pt, periodic step_*.pt
logs/layer_XX/train.jsonl train/val MSE vs step
outputs/layer_XX/figures/ spectrum.png, pca.png, loss.png
outputs/layer_XX/metrics/eval.json (d_{\mathrm{eff}}) and SWD vs standard-normal baseline

Evaluation

Held-out test hidden states are compared to an equal number of diffusion samples and to (\mathcal N(\hat\mu_{\mathrm{train}}, \hat\Sigma_{\mathrm{train}})):

  • covariance eigenvalue spectra and effective rank (d_{\mathrm{eff}}=(\sum\lambda_i)^2 / \sum\lambda_i^2)
  • PCA scatter in a basis fitted on the training set
  • sliced Wasserstein-2 over random 1-D projections

Smaller SWD and closer spectra indicate better recovery of (p_H).

Corrected-run protocol

configs/rankfix-pilot.yaml is a 2,500-step go/no-go run. It evaluates, checkpoints, and records diagnostics every 250 steps, uses 16 diagnostic samples, and does not stream additional training data. Proceed to configs/rankfix.yaml only if diagnostics remain finite, final and maximum normalized reverse RMS stay at or below 3, no hard-abort threshold is reached, and validation MSE is stable or improving. A nonfinite trajectory, RMS hard-fail, or persistently unhealthy diagnostic is a no-go.

The full run is 50,000 steps with the memo's cosine schedule and 8196-wide SiLU MLP, 500-step evaluation/checkpoint/diagnostic cadence, and streaming enabled. Both corrected configs reuse the frozen per-layer hidden-state stores but isolate normalizer, checkpoint, log, and output paths under runs/rankfix-pilot/ and runs/rankfix/ (nested as layer_XX).

Evaluation compares real and diffusion samples in the shared normalized space, plus an independent (\mathcal{N}(0,I)) baseline (not a fitted Gaussian). PCA is fit on training data at rank 32. Use validation only for tuning: if diagnostics are stable but high-timestep MSE remains poor, retry once at half the learning rate; if the RMS hard-fails twice, stop.

About

DDPM modeling of Olmo-3 hidden-state representations

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Used by

Contributors

Languages