Fix FactoredMatrix ellipsis indexing - #1849
yuanwuyuan9 wants to merge 3 commits into
Conversation
koriyoshi2041
left a comment
There was a problem hiding this comment.
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.
|
Hi @koriyoshi2041, thanks for catching this! Addressed in 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 |
|
Confirmed on |
jlarson4
left a comment
There was a problem hiding this comment.
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``). |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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.
| 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}" |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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, ...] |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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.
|
Confirmed on exact head |
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
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:
6affbe99make formatandmake check-formatpassedmypy .git diff --checkTests 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 withpytest.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:
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.