[Mlir-commits] [mlir] [mlir][SPIR-V] Add SPIRVToLLVM conversions for GL.FSign and GL.SSign (PR #206934)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed Jul 1 03:07:51 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-spirv
@llvm/pr-subscribers-mlir
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
---
Full diff: https://github.com/llvm/llvm-project/pull/206934.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp (+64)
- (modified) mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir (+54)
``````````diff
diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index c43415b27b1b3..eb86d0e5c417f 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -1564,6 +1564,68 @@ class SAbsPattern : public SPIRVToLLVMConversion<spirv::GLSAbsOp> {
}
};
+/// Converts `spirv.GL.FSign`/`spirv.GL.SSign` to a sign(x) sequence that maps
+/// the operand to -1/0/1 using two comparisons and two selects. The `isFloat`
+/// flag selects between floating-point and integer comparisons/constants.
+template <typename SPIRVOp, bool isFloat>
+class SignPattern : public SPIRVToLLVMConversion<SPIRVOp> {
+public:
+ using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
+
+ LogicalResult
+ matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto srcType = op.getType();
+ auto dstType = this->getTypeConverter()->convertType(srcType);
+ if (!dstType)
+ return rewriter.notifyMatchFailure(op, "type conversion failed");
+
+ Location loc = op.getLoc();
+ Value operand = adaptor.getOperand();
+ auto vecSrcType = dyn_cast<VectorType>(srcType);
+ Type i1Type = rewriter.getI1Type();
+ Type cmpType =
+ vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
+
+ Value zero, one, minusOne, gt, lt;
+ if constexpr (isFloat) {
+ zero = createFPConstant(loc, srcType, dstType, rewriter, 0.0);
+ one = createFPConstant(loc, srcType, dstType, rewriter, 1.0);
+ minusOne = createFPConstant(loc, srcType, dstType, rewriter, -1.0);
+ gt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
+ LLVM::FCmpPredicate::ogt, operand, zero);
+ lt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
+ LLVM::FCmpPredicate::olt, operand, zero);
+ } else {
+ auto intElemType = cast<IntegerType>(getElementTypeOrSelf(srcType));
+ IntegerAttr zeroAttr = rewriter.getIntegerAttr(intElemType, 0);
+ IntegerAttr oneAttr = rewriter.getIntegerAttr(intElemType, 1);
+ if (vecSrcType) {
+ zero = LLVM::ConstantOp::create(
+ rewriter, loc, dstType,
+ SplatElementsAttr::get(vecSrcType, zeroAttr));
+ one = LLVM::ConstantOp::create(
+ rewriter, loc, dstType,
+ SplatElementsAttr::get(vecSrcType, oneAttr));
+ } else {
+ zero = LLVM::ConstantOp::create(rewriter, loc, dstType, zeroAttr);
+ one = LLVM::ConstantOp::create(rewriter, loc, dstType, oneAttr);
+ }
+ minusOne = createConstantAllBitsSet(loc, srcType, dstType, rewriter);
+ gt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
+ LLVM::ICmpPredicate::sgt, operand, zero);
+ lt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
+ LLVM::ICmpPredicate::slt, operand, zero);
+ }
+
+ Value negOrZero =
+ LLVM::SelectOp::create(rewriter, loc, dstType, lt, minusOne, zero);
+ rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, gt, one,
+ negOrZero);
+ return success();
+ }
+};
+
class VariablePattern : public SPIRVToLLVMConversion<spirv::VariableOp> {
public:
using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
@@ -1917,6 +1979,8 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
InverseSqrtPattern, SAbsPattern, TanPattern, TanhPattern,
+ SignPattern<spirv::GLFSignOp, /*isFloat=*/true>,
+ SignPattern<spirv::GLSSignOp, /*isFloat=*/false>,
// OpenCL extended instruction set ops
DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
diff --git a/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
index ffa47efbf9213..99153da1e9509 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
@@ -361,3 +361,57 @@ spirv.func @asin_acos_atan(%arg0: f32, %arg1: vector<3xf16>) "None" {
%2 = spirv.GL.Atan %arg0 : f32
spirv.Return
}
+
+//===----------------------------------------------------------------------===//
+// spirv.GL.FSign
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @fsign_scalar
+spirv.func @fsign_scalar(%arg0: f32) "None" {
+ // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(0.000000e+00 : f32) : f32
+ // CHECK: %[[ONE:.*]] = llvm.mlir.constant(1.000000e+00 : f32) : f32
+ // CHECK: %[[MONE:.*]] = llvm.mlir.constant(-1.000000e+00 : f32) : f32
+ // CHECK: %[[GT:.*]] = llvm.fcmp "ogt" %{{.*}}, %[[ZERO]] : f32
+ // CHECK: %[[LT:.*]] = llvm.fcmp "olt" %{{.*}}, %[[ZERO]] : f32
+ // CHECK: %[[SEL0:.*]] = llvm.select %[[LT]], %[[MONE]], %[[ZERO]] : i1, f32
+ // CHECK: llvm.select %[[GT]], %[[ONE]], %[[SEL0]] : i1, f32
+ %0 = spirv.GL.FSign %arg0 : f32
+ spirv.Return
+}
+
+// CHECK-LABEL: @fsign_vector
+spirv.func @fsign_vector(%arg0: vector<4xf32>) "None" {
+ // CHECK: %[[GT:.*]] = llvm.fcmp "ogt" %{{.*}}, %{{.*}} : vector<4xf32>
+ // CHECK: %[[LT:.*]] = llvm.fcmp "olt" %{{.*}}, %{{.*}} : vector<4xf32>
+ // CHECK: %[[SEL0:.*]] = llvm.select %[[LT]], %{{.*}}, %{{.*}} : vector<4xi1>, vector<4xf32>
+ // CHECK: llvm.select %[[GT]], %{{.*}}, %[[SEL0]] : vector<4xi1>, vector<4xf32>
+ %0 = spirv.GL.FSign %arg0 : vector<4xf32>
+ spirv.Return
+}
+
+//===----------------------------------------------------------------------===//
+// spirv.GL.SSign
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @ssign_scalar
+spirv.func @ssign_scalar(%arg0: i32) "None" {
+ // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(0 : i32) : i32
+ // CHECK: %[[ONE:.*]] = llvm.mlir.constant(1 : i32) : i32
+ // CHECK: %[[MONE:.*]] = llvm.mlir.constant(-1 : i32) : i32
+ // CHECK: %[[GT:.*]] = llvm.icmp "sgt" %{{.*}}, %[[ZERO]] : i32
+ // CHECK: %[[LT:.*]] = llvm.icmp "slt" %{{.*}}, %[[ZERO]] : i32
+ // CHECK: %[[SEL0:.*]] = llvm.select %[[LT]], %[[MONE]], %[[ZERO]] : i1, i32
+ // CHECK: llvm.select %[[GT]], %[[ONE]], %[[SEL0]] : i1, i32
+ %0 = spirv.GL.SSign %arg0 : i32
+ spirv.Return
+}
+
+// CHECK-LABEL: @ssign_vector
+spirv.func @ssign_vector(%arg0: vector<4xi32>) "None" {
+ // CHECK: %[[GT:.*]] = llvm.icmp "sgt" %{{.*}}, %{{.*}} : vector<4xi32>
+ // CHECK: %[[LT:.*]] = llvm.icmp "slt" %{{.*}}, %{{.*}} : vector<4xi32>
+ // CHECK: %[[SEL0:.*]] = llvm.select %[[LT]], %{{.*}}, %{{.*}} : vector<4xi1>, vector<4xi32>
+ // CHECK: llvm.select %[[GT]], %{{.*}}, %[[SEL0]] : vector<4xi1>, vector<4xi32>
+ %0 = spirv.GL.SSign %arg0 : vector<4xi32>
+ spirv.Return
+}
``````````
</details>
https://github.com/llvm/llvm-project/pull/206934
More information about the Mlir-commits
mailing list