Skip to content

Fix FactoredMatrix ellipsis indexing - #1849

Open
yuanwuyuan9 wants to merge 3 commits into
TransformerLensOrg:devfrom
yuanwuyuan9:fix/factored-matrix-ellipsis-indexing
Open

yuanwuyuan9 wants to merge 3 commits into
TransformerLensOrg:devfrom
yuanwuyuan9:fix/factored-matrix-ellipsis-indexing

Conversation

@yuanwuyuan9

@yuanwuyuan9 yuanwuyuan9 commented Oct 2, 2026 •

Copy link
Copy Markdown

Description

FactoredMatrix[..., column] and ellipsis-based row/submatrix selections can index the hidden rank dimension, causing shape errors or incorrect results. This affects selections from QK/OV circuit matrices.

Expand ellipsis against the product's dimensions before mapping indices to the factors. Use the same axis-counting rule for expansion and factor selection: new axes consume none, while Boolean masks consume one axis per mask dimension. Preserve legacy byte-mask indexing and the existing integer row/column singleton axes, and reject multiple ellipses and excess indices. The factors are sliced directly; no dense product or new dependencies are required.

Fixes #1848

Type of change

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

Validation

The existing indexing test file now covers ellipsis with 0–2 leading dimensions, positive/negative row and column indices, submatrices, leading new axes, integer Tensor indices, Tensor/NumPy Boolean masks, legacy byte masks, zero-width ellipsis, invalid indices, and indexing without materializing AB.

Local CPU validation on macOS, Python 3.12.14 and PyTorch 2.11.0:

Check Result
54 new regression cases against original 6affbe99 Before: 38 failed, 16 passed; after: all passed within the directory run
Entire FactoredMatrix unit directory 199 passed, 1 existing skipped, 2 warnings; retries disabled
Tiny native QK/OV column selection Both pass after the fix; both asserted before it
Format checks make format and make check-format passed
mypy . No issues in 387 source files
git diff --check Passed

Tests used OMP_NUM_THREADS=1 MKL_NUM_THREADS=1. Before/after pytest runs report the same two SWIG deprecation warnings; legacy byte-mask deprecations are explicitly checked with pytest.warns. The existing skip concerns jaxtyping's handling of an inner-dimension mismatch constructor test. The full unit suite and HF integration tests have not been run for this change.

Checklist:

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

The docstring documents ellipsis and singleton-axis behavior. The unit-test checklist remains unchecked pending a broader suite run; all FactoredMatrix unit tests pass as reported above.

@koriyoshi2041 koriyoshi2041 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

One indexing form is still miscounted: nested Python boolean lists consume one input axis per list dimension in PyTorch, but index_ndim() treats every list as one axis.

m = FactoredMatrix(torch.randn(2, 3, 5, 2), torch.randn(2, 3, 2, 7))
mask = [[True, False, True], [False, True, False]]
m[mask, ..., 1]

At b7e6a786 this raises IndexError: Dimension out of range ... got 4; the equivalent dense selection m.AB[mask, ..., 1] returns shape (3, 5). The ellipsis gets one extra full slice because the two-dimensional list mask contributes 1 to consumed instead of 2.

I reproduced this against the exact PR head after the focused indexing file passed 68/68. A regression alongside the tensor/NumPy mask cases would protect the valid Python-sequence form too.

@yuanwuyuan9

Copy link
Copy Markdown
Author

Hi @koriyoshi2041, thanks for catching this! Addressed in 2e373523.

Boolean lists/tuples now consume dimensions according to their mask dimensionality, while integer sequences retain single-axis indexing. I also extended the mask coverage and added integer-sequence tests.

All 223 FactoredMatrix tests pass, as do formatting and mypy checks.

@koriyoshi2041

Copy link
Copy Markdown
Contributor

Confirmed on 2e373523: the original nested Python boolean-list case now matches the factored-matrix shape contract, and the added integer-sequence cases preserve single-axis consumption. The full tests/unit/factored_matrix suite passes locally (223 passed, 1 skipped), and the hosted format, typing, compatibility, coverage, benchmark, and notebook checks are green. Thanks for the quick fix.

@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 the quick turnaround on the nested-mask case. Sorry for the delay in my feedback, I have a couple additional notes below

"""Indexing - assumed to only apply to the leading dimensions."""
"""Index leading dimensions and matrix rows/columns without forming the product.

Ellipsis expands across the product's dimensions, excluding new axes (``None``).

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.

The docstring says ellipsis skips new axes, but a None at or after the row position still breaks the row/column mapping. m[..., None] and m[..., None, 1] raise, and m[0, 1, :, None, ...], which worked before, now raises. Can the row/column handling account for new axes the same way the ellipsis expansion does, with a test for a trailing None?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for catching this. Fixed in 25c64154 by preserving the row/column mapping when inserting new axes, without materializing the product. Added regressions for all three examples and multiple new axes. All 248 FactoredMatrix tests pass.

Comment thread transformer_lens/FactoredMatrix.py Outdated
consumed = sum(index_ndim(item) for item in idx)
if consumed > self.ndim:
raise ValueError(
f"{idx} is too long an index for a FactoredMatrix with shape {self.shape}"

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.

When the index is too long, the expansion adds no slices, so the final else already raises this same error. Is this early check needed?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed—the early check was redundant. Removed it in 25c64154; excess indices still raise through the final check, with additional coverage for indices containing None.

if numpy_mask:
mask = mask.numpy()
with pytest.warns(UserWarning, match="uint8"):
result = matrix[mask, ...]

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 leading-only selection lands correctly even if a byte mask is miscounted as one axis, so this test still passes when numpy uint8 is dropped from index_ndim. Could it also select a row or column?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good point. Added row and column selections for both Tensor and NumPy byte masks in 25c64154. I also verified that removing NumPy uint8 recognition makes both new NumPy cases fail.

@jlarson4 jlarson4 linked an issue Oct 5, 2026 that may be closed by this pull request
1 task done
@koriyoshi2041

Copy link
Copy Markdown
Contributor

Confirmed on exact head 25c64154: the full tests/unit/factored_matrix suite passes locally (248 passed, 1 skipped). I also checked six dense-vs-factored adversarial selections beyond the added regressions—three trailing new axes; new axes around row and column slices; a leading integer selection; a leading boolean-list selection; and mixed leading/matrix new axes—and all matched matrix.AB in shape and values. The three requested cases are covered, the redundant early length check is gone, and I did not find a new semantic regression in this pass.

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.

[Bug Report] FactoredMatrix ellipsis indexing selects the wrong factor axes

3 participants