Repository navigation
Conversation
nehal-a2z
marked this pull request as ready for review
September 28, 2026 13:57
jlarson4
reviewed
Sep 28, 2026
jlarson4
left a comment
Collaborator
There was a problem hiding this comment.
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) |
Collaborator
There was a problem hiding this comment.
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
force-pushed
the
nehal/fix-batched-logit-attribution
branch
from
October 5, 2026 02:40
b71a263 to
8b27e0a
Compare
jlarson4
approved these changes
Oct 5, 2026
jlarson4
left a comment
Collaborator
There was a problem hiding this comment.
This is approved! I had to resolve a merge conflict, but once the CI finishes, I will merge
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Description
Passing one target token per example to
logit_attrscan 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
Checklist
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