CoolFace
Modelpublic

achand45/gemma-3-12b-it-nla-L32

sourceHugging Facegemmaupdated 1mo agoView on Hugging Face
0likes
Model Card

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

StageMetricValue
AV SFTheld-out val perplexity3.783 (from 4.659 @ step 499)
AR SFTheld-out FVE, on gold explanations61.0% (MSE 0.0121 vs a 0.0310 baseline)
RL (GRPO, 400 steps)held-out FVE, on the AV's own explanations~68.6% (49.9% at step 0)

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: 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 ~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

Base modelgoogle/gemma-3-12b-it
Layerlayer_index = 32 — the output of block 32 (of 48), i.e. HF hidden_states[33]
d_model3840
AR depthar_num_layers = 33 — the reconstructor is the first 33 blocks, a truncation
Normalizationraw / unnormalized (norm: none); the AR's final RMSNorm is deliberately stripped (final_norm_stripped: true)
Loss-side scalingevery 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=32 hooks layers[32] and captures its output, which equals hidden_states[33] (index 0 is the embedding output). Verified numerically during data generation: worst cosine 0.9999927 over 5 rows against an independent output_hidden_states forward. Negative controls confirm the test discriminates — the neighbouring layers score hidden_states[32] 0.99746 and hidden_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.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, 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

python
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:

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

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

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:

  1. 1.The activations, by forwarding google/gemma-3-12b-it over detokenized_text_truncated and taking the block-32 residual stream at the final token.
  2. 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 SFTLoRA r=128 α=16 on all attn+MLP linears, lr 1e-4, batch 64, 3834 steps
AR SFTLoRA r=128 α=16 + value head, lr 2e-5, batch 64, 3834 steps, ar_num_layers=33
RLGRPO, 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 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.