[Mlir-commits] [mlir] [mlir][SPIR-V] Guard UMod canonicalization against zero divisor (PR #203513)

Arseniy Obolenskiy llvmlistbot at llvm.org
Fri Jun 12 04:50:22 PDT 2026


https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/203513

Chained `spirv.UMod` with a zero outer divisor reached `APInt::urem` which causes UB

>From 0072e41b6618cf7e536cccb95161963b2fcf1712 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Fri, 12 Jun 2026 13:48:31 +0200
Subject: [PATCH] [mlir][SPIR-V] Guard UMod canonicalization against zero
 divisor

Chained spirv.UMod with a zero outer divisor reached APInt::urem which causes UB
---
 .../SPIRV/IR/SPIRVCanonicalization.cpp        |  5 ++++
 .../SPIRV/Transforms/canonicalize.mlir        | 30 +++++++++++++++++++
 2 files changed, 35 insertions(+)

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