Gemma-3-12B-IT β€” Natural Language Autoencoder @ block 24

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 25-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 24 of 48, at half depth. It is one arm of a depth study; the sibling repos achand45/gemma-3-12b-it-nla-L32, -L40 and -L47 are the same pipeline, same data, same recipe at blocks 32, 40 and 47, so the four are directly comparable to each other.

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).

Stage Metric Value
AV SFT held-out val perplexity 3.635 (from 4.482 @ step 499)
AR SFT held-out FVE, on gold explanations 50.3% (MSE 0.0027 vs a 0.0054 baseline)
RL (GRPO, 400 steps) held-out FVE, on the AV's own explanations ~60.9% (45.7% at step 0)

FVE = fraction of activation variance explained, against a predict-the-mean baseline. Extraction rate was 100% at every eval β€” no format collapse β€” and KL from the SFT reference rose smoothly to ~1.21 with no runaway.

⚠️ Read the RL number as a band, not a point. Eval sampling runs at temperature 1.0 (eval_temperature: null falls back to temperature), 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 ~60.9% is the mean of the final ten evals (steps 300–390, range 59.9–62.1%); the single best eval was 62.1% at step 390. Differences under ~5 points in this table are not resolvable.

The gain is front-loaded β€” 45.7% β†’ 53.7% over the first 50 steps β€” but unlike this study's deeper arms the curve had not flattened at 400 steps: the last two evals (62.0% @380, 62.1% @390) are the run's two highest. The reported band is therefore a lower bound on what this recipe reaches at this depth, and iter_000400 is the checkpoint to prefer.

The baseline variance at this depth is small in absolute terms (0.0054, against 0.0310 at block 32), so FVE percentages are not on a common scale across layers even though they are on a common scale across the stages of this repo.

Extraction contract

Layer layer_index = 24 β€” the output of block 24 (of 48), i.e. HF hidden_states[25]
d_model 3840
AR depth ar_num_layers = 25 β€” the reconstructor is the first 25 blocks, a truncation
Normalization raw / unnormalized (norm: none); the AR's final RMSNorm is deliberately stripped (final_norm_stripped: true)
Loss-side scaling every row is rescaled to L2 norm √3840 = 61.9677, symmetrically on prediction and gold
Injection marker ㈜ (U+321C), token id 246566

⚠️ The layer convention is off-by-one relative to naive hidden_states[K] indexing. layer_index=24 hooks layers[24] and captures its output, which equals hidden_states[25] (index 0 is the embedding output). Verified numerically during data generation: worst cosine 0.9999985 over 5 rows against an independent output_hidden_states forward. Negative controls confirm the test discriminates β€” the neighbouring layers score hidden_states[24] 0.99924 and hidden_states[26] 0.99900.

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.yaml

The 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, which is not shipped in this repo (see Data provenance β€” it is reproducible from the source dataset with --cross-model regen).

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 -> activation

Inject 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 and is independent of the layer being read:

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-24 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.

peft must be <0.19 (e.g. 0.18.1) if you load adapters under torch.distributed: peft 0.19's set_peft_model_state_dict imports EmbeddingParallel from transformers.integrations.tensor_parallel, which does not exist in transformers 4.57.x, and the call is guarded by dist.is_initialized() β€” so it breaks only in distributed runs.

Data provenance

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). The L24 in that dataset's name refers to layer 24 of Qwen3-8B and is a coincidence β€” none of its activations were used here.

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:

  1. The activations, by forwarding google/gemma-3-12b-it over detokenized_text_truncated and taking the block-24 residual stream at the final token.
  2. 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.

AV SFT LoRA r=128 Ξ±=16 on all attn+MLP linears, lr 1e-4, batch 64, 3834 steps
AR SFT LoRA r=128 Ξ±=16 + value head, lr 2e-5, batch 64, 3834 steps, ar_num_layers=25
RL GRPO, 400 steps, batch 256 Γ— group 8, AV lr 1e-4 (r=128, rsLoRA) / AR lr 8e-5 (--ar-lora r=64), KL Ξ²=0.01 (k3), temp 1.0, max 256 new tokens

Both SFT stages together took 2 h 54 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, ~266 s/step, 29 h 39 m 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 β€” 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.

Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support

Model tree for achand45/gemma-3-12b-it-nla-L24

Adapter
(427)
this model
Adapters
1 model