1//===- SLPReductionUtils.cpp - SLP reduction match helpers ----------------===//
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#include "SLPReductionUtils.h"
10
11#include "SLPCostAnalysis.h"
12#include "SLPUtils.h"
13
14#include "llvm/ADT/STLExtras.h"
15#include "llvm/ADT/SmallBitVector.h"
16#include "llvm/Analysis/IVDescriptors.h"
17#include "llvm/Analysis/ValueTracking.h"
18#include "llvm/IR/Constants.h"
19#include "llvm/IR/DataLayout.h"
20#include "llvm/IR/IRBuilder.h"
21#include "llvm/IR/Instructions.h"
22#include "llvm/IR/Intrinsics.h"
23#include "llvm/IR/PatternMatch.h"
24#include "llvm/IR/Type.h"
25
26using namespace llvm;
27using namespace llvm::PatternMatch;
28
29namespace llvm::slpvectorizer {
30
31static bool matchRdxBop(Instruction *I, Value *&V0, Value *&V1) {
32 if (match(V: I, P: m_BinOp(L: m_Value(V&: V0), R: m_Value(V&: V1))))
33 return true;
34 if (match(V: I, P: m_FMaxNum(Op0: m_Value(V&: V0), Op1: m_Value(V&: V1))))
35 return true;
36 if (match(V: I, P: m_FMinNum(Op0: m_Value(V&: V0), Op1: m_Value(V&: V1))))
37 return true;
38 if (match(V: I, P: m_FMaximum(Op0: m_Value(V&: V0), Op1: m_Value(V&: V1))))
39 return true;
40 if (match(V: I, P: m_FMinimum(Op0: m_Value(V&: V0), Op1: m_Value(V&: V1))))
41 return true;
42 if (match(V: I, P: m_Intrinsic<Intrinsic::smax>(Ops: m_Value(V&: V0), Ops: m_Value(V&: V1))))
43 return true;
44 if (match(V: I, P: m_Intrinsic<Intrinsic::smin>(Ops: m_Value(V&: V0), Ops: m_Value(V&: V1))))
45 return true;
46 if (match(V: I, P: m_Intrinsic<Intrinsic::umax>(Ops: m_Value(V&: V0), Ops: m_Value(V&: V1))))
47 return true;
48 if (match(V: I, P: m_Intrinsic<Intrinsic::umin>(Ops: m_Value(V&: V0), Ops: m_Value(V&: V1))))
49 return true;
50 return false;
51}
52
53Instruction *getNonPhiOperand(Instruction *I, PHINode *Phi) {
54 Value *Op0 = nullptr;
55 Value *Op1 = nullptr;
56 if (!matchRdxBop(I, V0&: Op0, V1&: Op1))
57 return nullptr;
58 return dyn_cast<Instruction>(Val: Op0 == Phi ? Op1 : Op0);
59}
60
61bool isReductionCandidate(Instruction *I) {
62 bool IsSelect = match(V: I, P: m_Select(C: m_Value(), L: m_Value(), R: m_Value()));
63 Value *B0 = nullptr, *B1 = nullptr;
64 bool IsBinop = matchRdxBop(I, V0&: B0, V1&: B1);
65 return IsBinop || IsSelect;
66}
67
68Type *getBoolReduxWideTy(RecurKind RdxKind, Type *RootTy, Type *LeafTy) {
69 if ((RdxKind == RecurKind::And || RdxKind == RecurKind::Or) &&
70 RootTy->isIntegerTy(BitWidth: 1) && LeafTy->isIntegerTy() &&
71 !LeafTy->isIntegerTy(BitWidth: 1))
72 return LeafTy;
73 return nullptr;
74}
75
76BoolBitmask isBoolBitmaskRdx(
77 RecurKind RdxKind,
78 const SmallDenseMap<Value *, NarrowedLeafInfo> &NarrowedLeafShifts,
79 const DataLayout &DL) {
80 if (RdxKind != RecurKind::Or || DL.isBigEndian() ||
81 NarrowedLeafShifts.empty())
82 return BoolBitmask::None;
83 unsigned NumLeaves = NarrowedLeafShifts.size();
84 SmallBitVector Seen(NumLeaves);
85 bool NeedMask = false;
86 for (const auto &[V, L] : NarrowedLeafShifts) {
87 if (L.Shift >= NumLeaves || Seen.test(Idx: L.Shift))
88 return BoolBitmask::None;
89 Seen.set(L.Shift);
90 KnownBits Known = computeKnownBits(V, DL);
91 // The masked leaf must be known to be 0 or 1.
92 if ((L.Mask & ~Known.Zero).ugt(RHS: 1))
93 return BoolBitmask::None;
94 // The mask is redundant if it keeps all not-known-zero bits.
95 NeedMask |= !(Known.Zero | L.Mask).isAllOnes();
96 }
97 return NeedMask ? BoolBitmask::NeedMask : BoolBitmask::NoMask;
98}
99
100bool matchPackedFields(Value *V, unsigned MaxDepth,
101 SmallVectorImpl<Value *> &Fields,
102 SmallVectorImpl<Instruction *> &Chain) {
103 auto *PackTy = dyn_cast<IntegerType>(Val: V->getType());
104 if (!PackTy)
105 return false;
106 SmallVector<NarrowedLeafInfo> Leaves;
107 collectNarrowedLeaves(V, RdxOpcode: Instruction::Or, WideBW: PackTy->getBitWidth(), MaxDepth,
108 Leaves, ChainInsts&: Chain);
109 if (Leaves.size() < 2)
110 return false;
111 Type *FieldTy = Leaves.front().V->getType();
112 if (!FieldTy->isIntegerTy() ||
113 PackTy->getBitWidth() != Leaves.size() * FieldTy->getIntegerBitWidth())
114 return false;
115 llvm::sort(C&: Leaves, Comp: [](const NarrowedLeafInfo &A, const NarrowedLeafInfo &B) {
116 return A.Shift < B.Shift;
117 });
118 for (const auto &[Pos, L] : enumerate(First&: Leaves))
119 if (L.V->getType() != FieldTy || !L.Mask.isAllOnes() ||
120 L.Shift != Pos * FieldTy->getIntegerBitWidth())
121 return false;
122 append_range(
123 C&: Fields, R: map_range(C&: Leaves, F: [](const NarrowedLeafInfo &L) { return L.V; }));
124 return true;
125}
126
127Value *tryEmitBoolReduxBitcastCmp(IRBuilderBase &Builder,
128 const TargetTransformInfo &TTI,
129 RecurKind RdxKind, Value *Vec,
130 const Value *Root, FastMathFlags FMF,
131 const TTI::TargetCostKind CostKind) {
132 auto *VecTy = cast<FixedVectorType>(Val: Vec->getType());
133 unsigned VF = VecTy->getNumElements();
134 auto *I1VecTy = FixedVectorType::get(ElementType: Builder.getInt1Ty(), NumElts: VF);
135 DebugLoc DL = Builder.getCurrentDebugLocation();
136 Builder.SetCurrentDebugLocation(cast<Instruction>(Val: Root)->getDebugLoc());
137 Value *T = Builder.CreateTrunc(V: Vec, DestTy: I1VecTy);
138 Value *BC = Builder.CreateBitCast(V: T, DestTy: Builder.getIntNTy(N: VF));
139 CmpInst::Predicate Pred =
140 RdxKind == RecurKind::And ? CmpInst::ICMP_EQ : CmpInst::ICMP_NE;
141 Constant *RHS = RdxKind == RecurKind::And
142 ? Constant::getAllOnesValue(Ty: BC->getType())
143 : Constant::getNullValue(Ty: BC->getType());
144 Value *Res = Builder.CreateICmp(P: Pred, LHS: BC, RHS);
145 // The costs are evaluated from the emitted instructions; they are dropped
146 // if the wide reduction form is cheaper.
147 auto CastCost = [&](Value *V, unsigned Opcode, Type *SrcTy) {
148 auto *I = dyn_cast<Instruction>(Val: V);
149 if (!I)
150 return InstructionCost(0);
151 return TTI.getCastInstrCost(Opcode, Dst: I->getType(), Src: SrcTy,
152 CCH: TTI.getCastContextHint(I), CostKind, I);
153 };
154 InstructionCost BitcastCmpCost = CastCost(T, Instruction::Trunc, VecTy) +
155 CastCost(BC, Instruction::BitCast, I1VecTy);
156 if (auto *Cmp = dyn_cast<Instruction>(Val: Res))
157 BitcastCmpCost += TTI.getCmpSelInstrCost(
158 Opcode: Instruction::ICmp, ValTy: BC->getType(), /*CondTy=*/nullptr, VecPred: Pred, CostKind,
159 Op1Info: TTI.getOperandInfo(V: BC), Op2Info: TTI.getOperandInfo(V: RHS), I: Cmp);
160 if (BitcastCmpCost >=
161 getBoolReduxWideRdxCost(TTI, RdxKind, VecTy, Root, FMF, CostKind)) {
162 for (Value *V : {Res, BC, T})
163 if (auto *I = dyn_cast<Instruction>(Val: V))
164 I->eraseFromParent();
165 Builder.SetCurrentDebugLocation(DL);
166 return nullptr;
167 }
168 Builder.SetCurrentDebugLocation(DL);
169 return Res;
170}
171
172} // namespace llvm::slpvectorizer
173