Preserve zero-length dimensions during broadcast shape inference - #259
Open
Robertboy18 wants to merge 1 commit into
Open
Preserve zero-length dimensions during broadcast shape inference#259Robertboy18 wants to merge 1 commit into
Robertboy18 wants to merge 1 commit into
Conversation
Codecov Report✅ All modified and coverable lines are covered by tests. 🚀 New features to boost your workflow:
|
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
Hi! I was looking at einsum broadcasting semantics while formalizing some tensor shape behavior in Lean/TorchLean and noticed a discrepancy around zero-length dimensions.
find_output_shapecurrently uses the largest dimension for each output index. For a valid broadcast between dimensions0and1, that infers1, even though the broadcast result has dimension0. In a multi-operand contraction this incorrect intermediate shape can select a contraction that raises a shape mismatch, whilenumpy.einsumevaluates the same expression successfully.This change preserves the first non-singleton dimension, including
0, and adds parser-level and end-to-end regression coverage.I also have the corresponding small TorchLean/Lean formalization of the
0/1broadcast case and can attach it if that would be useful :)Tests
pytest -q opt_einsum/tests/test_parser.py opt_einsum/tests/test_edge_cases.py opt_einsum/tests/test_contract.py opt_einsum/tests/test_blas.py(7552 passed)uv run --extra test pytest -q(167 passed, 120 skipped)Status