[llvm] [NVPTX] Fold full-warp add reductions to redux.sync.add (PR #226345)
via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 24 19:56:17 PDT 2026
https://github.com/frost-miner created https://github.com/llvm/llvm-project/pull/226345
## Summary
Fold full-warp integer butterfly add reductions into `redux.sync.add` in the NVPTX IR peephole pass.
Compiler-generated full-warp reductions can be represented as a chain of five `shfl.sync.bfly.i32` operations and five integer additions. When the target supports `redux.sync.add`, this sequence can be replaced with a single reduction intrinsic, reducing the amount of generated code.
## Implementation
- Add warp reduction folding to `NVPTXIRPeephole`.
- Enable the transformation only when the subtarget supports both SM80 and PTX 7.0 or newer.
- Match five butterfly steps covering each lane-id bit exactly once.
- Require a full member mask, a clamp of `31`, and matching shuffle source operands.
- Accept different lane-mask orders and either operand ordering of the add.
- Replace the matched chain with `llvm.nvvm.redux.sync.add`.
## Tests
Added `nvptx-fold-warp-reduce.ll` to test supported and rejected warp-reduction folding cases on SM80/PTX70 and SM75. Extended `redux-sync.ll` to verify `redux.sync.add.s32` emission.
Both tests passed.
>From 42877677ffac69ca12fec98f5829dd2317b915ec Mon Sep 17 00:00:00 2001
From: frost-miner <2049180833 at qq.com>
Date: Fri, 25 Sep 2026 10:37:36 +0800
Subject: [PATCH] [NVPTX] Fold full-warp add reductions to redux.sync.add
---
llvm/lib/Target/NVPTX/NVPTX.h | 3 +
.../Target/NVPTX/NVPTXCodeGenPassBuilder.cpp | 2 +-
llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp | 118 ++++++-
llvm/lib/Target/NVPTX/NVPTXPassRegistry.def | 2 +-
llvm/lib/Target/NVPTX/NVPTXSubtarget.h | 3 +
.../CodeGen/NVPTX/nvptx-fold-warp-reduce.ll | 310 ++++++++++++++++++
llvm/test/CodeGen/NVPTX/redux-sync.ll | 20 ++
7 files changed, 452 insertions(+), 6 deletions(-)
create mode 100644 llvm/test/CodeGen/NVPTX/nvptx-fold-warp-reduce.ll
diff --git a/llvm/lib/Target/NVPTX/NVPTX.h b/llvm/lib/Target/NVPTX/NVPTX.h
index 6a58b55acd29d4..4bd21742892d72 100644
--- a/llvm/lib/Target/NVPTX/NVPTX.h
+++ b/llvm/lib/Target/NVPTX/NVPTX.h
@@ -159,7 +159,10 @@ class NVPTXImageOptimizerPass
};
class NVPTXIRPeepholePass : public OptionalPassInfoMixin<NVPTXIRPeepholePass> {
+ TargetMachine &TM;
+
public:
+ NVPTXIRPeepholePass(TargetMachine &TM) : TM(TM) {}
PreservedAnalyses run(Function &F, FunctionAnalysisManager &FAM);
};
diff --git a/llvm/lib/Target/NVPTX/NVPTXCodeGenPassBuilder.cpp b/llvm/lib/Target/NVPTX/NVPTXCodeGenPassBuilder.cpp
index d48b8f6f5fdc51..55ee1fad9a913d 100644
--- a/llvm/lib/Target/NVPTX/NVPTXCodeGenPassBuilder.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXCodeGenPassBuilder.cpp
@@ -234,7 +234,7 @@ void NVPTXCodeGenPassBuilder::addIRPasses(PassManagerWrapper &PMW) {
PMW);
addFunctionPass(NVPTXTagInvariantLoadsPass(), PMW);
if (!DisableNVPTXIRPeephole)
- addFunctionPass(NVPTXIRPeepholePass(), PMW);
+ addFunctionPass(NVPTXIRPeepholePass(TM), PMW);
}
if (ST.hasPTXASUnreachableBug()) {
diff --git a/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp b/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
index bd16c7213b1e79..c1c288b09814ab 100644
--- a/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXIRPeephole.cpp
@@ -20,17 +20,28 @@
// - fsub(a, fmul(b, c)) => fma(fneg(b), c, a)
// - fsub(fmul(a, b), fmul(c, d)) => fma(a, b, fneg(fmul(c, d)))
//
+// 2. Warp reduction folding (i32, redux.sync.add targets):
+// Transforms full-warp butterfly add reductions into redux.sync.add when
+// each lane-id bit is used exactly once. Supported pattern:
+// - add(%x, shfl.sync.bfly.i32(-1, %x, 1 << I, 31)), I = 0..4 in any order
+// => redux.sync.add(%x, -1)
+//
//===----------------------------------------------------------------------===//
+#include "NVPTXTargetMachine.h"
#include "NVPTXUtilities.h"
+#include "llvm/CodeGen/TargetPassConfig.h"
#include "llvm/IR/IRBuilder.h"
#include "llvm/IR/InstIterator.h"
#include "llvm/IR/Instructions.h"
#include "llvm/IR/Intrinsics.h"
+#include "llvm/IR/PatternMatch.h"
+#include "llvm/InitializePasses.h"
#define DEBUG_TYPE "nvptx-ir-peephole"
using namespace llvm;
+using namespace llvm::PatternMatch;
static bool tryFoldBinaryFMul(BinaryOperator *BI) {
Value *Op0 = BI->getOperand(0);
@@ -136,21 +147,117 @@ static bool foldFMA(Function &F) {
return Changed;
}
+// Fold a five-step butterfly add reduction into redux.sync.add.
+static bool tryFoldWarpReduceAdd(Instruction *Sum) {
+ // The chain is erased after replacing Sum. Keep users before their operands.
+ SmallVector<Instruction *, 10> Chain;
+ IntrinsicInst *FirstShfl = nullptr;
+ Value *Src = Sum;
+ unsigned SeenLanes = 0;
+
+ for (unsigned I = 0; I < 5; ++I) {
+ auto *Step = cast<Instruction>(Src);
+ // Intermediate partial sums must only feed the next step.
+ if (I != 0 && !Step->hasNUses(2))
+ return false;
+
+ Value *Op0, *Op1;
+ if (!match(Step, m_Add(m_Value(Op0), m_Value(Op1))))
+ return false;
+
+ uint64_t LaneMask = 0;
+ IntrinsicInst *Shfl = nullptr;
+ for (auto [Candidate, Other] : {std::pair(Op0, Op1), std::pair(Op1, Op0)}) {
+ auto *II = dyn_cast<IntrinsicInst>(Candidate);
+ const APInt *LM;
+ if (!II || II->getIntrinsicID() != Intrinsic::nvvm_shfl_sync_bfly_i32 ||
+ II->hasOperandBundles() || !II->hasOneUse() ||
+ !match(II->getArgOperand(0), m_AllOnes()) ||
+ II->getArgOperand(1) != Other ||
+ !match(II->getArgOperand(2), m_APInt(LM)) ||
+ !match(II->getArgOperand(3), m_SpecificInt(31)))
+ continue;
+ Src = Other;
+ LaneMask = LM->getZExtValue();
+ Shfl = II;
+ break;
+ }
+ if (!Shfl || !isPowerOf2_64(LaneMask) || LaneMask >= 32 ||
+ (SeenLanes & LaneMask))
+ return false;
+ SeenLanes |= LaneMask;
+
+ // Every step but the last must continue the chain.
+ if (I != 4 && !isa<Instruction>(Src))
+ return false;
+
+ Chain.push_back(Step);
+ Chain.push_back(Shfl);
+ FirstShfl = Shfl;
+ }
+
+ LLVM_DEBUG(dbgs() << "Found full-warp butterfly add reduction: " << *Sum
+ << "\n");
+
+ // Place the reduction at the first shuffle. The redux.sync.add wrap-around
+ // matches the add operations, so nsw/nuw flags are not preserved.
+ IRBuilder<> Builder(FirstShfl);
+ Value *Redux =
+ Builder.CreateIntrinsic(Intrinsic::nvvm_redux_sync_add, {},
+ {Src, Constant::getAllOnesValue(Src->getType())});
+ Redux->takeName(Sum);
+ Sum->replaceAllUsesWith(Redux);
+ for (Instruction *I : Chain)
+ I->eraseFromParent();
+ return true;
+}
+
+static bool foldWarpReductions(Function &F, const NVPTXSubtarget &ST) {
+ if (!ST.hasReduxSync())
+ return false;
+
+ // Folding erases instructions that may be later candidates.
+ SmallVector<WeakVH, 8> Candidates;
+ for (Instruction &I : instructions(F))
+ if (match(&I, m_c_Add(m_Value(),
+ m_Intrinsic<Intrinsic::nvvm_shfl_sync_bfly_i32>())))
+ Candidates.push_back(&I);
+
+ bool Changed = false;
+ for (WeakVH &V : Candidates)
+ if (auto *I = cast_or_null<Instruction>(V))
+ Changed |= tryFoldWarpReduceAdd(I);
+ return Changed;
+}
+
namespace {
struct NVPTXIRPeephole : public FunctionPass {
static char ID;
NVPTXIRPeephole() : FunctionPass(ID) {}
bool runOnFunction(Function &F) override;
+
+ void getAnalysisUsage(AnalysisUsage &AU) const override {
+ AU.addRequired<TargetPassConfig>();
+ }
};
} // namespace
char NVPTXIRPeephole::ID = 0;
-INITIALIZE_PASS(NVPTXIRPeephole, "nvptx-ir-peephole", "NVPTX IR Peephole",
- false, false)
+INITIALIZE_PASS_BEGIN(NVPTXIRPeephole, "nvptx-ir-peephole", "NVPTX IR Peephole",
+ false, false)
+INITIALIZE_PASS_DEPENDENCY(TargetPassConfig)
+INITIALIZE_PASS_END(NVPTXIRPeephole, "nvptx-ir-peephole", "NVPTX IR Peephole",
+ false, false)
-bool NVPTXIRPeephole::runOnFunction(Function &F) { return foldFMA(F); }
+bool NVPTXIRPeephole::runOnFunction(Function &F) {
+ auto &TM = getAnalysis<TargetPassConfig>().getTM<NVPTXTargetMachine>();
+ const NVPTXSubtarget &ST = TM.getSubtarget<NVPTXSubtarget>(F);
+ bool Changed = foldFMA(F);
+ Changed |= foldWarpReductions(F, ST);
+ return Changed;
+}
FunctionPass *llvm::createNVPTXIRPeepholePass() {
return new NVPTXIRPeephole();
@@ -158,7 +265,10 @@ FunctionPass *llvm::createNVPTXIRPeepholePass() {
PreservedAnalyses NVPTXIRPeepholePass::run(Function &F,
FunctionAnalysisManager &) {
- if (!foldFMA(F))
+ const NVPTXSubtarget &ST = TM.getSubtarget<NVPTXSubtarget>(F);
+ bool Changed = foldFMA(F);
+ Changed |= foldWarpReductions(F, ST);
+ if (!Changed)
return PreservedAnalyses::all();
PreservedAnalyses PA;
diff --git a/llvm/lib/Target/NVPTX/NVPTXPassRegistry.def b/llvm/lib/Target/NVPTX/NVPTXPassRegistry.def
index 118a34117773cb..b7c10309a5a5cb 100644
--- a/llvm/lib/Target/NVPTX/NVPTXPassRegistry.def
+++ b/llvm/lib/Target/NVPTX/NVPTXPassRegistry.def
@@ -45,7 +45,7 @@ FUNCTION_PASS("nvptx-alloca-hoisting", NVPTXAllocaHoistingPass())
FUNCTION_PASS("nvptx-atomic-lower", NVPTXAtomicLowerPass())
FUNCTION_PASS("nvptx-copy-byval-args", NVPTXCopyByValArgsPass())
FUNCTION_PASS("nvptx-image-optimizer", NVPTXImageOptimizerPass())
-FUNCTION_PASS("nvptx-ir-peephole", NVPTXIRPeepholePass())
+FUNCTION_PASS("nvptx-ir-peephole", NVPTXIRPeepholePass(*this))
FUNCTION_PASS("nvptx-lower-aggr-copies", NVPTXLowerAggrCopiesPass())
FUNCTION_PASS("nvptx-lower-alloca", NVPTXLowerAllocaPass())
FUNCTION_PASS("nvptx-lower-unreachable",
diff --git a/llvm/lib/Target/NVPTX/NVPTXSubtarget.h b/llvm/lib/Target/NVPTX/NVPTXSubtarget.h
index ee8adec2da060d..3dac32ae1a0e01 100644
--- a/llvm/lib/Target/NVPTX/NVPTXSubtarget.h
+++ b/llvm/lib/Target/NVPTX/NVPTXSubtarget.h
@@ -121,6 +121,9 @@ class NVPTXSubtarget : public NVPTXGenSubtargetInfo {
}
bool hasLocalVolatile() const { return hasFeature(NVPTX::PTX91); }
bool hasDotInstructions() const { return hasFeature(NVPTX::SM61); }
+ bool hasReduxSync() const {
+ return hasFeature(NVPTX::SM80) && hasFeature(NVPTX::PTX70);
+ }
bool hasCLMAD() const {
return hasFeature(NVPTX::SM80) && hasFeature(NVPTX::PTX93);
}
diff --git a/llvm/test/CodeGen/NVPTX/nvptx-fold-warp-reduce.ll b/llvm/test/CodeGen/NVPTX/nvptx-fold-warp-reduce.ll
new file mode 100644
index 00000000000000..98723fb860dd3e
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/nvptx-fold-warp-reduce.ll
@@ -0,0 +1,310 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt < %s -mcpu=sm_80 -mattr=+ptx70 -passes=nvptx-ir-peephole -S | FileCheck %s --check-prefixes=CHECK,SM80
+; RUN: opt < %s -mcpu=sm_75 -passes=nvptx-ir-peephole -S | FileCheck %s --check-prefixes=CHECK,SM75
+
+target triple = "nvptx64-nvidia-cuda"
+
+declare i32 @llvm.nvvm.shfl.sync.bfly.i32(i32, i32, i32, i32)
+declare i32 @use(i32)
+
+; The canonical five-step butterfly reduction from #224820.
+define i32 @reduce_add(i32 %x) {
+; SM80-LABEL: define i32 @reduce_add(
+; SM80-SAME: i32 [[X:%.*]]) #[[ATTR1:[0-9]+]] {
+; SM80-NEXT: [[ENTRY:.*:]]
+; SM80-NEXT: [[SUM:%.*]] = call i32 @llvm.nvvm.redux.sync.add(i32 [[X]], i32 -1)
+; SM80-NEXT: ret i32 [[SUM]]
+;
+; SM75-LABEL: define i32 @reduce_add(
+; SM75-SAME: i32 [[X:%.*]]) #[[ATTR1:[0-9]+]] {
+; SM75-NEXT: [[ENTRY:.*:]]
+; SM75-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 1, i32 31)
+; SM75-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; SM75-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; SM75-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; SM75-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; SM75-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; SM75-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; SM75-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; SM75-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 16, i32 31)
+; SM75-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S16]]
+; SM75-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 31)
+ %sum = add i32 %a8, %s16
+ ret i32 %sum
+}
+
+; The lane masks may be visited in any order, and the shuffle may be either
+; operand of the add. nsw/nuw flags are dropped.
+define i32 @reduce_add_reordered_commuted(i32 %x) {
+; SM80-LABEL: define i32 @reduce_add_reordered_commuted(
+; SM80-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; SM80-NEXT: [[ENTRY:.*:]]
+; SM80-NEXT: [[SUM:%.*]] = call i32 @llvm.nvvm.redux.sync.add(i32 [[X]], i32 -1)
+; SM80-NEXT: ret i32 [[SUM]]
+;
+; SM75-LABEL: define i32 @reduce_add_reordered_commuted(
+; SM75-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; SM75-NEXT: [[ENTRY:.*:]]
+; SM75-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 16, i32 31)
+; SM75-NEXT: [[A1:%.*]] = add nsw i32 [[S16]], [[X]]
+; SM75-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 1, i32 31)
+; SM75-NEXT: [[A2:%.*]] = add nuw i32 [[A1]], [[S1]]
+; SM75-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 8, i32 31)
+; SM75-NEXT: [[A3:%.*]] = add i32 [[S8]], [[A2]]
+; SM75-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A3]], i32 2, i32 31)
+; SM75-NEXT: [[A4:%.*]] = add i32 [[A3]], [[S2]]
+; SM75-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 4, i32 31)
+; SM75-NEXT: [[SUM:%.*]] = add nuw nsw i32 [[S4]], [[A4]]
+; SM75-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 16, i32 31)
+ %a1 = add nsw i32 %s16, %x
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 1, i32 31)
+ %a2 = add nuw i32 %a1, %s1
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 8, i32 31)
+ %a3 = add i32 %s8, %a2
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a3, i32 2, i32 31)
+ %a4 = add i32 %a3, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 4, i32 31)
+ %sum = add nsw nuw i32 %s4, %a4
+ ret i32 %sum
+}
+
+; The source of the reduction may have other uses.
+define i32 @reduce_add_src_multi_use(i32 %x) {
+; SM80-LABEL: define i32 @reduce_add_src_multi_use(
+; SM80-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; SM80-NEXT: [[ENTRY:.*:]]
+; SM80-NEXT: [[Y:%.*]] = mul i32 [[X]], 3
+; SM80-NEXT: [[SUM:%.*]] = call i32 @llvm.nvvm.redux.sync.add(i32 [[Y]], i32 -1)
+; SM80-NEXT: [[R:%.*]] = add i32 [[SUM]], [[Y]]
+; SM80-NEXT: ret i32 [[R]]
+;
+; SM75-LABEL: define i32 @reduce_add_src_multi_use(
+; SM75-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; SM75-NEXT: [[ENTRY:.*:]]
+; SM75-NEXT: [[Y:%.*]] = mul i32 [[X]], 3
+; SM75-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[Y]], i32 1, i32 31)
+; SM75-NEXT: [[A1:%.*]] = add i32 [[Y]], [[S1]]
+; SM75-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; SM75-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; SM75-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; SM75-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; SM75-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; SM75-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; SM75-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 16, i32 31)
+; SM75-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S16]]
+; SM75-NEXT: [[R:%.*]] = add i32 [[SUM]], [[Y]]
+; SM75-NEXT: ret i32 [[R]]
+;
+entry:
+ %y = mul i32 %x, 3
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %y, i32 1, i32 31)
+ %a1 = add i32 %y, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 31)
+ %sum = add i32 %a8, %s16
+ %r = add i32 %sum, %y
+ ret i32 %r
+}
+
+; Negative: only four steps reduce within groups of 16 lanes.
+define i32 @no_fold_four_steps(i32 %x) {
+; CHECK-LABEL: define i32 @no_fold_four_steps(
+; CHECK-SAME: i32 [[X:%.*]]) #[[ATTR1:[0-9]+]] {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 1, i32 31)
+; CHECK-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; CHECK-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; CHECK-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; CHECK-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; CHECK-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; CHECK-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; CHECK-NEXT: [[SUM:%.*]] = add i32 [[A4]], [[S8]]
+; CHECK-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %sum = add i32 %a4, %s8
+ ret i32 %sum
+}
+
+; Negative: lane mask 8 is repeated and 16 is missing.
+define i32 @no_fold_repeated_lane_mask(i32 %x) {
+; CHECK-LABEL: define i32 @no_fold_repeated_lane_mask(
+; CHECK-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 1, i32 31)
+; CHECK-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; CHECK-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; CHECK-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; CHECK-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; CHECK-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; CHECK-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; CHECK-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; CHECK-NEXT: [[S8B:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 8, i32 31)
+; CHECK-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S8B]]
+; CHECK-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s8b = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 8, i32 31)
+ %sum = add i32 %a8, %s8b
+ ret i32 %sum
+}
+
+; Negative: not all lanes participate.
+define i32 @no_fold_partial_member_mask(i32 %x) {
+; CHECK-LABEL: define i32 @no_fold_partial_member_mask(
+; CHECK-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 65535, i32 [[X]], i32 1, i32 31)
+; CHECK-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; CHECK-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; CHECK-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; CHECK-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; CHECK-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; CHECK-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; CHECK-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; CHECK-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 16, i32 31)
+; CHECK-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S16]]
+; CHECK-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 65535, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 31)
+ %sum = add i32 %a8, %s16
+ ret i32 %sum
+}
+
+; Negative: a clamp of 15 splits the warp into two 16-lane segments.
+define i32 @no_fold_segmented(i32 %x) {
+; CHECK-LABEL: define i32 @no_fold_segmented(
+; CHECK-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 1, i32 31)
+; CHECK-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; CHECK-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; CHECK-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; CHECK-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; CHECK-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; CHECK-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; CHECK-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; CHECK-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 16, i32 15)
+; CHECK-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S16]]
+; CHECK-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 15)
+ %sum = add i32 %a8, %s16
+ ret i32 %sum
+}
+
+; Negative: the partial sum %a4 is also used outside the chain.
+define i32 @no_fold_partial_sum_multi_use(i32 %x) {
+; CHECK-LABEL: define i32 @no_fold_partial_sum_multi_use(
+; CHECK-SAME: i32 [[X:%.*]]) #[[ATTR1]] {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 1, i32 31)
+; CHECK-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; CHECK-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; CHECK-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; CHECK-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; CHECK-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; CHECK-NEXT: [[TMP0:%.*]] = call i32 @use(i32 [[A4]])
+; CHECK-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A4]], i32 8, i32 31)
+; CHECK-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; CHECK-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 16, i32 31)
+; CHECK-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S16]]
+; CHECK-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ call i32 @use(i32 %a4)
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 31)
+ %sum = add i32 %a8, %s16
+ ret i32 %sum
+}
+
+; Negative: the shuffle reads a value other than the partial sum it is added to.
+define i32 @no_fold_mismatched_operand(i32 %x, i32 %y) {
+; CHECK-LABEL: define i32 @no_fold_mismatched_operand(
+; CHECK-SAME: i32 [[X:%.*]], i32 [[Y:%.*]]) #[[ATTR1]] {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[S1:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[X]], i32 1, i32 31)
+; CHECK-NEXT: [[A1:%.*]] = add i32 [[X]], [[S1]]
+; CHECK-NEXT: [[S2:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A1]], i32 2, i32 31)
+; CHECK-NEXT: [[A2:%.*]] = add i32 [[A1]], [[S2]]
+; CHECK-NEXT: [[S4:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A2]], i32 4, i32 31)
+; CHECK-NEXT: [[A4:%.*]] = add i32 [[A2]], [[S4]]
+; CHECK-NEXT: [[S8:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[Y]], i32 8, i32 31)
+; CHECK-NEXT: [[A8:%.*]] = add i32 [[A4]], [[S8]]
+; CHECK-NEXT: [[S16:%.*]] = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 [[A8]], i32 16, i32 31)
+; CHECK-NEXT: [[SUM:%.*]] = add i32 [[A8]], [[S16]]
+; CHECK-NEXT: ret i32 [[SUM]]
+;
+entry:
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %y, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 31)
+ %sum = add i32 %a8, %s16
+ ret i32 %sum
+}
diff --git a/llvm/test/CodeGen/NVPTX/redux-sync.ll b/llvm/test/CodeGen/NVPTX/redux-sync.ll
index 90b230850bd380..c8619e3de8c965 100644
--- a/llvm/test/CodeGen/NVPTX/redux-sync.ll
+++ b/llvm/test/CodeGen/NVPTX/redux-sync.ll
@@ -64,3 +64,23 @@ define i32 @redux_sync_or_b32(i32 %src, i32 %mask) {
%val = call i32 @llvm.nvvm.redux.sync.or(i32 %src, i32 %mask)
ret i32 %val
}
+
+; A full-warp butterfly add reduction is folded into redux.sync.add.
+declare i32 @llvm.nvvm.shfl.sync.bfly.i32(i32, i32, i32, i32)
+; CHECK-LABEL: .func{{.*}}butterfly_reduce_add
+define i32 @butterfly_reduce_add(i32 %x) {
+ ; CHECK-NOT: shfl.sync
+ ; CHECK: redux.sync.add.s32
+ ; CHECK-NOT: shfl.sync
+ %s1 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %x, i32 1, i32 31)
+ %a1 = add i32 %x, %s1
+ %s2 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a1, i32 2, i32 31)
+ %a2 = add i32 %a1, %s2
+ %s4 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a2, i32 4, i32 31)
+ %a4 = add i32 %a2, %s4
+ %s8 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a4, i32 8, i32 31)
+ %a8 = add i32 %a4, %s8
+ %s16 = call i32 @llvm.nvvm.shfl.sync.bfly.i32(i32 -1, i32 %a8, i32 16, i32 31)
+ %sum = add i32 %a8, %s16
+ ret i32 %sum
+}
More information about the llvm-commits
mailing list