[Mlir-commits] [mlir] [mlir][SPIR-V] Add SPIRVToLLVM conversion for VectorTimesScalar (PR #206949)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed Jul 1 04:16:54 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-spirv
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
---
Full diff: https://github.com/llvm/llvm-project/pull/206949.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp (+26)
- (modified) mlir/test/Conversion/SPIRVToLLVM/arithmetic-ops-to-llvm.mlir (+15)
``````````diff
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
+}
``````````
</details>
https://github.com/llvm/llvm-project/pull/206949
More information about the Mlir-commits
mailing list