[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