[Mlir-commits] [mlir] [mlir][SPIR-V] Add SPIRVToLLVM conversion for VectorTimesScalar (PR #206949)

Arseniy Obolenskiy llvmlistbot at llvm.org
Tue Jul 28 07:02:21 PDT 2026


https://github.com/aobolensk updated https://github.com/llvm/llvm-project/pull/206949

>From 86c221950fd16910955f66c6be5c59c7f76b99f0 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 1 Jul 2026 13:10:12 +0200
Subject: [PATCH 1/2] [mlir][SPIR-V] Add SPIRVToLLVM conversion for
 VectorTimesScalar

---
 .../Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp    | 26 +++++++++++++++++++
 .../SPIRVToLLVM/arithmetic-ops-to-llvm.mlir   | 15 +++++++++++
 2 files changed, 41 insertions(+)

diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index c43415b27b1b3..23007a58713bd 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -921,6 +921,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 {
+    auto srcType = op.getType();
+    auto 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.Load` and `spirv.Store` to LLVM dialect.
 template <typename SPIRVOp>
 class LoadStorePattern : public SPIRVToLLVMConversion<SPIRVOp> {
@@ -1825,6 +1850,7 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
       DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
       DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
       DirectConversionPattern<spirv::UModOp, LLVM::URemOp>,
+      VectorTimesScalarPattern,
 
       // 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..0739f377c008f 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
@@ -233,3 +233,18 @@ spirv.func @srem_vector(%arg0: vector<4xi32>, %arg1: vector<4xi32>) "None" {
   %0 = spirv.SRem %arg0, %arg1 : vector<4xi32>
   spirv.Return
 }
+
+//===----------------------------------------------------------------------===//
+// spirv.VectorTimesScalar
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @vector_times_scalar
+spirv.func @vector_times_scalar(%vector: vector<4xf32>, %scalar: f32) "None" {
+  // CHECK: %[[UNDEF:.*]] = llvm.mlir.poison : vector<4xf32>
+  // CHECK: llvm.insertelement %{{.*}}, %[[UNDEF]]
+  // CHECK-COUNT-2: llvm.insertelement
+  // CHECK: %[[BCAST:.*]] = llvm.insertelement
+  // CHECK: llvm.fmul %{{.*}}, %[[BCAST]] : vector<4xf32>
+  %0 = spirv.VectorTimesScalar %vector, %scalar : (vector<4xf32>, f32) -> vector<4xf32>
+  spirv.Return
+}

>From 88eb10f80b2515a6f9a9c2438edd5ec557d33577 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 28 Jul 2026 16:02:07 +0200
Subject: [PATCH 2/2] Address comments

---
 mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp  |  4 ++--
 .../SPIRVToLLVM/arithmetic-ops-to-llvm.mlir      | 16 +++++++++++-----
 2 files changed, 13 insertions(+), 7 deletions(-)

diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index 9c1f4511afcfd..99f9fe651d08f 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -930,8 +930,8 @@ class VectorTimesScalarPattern
   LogicalResult
   matchAndRewrite(spirv::VectorTimesScalarOp 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");
 
diff --git a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
index efc6dd952a328..6b16335c3f804 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir
@@ -239,12 +239,18 @@ spirv.func @srem_vector(%arg0: vector<4xi32>, %arg1: vector<4xi32>) "None" {
 //===----------------------------------------------------------------------===//
 
 // CHECK-LABEL: @vector_times_scalar
+//  CHECK-SAME: %[[VECTOR:.*]]: vector<4xf32>, %[[SCALAR:.*]]: f32
 spirv.func @vector_times_scalar(%vector: vector<4xf32>, %scalar: f32) "None" {
-  // CHECK: %[[UNDEF:.*]] = llvm.mlir.poison : vector<4xf32>
-  // CHECK: llvm.insertelement %{{.*}}, %[[UNDEF]]
-  // CHECK-COUNT-2: llvm.insertelement
-  // CHECK: %[[BCAST:.*]] = llvm.insertelement
-  // CHECK: llvm.fmul %{{.*}}, %[[BCAST]] : vector<4xf32>
+  // 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
 }



More information about the Mlir-commits mailing list