achand45/gemma-3-12b-it-nla-L32
Gemma-3-12B-IT — Natural Language Autoencoder @ block 32
A natural language autoencoder (NLA) for google/gemma-3-12b-it: a pair of models that compress a residual-stream activation into English and back.
- AV (verbalizer) — the base model with a LoRA, reading the activation as an injected marker token (norm-matched, à la Karvonen et al.) and writing an
<explanation>…</explanation>of it. - AR (reconstructor) — the base model plus a linear head mapping the explanation text back to the activation vector. At this depth the AR is a 33-block truncation of the 48-block stack (
ar_num_layers = layer_index + 1).
Trained with EasyNLA: SFT warm-start, then on-policy GRPO where the AV is rewarded by how well the AR reconstructs the activation from its words.
Block 32 of 48, at ~2/3 depth. It is one arm of a depth study; the sibling repos `achand45/gemma-3-12b-it-nla-L40` and `achand45/gemma-3-12b-it-nla-L47` are the same pipeline, same data, same recipe at blocks 40 and 47, so the three are directly comparable to each other. A separately-trained NLA exists at this same layer and base model — `kitft/nla-gemma3-12b-L32-av` / `-ar` — but it is not comparable to these numbers: corpus, truncation points and training budget differ.
Results
Held-out, doc-disjoint (every row of a held-out document is excluded from training — a row-level split leaks badly here because the corpus is row-shuffled and each document contributes ~10 rows).
FVE = fraction of activation variance explained, against a predict-the-mean baseline. Extraction rate stayed at 100% for effectively the whole run (a single eval dipped to 99%; no format collapse), and KL from the SFT reference rose smoothly to ~0.97 with no runaway.
⚠️ Read the RL number as a band, not a point. Eval sampling runs at temperature 1.0 (eval_temperature: nullfalls back totemperature), and repeated evals of identical weights spread over ~5 points — measured on the sibling L47 run's step-0 checkpoint: 31.5 / 26.3 / 27.9 / 28.3. The reported ~68.6% is the mean of the final ten evals (steps 300–390, range 67.5–69.6%); the single best eval was 69.6% at step 370. Differences under ~5 points in this table are not resolvable.
Most of the RL gain lands in the first ~50 steps (49.9% → 63.3%); from ~step 200 (67.6%) the curve is flat, the last 200 steps moving it ~1 point, inside the noise band, while KL kept climbing. If you only want the trained model, the late checkpoints are interchangeable within measurement error.
Extraction contract
⚠️ The layer convention is off-by-one relative to naivehidden_states[K]indexing.layer_index=32hookslayers[32]and captures its output, which equalshidden_states[33](index 0 is the embedding output). Verified numerically during data generation: worst cosine 0.9999927 over 5 rows against an independentoutput_hidden_statesforward. Negative controls confirm the test discriminates — the neighbouring layers scorehidden_states[32]0.99746 andhidden_states[34]0.99779, and the wrong token position (second to last) 0.98442.
Because magnitude is discarded by the per-row rescale, FVE here measures direction only.
Every checkpoint ships an nla_meta.yaml sidecar carrying this contract (marker token ids, prompt templates, scales). The trainers assert against it — the AR's depth is derived from extraction.layer_index + 1, not hardcoded.
Repository layout
av_sft/iter_*/ AV warm-start LoRA (final: iter_0003834)
ar_sft/iter_*/ AR warm-start LoRA + value head (final: iter_0003834)
merged/av_hf/ AV SFT merged to bf16 HF (regenerable: merge_lora_to_hf.py)
merged/ar_hf/ AR SFT merged to bf16 HF (+ value_head.safetensors)
rl_vllm/iter_*/ RL AV LoRA every 25 steps (final: iter_000400)
adapter_model.safetensors — the trained policy
reference/ — frozen SFT copy, the KL reference
rl_vllm/critic_latest/ RL co-trained AR reconstructor (full weights + value_head)
rl_vllm/optim_latest.pt, run_config.yaml, nla_meta.yamlThe run_config.yaml files have had this cluster's absolute paths rewritten to repo-relative ones; their ./data/rl_shuf.parquet refers to the regenerated activation parquet, published separately as `achand45/gemma-3-12b-it-nla-data` (config L32_rl).
To use the trained NLA you need `rl_vllm/iter_000400` + `rl_vllm/critic_latest`. The RL adapter is a LoRA on the raw base model (RL continued the SFT adapter via --av-adapter), so it does not require merged/av_hf. The intermediate iter_* snapshots are included for training-dynamics and ablation work.
Usage
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import PeftModel
from nla.models import NLACriticModel # pip install -e git+https://github.com/chand-ab/easy_nla
from nla.utils.arch_adapters import resolve_text_model
BASE = "google/gemma-3-12b-it"
REPO = "<local clone of this repo>"
tok = AutoTokenizer.from_pretrained(BASE)
base = AutoModelForCausalLM.from_pretrained(BASE, dtype=torch.bfloat16)
av = PeftModel.from_pretrained(
resolve_text_model(base), # <-- REQUIRED on Gemma-3, see below
f"{REPO}/rl_vllm/iter_000400",
) # verbalizer: activation -> text
ar = NLACriticModel.from_pretrained(
f"{REPO}/rl_vllm/critic_latest", torch_dtype=torch.bfloat16,
) # reconstructor: text -> activationInject the activation at the marker token with register_karvonen_hook from nla.utils — note it hooks decoder layer 1's residual output (layer_idx=1), which is the injection site for this pipeline:
from nla.utils import register_karvonen_hook
from nla.config import load_nla_config
cfg = load_nla_config(f"{REPO}/rl_vllm/nla_meta.yaml")
vref = [None] # set vref[0] to the activation per generation
register_karvonen_hook(av, vref, cfg.injection_token_id,
cfg.injection_left_neighbor_id,
cfg.injection_right_neighbor_id, layer_idx=1)scripts/show_nla_generations.py in the repo is the closest end-to-end example. Two of its defaults are Qwen-shaped, so pass --base-ckpt google/gemma-3-12b-it and --skip-rows 449846 (the RL run's eval_skip_rows, so you score held-out rows).
Expected warning when loading the critic
Loading rl_vllm/critic_latest (or merged/ar_hf) prints:
Some weights of Gemma3ForCausalLM were not initialized from the model checkpoint
... and are newly initialized: ['model.norm.weight']This is expected and harmless. The AR is trained with its final RMSNorm stripped (final_norm_stripped: true) so the value head sees the raw layer-32 residual, so the checkpoint genuinely has no model.norm.weight. NLACriticModel.from_pretrained replaces that module with nn.Identity() immediately after loading, discarding the randomly-initialized tensor.
⚠️ Gemma-3 gotcha: resolve the text model before loading the adapter
gemma-3-12b-it loads as a multimodal wrapper that nests the language model under model.language_model.*, while these adapters are keyed on the text model's own model.layers.*. Loading them onto the unresolved wrapper matches zero keys, and PEFT will silently random-initialize the policy instead of erroring — you get a fluent model that is not this NLA. Always pass the resolved text model (resolve_text_model, or base.model.language_model equivalently).
Sanity check: with one adapter attached, total params should be 12,289,797,888 (11,766,034,176 text + 523,763,712 LoRA). If you see ~12.73B, the vision tower is still attached and the adapter is on the wrong module tree.
merged/av_hf/config.json deliberately declares `Gemma3ForCausalLM`, not Gemma3ForConditionalGeneration, so that vLLM routes to its text gemma3 implementation rather than the gemma3_mm path.
peftmust be<0.19(e.g.0.18.1) if you load adapters undertorch.distributed: peft 0.19'sset_peft_model_state_dictimportsEmbeddingParallelfromtransformers.integrations.tensor_parallel, which does not exist in transformers 4.57.x, and the call is guarded bydist.is_initialized()— so it breaks only in distributed runs.
Data provenance
The parquets these models were trained on are published as `achand45/gemma-3-12b-it-nla-data` — configs L32_av_sft, L32_ar_sft, L32_rl.
Text, prompts, and gold explanations come from `ceselder/qwen3-8b-nla-L24-finefineweb-100k` (corpus: `m-a-p/FineFineWeb`, 100k docs; explanations by Claude Sonnet 4.6).
Those fields are model-agnostic and were reused unchanged: the explanations describe the source text, and the prompts store an <INJECT> placeholder rather than a literal marker.
Two things were regenerated for Gemma:
- The activations, by forwarding
google/gemma-3-12b-itoverdetokenized_text_truncatedand taking the block-32 residual stream at the final token. - The sidecar `tokens` block, because Gemma-3 does not share Qwen's tokenizer — the marker character maps to a different id (
246566) and the neighbour ids differ. Note this means row truncation points were inherited from Qwen tokenization, so they do not fall on Gemma token boundaries; this is shared identically across the arms of the study, but is a confound against externally-trained checkpoints.
Training rows: 247,261 (AV SFT) / 247,358 (AR SFT) / 499,846 (RL). One epoch of SFT each.
Training setup
8×A100-80GB.
Both SFT stages together took 3 h 02 m. RL ran as 4 data-parallel ranks with per-rank vLLM rollouts at tp=2 (--vllm-gpu-mem 0.26) and the critic offloaded to a partner GPU, ~275 s/step, 30.8 h wall clock.
License
These weights are a Model Derivative of google/gemma-3-12b-it and are distributed under, and subject to, the [Gemma Terms of Use](https://ai.google.dev/gemma/terms) — not the MIT licence of the training code. A copy of the Agreement is included in this repository as LICENSE, and the required notice as NOTICE:
Gemma is provided under and subject to the Gemma Terms of Use found at ai.google.dev/gemma/terms
Use restrictions. Your use of these weights is subject to the Gemma Prohibited Use Policy, incorporated by reference into the Terms (§3.2). If you redistribute these weights or anything derived from them, you must pass these restrictions on to your recipients as an enforceable provision, supply them a copy of the Agreement, and give notice that the weights are subject to those restrictions (§3.1).
Modification notice (§3.1). The files under merged/ are modified Gemma weights: google/gemma-3-12b-it with a LoRA merged in, and — for merged/ar_hf/ and rl_vllm/critic_latest/ — the final RMSNorm stripped and a value_head added. The adapters under av_sft/, ar_sft/ and rl_vllm/iter_*/ are new weights trained by us, not modified Gemma files, but they only function when applied to Gemma and are Model Derivatives on the same terms.
Other components, for clarity. The training code (`chand-ab/easy_nla`) is MIT. The source text, prompts and gold explanations come from `ceselder/qwen3-8b-nla-L24-finefineweb-100k` (Apache-2.0), itself derived from `m-a-p/FineFineWeb`. Neither licence displaces the Gemma Terms for the weights published here.
