feat(0831): Localize constant ReduceSum axes inside ONNX subfunctions - #1294
Open
vbaddi wants to merge 8 commits into
Open
feat(0831): Localize constant ReduceSum axes inside ONNX subfunctions#1294vbaddi wants to merge 8 commits into
vbaddi wants to merge 8 commits into
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 |
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: