[Mlir-commits] [mlir] [mlir][SPIR-V] Convert math.fpowi to spirv.GL.Pow (PR #200563)

Arseniy Obolenskiy llvmlistbot at llvm.org
Sat May 30 03:56:51 PDT 2026


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

>From 3b7497f18659576ad1aa591e12c15758df418877 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Sat, 30 May 2026 12:50:16 +0200
Subject: [PATCH] [mlir][SPIR-V] Convert math.fpowi to spirv.GL.Pow

---
 .../Conversion/MathToSPIRV/MathToSPIRV.cpp    | 101 +++++++++++++-----
 .../MathToSPIRV/math-to-gl-spirv.mlir         |  34 ++++++
 2 files changed, 106 insertions(+), 29 deletions(-)

diff --git a/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp b/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp
index 674fb1e586701..0489138285b0d 100644
--- a/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp
+++ b/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp
@@ -463,6 +463,48 @@ struct PowIOpPattern final : public OpConversionPattern<math::FPowIOp> {
   }
 };
 
+/// Converts math.fpowi to GLSL SPIR-V ops. GL has no integer-power op, so the
+/// exponent is converted to float and lowered through spirv.GL.Pow. As GL.Pow
+/// is undefined for a negative base, the base is made positive and the result
+/// is negated when the base is negative and the exponent is odd.
+struct PowIOpGLPattern final : public OpConversionPattern<math::FPowIOp> {
+  using Base::Base;
+
+  LogicalResult
+  matchAndRewrite(math::FPowIOp op, OpAdaptor adaptor,
+                  ConversionPatternRewriter &rewriter) const override {
+    if (LogicalResult res = checkSourceOpTypes(rewriter, op); failed(res))
+      return res;
+
+    Type dstType = getTypeConverter()->convertType(op.getType());
+    if (!dstType)
+      return failure();
+
+    Location loc = op.getLoc();
+    Value base = adaptor.getLhs();
+    Value power = adaptor.getRhs();
+
+    Value expFloat =
+        spirv::ConvertSToFOp::create(rewriter, loc, dstType, power);
+    Value abs = spirv::GLFAbsOp::create(rewriter, loc, base);
+    Value pow = spirv::GLPowOp::create(rewriter, loc, abs, expFloat);
+
+    Value zeroF = spirv::ConstantOp::getZero(dstType, loc, rewriter);
+    Value lessThan = spirv::FOrdLessThanOp::create(rewriter, loc, base, zeroF);
+
+    Type powerType = power.getType();
+    Value oneI = spirv::ConstantOp::getOne(powerType, loc, rewriter);
+    Value lowBit = spirv::BitwiseAndOp::create(rewriter, loc, power, oneI);
+    Value isOdd = spirv::IEqualOp::create(rewriter, loc, lowBit, oneI);
+
+    Value shouldNegate =
+        spirv::LogicalAndOp::create(rewriter, loc, lessThan, isOdd);
+    Value negate = spirv::FNegateOp::create(rewriter, loc, pow);
+    rewriter.replaceOpWithNewOp<spirv::SelectOp>(op, shouldNegate, negate, pow);
+    return success();
+  }
+};
+
 /// Converts math.round to GLSL SPIRV extended ops.
 struct RoundOpPattern final : public OpConversionPattern<math::RoundOp> {
   using Base::Base;
@@ -530,35 +572,36 @@ void populateMathToSPIRVPatterns(const SPIRVTypeConverter &typeConverter,
           typeConverter, patterns.getContext());
 
   // GLSL patterns
-  patterns.add<
-      CountLeadingZerosPattern, Log1pOpPattern<spirv::GLLogOp>, Log10OpPattern,
-      ExpM1OpPattern<spirv::GLExpOp>, PowFOpPattern, RoundOpPattern,
-      CheckedElementwiseOpPattern<math::AbsFOp, spirv::GLFAbsOp>,
-      CheckedElementwiseOpPattern<math::AbsIOp, spirv::GLSAbsOp>,
-      CheckedElementwiseOpPattern<math::AtanOp, spirv::GLAtanOp>,
-      CheckedElementwiseOpPattern<math::CeilOp, spirv::GLCeilOp>,
-      CheckedElementwiseOpPattern<math::ClampFOp, spirv::GLFClampOp>,
-      CheckedElementwiseOpPattern<math::CosOp, spirv::GLCosOp>,
-      CheckedElementwiseOpPattern<math::ExpOp, spirv::GLExpOp>,
-      CheckedElementwiseOpPattern<math::Exp2Op, spirv::GLExp2Op>,
-      CheckedElementwiseOpPattern<math::FloorOp, spirv::GLFloorOp>,
-      CheckedElementwiseOpPattern<math::FmaOp, spirv::GLFmaOp>,
-      CheckedElementwiseOpPattern<math::LogOp, spirv::GLLogOp>,
-      CheckedElementwiseOpPattern<math::Log2Op, spirv::GLLog2Op>,
-      CheckedElementwiseOpPattern<math::RoundEvenOp, spirv::GLRoundEvenOp>,
-      CheckedElementwiseOpPattern<math::RsqrtOp, spirv::GLInverseSqrtOp>,
-      CheckedElementwiseOpPattern<math::SinOp, spirv::GLSinOp>,
-      CheckedElementwiseOpPattern<math::SqrtOp, spirv::GLSqrtOp>,
-      CheckedElementwiseOpPattern<math::TanhOp, spirv::GLTanhOp>,
-      CheckedElementwiseOpPattern<math::TanOp, spirv::GLTanOp>,
-      CheckedElementwiseOpPattern<math::AsinOp, spirv::GLAsinOp>,
-      CheckedElementwiseOpPattern<math::AcosOp, spirv::GLAcosOp>,
-      CheckedElementwiseOpPattern<math::SinhOp, spirv::GLSinhOp>,
-      CheckedElementwiseOpPattern<math::CoshOp, spirv::GLCoshOp>,
-      CheckedElementwiseOpPattern<math::AsinhOp, spirv::GLAsinhOp>,
-      CheckedElementwiseOpPattern<math::AcoshOp, spirv::GLAcoshOp>,
-      CheckedElementwiseOpPattern<math::AtanhOp, spirv::GLAtanhOp>>(
-      typeConverter, patterns.getContext());
+  patterns
+      .add<CountLeadingZerosPattern, Log1pOpPattern<spirv::GLLogOp>,
+           Log10OpPattern, ExpM1OpPattern<spirv::GLExpOp>, PowFOpPattern,
+           PowIOpGLPattern, RoundOpPattern,
+           CheckedElementwiseOpPattern<math::AbsFOp, spirv::GLFAbsOp>,
+           CheckedElementwiseOpPattern<math::AbsIOp, spirv::GLSAbsOp>,
+           CheckedElementwiseOpPattern<math::AtanOp, spirv::GLAtanOp>,
+           CheckedElementwiseOpPattern<math::CeilOp, spirv::GLCeilOp>,
+           CheckedElementwiseOpPattern<math::ClampFOp, spirv::GLFClampOp>,
+           CheckedElementwiseOpPattern<math::CosOp, spirv::GLCosOp>,
+           CheckedElementwiseOpPattern<math::ExpOp, spirv::GLExpOp>,
+           CheckedElementwiseOpPattern<math::Exp2Op, spirv::GLExp2Op>,
+           CheckedElementwiseOpPattern<math::FloorOp, spirv::GLFloorOp>,
+           CheckedElementwiseOpPattern<math::FmaOp, spirv::GLFmaOp>,
+           CheckedElementwiseOpPattern<math::LogOp, spirv::GLLogOp>,
+           CheckedElementwiseOpPattern<math::Log2Op, spirv::GLLog2Op>,
+           CheckedElementwiseOpPattern<math::RoundEvenOp, spirv::GLRoundEvenOp>,
+           CheckedElementwiseOpPattern<math::RsqrtOp, spirv::GLInverseSqrtOp>,
+           CheckedElementwiseOpPattern<math::SinOp, spirv::GLSinOp>,
+           CheckedElementwiseOpPattern<math::SqrtOp, spirv::GLSqrtOp>,
+           CheckedElementwiseOpPattern<math::TanhOp, spirv::GLTanhOp>,
+           CheckedElementwiseOpPattern<math::TanOp, spirv::GLTanOp>,
+           CheckedElementwiseOpPattern<math::AsinOp, spirv::GLAsinOp>,
+           CheckedElementwiseOpPattern<math::AcosOp, spirv::GLAcosOp>,
+           CheckedElementwiseOpPattern<math::SinhOp, spirv::GLSinhOp>,
+           CheckedElementwiseOpPattern<math::CoshOp, spirv::GLCoshOp>,
+           CheckedElementwiseOpPattern<math::AsinhOp, spirv::GLAsinhOp>,
+           CheckedElementwiseOpPattern<math::AcoshOp, spirv::GLAcoshOp>,
+           CheckedElementwiseOpPattern<math::AtanhOp, spirv::GLAtanhOp>>(
+          typeConverter, patterns.getContext());
 
   // OpenCL patterns
   patterns.add<
diff --git a/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir b/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir
index 9dee79914b0dc..c1a67c25084c1 100644
--- a/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir
+++ b/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir
@@ -293,6 +293,40 @@ func.func @powf_const_mixed_int_exp_vector(%lhs: vector<4xf32>) -> vector<4xf32>
   return %0: vector<4xf32>
 }
 
+// CHECK-LABEL: @fpowi_scalar
+//  CHECK-SAME: (%[[BASE:.+]]: f32, %[[POW:.+]]: i32)
+func.func @fpowi_scalar(%base: f32, %power: i32) -> f32 {
+  // CHECK: %[[EXP:.+]] = spirv.ConvertSToF %[[POW]] : i32 to f32
+  // CHECK: %[[ABS:.+]] = spirv.GL.FAbs %[[BASE]] : f32
+  // CHECK: %[[POWF:.+]] = spirv.GL.Pow %[[ABS]], %[[EXP]] : f32
+  // CHECK: %[[F0:.+]] = spirv.Constant 0.000000e+00 : f32
+  // CHECK: %[[LT:.+]] = spirv.FOrdLessThan %[[BASE]], %[[F0]] : f32
+  // CHECK: %[[I1:.+]] = spirv.Constant 1 : i32
+  // CHECK: %[[AND:.+]] = spirv.BitwiseAnd %[[POW]], %[[I1]] : i32
+  // CHECK: %[[ODD:.+]] = spirv.IEqual %[[AND]], %[[I1]] : i32
+  // CHECK: %[[NEG_C:.+]] = spirv.LogicalAnd %[[LT]], %[[ODD]] : i1
+  // CHECK: %[[NEG:.+]] = spirv.FNegate %[[POWF]] : f32
+  // CHECK: %[[SEL:.+]] = spirv.Select %[[NEG_C]], %[[NEG]], %[[POWF]] : i1, f32
+  %0 = math.fpowi %base, %power : f32, i32
+  // CHECK: return %[[SEL]]
+  return %0 : f32
+}
+
+// CHECK-LABEL: @fpowi_vector
+func.func @fpowi_vector(%base: vector<4xf32>, %power: vector<4xi32>) -> vector<4xf32> {
+  // CHECK: spirv.ConvertSToF %{{.*}} : vector<4xi32> to vector<4xf32>
+  // CHECK: spirv.GL.FAbs %{{.*}} : vector<4xf32>
+  // CHECK: spirv.GL.Pow %{{.*}} : vector<4xf32>
+  // CHECK: spirv.FOrdLessThan %{{.*}} : vector<4xf32>
+  // CHECK: spirv.BitwiseAnd %{{.*}} : vector<4xi32>
+  // CHECK: spirv.IEqual %{{.*}} : vector<4xi32>
+  // CHECK: spirv.LogicalAnd %{{.*}} : vector<4xi1>
+  // CHECK: spirv.FNegate %{{.*}} : vector<4xf32>
+  // CHECK: spirv.Select %{{.*}} : vector<4xi1>, vector<4xf32>
+  %0 = math.fpowi %base, %power : vector<4xf32>, vector<4xi32>
+  return %0 : vector<4xf32>
+}
+
 // CHECK-LABEL: @round_scalar
 func.func @round_scalar(%x: f32) -> f32 {
   // CHECK: %[[ZERO:.+]] = spirv.Constant 0.000000e+00



More information about the Mlir-commits mailing list