LLVM 24.0.0git
ExpandReductions.cpp
Go to the documentation of this file.
1//===- ExpandReductions.cpp - Expand reduction intrinsics -----------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This pass implements IR expansion for reduction intrinsics, allowing targets
10// to enable the intrinsics until just before codegen.
11//
12//===----------------------------------------------------------------------===//
13
17#include "llvm/CodeGen/Passes.h"
18#include "llvm/IR/Dominators.h"
19#include "llvm/IR/IRBuilder.h"
22#include "llvm/IR/Intrinsics.h"
24#include "llvm/Pass.h"
26
27using namespace llvm;
28
29namespace {
30
31bool expandReductions(Function &F, const TargetTransformInfo *TTI,
32 DominatorTree *DT, LoopInfo *LI) {
33 bool Changed = false;
35 for (auto &I : instructions(F)) {
36 if (auto *II = dyn_cast<IntrinsicInst>(&I)) {
37 switch (II->getIntrinsicID()) {
38 default:
39 break;
40 case Intrinsic::vector_reduce_fadd:
41 case Intrinsic::vector_reduce_fmul:
42 case Intrinsic::vector_reduce_add:
43 case Intrinsic::vector_reduce_mul:
44 case Intrinsic::vector_reduce_and:
45 case Intrinsic::vector_reduce_or:
46 case Intrinsic::vector_reduce_xor:
47 case Intrinsic::vector_reduce_smax:
48 case Intrinsic::vector_reduce_smin:
49 case Intrinsic::vector_reduce_umax:
50 case Intrinsic::vector_reduce_umin:
51 case Intrinsic::vector_reduce_fmax:
52 case Intrinsic::vector_reduce_fmin:
53 case Intrinsic::vector_reduce_fmaximum:
54 case Intrinsic::vector_reduce_fminimum:
55 case Intrinsic::vector_reduce_fmaximumnum:
56 case Intrinsic::vector_reduce_fminimumnum: {
57 // Only expand if the target doesn't support this operation natively.
58 if (TTI->shouldExpandReduction(II))
59 Worklist.push_back(II);
60 break;
61 }
62 }
63 }
64 }
65
66 for (auto *II : Worklist) {
67 FastMathFlags FMF = II->getFastMathFlagsOrNone();
68 Intrinsic::ID ID = II->getIntrinsicID();
71 TTI->getPreferredExpandedReductionShuffle(II);
72
73 Value *Rdx = nullptr;
74 IRBuilder<> Builder(II);
75 IRBuilder<>::FastMathFlagGuard FMFGuard(Builder);
76 Builder.setFastMathFlags(FMF);
77 switch (ID) {
78 default:
79 llvm_unreachable("Unexpected intrinsic!");
80 case Intrinsic::vector_reduce_fadd:
81 case Intrinsic::vector_reduce_fmul: {
82 // FMFs must be attached to the call, otherwise it's an ordered reduction
83 // and it can't be handled by generating a shuffle sequence.
84 Value *Acc = II->getArgOperand(0);
85 Value *Vec = II->getArgOperand(1);
86 unsigned RdxOpcode = getArithmeticReductionInstruction(ID);
87 if (isa<ScalableVectorType>(Vec->getType())) {
88 Rdx = expandReductionViaLoop(Builder, Vec, RdxOpcode, Acc, DT, LI);
89 break;
90 }
91 if (!FMF.allowReassoc())
92 Rdx = getOrderedReduction(Builder, Acc, Vec, RdxOpcode, RK);
93 else {
94 if (!isPowerOf2_32(
95 cast<FixedVectorType>(Vec->getType())->getNumElements()))
96 continue;
97 Rdx = getShuffleReduction(Builder, Vec, RdxOpcode, RS, RK);
98 Rdx = Builder.CreateBinOp((Instruction::BinaryOps)RdxOpcode, Acc, Rdx,
99 "bin.rdx");
100 }
101 break;
102 }
103 case Intrinsic::vector_reduce_and:
104 case Intrinsic::vector_reduce_or: {
105 // Canonicalize logical or/and reductions:
106 // Or reduction for i1 is represented as:
107 // %val = bitcast <ReduxWidth x i1> to iReduxWidth
108 // %res = cmp ne iReduxWidth %val, 0
109 // And reduction for i1 is represented as:
110 // %val = bitcast <ReduxWidth x i1> to iReduxWidth
111 // %res = cmp eq iReduxWidth %val, 11111
112 Value *Vec = II->getArgOperand(0);
113 auto *FTy = cast<FixedVectorType>(Vec->getType());
114 unsigned NumElts = FTy->getNumElements();
115 if (!isPowerOf2_32(NumElts))
116 continue;
117
118 if (FTy->getElementType() == Builder.getInt1Ty()) {
119 Rdx = Builder.CreateBitCast(Vec, Builder.getIntNTy(NumElts));
120 if (ID == Intrinsic::vector_reduce_and) {
121 Rdx = Builder.CreateICmpEQ(
123 } else {
124 assert(ID == Intrinsic::vector_reduce_or && "Expected or reduction.");
125 Rdx = Builder.CreateIsNotNull(Rdx);
126 }
127 break;
128 }
129 unsigned RdxOpcode = getArithmeticReductionInstruction(ID);
130 Rdx = getShuffleReduction(Builder, Vec, RdxOpcode, RS, RK);
131 break;
132 }
133 case Intrinsic::vector_reduce_add:
134 case Intrinsic::vector_reduce_mul:
135 case Intrinsic::vector_reduce_xor:
136 case Intrinsic::vector_reduce_smax:
137 case Intrinsic::vector_reduce_smin:
138 case Intrinsic::vector_reduce_umax:
139 case Intrinsic::vector_reduce_umin: {
140 Value *Vec = II->getArgOperand(0);
141 unsigned RdxOpcode = getArithmeticReductionInstruction(ID);
142 if (isa<ScalableVectorType>(Vec->getType())) {
143 Type *EltTy = Vec->getType()->getScalarType();
144 Value *Ident = getReductionIdentity(ID, EltTy, FMF);
145 Rdx = expandReductionViaLoop(Builder, Vec, RdxOpcode, Ident, DT, LI);
146 break;
147 }
148 if (!isPowerOf2_32(
149 cast<FixedVectorType>(Vec->getType())->getNumElements()))
150 continue;
151 Rdx = getShuffleReduction(Builder, Vec, RdxOpcode, RS, RK);
152 break;
153 }
154 case Intrinsic::vector_reduce_fmax:
155 case Intrinsic::vector_reduce_fmin: {
156 // We require "nnan" to use a shuffle reduction; "nsz" is implied by the
157 // semantics of the reduction.
158 Value *Vec = II->getArgOperand(0);
159 if (!isPowerOf2_32(
160 cast<FixedVectorType>(Vec->getType())->getNumElements()) ||
161 !FMF.noNaNs())
162 continue;
163 unsigned RdxOpcode = getArithmeticReductionInstruction(ID);
164 Rdx = getShuffleReduction(Builder, Vec, RdxOpcode, RS, RK);
165 break;
166 }
167 case Intrinsic::vector_reduce_fmaximum:
168 case Intrinsic::vector_reduce_fminimum:
169 case Intrinsic::vector_reduce_fmaximumnum:
170 case Intrinsic::vector_reduce_fminimumnum: {
171 Value *Vec = II->getArgOperand(0);
172 if (!isPowerOf2_32(
173 cast<FixedVectorType>(Vec->getType())->getNumElements()))
174 continue;
175 unsigned RdxOpcode = getArithmeticReductionInstruction(ID);
176 Rdx = getShuffleReduction(Builder, Vec, RdxOpcode, RS, RK);
177 break;
178 }
179 }
180 II->replaceAllUsesWith(Rdx);
181 II->eraseFromParent();
182 Changed = true;
183 }
184 return Changed;
185}
186
187class ExpandReductions : public FunctionPass {
188public:
189 static char ID;
190 ExpandReductions() : FunctionPass(ID) {}
191
192 bool runOnFunction(Function &F) override {
193 const auto *TTI =&getAnalysis<TargetTransformInfoWrapperPass>().getTTI(F);
194 auto *DTWP = getAnalysisIfAvailable<DominatorTreeWrapperPass>();
195 auto *LIWP = getAnalysisIfAvailable<LoopInfoWrapperPass>();
196 auto *DT = DTWP ? &DTWP->getDomTree() : nullptr;
197 auto *LI = LIWP ? &LIWP->getLoopInfo() : nullptr;
198 return expandReductions(F, TTI, DT, LI);
199 }
200
201 void getAnalysisUsage(AnalysisUsage &AU) const override {
202 AU.addRequired<TargetTransformInfoWrapperPass>();
203 AU.addPreserved<DominatorTreeWrapperPass>();
204 AU.addPreserved<LoopInfoWrapperPass>();
205 }
206};
207}
208
209char ExpandReductions::ID;
210INITIALIZE_PASS_BEGIN(ExpandReductions, "expand-reductions",
211 "Expand reduction intrinsics", false, false)
213INITIALIZE_PASS_END(ExpandReductions, "expand-reductions",
214 "Expand reduction intrinsics", false, false)
215
217 return new ExpandReductions();
218}
219
222 const auto &TTI = AM.getResult<TargetIRAnalysis>(F);
224 auto *LI = AM.getCachedResult<LoopAnalysis>(F);
225 if (!expandReductions(F, &TTI, DT, LI))
226 return PreservedAnalyses::all();
230 return PA;
231}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
Expand Atomic instructions
static bool runOnFunction(Function &F, bool PostInlining)
#define F(x, y, z)
Definition MD5.cpp:54
#define I(x, y, z)
Definition MD5.cpp:57
uint64_t IntrinsicInst * II
#define INITIALIZE_PASS_DEPENDENCY(depName)
Definition PassSupport.h:42
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
Definition PassSupport.h:44
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
Definition PassSupport.h:39
This pass exposes codegen information to IR-level passes.
PassT::Result * getCachedResult(IRUnitT &IR) const
Get the cached result of an analysis pass for a given IR unit.
PassT::Result & getResult(IRUnitT &IR, ExtraArgTs... ExtraArgs)
Get the result of an analysis pass for a given IR unit.
AnalysisUsage & addRequired()
AnalysisUsage & addPreserved()
Add the specified Pass class to the set of analyses preserved by this pass.
static LLVM_ABI Constant * getAllOnesValue(Type *Ty)
Analysis pass which computes a DominatorTree.
Definition Dominators.h:241
Concrete subclass of DominatorTreeBase that is used to compute a normal dominator tree.
Definition Dominators.h:122
LLVM_ABI PreservedAnalyses run(Function &F, FunctionAnalysisManager &AM)
Convenience struct for specifying and reasoning about fast-math flags.
Definition FMF.h:23
bool allowReassoc() const
Flag queries.
Definition FMF.h:64
bool noNaNs() const
Definition FMF.h:65
FunctionPass class - This class is used to implement most global optimizations.
Definition Pass.h:314
This provides a uniform API for creating instructions and inserting them into a basic block: either a...
Definition IRBuilder.h:2903
Analysis pass that exposes the LoopInfo for a function.
Definition LoopInfo.h:594
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
PreservedAnalyses & preserve()
Mark an analysis as preserved.
Definition Analysis.h:132
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
Analysis pass providing the TargetTransformInfo.
Wrapper pass for TargetTransformInfo.
This pass provides access to the codegen interfaces that are needed for IR-level transformations.
The instances of the Type class are immutable: once they are created, they are never changed.
Definition Type.h:46
Type * getScalarType() const
If this is a vector type, return the element type, otherwise return 'this'.
Definition Type.h:368
LLVM Value Representation.
Definition Value.h:75
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:255
Changed
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
This is an optimization pass for GlobalISel generic memory operations.
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
LLVM_ABI Value * getReductionIdentity(Intrinsic::ID RdxID, Type *Ty, FastMathFlags FMF)
Given information about an @llvm.vector.reduce.
LLVM_ABI unsigned getArithmeticReductionInstruction(Intrinsic::ID RdxID)
Returns the arithmetic instruction opcode used when expanding a reduction.
constexpr bool isPowerOf2_32(uint32_t Value)
Return true if the argument is a power of two > 0.
Definition MathExtras.h:280
LLVM_ABI Value * getShuffleReduction(IRBuilderBase &Builder, Value *Src, unsigned Op, TargetTransformInfo::ReductionShuffle RS, RecurKind MinMaxKind=RecurKind::None)
Generates a vector reduction using shufflevectors to reduce the value.
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
Definition Casting.h:547
TargetTransformInfo TTI
RecurKind
These are the kinds of recurrences that we support.
LLVM_ABI FunctionPass * createExpandReductionsPass()
This pass expands the reduction intrinsics into sequences of shuffles.
LLVM_ABI Value * expandReductionViaLoop(IRBuilderBase &Builder, Value *Vec, unsigned RdxOpcode, Value *Acc, DominatorTree *DT=nullptr, LoopInfo *LI=nullptr)
Expand a scalable vector reduction into a runtime loop that applies RdxOpcode element by element,...
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
LLVM_ABI RecurKind getMinMaxReductionRecurKind(Intrinsic::ID RdxID)
Returns the recurence kind used when expanding a min/max reduction.
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
LLVM_ABI Value * getOrderedReduction(IRBuilderBase &Builder, Value *Acc, Value *Src, unsigned Op, RecurKind MinMaxKind=RecurKind::None)
Generates an ordered vector reduction using extracts to reduce the value.