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?
This branch has not been deployed
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. 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
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