Hacktoberfest 2026: los issues que los mantenedores marcaron para octubre, abiertos y aptos para principiantes. Explorar issues de Hacktoberfest

[Proposal] Backward Lens: support gated MLP gate/ up/ down gradient factors

Abierto
#1,832 0 comentarios 0 reacciones 1 asignado Ver en GitHub

Los mantenedores suelen responder en 1 día

@janmenjayap ya está trabajando en esto.

Desde el 1/10/2026.

  • #1833 de @janmenjayap — abierto

Evaluación

Este issue todavía no se ha evaluado.

Descripción

complexity-moderate enhancement TransformerBridge

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 input x).
  • gate.output_gradients = dL/dg: d_mlp width, not projectable.
  • in.output_gradients = dL/du: d_mlp width, not projectable.
  • out.forward_inputs = h: d_mlp width, not projectable.
  • out.output_gradients = dL/dy: residual width, projectable (the "shift" direction, exactly
    GPT-2's FF2 treatment).

Two consequences drive the contract:

  1. Down (out) mirrors GPT-2 FF2 exactly: project its residual-width output gradient. Nothing new
    conceptually.
  2. Gate and up share one imprint direction. Because gate and up both consume the same residual
    input x, 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 distinct d_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)
    gains gate_projection: BackwardLensMatrixResult | None = None as a trailing, default-None field.
    Dense models (GPT-2, Pythia) leave it None; gated models populate it.
  • BackwardLensMatrixResult is unchanged: it already carries factors, 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_projection and input_projection
    both with projected_factor="forward_inputs" (the shared imprint x), and output_projection with
    projected_factor="output_gradients" (the shift dL/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 one torch.autograd.grad call.
  • Add gate_projection to the public layer result; project gate/up (shared imprint) and down (shift).
  • Accept gated TransformerBridge models 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) and Qwen2-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.5B is 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_projection and input_projection produce identical projected logits for a gated model
    (shared residual input x), 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.grad call 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.md and demos/Backward_Lens_Demo.ipynb state 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.py exists in the component set). If mlp.gate is absent but the MLP is still
    gated via a fused projection, the three-linear discovery will not find a separate gate LinearBridge.
    Boundary: only support MLPs exposing a distinct gate LinearBridge; reject fused-gate MLPs with
    a clear error (own follow-up). Verify Qwen2-0.5B uses distinct gate_proj/up_proj (it does:
    _gated_mlp() wires gate/in/out).
  • Tokenizer on tiny-random-llama-2. The full analyze() path needs model.to_tokens and a
    single-token target, which requires a tokenizer. If boot_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 make Qwen2-0.5B carry
    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-derives act_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_projection rather 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
    vendor shacharKZ/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

Abrir en Codespaces

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

  1. Lee el issue completo y luego la guía de contribución del proyecto.
  2. Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
  3. Haz un fork del repositorio y trabaja en una rama.
  4. Abre un pull request que haga referencia al número del issue.

Más de TransformerLensOrg/TransformerLens

Todos los issues de TransformerLensOrg/TransformerLens

Issues similares

Más issues de Python

Recibe los nuevos issues en tu correo

Un resumen breve de issues de GitHub para principiantes.