[Relax] Preserve out_dtype in AdjustMatmulOrder - #20296
Merged
Merged
Conversation
Preserve the original outer matmul output dtype in both cost-based and compile-time grouping rewrites. Add regression coverage for both reassociation directions, transposed inner matmuls, inferred output dtype, and idempotence. Fixes apache#20202. Validation: 41 AdjustMatmulOrder tests passed; LLVM VM float32 output checks passed in both directions; pre-commit passed on changed files. Before the fix, all six explicit float32 regression cases failed.
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.
Fixes #20202.
AdjustMatmulOrderrebuilt reassociated outer matmuls without_dtype=void, which could change the result dtype of the originalprogram. For example, a
float16matmul chain whose outer matmul explicitlyreturns
float32was rewritten to returnfloat16.This patch preserves the original outer
MatmulAttrs::out_dtypewhenconstructing the replacement outer matmul. Newly introduced inner matmuls
continue to infer their output dtype.
Tests:
AdjustMatmulOrdertest file: 41 passed.pre-commitandgit diff --check: passed.