Use the ln2 scale for neurons in fused LN+projection resid decomposition - #1847
Zhuoxi2000 wants to merge 2 commits into
Conversation
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
left a comment
There was a problem hiding this comment.
Thanks for putting this together @Zhuoxi2000! Looks great, just one small note below
| ``[..., 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. |
There was a problem hiding this comment.
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.
|
Good catch, thanks! Done in 6536628. With New test: |
|
Rechecked the changed head |
Base branch:
devDescription
get_full_resid_decomposition(mlp_input=True, apply_ln=True, project_output_onto=p)normalized the head, embed, pos_embed and bias rows with theln2scale, through_ln_then_project/apply_ln_to_stack(mlp_input=...). The neuron rows instead went through the fused pathstack_neuron_results->_stack_neuron_results_apply_ln_projected, which always used theln1scale becausestack_neuron_resultshad nomlp_input. So the projected decomposition did not equal the unprojected one projected ontop, and it no longer summed toLN2(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_resultsgainsmlp_input: bool = Falseas its last parameter, so existing callers are unchanged. Withapply_ln=Trueit selects the LN scale of the input tolayer'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;apply_ln_to_stackcall on the fusedincl_remainderbranch, so the remainder stays consistent with the neurons;apply_ln_to_stackcall on the unfused branch.get_full_resid_decompositionpasses itsmlp_inputtostack_neuron_results.mlp_inputis documented onstack_neuron_results.Out of scope: the set of neurons stacked and the
incl_remainderbase (resid_post[layer - 1]) are unchanged.mlp_inputonly selects the LN scale, as inapply_ln_to_stack.This branch is rebased onto
devafter #1844. Both PRs touchstack_neuron_results, but the changes are independent: #1844 normalizesneuron_slice, and this PR threadsmlp_input.Fixes #1846
Type of change
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 equalsunprojected @ pand sums toapply_ln_to_stack(resid_mid, layer, mlp_input=True) @ p.test_stack_neuron_results_projected_ln_honours_mlp_input(with and withoutincl_remainder) checks that the fused and unfusedstack_neuron_results(mlp_input=True, apply_ln=True)paths agree.Before:
transformer_lens/ActivationCache.pyfromdev@ 0c254e6, with this branch's tests.The 4
get_full_resid_decompositioncases fail on the values. The 4stack_neuron_resultscases fail withTypeError: got an unexpected keyword argument 'mlp_input'because the parameter does not exist ondevyet. As a check that the second test covers the remainder branch, removingmlp_inputonly from theincl_remainderapply_ln_to_stackcall makes bothwith-remaindercases fail.After (this branch):
Real model (TinyStories-1M,
boot_transformers+enable_compatibility_mode(), layer 3,pos_slice=-1, randomp). The target isapply_ln_to_stack(resid_mid[:, -1][None], 3, mlp_input=True, pos_slice=-1)[0] @ p:dev@ 0c254e6max |projected - unprojected @ p|max |projected.sum(0) - target|max |(unprojected @ p).sum(0) - target|Note on the repro in #1846: its
targetline passes the fullresid_midwithpos_slice=-1, sotargetcovers every position, while the decomposition covers only the last one. As written, ondevits second and third prints give 6.90 and 11.67 instead of the 6.00 and 3.1e-06 shown there. Slicingresid_midto 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): themake check-formatcommands (pycln / isort / black, repo-wide) are clean, anduv run mypy .reports "Success: no issues found in 387 source files". black skipped.ipynbfiles because the Jupyter extra is not installed in thedev-only environment. The acceptance tests (tests/acceptance/test_activation_cache.py) were not run because they need distilgpt2, which is not cached locally.Checklist:
AI assistance: this change was drafted with an AI coding assistant (Claude) and verified locally with the tests above.