[Mlir-commits] [mlir] d426cca - [mlir][SPIR-V] Guard UMod canonicalization against zero divisor (#203513)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Fri Jun 12 07:15:45 PDT 2026
Author: Arseniy Obolenskiy
Date: 2026-06-12T16:15:40+02:00
New Revision: d426cca8835e109069f60630f886574af273e803
URL: https://github.com/llvm/llvm-project/commit/d426cca8835e109069f60630f886574af273e803
DIFF: https://github.com/llvm/llvm-project/commit/d426cca8835e109069f60630f886574af273e803.diff
LOG: [mlir][SPIR-V] Guard UMod canonicalization against zero divisor (#203513)
Chained `spirv.UMod` with a zero outer divisor reached `APInt::urem`
which causes UB
Added:
Modified:
mlir/lib/Dialect/SPIRV/IR/SPIRVCanonicalization.cpp
mlir/test/Dialect/SPIRV/Transforms/canonicalize.mlir
Removed:
################################################################################
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>
+}
+
// -----
//===----------------------------------------------------------------------===//
More information about the Mlir-commits
mailing list