Instructions to use Yooniel/gemma-3-12b-it-nla-L24-wildchat with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use Yooniel/gemma-3-12b-it-nla-L24-wildchat with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
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-itloads as a multimodal wrapper nesting the language model undermodel.language_model.*, while these adapters are keyed onmodel.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, andNLACriticModelreplaces that module withnn.Identity()immediately after loading.
peftmust be<0.19(e.g.0.18.1) if you load adapters undertorch.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
- -
Model tree for Yooniel/gemma-3-12b-it-nla-L24-wildchat
Base model
google/gemma-3-12b-pt