feat(0831): Localize constant ReduceSum axes inside ONNX subfunctions - #1294
Merged
Merged
Conversation
Contributor
Author
|
CI-Ready |
Contributor
Author
|
CI-Ready |
Replace QEff einsum-based reduction workarounds with equivalent torch.sum()
forms now that ONNX subfunction ReduceSum axes are localized as constants.
Add a narrow ONNX transform that detects ReduceSum nodes inside FunctionProto
bodies where the axes input was promoted to a function formal input. When every
top-level call site passes the same compile-time integer constant for that
formal input, the transform inserts a local Constant node inside the function,
rewires ReduceSum to use it, removes the formal function input, and removes the
matching actual argument from each call site.
This keeps PyTorch code such as:
(query * query).sum(dim=-1, keepdim=True)
lowering to Mul + ReduceSum instead of requiring an einsum workaround, while
still presenting ReduceSum axes to the compiler as a compile-time constant
inside the ONNX FunctionProto.
Signed-off-by: vbaddi <vbaddi@qti.qualcomm.com>
Signed-off-by: vbaddi <vbaddi@qti.qualcomm.com>
Signed-off-by: Kushal Dulla <kdulla@qti.qualcomm.com> (cherry picked from commit 66b459b) Signed-off-by: ochougul <ochougul@qti.qualcomm.com>
Signed-off-by: ochougul <ochougul@qti.qualcomm.com> (cherry picked from commit c827393) Signed-off-by: ochougul <ochougul@qti.qualcomm.com>
Signed-off-by: ochougul <ochougul@qti.qualcomm.com>
Signed-off-by: vbaddi <vbaddi@qti.qualcomm.com>
Signed-off-by: Mohit Soni <mohisoni@qti.qualcomm.com>
Contributor
|
CI-Ready |
Contributor
|
CI-Ready |
Contributor
|
CI-Ready |
1 similar comment
|
CI-Ready |
athavale-shivani
added a commit
to vaibverm/efficient-transformers-blocking-techniques
that referenced
this pull request
Sep 15, 2026
Extends the existing PR quic#1294 guard to the remaining call sites that hit the same issue: test_full_*, test_few_* (main loop + qkv_paged block), and test_dummy_*'s hqkv_paged block. Signed-off-by: Shivani Athavale <athavale@qti.qualcomm.com>
athavale-shivani
pushed a commit
to vaibverm/efficient-transformers-blocking-techniques
that referenced
this pull request
Sep 16, 2026
…quic#1294) ## Summary - Added `LocalizeFunctionReduceSumAxesTransform` for ONNX subfunction exports. - Replaced the previous `torch.einsum(... reduction ...)` workarounds back to equivalent `.sum(...)` forms across QEff modeling/MoE/blocking code. - Enabled the ReduceSum axes localization only when `use_onnx_subfunctions=True`. We previously used `einsum` in several reduction patterns to avoid ONNX subfunction export promoting constant `ReduceSum` axes into `FunctionProto` inputs. That workaround avoided compiler failures, but it could lower into heavier performance dip's due to lowering of BMM instead of the intended elementwise multiply plus ReduceSum path. The desired PyTorch source should be able to stay as: ```python (query * query).sum(dim=-1, keepdim=True) and export as: Mul -> ReduceSum <- Constant([-1]) -- inside the ONNX function body. ``` ## ONNX Transform Details LocalizeFunctionReduceSumAxesTransform it rewrites a function input when: - a node inside an ONNX FunctionProto is ReduceSum - the ReduceSum axes input is one of the function formal inputs - every top-level call site passes a compile-time constant for that formal input - the value is a valid integer scalar or 1-D axes tensor When eligible, the transform: - inserts a local Constant node inside the function body - rewires each matching ReduceSum to use that local constant - removes the axes formal input from the function signature --------- Signed-off-by: vbaddi <vbaddi@qti.qualcomm.com> Signed-off-by: Kushal Dulla <kdulla@qti.qualcomm.com> Signed-off-by: ochougul <ochougul@qti.qualcomm.com> Signed-off-by: Mohit Soni <mohisoni@qti.qualcomm.com> Co-authored-by: Kushal Dulla <kdulla@qti.qualcomm.com> Co-authored-by: ochougul <ochougul@qti.qualcomm.com> Co-authored-by: Mohit Soni <mohisoni@qti.qualcomm.com> Co-authored-by: Onkar Chougule <168134249+ochougul@users.noreply.github.com>
athavale-shivani
added a commit
to vaibverm/efficient-transformers-blocking-techniques
that referenced
this pull request
Sep 16, 2026
Extends the existing PR quic#1294 guard to the remaining call sites that hit the same issue: test_full_*, test_few_* (main loop + qkv_paged block), and test_dummy_*'s hqkv_paged block. Signed-off-by: Shivani Athavale <athavale@qti.qualcomm.com>
athavale-shivani
added a commit
to vaibverm/efficient-transformers-blocking-techniques
that referenced
this pull request
Sep 16, 2026
Extends the existing PR quic#1294 guard to the remaining call sites that hit the same issue: test_full_*, test_few_* (main loop + qkv_paged block), and test_dummy_*'s hqkv_paged block. Signed-off-by: Shivani Athavale <athavale@qti.qualcomm.com>
athavale-shivani
added a commit
to vaibverm/efficient-transformers-blocking-techniques
that referenced
this pull request
Sep 18, 2026
Extends the existing PR quic#1294 guard to the remaining call sites that hit the same issue: test_full_*, test_few_* (main loop + qkv_paged block), and test_dummy_*'s hqkv_paged block. Signed-off-by: Shivani Athavale <athavale@qti.qualcomm.com>
quic-rishinr
pushed a commit
to vaibverm/efficient-transformers-blocking-techniques
that referenced
this pull request
Sep 19, 2026
Extends the existing PR quic#1294 guard to the remaining call sites that hit the same issue: test_full_*, test_few_* (main loop + qkv_paged block), and test_dummy_*'s hqkv_paged block. Signed-off-by: Shivani Athavale <athavale@qti.qualcomm.com>
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.
Summary
LocalizeFunctionReduceSumAxesTransformfor ONNX subfunction exports.torch.einsum(... reduction ...)workarounds back to equivalent.sum(...)forms across QEff modeling/MoE/blocking code.use_onnx_subfunctions=True.We previously used
einsumin several reduction patterns to avoid ONNX subfunction export promoting constantReduceSumaxes intoFunctionProtoinputs. That workaround avoided compiler failures, but it could lower into heavier performance dip's due to lowering of BMM instead of the intended elementwise multiply plus ReduceSum path.The desired PyTorch source should be able to stay as:
ONNX Transform Details
LocalizeFunctionReduceSumAxesTransform it rewrites a function input when:
When eligible, the transform: