[Mlir-commits] [mlir] [mlir][SPIR-V] Guard UMod canonicalization against zero divisor (PR #203513)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jun 12 04:51:08 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
Chained `spirv.UMod` with a zero outer divisor reached `APInt::urem` which causes UB
---
Full diff: https://github.com/llvm/llvm-project/pull/203513.diff
2 Files Affected:
- (modified) mlir/lib/Dialect/SPIRV/IR/SPIRVCanonicalization.cpp (+5)
- (modified) mlir/test/Dialect/SPIRV/Transforms/canonicalize.mlir (+30)
``````````diff
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVCanonicalization.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVCanonicalization.cpp
index acfa6cdf85052..2d5c4d7d3fd0e 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVCanonicalization.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVCanonicalization.cpp
@@ -314,9 +314,14 @@ struct UModSimplification final : OpRewritePattern<spirv::UModOp> {
bool isApplicable = false;
if (auto prevInt = dyn_cast<IntegerAttr>(prevValue)) {
auto currInt = cast<IntegerAttr>(currValue);
+ if (currInt.getValue().isZero())
+ return failure();
isApplicable = prevInt.getValue().urem(currInt.getValue()) == 0;
} else if (auto prevVec = dyn_cast<DenseElementsAttr>(prevValue)) {
auto currVec = cast<DenseElementsAttr>(currValue);
+ if (llvm::any_of(currVec.getValues<APInt>(),
+ [](const APInt &curr) { return curr.isZero(); }))
+ return failure();
isApplicable = llvm::all_of(llvm::zip_equal(prevVec.getValues<APInt>(),
currVec.getValues<APInt>()),
[](const auto &pair) {
diff --git a/mlir/test/Dialect/SPIRV/Transforms/canonicalize.mlir b/mlir/test/Dialect/SPIRV/Transforms/canonicalize.mlir
index e49372ae91aed..94d9c53db0bbc 100644
--- a/mlir/test/Dialect/SPIRV/Transforms/canonicalize.mlir
+++ b/mlir/test/Dialect/SPIRV/Transforms/canonicalize.mlir
@@ -1063,6 +1063,36 @@ func.func @umod_vector_fail_2_fold(%arg0: vector<4xi32>) -> (vector<4xi32>, vect
return %0, %1: vector<4xi32>, vector<4xi32>
}
+// CHECK-LABEL: @umod_fail_div_0_fold
+// CHECK-SAME: (%[[ARG:.*]]: i32)
+func.func @umod_fail_div_0_fold(%arg0: i32) -> (i32, i32) {
+ // CHECK: %[[CONST0:.*]] = spirv.Constant 0
+ // CHECK: %[[CONST32:.*]] = spirv.Constant 32
+ %const1 = spirv.Constant 32 : i32
+ %0 = spirv.UMod %arg0, %const1 : i32
+ // CHECK: %[[UMOD0:.*]] = spirv.UMod %[[ARG]], %[[CONST32]]
+ %const2 = spirv.Constant 0 : i32
+ %1 = spirv.UMod %0, %const2 : i32
+ // CHECK: %[[UMOD1:.*]] = spirv.UMod %[[UMOD0]], %[[CONST0]]
+ // CHECK: return %[[UMOD0]], %[[UMOD1]]
+ return %0, %1: i32, i32
+}
+
+// CHECK-LABEL: @umod_vector_fail_div_0_fold
+// CHECK-SAME: (%[[ARG:.*]]: vector<4xi32>)
+func.func @umod_vector_fail_div_0_fold(%arg0: vector<4xi32>) -> (vector<4xi32>, vector<4xi32>) {
+ // CHECK: %[[CONST0:.*]] = spirv.Constant dense<[4, 0, 4, 0]> : vector<4xi32>
+ // CHECK: %[[CONST32:.*]] = spirv.Constant dense<32> : vector<4xi32>
+ %const1 = spirv.Constant dense<32> : vector<4xi32>
+ %0 = spirv.UMod %arg0, %const1 : vector<4xi32>
+ // CHECK: %[[UMOD0:.*]] = spirv.UMod %[[ARG]], %[[CONST32]]
+ %const2 = spirv.Constant dense<[4, 0, 4, 0]> : vector<4xi32>
+ %1 = spirv.UMod %0, %const2 : vector<4xi32>
+ // CHECK: %[[UMOD1:.*]] = spirv.UMod %[[UMOD0]], %[[CONST0]]
+ // CHECK: return %[[UMOD0]], %[[UMOD1]]
+ return %0, %1: vector<4xi32>, vector<4xi32>
+}
+
// -----
//===----------------------------------------------------------------------===//
``````````
</details>
https://github.com/llvm/llvm-project/pull/203513
More information about the Mlir-commits
mailing list