[Mlir-commits] [mlir] [mlir][SPIR-V] Add SPIRVToLLVM conversions for GL.FMix and CL.mix (PR #206935)

Arseniy Obolenskiy llvmlistbot at llvm.org
Mon Jul 6 02:15:42 PDT 2026


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

>From 4fa5c4ca4a5a4333743f1fc3afa38979d20a176a Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 1 Jul 2026 11:54:35 +0200
Subject: [PATCH 1/3] [mlir][SPIR-V] Add SPIRVToLLVM conversions for GL.FMix
 and CL.mix

---
 .../Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp    | 27 +++++++++++++++++++
 .../SPIRVToLLVM/cl-ops-to-llvm.mlir           | 20 ++++++++++++++
 .../SPIRVToLLVM/gl-ops-to-llvm.mlir           | 20 ++++++++++++++
 3 files changed, 67 insertions(+)

diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index c43415b27b1b3..c3fff3b55587c 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -1564,6 +1564,31 @@ class SAbsPattern : public SPIRVToLLVMConversion<spirv::GLSAbsOp> {
   }
 };
 
+/// Converts `spirv.GL.FMix`/`spirv.CL.mix` to `fma(a, y - x, x)`, the linear
+/// blend `x * (1 - a) + y * a`. Operands are taken positionally because the GL
+/// and CL ops name their third operand differently (`a` vs. `z`).
+template <typename SPIRVOp>
+class MixPattern : public SPIRVToLLVMConversion<SPIRVOp> {
+public:
+  using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
+
+  LogicalResult
+  matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    auto dstType = this->getTypeConverter()->convertType(op.getType());
+    if (!dstType)
+      return rewriter.notifyMatchFailure(op, "type conversion failed");
+
+    Location loc = op.getLoc();
+    Value x = adaptor.getOperands()[0];
+    Value y = adaptor.getOperands()[1];
+    Value a = adaptor.getOperands()[2];
+    Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
+    rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
+    return success();
+  }
+};
+
 class VariablePattern : public SPIRVToLLVMConversion<spirv::VariableOp> {
 public:
   using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
@@ -1917,6 +1942,7 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
       DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
       DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
       InverseSqrtPattern, SAbsPattern, TanPattern, TanhPattern,
+      MixPattern<spirv::GLFMixOp>,
 
       // OpenCL extended instruction set ops
       DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
@@ -1950,6 +1976,7 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
       DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
       DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
       DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>,
+      MixPattern<spirv::CLMixOp>,
 
       // Logical ops
       DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
diff --git a/mlir/test/Conversion/SPIRVToLLVM/cl-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/cl-ops-to-llvm.mlir
index f0568e1035cd9..a986be42a5437 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/cl-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/cl-ops-to-llvm.mlir
@@ -99,3 +99,23 @@ spirv.func @cl_integer(%arg0: i32, %arg1: i32) "None" {
   %3 = spirv.CL.u_min %arg0, %arg1 : i32
   spirv.Return
 }
+
+//===----------------------------------------------------------------------===//
+// spirv.CL.mix
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @mix_scalar
+spirv.func @mix_scalar(%x: f32, %y: f32, %a: f32) "None" {
+  // CHECK: %[[DIFF:.*]] = llvm.fsub %{{.*}}, %{{.*}} : f32
+  // CHECK: llvm.intr.fma(%{{.*}}, %[[DIFF]], %{{.*}}) : (f32, f32, f32) -> f32
+  %0 = spirv.CL.mix %x, %y, %a : f32
+  spirv.Return
+}
+
+// CHECK-LABEL: @mix_vector
+spirv.func @mix_vector(%x: vector<4xf32>, %y: vector<4xf32>, %a: vector<4xf32>) "None" {
+  // CHECK: %[[DIFF:.*]] = llvm.fsub %{{.*}}, %{{.*}} : vector<4xf32>
+  // CHECK: llvm.intr.fma(%{{.*}}, %[[DIFF]], %{{.*}}) : (vector<4xf32>, vector<4xf32>, vector<4xf32>) -> vector<4xf32>
+  %0 = spirv.CL.mix %x, %y, %a : vector<4xf32>
+  spirv.Return
+}
diff --git a/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
index ffa47efbf9213..f405ef48c6fe1 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
@@ -361,3 +361,23 @@ spirv.func @asin_acos_atan(%arg0: f32, %arg1: vector<3xf16>) "None" {
   %2 = spirv.GL.Atan %arg0 : f32
   spirv.Return
 }
+
+//===----------------------------------------------------------------------===//
+// spirv.GL.FMix
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: @fmix_scalar
+spirv.func @fmix_scalar(%x: f32, %y: f32, %a: f32) "None" {
+  // CHECK: %[[DIFF:.*]] = llvm.fsub %{{.*}}, %{{.*}} : f32
+  // CHECK: llvm.intr.fma(%{{.*}}, %[[DIFF]], %{{.*}}) : (f32, f32, f32) -> f32
+  %0 = spirv.GL.FMix %x : f32, %y : f32, %a : f32 -> f32
+  spirv.Return
+}
+
+// CHECK-LABEL: @fmix_vector
+spirv.func @fmix_vector(%x: vector<4xf32>, %y: vector<4xf32>, %a: vector<4xf32>) "None" {
+  // CHECK: %[[DIFF:.*]] = llvm.fsub %{{.*}}, %{{.*}} : vector<4xf32>
+  // CHECK: llvm.intr.fma(%{{.*}}, %[[DIFF]], %{{.*}}) : (vector<4xf32>, vector<4xf32>, vector<4xf32>) -> vector<4xf32>
+  %0 = spirv.GL.FMix %x : vector<4xf32>, %y : vector<4xf32>, %a : vector<4xf32> -> vector<4xf32>
+  spirv.Return
+}

>From 5512c6003aa7212a15226316a7ca9c6aef850eb0 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Mon, 6 Jul 2026 06:56:23 +0200
Subject: [PATCH 2/3] Address comments

---
 .../Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp    | 45 ++++++++++++++-----
 .../SPIRVToLLVM/gl-ops-to-llvm.mlir           | 14 ++++--
 2 files changed, 44 insertions(+), 15 deletions(-)

diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index 24fdccb55ffaa..2fcc14777a5eb 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -1584,18 +1584,42 @@ class FractPattern : public SPIRVToLLVMConversion<spirv::GLFractOp> {
   }
 };
 
-/// Converts `spirv.GL.FMix`/`spirv.CL.mix` to `fma(a, y - x, x)`, the linear
-/// blend `x * (1 - a) + y * a`. Operands are taken positionally because the GL
-/// and CL ops name their third operand differently (`a` vs. `z`).
-template <typename SPIRVOp>
-class MixPattern : public SPIRVToLLVMConversion<SPIRVOp> {
+/// Converts `spirv.GL.FMix` to `x * (1 - a) + y * a` as specified by
+/// GL.std.450.
+class GLFMixPattern : public SPIRVToLLVMConversion<spirv::GLFMixOp> {
 public:
-  using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
+  using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
 
   LogicalResult
-  matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
+  matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
                   ConversionPatternRewriter &rewriter) const override {
-    auto dstType = this->getTypeConverter()->convertType(op.getType());
+    Type dstType = getTypeConverter()->convertType(op.getType());
+    if (!dstType)
+      return rewriter.notifyMatchFailure(op, "type conversion failed");
+
+    Location loc = op.getLoc();
+    Value x = adaptor.getX();
+    Value y = adaptor.getY();
+    Value a = adaptor.getA();
+    Value one = createFPConstant(loc, op.getType(), dstType, rewriter, 1.0);
+    Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
+    Value lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
+    Value rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
+    rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType, lhs, rhs);
+    return success();
+  }
+};
+
+/// Converts `spirv.CL.mix` to `fma(a, y - x, x)`. The OpenCL spec defines
+/// mix as `x + (y - x) * a` and explicitly permits FMA contractions.
+class CLMixPattern : public SPIRVToLLVMConversion<spirv::CLMixOp> {
+public:
+  using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
+
+  LogicalResult
+  matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    Type dstType = getTypeConverter()->convertType(op.getType());
     if (!dstType)
       return rewriter.notifyMatchFailure(op, "type conversion failed");
 
@@ -1991,7 +2015,7 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
       DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
       DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
       InverseSqrtPattern, SAbsPattern, TanPattern, TanhPattern, FractPattern,
-      MixPattern<spirv::GLFMixOp>,
+      GLFMixPattern,
 
       // OpenCL extended instruction set ops
       DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
@@ -2024,8 +2048,7 @@ void mlir::populateSPIRVToLLVMConversionPatterns(
       DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
       DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
       DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
-      DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>,
-      MixPattern<spirv::CLMixOp>,
+      DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
 
       // Logical ops
       DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
diff --git a/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
index beff049c7ebc8..caecf1f9f8d32 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/gl-ops-to-llvm.mlir
@@ -413,16 +413,22 @@ spirv.func @fract(%arg0: f32, %arg1: vector<3xf16>) "None" {
 
 // CHECK-LABEL: @fmix_scalar
 spirv.func @fmix_scalar(%x: f32, %y: f32, %a: f32) "None" {
-  // CHECK: %[[DIFF:.*]] = llvm.fsub %{{.*}}, %{{.*}} : f32
-  // CHECK: llvm.intr.fma(%{{.*}}, %[[DIFF]], %{{.*}}) : (f32, f32, f32) -> f32
+  // CHECK: %[[ONE:.*]] = llvm.mlir.constant(1.000000e+00 : f32) : f32
+  // CHECK: %[[ONE_MINUS_A:.*]] = llvm.fsub %[[ONE]], %{{.*}} : f32
+  // CHECK: %[[LHS:.*]] = llvm.fmul %{{.*}}, %[[ONE_MINUS_A]] : f32
+  // CHECK: %[[RHS:.*]] = llvm.fmul %{{.*}}, %{{.*}} : f32
+  // CHECK: llvm.fadd %[[LHS]], %[[RHS]] : f32
   %0 = spirv.GL.FMix %x : f32, %y : f32, %a : f32 -> f32
   spirv.Return
 }
 
 // CHECK-LABEL: @fmix_vector
 spirv.func @fmix_vector(%x: vector<4xf32>, %y: vector<4xf32>, %a: vector<4xf32>) "None" {
-  // CHECK: %[[DIFF:.*]] = llvm.fsub %{{.*}}, %{{.*}} : vector<4xf32>
-  // CHECK: llvm.intr.fma(%{{.*}}, %[[DIFF]], %{{.*}}) : (vector<4xf32>, vector<4xf32>, vector<4xf32>) -> vector<4xf32>
+  // CHECK: %[[ONE:.*]] = llvm.mlir.constant(dense<1.000000e+00> : vector<4xf32>) : vector<4xf32>
+  // CHECK: %[[ONE_MINUS_A:.*]] = llvm.fsub %[[ONE]], %{{.*}} : vector<4xf32>
+  // CHECK: %[[LHS:.*]] = llvm.fmul %{{.*}}, %[[ONE_MINUS_A]] : vector<4xf32>
+  // CHECK: %[[RHS:.*]] = llvm.fmul %{{.*}}, %{{.*}} : vector<4xf32>
+  // CHECK: llvm.fadd %[[LHS]], %[[RHS]] : vector<4xf32>
   %0 = spirv.GL.FMix %x : vector<4xf32>, %y : vector<4xf32>, %a : vector<4xf32> -> vector<4xf32>
   spirv.Return
 }

>From 6b8d06e38a7791aced5f7adb9f043ab91ac3f142 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Mon, 6 Jul 2026 11:15:31 +0200
Subject: [PATCH 3/3] Apply comment

---
 mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp | 6 +++---
 1 file changed, 3 insertions(+), 3 deletions(-)

diff --git a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
index 2fcc14777a5eb..d1c8822bf0f05 100644
--- a/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
+++ b/mlir/lib/Conversion/SPIRVToLLVM/SPIRVToLLVM.cpp
@@ -1624,9 +1624,9 @@ class CLMixPattern : public SPIRVToLLVMConversion<spirv::CLMixOp> {
       return rewriter.notifyMatchFailure(op, "type conversion failed");
 
     Location loc = op.getLoc();
-    Value x = adaptor.getOperands()[0];
-    Value y = adaptor.getOperands()[1];
-    Value a = adaptor.getOperands()[2];
+    Value x = adaptor.getX();
+    Value y = adaptor.getY();
+    Value a = adaptor.getZ();
     Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
     rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
     return success();



More information about the Mlir-commits mailing list