Skip to content

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

Open
nehal-a2z wants to merge 1 commit into
TransformerLensOrg:devfrom
nehal-a2z:nehal/fix-batched-logit-attribution
Open

nehal-a2z wants to merge 1 commit into
TransformerLensOrg:devfrom
nehal-a2z:nehal/fix-batched-logit-attribution

Conversation

@nehal-a2z

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

96 focused tests pass with runtime jaxtyping enabled. The tests use native synthetic TransformerBridge models and compare batched results with separate examples across LN/RMS, slicing, and logit differences. On the original source, 24 regression cases fail.

Tested on CPU with PyTorch 2.7.1 and Transformers 4.57.6; the full pinned dev environment and pretrained models weren't exercised. The limited changed-file mypy run reports the same two existing shape-annotation diagnostics on both the original and patched source, so this isn't a typecheck pass.

Sent using @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?

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.

2 participants