Skip to content

fix: align per-example logit attribution across positions - #1835

Merged
jlarson4 merged 3 commits into
TransformerLensOrg:devfrom
nehal-a2z:nehal/fix-batched-logit-attribution
Oct 6, 2026
Merged

jlarson4 merged 3 commits into
TransformerLensOrg:devfrom
nehal-a2z:nehal/fix-batched-logit-attribution

Conversation

@nehal-a2z

@nehal-a2z nehal-a2z commented Sep 28, 2026 •

Copy link
Copy Markdown

Description

Passing one target token per example to logit_attrs can silently change the attribution when the residual stack retains multiple positions. The [batch, d_model] directions broadcast against the position axis: equal batch and position sizes produce wrong values, while unequal sizes can raise.

This adds a position axis to per-example directions when both batch and position axes remain. Batched attribution now matches analyzing each example separately. A 1-D vector is treated as per-example only when its length matches the retained batch size; otherwise it keeps its shared per-position meaning. Scalar targets, explicit target grids, batchless inputs, and scalar slices keep their behavior. The docstring clarifies that target tensors already match the selected slices. No new dependencies.

Type of change

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

Checklist

  • Commented the shape-sensitive operation and updated its docstring.
  • Added regression coverage without rewriting existing interface tests.
  • New and adjacent cache/DLA unit tests pass locally.
  • Changed-file Black, isort, pycln, and diff checks pass.
  • Full repository typecheck and test suite.

Validation

226 focused cache/DLA tests pass with runtime jaxtyping on current dev (6affbe99). Tests cover LN/RMS, batch/position slicing, logit differences, and single-prompt/shared per-position vectors. The new compatibility cases fail 8/8 on the previous PR head and pass with the batch-length guard. Original batched-attribution regressions also remain covered.

Changed-file Black, isort, pycln and diff checks pass. The scoped mypy run with imports followed reports the same 10 diagnostics in unchanged tests on both current dev and this patch; none are in the new tests. This is not a full typecheck pass.

Tested on CPU with PyTorch 2.7.1 and Transformers 4.57.6 in the existing personal environment. Full pinned-environment and pretrained-model integration tests were not run locally; CI runs on the pushed revision.

Sent from Autobox

@nehal-a2z
nehal-a2z marked this pull request as ready for review September 28, 2026 13:57

@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 this. Checking batched results against separate examples is a solid way to test the fix. One comment on how 1-D target tensors are read.

and logit_directions.ndim == 2
):
# Per-example directions must broadcast over positions, not across the batch.
logit_directions = logit_directions.unsqueeze(-2)

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.

A 1-D target tensor whose length doesn't match the batch size used to mean one token per position, and this branch now treats it as per-example. On a single-prompt cache that silently returns a [component, pos, pos] grid, and on larger batches it raises. Could the new branch tell the two apart, and add a single-prompt case to the layouts test?

@nehal-a2z
nehal-a2z force-pushed the nehal/fix-batched-logit-attribution branch from b71a263 to 8b27e0a Compare October 5, 2026 02:40

@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.

This is approved! I had to resolve a merge conflict, but once the CI finishes, I will merge

@jlarson4
jlarson4 merged commit 56ce454 into TransformerLensOrg:dev Oct 6, 2026
27 checks passed
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.

2 participants