[Mlir-commits] [mlir] [mlir][SPIR-V] Add SPIRVToLLVM conversion for SNegate (PR #206950)
Arseniy Obolenskiy
llvmlistbot at llvm.org
Fri Jul 10 04:45:42 PDT 2026
https://github.com/aobolensk updated https://github.com/llvm/llvm-project/pull/206950
>From 187cd538f98c32b897f1d8839a38a3b6915917a1 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 1 Jul 2026 13:12:30 +0200
Subject: [PATCH 1/2] [mlir][SPIR-V] Add SPIRVToLLVM conversion for SNegate
---
.../Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp | 31 ++++++++++++++++++-
.../SPIRVToLLVM/arithmetic-ops-to-llvm.mlir | 20 ++++++++++++
2 files changed, 50 insertions(+), 1 deletion(-)
diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index c43415b27b1b3..610c052b64ea7 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -921,6 +921,35 @@ class InverseSqrtPattern
}
};
+/// Converts `spirv.SNegate` to `0 - x`.
+class SNegatePattern : public SPIRVToLLVMConversion<spirv::SNegateOp> {
+public:
+ using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
+
+ LogicalResult
+ matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto srcType = op.getType();
+ auto dstType = getTypeConverter()->convertType(srcType);
+ if (!dstType)
+ return rewriter.notifyMatchFailure(op, "type conversion failed");
+
+ Location loc = op.getLoc();
+ auto vecSrcType = dyn_cast<VectorType>(srcType);
+ IntegerAttr zeroAttr = rewriter.getIntegerAttr(
+ cast<IntegerType>(getElementTypeOrSelf(srcType)), 0);
+ Value zero;
+ if (vecSrcType)
+ zero = LLVM::ConstantOp::create(
+ rewriter, loc, dstType, SplatElementsAttr::get(vecSrcType, zeroAttr));
+ else
+ zero = LLVM::ConstantOp::create(rewriter, loc, dstType, zeroAttr);
+ rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
+ adaptor.getOperand());
+ return success();
+ }
+};
+
/// Converts `spirv.Load` and `spirv.Store` to LLVM dialect.
template <typename SPIRVOp>
class LoadStorePattern : public SPIRVToLLVMConversion<SPIRVOp> {
@@ -1824,7 +1853,7 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
- DirectConversionPattern<spirv::UModOp, LLVM::URemOp>,
+ DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, SNegatePattern,
// Bitwise ops
BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
diff --git a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
index dbbf8610afb4d..ee9c06fbb10e1 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
@@ -233,3 +233,23 @@ spirv.func @srem_vector(%arg0: vector<4xi32>, %arg1: vector<4xi32>) "None" {
%0 = spirv.SRem %arg0, %arg1 : vector<4xi32>
spirv.Return
}
+
+//===----------------------------------------------------------------------===//
+// spirv.SNegate
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @snegate_scalar
+spirv.func @snegate_scalar(%arg0: i32) "None" {
+ // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(0 : i32) : i32
+ // CHECK: llvm.sub %[[ZERO]], %{{.*}} : i32
+ %0 = spirv.SNegate %arg0 : i32
+ spirv.Return
+}
+
+// CHECK-LABEL: @snegate_vector
+spirv.func @snegate_vector(%arg0: vector<4xi32>) "None" {
+ // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(dense<0> : vector<4xi32>) : vector<4xi32>
+ // CHECK: llvm.sub %[[ZERO]], %{{.*}} : vector<4xi32>
+ %0 = spirv.SNegate %arg0 : vector<4xi32>
+ spirv.Return
+}
>From 204b7a856af639ca4bb46501848cabdcef02e393 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Fri, 10 Jul 2026 13:45:23 +0200
Subject: [PATCH 2/2] Address comments
---
.../Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp | 34 +++++++++----------
1 file changed, 17 insertions(+), 17 deletions(-)
diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index 96acea0f4f127..a8bee01b08d26 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -90,17 +90,22 @@ static IntegerAttr minusOneIntegerAttribute(Type type, Builder builder) {
return builder.getIntegerAttr(integerType, -1);
}
+/// Creates `llvm.mlir.constant` with a scalar or vector integer value,
+/// broadcasting `scalarAttr` across the vector if `srcType` is a vector.
+static Value createIntegerConstant(Location loc, Type srcType, Type dstType,
+ PatternRewriter &rewriter,
+ IntegerAttr scalarAttr) {
+ if (auto vecType = dyn_cast<VectorType>(srcType))
+ return LLVM::ConstantOp::create(
+ rewriter, loc, dstType, SplatElementsAttr::get(vecType, scalarAttr));
+ return LLVM::ConstantOp::create(rewriter, loc, dstType, scalarAttr);
+}
+
/// Creates `llvm.mlir.constant` with all bits set for the given type.
static Value createConstantAllBitsSet(Location loc, Type srcType, Type dstType,
PatternRewriter &rewriter) {
- if (isa<VectorType>(srcType)) {
- return LLVM::ConstantOp::create(
- rewriter, loc, dstType,
- SplatElementsAttr::get(cast<ShapedType>(srcType),
- minusOneIntegerAttribute(srcType, rewriter)));
- }
- return LLVM::ConstantOp::create(rewriter, loc, dstType,
- minusOneIntegerAttribute(srcType, rewriter));
+ return createIntegerConstant(loc, srcType, dstType, rewriter,
+ minusOneIntegerAttribute(srcType, rewriter));
}
/// Creates `llvm.mlir.constant` with a floating-point scalar or vector value.
@@ -929,21 +934,16 @@ class SNegatePattern : public SPIRVToLLVMConversion<spirv::SNegateOp> {
LogicalResult
matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
- auto srcType = op.getType();
- auto dstType = getTypeConverter()->convertType(srcType);
+ Type srcType = op.getType();
+ Type dstType = getTypeConverter()->convertType(srcType);
if (!dstType)
return rewriter.notifyMatchFailure(op, "type conversion failed");
Location loc = op.getLoc();
- auto vecSrcType = dyn_cast<VectorType>(srcType);
IntegerAttr zeroAttr = rewriter.getIntegerAttr(
cast<IntegerType>(getElementTypeOrSelf(srcType)), 0);
- Value zero;
- if (vecSrcType)
- zero = LLVM::ConstantOp::create(
- rewriter, loc, dstType, SplatElementsAttr::get(vecSrcType, zeroAttr));
- else
- zero = LLVM::ConstantOp::create(rewriter, loc, dstType, zeroAttr);
+ Value zero =
+ createIntegerConstant(loc, srcType, dstType, rewriter, zeroAttr);
rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
adaptor.getOperand());
return success();
More information about the Mlir-commits
mailing list