Skip to content

Use the ln2 scale for neurons in fused LN+projection resid decomposition - #1847

Open
Zhuoxi2000 wants to merge 2 commits into
TransformerLensOrg:devfrom
Zhuoxi2000:fix-resid-decomp-mlp-input-ln-scale
Open

Zhuoxi2000 wants to merge 2 commits into
TransformerLensOrg:devfrom
Zhuoxi2000:fix-resid-decomp-mlp-input-ln-scale

Conversation

@Zhuoxi2000

Copy link
Copy Markdown

Base branch: dev

Description

get_full_resid_decomposition(mlp_input=True, apply_ln=True, project_output_onto=p) normalized the head, embed, pos_embed and bias rows with the ln2 scale, through _ln_then_project / apply_ln_to_stack(mlp_input=...). The neuron rows instead went through the fused path stack_neuron_results -> _stack_neuron_results_apply_ln_projected, which always used the ln1 scale because stack_neuron_results had no mlp_input. So the projected decomposition did not equal the unprojected one projected onto p, and it no longer summed to LN2(resid_mid) @ p. On TinyStories-1M at layer 3 the max error was 2.5. The fused path came in with #1300.

Changes:

  • stack_neuron_results gains mlp_input: bool = False as its last parameter, so existing callers are unchanged. With apply_ln=True it selects the LN scale of the input to layer's MLP (ln2) instead of its attention input (ln1). It is passed to:
    • _stack_neuron_results_apply_ln_projected, and from there to _get_cached_ln_scale;
    • the apply_ln_to_stack call on the fused incl_remainder branch, so the remainder stays consistent with the neurons;
    • the apply_ln_to_stack call on the unfused branch.
  • get_full_resid_decomposition passes its mlp_input to stack_neuron_results.
  • The helper's docstring note that it "always uses the ln1 scale" is replaced, and mlp_input is documented on stack_neuron_results.

Out of scope: the set of neurons stacked and the incl_remainder base (resid_post[layer - 1]) are unchanged. mlp_input only selects the LN scale, as in apply_ln_to_stack.

This branch is rebased onto dev after #1844. Both PRs touch stack_neuron_results, but the changes are independent: #1844 normalizes neuron_slice, and this PR threads mlp_input.

Fixes #1846

Type of change

  • Bug fix (non-breaking change which fixes an issue)

Test evidence

Two regression tests were added at the end of tests/unit/test_activation_cache.py, both using the existing LN/RMS module fixture (2 layers, d_model=16, built natively, no download):

  • test_full_resid_decomposition_projected_ln_uses_mlp_input_scale (vector and matrix projection) checks that the projected decomposition equals unprojected @ p and sums to apply_ln_to_stack(resid_mid, layer, mlp_input=True) @ p.
  • test_stack_neuron_results_projected_ln_honours_mlp_input (with and without incl_remainder) checks that the fused and unfused stack_neuron_results(mlp_input=True, apply_ln=True) paths agree.

Before: transformer_lens/ActivationCache.py from dev @ 0c254e6, with this branch's tests.

$ uv run pytest tests/unit/test_activation_cache.py -q
FAILED ...test_full_resid_decomposition_projected_ln_uses_mlp_input_scale[LN-vector]
  Mismatched elements: 767 / 900 (85.2%)
  Greatest absolute difference: 0.4491734504699707 at index (23, 0, 0) (up to 1e-05 allowed)
FAILED ...[LN-matrix] / [RMS-vector] / [RMS-matrix]   (max abs diff 0.514 / 0.441 / 0.666)
FAILED ...test_stack_neuron_results_projected_ln_honours_mlp_input[LN|RMS-neurons|with-remainder]
8 failed, 108 passed

The 4 get_full_resid_decomposition cases fail on the values. The 4 stack_neuron_results cases fail with TypeError: got an unexpected keyword argument 'mlp_input' because the parameter does not exist on dev yet. As a check that the second test covers the remainder branch, removing mlp_input only from the incl_remainder apply_ln_to_stack call makes both with-remainder cases fail.

After (this branch):

$ uv run pytest tests/unit/test_activation_cache.py -q
116 passed

Real model (TinyStories-1M, boot_transformers + enable_compatibility_mode(), layer 3, pos_slice=-1, random p). The target is apply_ln_to_stack(resid_mid[:, -1][None], 3, mlp_input=True, pos_slice=-1)[0] @ p:

dev @ 0c254e6 this PR
max |projected - unprojected @ p| 2.51 9.5e-07
max |projected.sum(0) - target| 6.00 3.1e-06
max |(unprojected @ p).sum(0) - target| 3.1e-06 3.1e-06

Note on the repro in #1846: its target line passes the full resid_mid with pos_slice=-1, so target covers every position, while the decomposition covers only the last one. As written, on dev its second and third prints give 6.90 and 11.67 instead of the 6.00 and 3.1e-06 shown there. Slicing resid_mid to the last position, as in the table above, reproduces the numbers in the issue. The first print (2.51) is unaffected.

uv run pytest transformer_lens/ActivationCache.py --doctest-modules: 5 passed.

Lint, using the versions from uv.lock (black 23.12.1, isort 5.8.0, pycln 2.5.0, mypy 1.17.0): the make check-format commands (pycln / isort / black, repo-wide) are clean, and uv run mypy . reports "Success: no issues found in 387 source files". black skipped .ipynb files because the Jupyter extra is not installed in the dev-only environment. The acceptance tests (tests/acceptance/test_activation_cache.py) were not run because they need distilgpt2, which is not cached locally.

Checklist:

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

AI assistance: this change was drafted with an AI coding assistant (Claude) and verified locally with the tests above.

get_full_resid_decomposition(mlp_input=True, apply_ln=True,
project_output_onto=p) normalised head, embed and bias rows with the ln2
scale but routed neurons through the fused LN+projection path, which
always used the ln1 scale. The result therefore differed from the
unprojected decomposition projected onto p and no longer summed to
LN2(resid_mid) @ p.

Expose mlp_input on stack_neuron_results (default False, so existing
callers are unchanged), thread it into the fused helper and both
apply_ln_to_stack calls, and pass it from get_full_resid_decomposition.

Fixes TransformerLensOrg#1846

@jlarson4 jlarson4 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for putting this together @Zhuoxi2000! Looks great, just one small note below

Comment thread transformer_lens/ActivationCache.py Outdated
``[..., d_mlp, d_model]`` intermediate is still never materialized.
mlp_input:
With ``apply_ln=True``, normalize with the LN scale of the input to ``layer``'s MLP
(ln2) instead of its attention (ln1). The neurons stacked are unchanged.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

With mlp_input=True, the remainder still fills up to resid_pre, so the stack sums to the attention input scaled by ln2, which neither LayerNorm ever sees. Elsewhere in this class mlp_input moves the target to resid_mid (see decompose_resid and accumulated_resid). Can the remainder follow that convention, so the stack sums to the normalized MLP input?

…t=True

Follow the decompose_resid and accumulated_resid convention: with mlp_input=True the stack is the input to layer's MLP, so the remainder fills it up to resid_mid instead of resid_pre, and the stack sums to the (normalized) MLP input.
@Zhuoxi2000

Copy link
Copy Markdown
Author

Good catch, thanks! Done in 6536628. With mlp_input=True, the remainder now fills the stack up to resid_mid instead of resid_pre, following decompose_resid / accumulated_resid. The change covers all three remainder paths: the fused LN+projection path, the unfused path, and the no-neurons case. The mlp_input docstring says this now.

New test: test_stack_neuron_results_mlp_input_remainder_fills_to_resid_mid, parameterized over apply_ln and the projection, on both the LN and RMS fixtures. It checks that the stack sums to resid_mid, to LN2(resid_mid), or to their projection. All 8 cases fail on the previous commit and pass now. test_activation_cache.py passes (124), and isort/black/mypy are clean.

@koriyoshi2041

koriyoshi2041 commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

Rechecked the changed head 6536628b locally. The full activation-cache unit file passes (124/124), and the new remainder test exercises raw/LN × projected/unprojected paths. I also traced the three implementation branches: each now derives its remainder from resid_mid[layer] when mlp_input=True, while the existing mlp_input=False target remains resid_post[layer - 1]. This closes the semantic gap from the earlier review without changing the neuron stack itself.

This branch has not been deployed

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants