[Proposal] Backward Lens: support gated MLP gate/ up/ down gradient factors
Los mantenedores suelen responder en 1 día
Evaluación
Este issue todavía no se ha evaluado.
Descripción
Proposal
Backward Lens (Katz, Belinkov, Geva, Wolf 2024) factorizes an MLP weight gradient into a sum of
token-position outer products and projects the residual-width factors into the vocabulary. PR #1723
shipped the two-matrix dense case (GPT-2 c_fc/c_proj); #1777 generalizes that to any dense
LinearBridge MLP (adds Pythia/GPT-NeoX). This proposal adds gated (SwiGLU-style) MLPs, covering
the Llama/Mistral/Qwen2/Gemma family, whose MLP block has three linear projections instead of two.
The factorization math is unchanged per linear; the new work is a three-matrix result contract, the
correct choice of which factor is vocabulary-projectable for each matrix, and honest documentation of
why the gate and up projections coincide in the vocabulary while the down projection does not.
Only residual-width (d_model) factors can pass through ln_final + unembed
(_project_residual_factors enforces factors.shape[-1] == d_model,
transformer_lens/tools/analysis/backward_lens.py:421-424).
Classifying the six gated factors by width:
gate.forward_inputs = x: residual width, projectable (the "imprint" direction, the same
role GPT-2's FF1 input plays).in.forward_inputs = x: residual width, projectable, and it is the same tensor as
gate.forward_inputs(both read the MLP block inputx).gate.output_gradients = dL/dg:d_mlpwidth, not projectable.in.output_gradients = dL/du:d_mlpwidth, not projectable.out.forward_inputs = h:d_mlpwidth, not projectable.out.output_gradients = dL/dy: residual width, projectable (the "shift" direction, exactly
GPT-2's FF2 treatment).
Two consequences drive the contract:
- Down (
out) mirrors GPT-2 FF2 exactly: project its residual-width output gradient. Nothing new
conceptually. - Gate and up share one imprint direction. Because
gateandupboth consume the same residual
inputx, their vocabulary projections are identical by construction. The gate-vs-up
distinction is real but lives in (a) the two independent weight-gradient reconstructions and (b) the
two distinctd_mlp-width output gradients, not in the projected vocabulary. This must be
stated plainly so users do not read a nonexistent difference into two identical token tables. This
is the honest form of the acceptance criterion "documentation explicitly states why gate/up cannot
be treated identically to down/output": they are identical in the projection dimension and
different in the reconstruction dimension.
Keep the merged two-matrix schema and add one optional third matrix, so dense models stay
byte-for-byte compatible:
BackwardLensLayerResult(transformer_lens/tools/analysis/backward_lens.py:183-189)
gainsgate_projection: BackwardLensMatrixResult | None = Noneas a trailing, default-Nonefield.
Dense models (GPT-2, Pythia) leave itNone; gated models populate it.BackwardLensMatrixResultis unchanged: it already carriesfactors,projected_factor,
rankings, norms, target ranks, and optional full logits for any single linear. Its docstring's
"one GPT-2 MLP weight matrix" wording (already broadened to "dense MLP" by #1777) becomes "one MLP
weight matrix".- For gated layers,
analyze()produces three matrix results:gate_projectionandinput_projection
both withprojected_factor="forward_inputs"(the shared imprintx), andoutput_projectionwith
projected_factor="output_gradients"(the shiftdL/dy). All three carry independent reconstructions.
Motivation
A gated MLP block computes, for residual-width input x (shape [position, d_model]):
g = gate(x) # gate_proj: [position, d_mlp] (gate pre-activation)
u = up(x) # up_proj/in: [position, d_mlp] (linear component)
h = act_fn(g) * u # [position, d_mlp] (gated hidden)
y = down(h) # down_proj/out: [position, d_model] (block output)
This is exactly the structure documented on GatedMLPBridge
(transformer_lens/model_bridge/generalized_components/gated_mlp.py:55-62):
output = down_proj(act_fn(gate_proj(x)) * up_proj(x)), with submodules named gate (gate_proj),
in (up_proj), and out (down_proj).
Each of the three linears factorizes its own weight gradient the same way GPT-2's two do: a linear's
weight gradient is the sum of per-position outer products of its captured forward input and its
output-side VJP:
| Matrix | Forward input | Output VJP | grad_W (autograd) |
|---|---|---|---|
gate (gate_proj) |
x [pos, d_model] |
dL/dg [pos, d_mlp] |
sum_i outer(x_i, dL/dg_i) |
in (up_proj) |
x [pos, d_model] |
dL/du [pos, d_mlp] |
sum_i outer(x_i, dL/du_i) |
out (down_proj) |
h [pos, d_mlp] |
dL/dy [pos, d_model] |
sum_i outer(h_i, dL/dy_i) |
The existing _build_linear_gradient_factors
(transformer_lens/tools/analysis/backward_lens.py:296-355) already
performs this reconstruction for an arbitrary linear given (forward_inputs, output_gradients, weight_gradient, weight_layout). For Llama/Qwen2 the three projections are torch.nn.Linear, so the
Bridge weight-layout oracle weight_layout_in_out
(transformer_lens/model_bridge/generalized_components/mlp.py:15-29)
returns False -> "out_in" for all three, which #1777 already threads through capture. No new
factorization math is required: the change is capturing three linears instead of two and carrying a
third factor through the result.
Pitch
Yes, via the existing hook mechanism. GatedMLPBridge.forward in non-compatibility mode calls the HF
MLP as an opaque forward
(transformer_lens/model_bridge/generalized_components/gated_mlp.py:136-145),
and its docstring warns that the bridge's own intermediate hooks fire only in compatibility mode.
That warning is about GatedMLPBridge.forward explicitly calling self.gate.hook_out(...); it does
not govern the gate/in/out LinearBridge submodules' own hooks. Those submodules are
installed into the HF module tree by replace_remote_component
(transformer_lens/model_bridge/component_setup.py:99,
transformer_lens/model_bridge/component_setup.py:253-254) via setattr, so when
the opaque HF LlamaMLP.forward runs self.gate_proj(x) / self.up_proj(x) / self.down_proj(h), it
calls the LinearBridge wrappers, whose forward fires hook_in/hook_out. This is the identical
mechanism the merged GPT-2 capture already relies on for c_fc/c_proj. The interaction is
load-bearing and the docstring is easy to misread, so it is verified empirically by a test that
asserts all six gate/in/out hook_in/hook_out hooks fire exactly once during one gated
forward.
In scope:
- Capture, reconstruct, and independently validate all three gated weight gradients
(gate/up/down) via one grad-enabled forward and exactly onetorch.autograd.gradcall. - Add
gate_projectionto the public layer result; project gate/up (shared imprint) and down (shift). - Accept gated
TransformerBridgemodels in the public guard (remove the current gated rejection at
transformer_lens/tools/analysis/backward_lens.py:580-581, which
becomes the Issue-1-renamed guard by the time this lands). tiny-random-llama-2(structural, no Hub auth) andQwen2-0.5B(required real gated-model
reconstruction/projection); both are already in the CI cache
(tests/AGENTS.md:64).- Preserve every state-safety guarantee already asserted for GPT-2: weights,
.grad,requires_grad,
train/eval mode, existing hooks, CPU/CUDA/MPS RNG; temporary-hook cleanup on success and failure. - Keep GPT-2 and dense-MLP (Pythia) behavior byte-for-byte compatible.
Explicitly out of scope (each its own follow-up):
- Batched prompts and multi-token target losses.
- Non-SwiGLU exotic gating (MoE routing, GLU variants with fused gate+up matrices such as
joint_gate_up_mlp), recurrent/channel-mix MLPs, and any MLP whose linears the Bridge cannot orient.
Reject these with a clear error; do not silently mis-factor them. - Llama2-7B or any gated-access / large model in ordinary CI.
Qwen2-0.5Bis the required real model. - Weight editing / causal claims (separate editing discussion).
| Model | Role | Reason |
|---|---|---|
gpt2 (openai-community/gpt2) |
Required: regression | Dense Conv1D; must stay byte-for-byte compatible (gate_projection is None) |
EleutherAI/pythia-70m |
Required: regression | Dense nn.Linear; confirms dense path untouched |
trl-internal-testing/tiny-random-llama-2 |
Required: structural | Cached, no Hub auth; gate discovery and three-linear validation |
Qwen2-0.5B |
Required: real gated reconstruction/projection | Cached, ungated access, real tokenizer, SwiGLU gated MLP (_gated_mlp()); the full analyze() path runs here |
Llama2-7B |
Out of scope | Paper-scale only; gated access + memory; never in CI |
- On
Qwen2-0.5B,BackwardLens(model).analyze(...)returns, per requested layer, three matrix
results (gate_projection,input_projection,output_projection), and all three weight
gradients reconstruct independently against autograd within the GPT-2 tolerance bands
(atol=2e-6,rtol=2e-5;absolute_reconstruction_error <= 2e-6,
relative_reconstruction_error <= 2e-5, matching the current assertions in
tests/integration/test_backward_lens.py). gate_projectionandinput_projectionproduce identical projected logits for a gated model
(shared residual inputx), asserted by a test, and this is documented as expected.- GPT-2 and Pythia results are unchanged:
gate_projection is None, identical shapes, identical
reconstruction identities, identical public API surface. - No new required model download in CI beyond the already-cached
tiny-random-llama-2/Qwen2-0.5B. - Exactly one forward pass and one
torch.autograd.gradcall on the gated path, verified by test; all
state-preservation and hook-cleanup tests pass on the gated bridge too. docs/source/content/backward_lens.mdanddemos/Backward_Lens_Demo.ipynbstate the gated
gate/up/down contract, including the gate/up projection coincidence and the down shift, without
implying full architecture generality (no batched/multi-token/MoE/fused-gate support).
Alternatives
This deliberately does not rename input_projection/output_projection to up/down. The names
already map cleanly (in->up, out->down), renaming would break the merged public API for no gain, and
adding gate_projection is sufficient to "distinguish which matrix a factor belongs to."
Additional context
- Fused gate+up matrices. Some architectures store gate and up as one matrix
(joint_gate_up_mlp.pyexists in the component set). Ifmlp.gateis absent but the MLP is still
gated via a fused projection, the three-linear discovery will not find a separate gateLinearBridge.
Boundary: only support MLPs exposing a distinctgateLinearBridge; reject fused-gate MLPs with
a clear error (own follow-up). VerifyQwen2-0.5Buses distinctgate_proj/up_proj(it does:
_gated_mlp()wiresgate/in/out). - Tokenizer on
tiny-random-llama-2. The fullanalyze()path needsmodel.to_tokensand a
single-token target, which requires a tokenizer. Ifboot_transformers("tiny-random-llama-2")has no
usable tokenizer, keep tiny-llama for structural tests only (discovery, gate acceptance, three
linears, state-safety via the private capture with explicit token ids) and makeQwen2-0.5Bcarry
the reconstruction + projection assertions. - Activation function correctness is irrelevant to reconstruction. Because each linear's gradient is
captured by autograd on the real forward, the implementation never re-derivesact_fn; the elementwise gate is
inside the opaque forward and the captured VJPs already account for it.resolve_activation_fn
(transformer_lens/model_bridge/generalized_components/gated_mlp.py:15-46)
is only used by compatibility-mode processed weights, which Backward Lens forbids; do not depend on
it. - Naming lock-in. Adding
gate_projectionrather than renaming preserves the merged API and leaves
room for future batching support to extend shapes without touching matrix identity. - Licensing/scientific framing unchanged. Vocabulary readability is diagnostic, not causal; do not
vendorshacharKZ/BackwardLens; use in-repo algebraic/projection oracles; do not promise GPT-2
behavior transfers to gated models.
Checklist
- I have checked that there is no similar issue in the repo (required)
Suggested labels: enhancement, TransformerBridge, complexity-moderate
- Lenguaje dominante
- Python
- Estrellas
- 3.9k
- Forks
- 708
- Merge medio
- 1 d 18 h
- PR fusionados (30 d)
- 65
Preparar el entorno
Inicia el contenedor de desarrollo del proyecto en tu navegador, con tu propia cuenta de GitHub.
- Sin Dockerfile ni archivo de Docker Compose
- Tiene una plantilla de pull request
- Sin guía de contribución
Primeros pasos
- Lee el issue completo y luego la guía de contribución del proyecto.
- Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
- Haz un fork del repositorio y trabaja en una rama.
- Abre un pull request que haga referencia al número del issue.
Más de TransformerLensOrg/TransformerLens
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 62/100
TransformerLensOrg/TransformerLens#1868 ·
Los mantenedores suelen responder en 1 día
-
[Bug Report] run_with_hooks(remove_batch_dim=True) raises AttributeError when a hook returns NonePosiblemente ocupada @Mudassiruddin7 la tomó hace 1 día. Abiertobug complexity-simple TransformerBridge
TransformerLensOrg/TransformerLens#1863 · 1 asignado ·
Los mantenedores suelen responder en 1 día
-
[Bug Report] `remove_batch_dim=True` with batch `size > 1`: the three caching paths disagreePosiblemente ocupada @Mudassiruddin7 la tomó hace 1 día. Abiertobug complexity-simple TransformerBridge
TransformerLensOrg/TransformerLens#1858 · 3 comentarios · 1 asignado ·
Los mantenedores suelen responder en 1 día
-
[Bug Report] `get_caching_hooks` fires for alias names in `names_filter` but caches only under the canonical keyPosiblemente ocupada @Mudassiruddin7 la tomó hace 1 día. Abiertobug complexity-moderate TransformerBridge
TransformerLensOrg/TransformerLens#1857 · 2 comentarios · 1 asignado ·
Los mantenedores suelen responder en 1 día
-
[Bug Report] get_caching_hooks slices the wrong axis when remove_batch_dim=True is combined with pos_slicePosiblemente ocupada @JoeyTan21 la tomó hace 2 días. Abierto
TransformerLensOrg/TransformerLens#1856 · 1 comentario · 1 asignado ·
Los mantenedores suelen responder en 1 día
Todos los issues de TransformerLensOrg/TransformerLens
Issues similares
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
Los mantenedores suelen responder en 1 día
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 76/100
rpm-software-management/mock#1824 ·
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 75/100
jpata/particleflow#520 ·
Los mantenedores suelen responder en 1 día
-
bug good first issue hacktoberfest
Dificultad 1/5 Menos de una hora Aptitud para principiantes 78/100
gridhead/gi-loadouts#699 ·
Los mantenedores suelen responder en 13 días
-
Dificultad 1/5 Menos de una hora Aptitud para principiantes 86/100
FinanceFlash/unvibecode#206 ·
Los mantenedores suelen responder en 1 día