Instructions to use achand45/gemma-3-12b-it-nla-L24 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use achand45/gemma-3-12b-it-nla-L24 with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
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: 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 ~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=24hookslayers[24]and captures its output, which equalshidden_states[25](index 0 is the embedding output). Verified numerically during data generation: worst cosine 0.9999985 over 5 rows against an independentoutput_hidden_statesforward. Negative controls confirm the test discriminates β the neighbouring layers scorehidden_states[24]0.99924 andhidden_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.
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
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:
- The activations, by forwarding
google/gemma-3-12b-itoverdetokenized_text_truncatedand taking the block-24 residual stream at the final token. - The sidecar
tokensblock, 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
- -