maxact-fast / DESIGN.md
ceselder's picture
maxact-fast: clean repo — DESIGN + pretraining pipeline (embed/cluster/probe-cache/build-data/pretrain), Dr. GRPO RL to follow
8505f8e
|
Raw
History Blame Contribute Delete
3.89 kB

maxact-fast

Train a model to generate text that maximally activates an arbitrary linear direction in its own residual stream, conditioned activation-oracle style (direction injected at layer 1, maximized at layer L). The bet: if we pretrain on a huge, diverse set of directions — not just SAE features — the model learns a general residual-direction→text inverter that transfers to SAE features for free (cross-uplift).

The direction firehose (why this beats SAE-only conditioning)

SAE conditioning caps you at ~65k features and inherits the SAE's blind spots. Instead we mint an effectively unlimited, diverse supply of concept directions straight from a corpus:

  1. Embed a large corpus (UltraFineWeb, streamed) with a fast embedding model (BGE) — cheap.
  2. Cluster the embeddings into K≈10k clusters (faiss k-means, GPU). Each cluster ≈ a concept.
  3. For a training example, sample cluster A vs cluster B, train a quick linear probe on the policy model's layer-L residuals to separate A from B. The probe weight vector w is the conditioning direction — same space as an SAE encoder column, so it drops straight into the existing inject@1 / read@L machinery.
  4. SFT target = the few real corpus texts closest to cluster A's centroid. "Given direction w, write text like this."
  5. Pretrain on millions of these (shuffled). Then RL (Dr. GRPO, no KL) on cluster-probes and/or SAE features, and measure whether cluster-pretraining lifts SAE-feature performance.

Probe residuals are cached once per cluster (the only expensive Qwen3 pass), so step 3-4 mint millions of (direction, targets) examples cheaply.

Objective & injection (unchanged from the activation-oracle line)

  • Inject unit(w) * ‖resid‖ * coeff into the layer-1 residual at a ? marker (norm-matched add).
  • Reward = max over token positions of (x_L − μ)·unit(w), measured on the STANDALONE re-tokenized generation through the clean base model (no adapter, no injection). Same fair protocol as before.

RL = Dr. GRPO, done right, no KL

  • advantage = r − group_meanno /std (the removed Dr. GRPO difficulty bias).
  • loss = Σ_tokens(−min(ratio·A, clip(ratio)·A) · mask) / total_tokens_in_batch (global constant normalizer, no per-response length bias).
  • No KL term. Stability comes from lr + grad-clip + fluency/distinct gates (optional), not a leash.
  • Rollouts via vllm-lens (SteeringVector(norm_match=True) == our injection), ~5× vs HF generate.
  • Reuse vLLM generation logprobs as old_logps (one on-policy update/batch → ratio≈1, no extra fwd).

Pipeline stages (each a CLI in scripts/)

  1. embed_cluster — stream corpus → BGE embeddings → faiss k-means K clusters → cluster assignments.
  2. cache_resids — per cluster, run M member texts through Qwen3, cache layer-L mean residuals.
  3. build_data — sample cluster pairs, train probes on cached resids, pick centroid targets → (direction, target_text) shards. Millions.
  4. pretrain — SFT the generator (AO/LoRA) on the direction→text firehose. Massive.
  5. finetune — final polish on per-feature MAX-activating CORPUS examples (cleanest targets).
  6. rl — Dr. GRPO (no KL), vllm-lens rollouts, on cluster-probes and/or SAE features.
  7. eval — beat-corpus / normalized-activation on held-out SAE features (cross-uplift test).
  8. interface — tiny gradio app: type/pick a direction, watch generations + live activation.

Speed principles

bf16 + flash-attn everywhere; faiss-GPU k-means; HF streaming (never materialize the corpus); cache the one expensive Qwen3 residual pass; vLLM-lens rollouts; batchmax forwards on H200; torch.compile the training step. Target: pretraining data mine + pretrain in a day, not a week.