Gemma-3-12B-IT NLA @ block 24 — WildChat-adapted

A natural language autoencoder (NLA) for google/gemma-3-12b-it at block 24, adapted from web text to conversational activations with 400 further steps of on-policy GRPO on WildChat-1M.

This is a continuation, not a from-scratch model. It resumes achand45/gemma-3-12b-it-nla-L24 at rl_vllm/iter_000400 and keeps training on chat activations. The AV is the same LoRA lineage (r=128, rsLoRA, continued via --av-adapter); the AR is that release's co-trained critic, further co-trained here.

It is the half-depth arm of a two-layer study. The sibling Yooniel/gemma-3-12b-it-nla-L32-wildchat is the same procedure at block 32, run against the same corpus file and therefore the same held-out documents.

Why

An NLA trained on FineFineWeb is measured on web-text activations. Chat activations are a different distribution, and the parent model loses a lot of ground on them. This run quantifies that drop at half depth and how much of it RL recovers.

FineFineWeb WildChat
parent AV @ iter_000400 ~60.9% (their number) 46.5%
this model, +400 steps on chat — 55.6%

So the parent drops ~14pp when pointed at chat activations, and 400 steps of WildChat GRPO recover +9.1pp — roughly two-thirds of the transfer gap.

The depth comparison

Run side by side with block 32, the two depths gained almost exactly the same amount. The whole difference between them is where they start.

start (step 400) final (step 790) gain
block 24 (this model) 46.5% 55.6% +9.1pp
block 32 (sibling) 55.1% 64.3% +9.2pp

Block 24's parent transfers to chat activations ~8.6pp worse, and 400 WildChat steps close none of that gap — they shift both curves up in parallel. Note the coincidence that block 24 ends at 55.6%, essentially where block 32 began.

⚠️ FVE is not on a common scale across depths. The predict-the-mean baseline is 0.0049 here against 0.0299 at block 32 — block 24's residual stream is ~6× more concentrated, so a given FVE percentage is not the same amount of recovered structure at the two depths. The gains (+9.1 vs +9.2pp) are each measured against their own baseline and are the comparable quantity; the absolute levels are not.

Results

Held-out, doc-disjoint (every row of a held-out document is excluded from training; the corpus is row-shuffled and each document contributes ~10 rows, so a row-level split leaks).

Metric Value
held-out FVE @ step 400 (start) 46.5%
held-out FVE @ step 790 (final) 55.6%
best single eval 55.6% @ step 790 (the final eval)
mean of final 10 evals 54.7% (range 53.9–55.6)
extraction rate 100% on all 40 evals

Predict-the-mean baseline on this eval set: 0.0049 (train-set baseline 0.0055).

⚠️ Read these as a band. Eval sampling runs at temperature 1.0, and the parent release measured repeated evals of identical weights spreading ~5 points. The +9.1pp gain clears that; differences under ~5 points here do not.

The curve is front-loaded — 46.5% → 51.1% over the first 50 steps — then climbs slowly. Evals 600–690 averaged 54.0% and evals 700–790 averaged 54.7%, so the last 100 steps moved +0.7pp: flattened, but not fully stopped. The best eval landing on the very last step is inside the noise band, not evidence of remaining headroom.

Entropy rose 1.76 → ~2.1 and settled; KL from the SFT reference rose smoothly 0.99 → ~1.09. 400 steps in 16h58m on 4×H100 at ~152 s/step. No non-finite gradients, no format collapse, no OOMs.

⚠️ Not a like-for-like continuation of the parent

The parent ran GRPO at batch_prompts: 256; this run used 128. Everything else matches — lr 1e-4, critic lr 8e-5, KL β=0.01 (k3), group 8, temp 1.0, 256 max new tokens. So this run saw 51,200 prompts against the parent's 102,400 per equivalent step count, with noisier per-step gradients. Absolute FVE here is not directly comparable to the parent's published figures. It is comparable to the block-32 sibling, which used the same 128.

Extraction contract

Unchanged from the parent — the sidecar (rl_wildchat/nla_meta.yaml) is authoritative and the trainers assert against it.

Base model google/gemma-3-12b-it
Layer layer_index = 24 — output of block 24 of 48, i.e. HF hidden_states[25]
d_model 3840
AR depth ar_num_layers = 25, final RMSNorm stripped
Normalization raw (norm: none); loss rescales each row to L2 norm √3840 = 61.9677
Injection marker ㈜ (U+321C), token id 246566

Because magnitude is discarded by the per-row rescale, FVE measures direction only.

Training data

The same 20,000 WildChat-1M conversations used for the block-32 sibling (English, non-toxic, ≥400 chars, selected from 45,560 scanned), rendered with the Gemma-3 chat template and forwarded to get block-24 residuals at 10 sampled token positions each → 200,000 rows, ~150,000 after the doc-disjoint holdout.

The two depths point at the same corpus file on disk, deliberately. Document ids are derived as {corpus path}:{split}:{row index} and the holdout is a hash of that id, so a byte-identical copy at a different path would have produced a different held-out set and made the depths non-comparable.

Reproducibility detail worth knowing: apply_chat_template emits a leading <bos>, and the extraction path tokenizes with add_special_tokens=True, which adds another — giving [2, 2, 105, ...]. The template's BOS is stripped before extraction so each document carries exactly one. Position sampling skips special tokens, so double-BOS would not corrupt the vectors directly, but every activation would be conditioned on a prefix the model never sees at inference.

The activation parquet is not published here. It carries detokenized_text_truncated — verbatim WildChat conversations, which are real user–chatbot logs. Redistributing that is a separate decision under WildChat's own terms, not something this model card covers.

Contents

rl_wildchat/iter_000425 … iter_000800/   AV LoRA every 25 steps (16 checkpoints)
                             adapter_model.safetensors  the trained policy
                             reference/                 frozen SFT copy, the KL reference
rl_wildchat/critic_latest/  AR reconstructor (full weights + value_head)
rl_wildchat/nla_meta.yaml, run_config.yaml, optim_latest.pt

To use the NLA you need one iter_* adapter + critic_latest. Use iter_000800, the final step. Evals run every 10 steps but checkpoints save every 25, so the best eval (55.6% @790) has no exact checkpoint of its own; the final 10 evals span 53.9–55.6%, well inside the noise band, so 800 is the defensible pick.

The full every-25-step series is included for training-dynamics and ablation work. The RL adapter is a LoRA on the raw base model (RL continued the SFT adapter via --av-adapter), so no merged AV is required.

Usage

Identical to the parent — same contract, same gotchas.

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_wildchat/iter_000800",
)
ar = NLACriticModel.from_pretrained(
    f"{REPO}/rl_wildchat/critic_latest", torch_dtype=torch.bfloat16,
)

Inject the activation at the marker with register_karvonen_hook from nla.utils; it hooks decoder layer 1's residual output (layer_idx=1).

⚠️ Gemma-3 gotcha. gemma-3-12b-it loads as a multimodal wrapper nesting the language model under model.language_model.*, while these adapters are keyed on model.layers.*. Loading onto the unresolved wrapper matches zero keys and PEFT silently random-initializes instead of erroring — you get a fluent model that is not this NLA. Always pass the resolved text model. Sanity check: with one adapter attached, total params should be 12,289,797,888.

Loading the critic prints ... newly initialized: ['model.norm.weight']. Expected — the AR is trained with its final RMSNorm stripped, and NLACriticModel replaces that module with nn.Identity() immediately after loading.

peft must be <0.19 (e.g. 0.18.1) if you load adapters under torch.distributed.

Attribution

  • Parent NLA: achand45/gemma-3-12b-it-nla-L24 — the SFT warm-start and the first 400 RL steps are theirs; this repo continues them.
  • Training code: EasyNLA (MIT).
  • Conversations: allenai/WildChat-1M, subject to its own terms.
  • Method: Natural Language Autoencoders (Anthropic, 2026); Karvonen et al. norm-matched activation injection.

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 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). rl_wildchat/critic_latest/ is modified Gemma weights: google/gemma-3-12b-it truncated to its first 25 blocks, with the final RMSNorm stripped and a value_head added, further trained by us. The adapters under rl_wildchat/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.

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

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

Adapter
(1)
this model

Collection including Yooniel/gemma-3-12b-it-nla-L24-wildchat