[Mlir-commits] [mlir] [mlir][MathToSPIRV] Allow math.cttz lowering for non-i32 integer widths (PR #206400)

Arseniy Obolenskiy llvmlistbot at llvm.org
Sun Jun 28 22:07:11 PDT 2026


https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/206400

None

>From ce03d34e0d8e9adfcd2c210e49da949357db5848 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Mon, 29 Jun 2026 06:57:26 +0200
Subject: [PATCH] [mlir][MathToSPIRV] Allow math.cttz lowering for non-i32
 integer widths

---
 .../Conversion/MathToSPIRV/MathToSPIRV.cpp    | 16 +++++++++-----
 .../MathToSPIRV/math-to-gl-spirv.mlir         | 22 +++++++++++++++++++
 mlir/test/Target/SPIRV/gl-ops.mlir            | 21 +++++++++++++++++-
 3 files changed, 53 insertions(+), 6 deletions(-)

diff --git a/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp b/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp
index 61dddaf5155d3..892996f0ea840 100644
--- a/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp
+++ b/mlir/lib/Conversion/MathToSPIRV/MathToSPIRV.cpp
@@ -259,19 +259,25 @@ struct CountTrailingZerosPattern final
     if (!type)
       return failure();
 
+    auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
+    if (!typeConverter.getTargetEnv().allows(spirv::Capability::Shader))
+      return rewriter.notifyMatchFailure(countOp, "requires Shader capability");
+
     unsigned bitwidth = 0;
     if (isa<IntegerType>(type))
       bitwidth = type.getIntOrFloatBitWidth();
     else if (auto vectorType = dyn_cast<VectorType>(type))
       bitwidth = vectorType.getElementTypeBitWidth();
-    if (bitwidth != 32)
-      return failure();
 
     Location loc = countOp.getLoc();
     Value input = adaptor.getOperand();
-    Value val0 = getScalarOrVectorI32Constant(type, 0, rewriter, loc);
-    Value valBitwidth =
-        getScalarOrVectorI32Constant(type, bitwidth, rewriter, loc);
+    Value val0 = spirv::ConstantOp::getZero(type, loc, rewriter);
+    Type elemType = getElementTypeOrSelf(type);
+    Attribute bwAttr = IntegerAttr::get(elemType, bitwidth);
+    Attribute bwSplat = bwAttr;
+    if (auto vecType = dyn_cast<VectorType>(type))
+      bwSplat = SplatElementsAttr::get(vecType, bwAttr);
+    Value valBitwidth = spirv::ConstantOp::create(rewriter, loc, type, bwSplat);
 
     Value lsb = spirv::GLFindILsbOp::create(rewriter, loc, input);
     Value isZero = spirv::IEqualOp::create(rewriter, loc, input, val0);
diff --git a/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir b/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir
index 45e7bb01edc0b..45c921f0f7641 100644
--- a/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir
+++ b/mlir/test/Conversion/MathToSPIRV/math-to-gl-spirv.mlir
@@ -442,6 +442,28 @@ func.func @ctlz_vector2(%val: vector<2xi16>) -> vector<2xi16> {
   return %0 : vector<2xi16>
 }
 
+// CHECK-LABEL: @cttz_scalar_i64
+//  CHECK-SAME: (%[[VAL:.+]]: i64)
+func.func @cttz_scalar_i64(%val: i64) -> i64 {
+  // CHECK-DAG: %[[V0:.+]] = spirv.Constant 0 : i64
+  // CHECK-DAG: %[[V64:.+]] = spirv.Constant 64 : i64
+  // CHECK: %[[LSB:.+]] = spirv.GL.FindILsb %[[VAL]] : i64
+  // CHECK: %[[CMP:.+]] = spirv.IEqual %[[VAL]], %[[V0]] : i64
+  // CHECK: %[[R:.+]] = spirv.Select %[[CMP]], %[[V64]], %[[LSB]] : i1, i64
+  // CHECK: return %[[R]]
+  %0 = math.cttz %val : i64
+  return %0 : i64
+}
+
+// CHECK-LABEL: @cttz_vector_i16
+func.func @cttz_vector_i16(%val: vector<2xi16>) -> vector<2xi16> {
+  // CHECK: spirv.GL.FindILsb %{{.+}} : vector<2xi16>
+  // CHECK: spirv.IEqual
+  // CHECK: spirv.Select
+  %0 = math.cttz %val : vector<2xi16>
+  return %0 : vector<2xi16>
+}
+
 } // end module
 
 // -----
diff --git a/mlir/test/Target/SPIRV/gl-ops.mlir b/mlir/test/Target/SPIRV/gl-ops.mlir
index bc17ff82cc441..c643535cd5fbf 100644
--- a/mlir/test/Target/SPIRV/gl-ops.mlir
+++ b/mlir/test/Target/SPIRV/gl-ops.mlir
@@ -5,7 +5,7 @@
 // RUN: %if spirv-tools %{ mlir-translate --no-implicit-module --serialize-spirv --split-input-file --spirv-save-validation-files-with-prefix=%t/module %s %}
 // RUN: %if spirv-tools %{ spirv-val %t %}
 
-spirv.module Logical GLSL450 requires #spirv.vce<v1.0, [Shader, Linkage], []> {
+spirv.module Logical GLSL450 requires #spirv.vce<v1.0, [Shader, Linkage, Int16, Int64], []> {
   spirv.func @math(%arg0 : f32, %arg1 : f32, %arg2 : i32) "None" {
     // CHECK: {{%.*}} = spirv.GL.Exp {{%.*}} : f32
     %0 = spirv.GL.Exp %arg0 : f32
@@ -119,6 +119,25 @@ spirv.module Logical GLSL450 requires #spirv.vce<v1.0, [Shader, Linkage], []> {
     %2 = spirv.GL.FindILsb %arg0 : i32
     spirv.Return
   }
+
+  spirv.func @findilsb_i64(%arg0 : i64) "None" {
+    // CHECK: spirv.GL.FindILsb {{%.*}} : i64
+    %0 = spirv.GL.FindILsb %arg0 : i64
+    spirv.Return
+  }
+
+  spirv.func @findilsb_i16(%arg0 : i16) "None" {
+    // CHECK: spirv.GL.FindILsb {{%.*}} : i16
+    %0 = spirv.GL.FindILsb %arg0 : i16
+    spirv.Return
+  }
+
+  spirv.func @findilsb_vector_i64(%arg0 : vector<2xi64>) "None" {
+    // CHECK: spirv.GL.FindILsb {{%.*}} : vector<2xi64>
+    %0 = spirv.GL.FindILsb %arg0 : vector<2xi64>
+    spirv.Return
+  }
+
   spirv.func @findsmsb(%arg0 : i32) "None" {
     // CHECK: spirv.GL.FindSMsb {{%.*}} : i32
     %2 = spirv.GL.FindSMsb %arg0 : i32



More information about the Mlir-commits mailing list