[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