Skip to content

[Proposal] Backward Lens: generalize Backward Lens beyond GPT-2 (dense MLP families) #1777

Description

@janmenjayap

Related issues/PRs:

  • #1686 — original Backward Lens proposal; closed by #1723, which shipped the GPT-2-only vertical slice this issue generalizes.
  • #1751 — TransformerBridge NeoX/Pythia unembed resolved to stale embed_out instead of lm_head (transformers >= 5.13), blocking EleutherAI/pythia-70m from booting at all; fixed by #1752. This is a hard prerequisite for the Pythia-70m integration target below.

Proposal

Replace Backward Lens's hard-coded GPT-2/Conv1D requirement with a capability
contract over TransformerBridge
, so any decoder-only model whose MLP exposes dense
(non-gated) input/output linear components with resolvable original weights and a known
weight-storage orientation is supported — not just GPT-2. Ship
EleutherAI/pythia-70m (GPT-NeoX, torch.nn.Linear MLP) as the second required
integration-tested family, while keeping existing GPT-2 behavior byte-for-byte
compatible.


Motivation

Backward Lens (Katz, Belinkov, Geva, Wolf 2024) captures MLP weight-gradient factors and
projects residual-width factors into the vocabulary. #1686
requested this tool, and PR #1723
shipped it as a deliberately narrow GPT-2 vertical slice. I'm frustrated that the tool
now sits behind an artificial wall: the math it implements does not depend on GPT-2 at
all, yet the current code hard-rejects every non-GPT-2 model before that math ever runs.

The math is architecture-neutral for any dense (non-gated) MLP:

  • FF1 (input projection): grad_W_in = Σ_i outer(x_i, grad_pre_i)
  • FF2 (output projection): grad_W_out = Σ_i outer(hidden_i, grad_out_i)

Both are just sums of per-position outer products of the linear's captured forward input
and its output-side VJP. The only architecture-specific fact is weight storage
orientation
: GPT-2 stores Conv1D weights as [in, out] (weight_layout="in_out"),
whereas standard torch.nn.Linear models store [out, in] (weight_layout="out_in").
The tool already models both layouts — WeightLayout = Literal["in_out", "out_in"]
(backward_lens.py:19) and
_build_linear_gradient_factors transposes correctly for each
(backward_lens.py:296-355) —
but the capture path never exercises the out_in branch because it refuses non-GPT-2
models before it gets there.

Three coupling points reject non-GPT-2 models even though nothing below them needs
GPT-2:

  1. _require_raw_gpt2_bridge (backward_lens.py:549-586) —
    isinstance(model.adapter, GPT2ArchitectureAdapter) raises NotImplementedError for
    every other adapter. This is a class-name check, exactly what this follow-up asks us
    to remove.
  2. _get_gpt2_mlp_projections (backward_lens.py:589-635) —
    requires isinstance(projection.original_component, Conv1D) and hard-codes the two
    expected shapes as (d_model, d_mlp) and (d_mlp, d_model) (i.e. the [in, out]
    layout). A torch.nn.Linear model would store (d_mlp, d_model) and
    (d_model, d_mlp) respectively and fail the Conv1D check first.
  3. Hard-coded weight_layout="in_out" in _capture_gpt2_mlp_gradient_factors for
    both projections (backward_lens.py:798
    and backward_lens.py:810).
    For torch.nn.Linear families this must become "out_in", derived from the Bridge
    component rather than assumed.

Everything downstream — _build_linear_gradient_factors, _project_residual_factors
(which already uses live model.W_U / model.ln_final / model.unembed), rankings,
norms, and the whole result schema — is already model-agnostic.

This also depends on #1751:
the GPT-NeoX adapter used by Pythia previously resolved unembedding to a stale
embed_out attribute instead of lm_head (renamed upstream in transformers >= 5.13),
which blocked EleutherAI/pythia-70m from booting via TransformerBridge at all.
PR #1752 fixed that,
which is what makes Pythia-70m viable as the second integration target here.


Pitch

Replace the adapter-class and Conv1D checks with a contract expressed purely in terms
of TransformerBridge component capabilities. A model is supported when, for each
requested layer:

  • The block exposes a dense (non-gated) MLP bridge — its gate subcomponent is absent /
    None (the current gated-MLP rejection stays, deferred to Issue 2 — see
    issue-2-gated-mlp/).
  • The MLP exposes input and output linear components (in / out) that are
    LinearBridge instances with a resolvable original weight Parameter that is
    floating-point and requires_grad.
  • Each linear's weight orientation is known from the Bridge component
    ("in_out" vs "out_in"), not inferred from adapter class or model-name string.
  • Weight shapes are validated against (d_model, d_mlp) / (d_mlp, d_model) in the
    orientation the component declares
    , rather than assuming the GPT-2 [in, out] layout.
  • The model exposes the standard blocks, ln_final, and unembed components and a
    tokenizer, is a raw single-device Bridge, and has neither compatibility mode nor
    processed weights enabled (all existing guards retained unchanged).

Open design question (resolved in the implementation plan): phrasing this as
reading WeightLayout.IN_OUT / WeightLayout.OUT_IN "from the Bridge component" is
aspirational — in the current code WeightLayout is a private Literal in
backward_lens.py, not a Bridge concept. Whether the Bridge already exposes
orientation, or whether we derive it (e.g. Conv1D ⇒ in_out, nn.Linear ⇒
out_in) via a small capability probe that stays out of the analysis math, is settled
in the implementation plan based on the Bridge-component capability map.

Scope — in:

  • Replace the three GPT-2 coupling points above with the capability contract.
  • Keep existing GPT-2 public behavior byte-for-byte compatible — same shapes, same
    reconstruction identities, same public API surface and dataclass names.
  • Add EleutherAI/pythia-70m as the second required integration-tested family
    (small, cached, GPT-NeoX adapter, standard dense nn.Linear MLP), unblocked by
    #1752.
  • Keep projection using each model's own live final normalization and unembedding (no
    hard-coded ln_final / W_U path assumptions — already true, must stay true).
  • Preserve every state-safety guarantee: weights, .grad, requires_grad, train/eval
    mode, existing hooks, and CPU/CUDA/MPS RNG; temporary-hook cleanup on success and
    failure; one forward pass and exactly one torch.autograd.grad call.

Scope — explicitly out (each is its own follow-up):

  • Batched prompts and multi-token target losses (Issue 3).
  • Gated MLPs — Llama/SwiGLU gate/up/down math (Issue 2, see
    issue-2-gated-mlp/). The dense-MLP gate rejection stays in
    place.
  • EleutherAI/gpt-neo-125M, GPT-J, OPT, etc. — allowed to work if they satisfy the
    contract, but only Pythia-70M is a required tested family here. Neo-125M is optional
    and lower priority.

Acceptance criteria:

  • GPT-2 and Pythia-70M both reconstruct both MLP weight gradients at selected layers
    within the tolerance bands documented for GPT-2 today: absolute reconstruction error
    ≤ 2e-6 and relative reconstruction error ≤ 2e-5 (the bands asserted in
    tests/integration/test_backward_lens.py,
    atol=2e-6 / rtol=2e-5).
  • No model-name or class-name conditionals remain in the capture/validation path —
    only Bridge component capabilities. In particular, the
    isinstance(model.adapter, GPT2ArchitectureAdapter) and
    isinstance(..., Conv1D)-only checks are gone.
  • Unit tests remain synthetic / model-free where possible; integration tests use cached
    models only (no network download in CI beyond what is already cached).
  • Existing state-preservation and cleanup tests remain green, extended to cover the
    out_in (nn.Linear) path.
  • docs/source/content/backward_lens.md and
    demos/Backward_Lens_Demo.ipynb are updated to
    state the generalized dense-MLP contract without implying full architecture
    generality (still no gated / batched / multi-token support). Today the docs say
    "focused GPT-2 implementation" and the requirements section names GPT2ArchitectureAdapter
    and Conv1D explicitly (backward_lens.md:9-11, 154-164).

Model matrix:

Model Role Reason
gpt2 (openai-community/gpt2) Required — existing Conv1D / in_out; must stay byte-for-byte compatible
EleutherAI/pythia-70m Required — new Cached, small, GPT-NeoX adapter, standard dense nn.Linear / out_in MLP; unblocked by #1752
EleutherAI/gpt-neo-125M Optional Lower priority once Pythia passes
GPT-J, OPT, … Best-effort Should work if the contract holds; not required tests

Alternatives

  • Keep GPT-2-only and file per-model exceptions. Rejected: this repeats the same
    isinstance pattern for every new family and never removes the underlying coupling;
    it doesn't scale past two or three architectures.
  • Infer orientation from model name or adapter class instead of a capability probe.
    Rejected: it's exactly the class-name conditional this proposal is trying to remove,
    and it silently breaks for any adapter that mixes Conv1D and nn.Linear MLPs.
  • Wait for a first-class WeightLayout concept on TransformerBridge before doing
    anything.
    Deferred, not rejected: if the Bridge doesn't already expose orientation,
    the implementation plan allows a narrow, local capability probe now rather than
    blocking this issue on a larger Bridge API change.

Additional context

  • Do not promise GPT-2 behavior transfers to larger or gated models.
  • Vocabulary readability is a diagnostic, not causal evidence — unchanged framing.
  • No dependency on or vendoring of shacharKZ/BackwardLens; use in-repo algebraic /
    projection oracles for tests.

Risks:

  • Naming lock-in: the public schema uses GPT-2 input_projection / output_projection
    (FF1/FF2) naming. That is fine for dense MLPs across families but will need the
    gate/up/down distinction in Issue 2; this proposal must not paint the schema into a
    corner. It does not — dense in/out maps cleanly onto every target family here.
  • Orientation source of truth: if the Bridge does not already carry weight
    orientation, we must derive it without reintroducing a class-name check in the math path.
    Resolved in the implementation plan.
  • Pythia tokenizer / BOS differences: Pythia's tokenizer and BOS behavior differ from
    GPT-2; the single-token-target and prompt_token_ids-as-source-of-truth contracts must
    be re-verified per model, not assumed.

Checklist

  • I have checked that there is no similar issue in the repo (required)

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions