[Mlir-commits] [mlir] 9876e04 - [mlir][SPIR-V] Add SPIRVToLLVM conversion for VectorTimesScalar (#206949)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Tue Jul 28 08:23:44 PDT 2026
Author: Arseniy Obolenskiy
Date: 2026-07-28T17:23:39+02:00
New Revision: 9876e0469e315b767c4ba114e7f51207d198e762
URL: https://github.com/llvm/llvm-project/commit/9876e0469e315b767c4ba114e7f51207d198e762
DIFF: https://github.com/llvm/llvm-project/commit/9876e0469e315b767c4ba114e7f51207d198e762.diff
LOG: [mlir][SPIR-V] Add SPIRVToLLVM conversion for VectorTimesScalar (#206949)
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 2982061c957e0..99f9fe651d08f 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -919,6 +919,31 @@ class InverseSqrtPattern
}
};
+/// Converts `spirv.VectorTimesScalar` to a broadcast of the scalar followed by
+/// an `llvm.fmul`.
+class VectorTimesScalarPattern
+ : public SPIRVToLLVMConversion<spirv::VectorTimesScalarOp> {
+public:
+ using SPIRVToLLVMConversion<
+ spirv::VectorTimesScalarOp>::SPIRVToLLVMConversion;
+
+ LogicalResult
+ matchAndRewrite(spirv::VectorTimesScalarOp 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");
+
+ unsigned numElements = op.getVector().getType().getNumElements();
+ Value broadcasted = broadcast(op.getLoc(), adaptor.getScalar(), numElements,
+ *getTypeConverter(), rewriter);
+ rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getVector(),
+ broadcasted);
+ return success();
+ }
+};
+
/// Converts `spirv.SNegate` to `0 - x`.
class SNegatePattern : public SPIRVToLLVMConversion<spirv::SNegateOp> {
public:
@@ -1946,7 +1971,8 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
- DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, SNegatePattern,
+ DirectConversionPattern<spirv::UModOp, LLVM::URemOp>,
+ VectorTimesScalarPattern, 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 ee9c06fbb10e1..6b16335c3f804 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
@@ -234,6 +234,27 @@ spirv.func @srem_vector(%arg0: vector<4xi32>, %arg1: vector<4xi32>) "None" {
spirv.Return
}
+//===----------------------------------------------------------------------===//
+// spirv.VectorTimesScalar
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @vector_times_scalar
+// CHECK-SAME: %[[VECTOR:.*]]: vector<4xf32>, %[[SCALAR:.*]]: f32
+spirv.func @vector_times_scalar(%vector: vector<4xf32>, %scalar: f32) "None" {
+ // CHECK: %[[BCAST0:.*]] = llvm.mlir.poison : vector<4xf32>
+ // CHECK: %[[ZERO:.*]] = llvm.mlir.constant(0 : i32) : i32
+ // CHECK: %[[BCAST1:.*]] = llvm.insertelement %[[SCALAR]], %[[BCAST0]][%[[ZERO]] : i32] : vector<4xf32>
+ // CHECK: %[[ONE:.*]] = llvm.mlir.constant(1 : i32) : i32
+ // CHECK: %[[BCAST2:.*]] = llvm.insertelement %[[SCALAR]], %[[BCAST1]][%[[ONE]] : i32] : vector<4xf32>
+ // CHECK: %[[TWO:.*]] = llvm.mlir.constant(2 : i32) : i32
+ // CHECK: %[[BCAST3:.*]] = llvm.insertelement %[[SCALAR]], %[[BCAST2]][%[[TWO]] : i32] : vector<4xf32>
+ // CHECK: %[[THREE:.*]] = llvm.mlir.constant(3 : i32) : i32
+ // CHECK: %[[BCAST4:.*]] = llvm.insertelement %[[SCALAR]], %[[BCAST3]][%[[THREE]] : i32] : vector<4xf32>
+ // CHECK: llvm.fmul %[[VECTOR]], %[[BCAST4]] : vector<4xf32>
+ %0 = spirv.VectorTimesScalar %vector, %scalar : (vector<4xf32>, f32) -> vector<4xf32>
+ spirv.Return
+}
+
//===----------------------------------------------------------------------===//
// spirv.SNegate
//===----------------------------------------------------------------------===//
More information about the Mlir-commits
mailing list