[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