[Mlir-commits] [mlir] 4fbad31 - [mlir][SPIR-V] Add SPIRVToLLVM conversions for FMod and SMod (#206933)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Fri Jul 31 01:36:20 PDT 2026


Author: Arseniy Obolenskiy
Date: 2026-07-31T10:36:15+02:00
New Revision: 4fbad31a8adf437e1d0f1307ff4c26a248462b09

URL: https://github.com/llvm/llvm-project/commit/4fbad31a8adf437e1d0f1307ff4c26a248462b09
DIFF: https://github.com/llvm/llvm-project/commit/4fbad31a8adf437e1d0f1307ff4c26a248462b09.diff

LOG: [mlir][SPIR-V] Add SPIRVToLLVM conversions for FMod and SMod (#206933)

Added: 
    

Modified: 
    mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
    mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir

Removed: 
    


################################################################################
diff  --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index c272a91f0ed76..63f9188a4e464 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -1032,6 +1032,79 @@ class ClampPattern : public SPIRVToLLVMConversion<SPIRVOp> {
   }
 };
 
+/// Converts `spirv.FMod` to `x - y * floor(x / y)`. The SPIR-V op requires the
+/// result to take the sign of the divisor, whereas `llvm.frem` keeps the sign
+/// of the dividend, so `frem` cannot be used directly.
+class FModPattern : public SPIRVToLLVMConversion<spirv::FModOp> {
+public:
+  using SPIRVToLLVMConversion<spirv::FModOp>::SPIRVToLLVMConversion;
+
+  LogicalResult
+  matchAndRewrite(spirv::FModOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type dstType = getTypeConverter()->convertType(op.getType());
+    if (!dstType)
+      return rewriter.notifyMatchFailure(op, "type conversion failed");
+
+    Location loc = op.getLoc();
+    Value lhs = adaptor.getOperand1();
+    Value rhs = adaptor.getOperand2();
+    Value div = LLVM::FDivOp::create(rewriter, loc, dstType, lhs, rhs);
+    Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, div);
+    Value scaled = LLVM::FMulOp::create(rewriter, loc, dstType, rhs, floored);
+    rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, lhs, scaled);
+    return success();
+  }
+};
+
+/// Converts `spirv.SMod` to a signed remainder corrected to take the sign of
+/// the divisor. `llvm.srem` keeps the sign of the dividend, so the result is
+/// adjusted by adding the divisor when the remainder is non-zero and its sign
+/// 
diff ers from the divisor's.
+class SModPattern : public SPIRVToLLVMConversion<spirv::SModOp> {
+public:
+  using SPIRVToLLVMConversion<spirv::SModOp>::SPIRVToLLVMConversion;
+
+  LogicalResult
+  matchAndRewrite(spirv::SModOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type srcType = op.getType();
+    Type dstType = getTypeConverter()->convertType(srcType);
+    if (!dstType)
+      return rewriter.notifyMatchFailure(op, "type conversion failed");
+
+    Location loc = op.getLoc();
+    Value lhs = adaptor.getOperand1();
+    Value rhs = adaptor.getOperand2();
+    Type i1Type = rewriter.getI1Type();
+    auto vecSrcType = dyn_cast<VectorType>(srcType);
+    Type cmpType =
+        vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
+
+    Value rem = LLVM::SRemOp::create(rewriter, loc, dstType, lhs, rhs);
+    IntegerAttr zeroAttr = rewriter.getIntegerAttr(
+        cast<IntegerType>(getElementTypeOrSelf(srcType)), 0);
+    Value zero =
+        createIntegerConstant(loc, srcType, dstType, rewriter, zeroAttr);
+
+    Value remNonZero = LLVM::ICmpOp::create(rewriter, loc, cmpType,
+                                            LLVM::ICmpPredicate::ne, rem, zero);
+    Value remNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
+                                        LLVM::ICmpPredicate::slt, rem, zero);
+    Value rhsNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
+                                        LLVM::ICmpPredicate::slt, rhs, zero);
+    Value signMismatch =
+        LLVM::XOrOp::create(rewriter, loc, cmpType, remNeg, rhsNeg);
+    Value needsAdjust =
+        LLVM::AndOp::create(rewriter, loc, cmpType, remNonZero, signMismatch);
+
+    Value adjusted = LLVM::AddOp::create(rewriter, loc, dstType, rem, rhs);
+    rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, needsAdjust,
+                                                adjusted, rem);
+    return success();
+  }
+};
+
 /// Converts `spirv.Load` and `spirv.Store` to LLVM dialect.
 template <typename SPIRVOp>
 class LoadStorePattern : public SPIRVToLLVMConversion<SPIRVOp> {
@@ -2065,8 +2138,8 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
       DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
       DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
       DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
-      DirectConversionPattern<spirv::UModOp, LLVM::URemOp>,
-      VectorTimesScalarPattern, SNegatePattern,
+      DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, FModPattern,
+      SModPattern, VectorTimesScalarPattern, SNegatePattern,
       ArithmeticWithOverflowPattern<spirv::IAddCarryOp,
                                     LLVM::UAddWithOverflowOp>,
       ArithmeticWithOverflowPattern<spirv::ISubBorrowOp,

diff  --git a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
index 64efbdf927c68..14ddfd64fefa9 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
@@ -294,6 +294,60 @@ spirv.func @srem_vector(%arg0: vector<4xi32>, %arg1: vector<4xi32>) "None" {
   spirv.Return
 }
 
+//===----------------------------------------------------------------------===//
+// spirv.FMod
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @fmod_scalar
+spirv.func @fmod_scalar(%arg0: f32, %arg1: f32) "None" {
+  // CHECK: %[[DIV:.*]] = llvm.fdiv %{{.*}}, %{{.*}} : f32
+  // CHECK: %[[FLOOR:.*]] = llvm.intr.floor(%[[DIV]]) : (f32) -> f32
+  // CHECK: %[[MUL:.*]] = llvm.fmul %{{.*}}, %[[FLOOR]] : f32
+  // CHECK: llvm.fsub %{{.*}}, %[[MUL]] : f32
+  %0 = spirv.FMod %arg0, %arg1 : f32
+  spirv.Return
+}
+
+// CHECK-LABEL: @fmod_vector
+spirv.func @fmod_vector(%arg0: vector<4xf32>, %arg1: vector<4xf32>) "None" {
+  // CHECK: %[[DIV:.*]] = llvm.fdiv %{{.*}}, %{{.*}} : vector<4xf32>
+  // CHECK: %[[FLOOR:.*]] = llvm.intr.floor(%[[DIV]]) : (vector<4xf32>) -> vector<4xf32>
+  // CHECK: %[[MUL:.*]] = llvm.fmul %{{.*}}, %[[FLOOR]] : vector<4xf32>
+  // CHECK: llvm.fsub %{{.*}}, %[[MUL]] : vector<4xf32>
+  %0 = spirv.FMod %arg0, %arg1 : vector<4xf32>
+  spirv.Return
+}
+
+//===----------------------------------------------------------------------===//
+// spirv.SMod
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @smod_scalar
+spirv.func @smod_scalar(%arg0: i32, %arg1: i32) "None" {
+  // CHECK: %[[REM:.*]] = llvm.srem %{{.*}}, %{{.*}} : i32
+  // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(0 : i32) : i32
+  // CHECK: %[[NZ:.*]] = llvm.icmp "ne" %[[REM]], %[[ZERO]] : i32
+  // CHECK: %[[RNEG:.*]] = llvm.icmp "slt" %[[REM]], %[[ZERO]] : i32
+  // CHECK: %[[DNEG:.*]] = llvm.icmp "slt" %{{.*}}, %[[ZERO]] : i32
+  // CHECK: %[[XOR:.*]] = llvm.xor %[[RNEG]], %[[DNEG]] : i1
+  // CHECK: %[[ADJ:.*]] = llvm.and %[[NZ]], %[[XOR]] : i1
+  // CHECK: %[[ADD:.*]] = llvm.add %[[REM]], %{{.*}} : i32
+  // CHECK: llvm.select %[[ADJ]], %[[ADD]], %[[REM]] : i1, i32
+  %0 = spirv.SMod %arg0, %arg1 : i32
+  spirv.Return
+}
+
+// CHECK-LABEL: @smod_vector
+spirv.func @smod_vector(%arg0: vector<4xi32>, %arg1: vector<4xi32>) "None" {
+  // CHECK: %[[REM:.*]] = llvm.srem %{{.*}}, %{{.*}} : vector<4xi32>
+  // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(dense<0> : vector<4xi32>) : vector<4xi32>
+  // CHECK: %[[NZ:.*]] = llvm.icmp "ne" %[[REM]], %[[ZERO]] : vector<4xi32>
+  // CHECK: %[[ADD:.*]] = llvm.add %[[REM]], %{{.*}} : vector<4xi32>
+  // CHECK: llvm.select %{{.*}}, %[[ADD]], %[[REM]] : vector<4xi1>, vector<4xi32>
+  %0 = spirv.SMod %arg0, %arg1 : vector<4xi32>
+  spirv.Return
+}
+
 //===----------------------------------------------------------------------===//
 // spirv.VectorTimesScalar
 //===----------------------------------------------------------------------===//


        


More information about the Mlir-commits mailing list