[Mlir-commits] [mlir] [mlir][gpu][nvvm] Add subgroup_reduce shuffle fallback for clustered and non-i32 cases (PR #209098)
Johnny Lin
llvmlistbot at llvm.org
Mon Jul 13 00:18:00 PDT 2026
https://github.com/johnny19436 updated https://github.com/llvm/llvm-project/pull/209098
>From 4fab19a557f8eefcfc40fdb5263ded9ca27c31f6 Mon Sep 17 00:00:00 2001
From: johnny19436 <johnny194369672 at gmail.com>
Date: Mon, 13 Jul 2026 15:19:33 +0800
Subject: [PATCH] [mlir][gpu][nvvm] Add subgroup_reduce shuffle fallback for
clustered and non-i32 cases
---
.../GPUToNVVM/LowerGpuOpsToNVVMOps.cpp | 35 +++++++++++++++++++
.../Conversion/GPUToNVVM/gpu-to-nvvm.mlir | 22 ++++++++++++
2 files changed, 57 insertions(+)
diff --git a/mlir/lib/Conversion/GPUToNVVM/LowerGpuOpsToNVVMOps.cpp b/mlir/lib/Conversion/GPUToNVVM/LowerGpuOpsToNVVMOps.cpp
index 80420c26537c3..59b150deb546d 100644
--- a/mlir/lib/Conversion/GPUToNVVM/LowerGpuOpsToNVVMOps.cpp
+++ b/mlir/lib/Conversion/GPUToNVVM/LowerGpuOpsToNVVMOps.cpp
@@ -96,6 +96,14 @@ convertToNVVMReductionKind(gpu::AllReduceOperation mode) {
return std::nullopt;
}
+static bool canLowerSubgroupReduceToNVVMRedux(gpu::SubgroupReduceOp op) {
+ if (op.getClusterSize() || !op.getUniform())
+ return false;
+ if (!op.getValue().getType().isInteger(32))
+ return false;
+ return convertToNVVMReductionKind(op.getOp()).has_value();
+}
+
static constexpr llvm::StringLiteral kNVVMNamedBarrierIdPrefix =
"__named_barrier_id";
static constexpr int32_t kNVVMFirstNamedBarrierId = 1;
@@ -529,6 +537,33 @@ struct LowerGpuOpsToNVVMOpsPass final
// ops which need to be lowered further, which is not supported by a
// single conversion pass.
{
+ // Lower subgroup reductions that cannot use nvvm.redux to shuffles
+ // before conversion. Keep redux-compatible cases untouched so the
+ // dedicated conversion pattern still applies.
+ SmallVector<Operation *> subgroupReduceOpsToLower;
+ m.walk([&](gpu::SubgroupReduceOp op) {
+ if (!op.getUniform())
+ return;
+ if (!this->hasRedux || !canLowerSubgroupReduceToNVVMRedux(op))
+ subgroupReduceOpsToLower.push_back(op.getOperation());
+ });
+ if (!subgroupReduceOpsToLower.empty()) {
+ RewritePatternSet subgroupReducePatterns(m.getContext());
+ populateGpuBreakDownSubgroupReducePatterns(
+ subgroupReducePatterns, /*maxShuffleBitwidth=*/kNVVMWarpSize);
+ populateGpuLowerSubgroupReduceToShufflePatterns(
+ subgroupReducePatterns,
+ /*subgroupSize=*/kNVVMWarpSize,
+ /*shuffleBitwidth=*/kNVVMWarpSize);
+ populateGpuLowerClusteredSubgroupReduceToShufflePatterns(
+ subgroupReducePatterns,
+ /*subgroupSize=*/kNVVMWarpSize,
+ /*shuffleBitwidth=*/kNVVMWarpSize);
+ if (failed(applyOpPatternsGreedily(subgroupReduceOpsToLower,
+ std::move(subgroupReducePatterns))))
+ return signalPassFailure();
+ }
+
RewritePatternSet patterns(m.getContext());
populateGpuRewritePatterns(patterns);
// Transform N-D vector.from_elements to 1-D vector.from_elements before
diff --git a/mlir/test/Conversion/GPUToNVVM/gpu-to-nvvm.mlir b/mlir/test/Conversion/GPUToNVVM/gpu-to-nvvm.mlir
index b96069ac41a44..8c76285b3e47e 100644
--- a/mlir/test/Conversion/GPUToNVVM/gpu-to-nvvm.mlir
+++ b/mlir/test/Conversion/GPUToNVVM/gpu-to-nvvm.mlir
@@ -739,6 +739,28 @@ gpu.module @test_module_30 {
}
}
+gpu.module @test_module_30_fallback {
+ // CHECK-LABEL: func @subgroup_reduce_add_f32_fallback
+ gpu.func @subgroup_reduce_add_f32_fallback(%arg0 : f32, %buf : memref<f32>) {
+ // CHECK-NOT: nvvm.redux.sync
+ // CHECK: nvvm.shfl.sync bfly
+ %result = gpu.subgroup_reduce add %arg0 uniform {} : (f32) -> (f32)
+ memref.store %result, %buf[] : memref<f32>
+ gpu.return
+ }
+
+ // CHECK-LABEL: func @subgroup_reduce_clustered_i32_fallback
+ gpu.func @subgroup_reduce_clustered_i32_fallback(%arg0 : i32,
+ %buf : memref<i32>) {
+ // CHECK-NOT: nvvm.redux.sync
+ // CHECK: nvvm.shfl.sync bfly
+ %result = gpu.subgroup_reduce add %arg0 uniform cluster(size = 8) :
+ (i32) -> (i32)
+ memref.store %result, %buf[] : memref<i32>
+ gpu.return
+ }
+}
+
gpu.module @test_module_31 {
// CHECK: llvm.func @__nv_fmodf(f32, f32) -> f32
// CHECK: llvm.func @__nv_fmod(f64, f64) -> f64
More information about the Mlir-commits
mailing list