1//===- Legality.h -----------------------------------------------*- C++ -*-===//
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// Legality checks for the Sandbox Vectorizer.
10//
11
12#ifndef LLVM_TRANSFORMS_VECTORIZE_SANDBOXVECTORIZER_LEGALITY_H
13#define LLVM_TRANSFORMS_VECTORIZE_SANDBOXVECTORIZER_LEGALITY_H
14
15#include "llvm/ADT/ArrayRef.h"
16#include "llvm/Analysis/ScalarEvolution.h"
17#include "llvm/Analysis/TargetTransformInfo.h"
18#include "llvm/IR/DataLayout.h"
19#include "llvm/Support/Casting.h"
20#include "llvm/Support/Compiler.h"
21#include "llvm/Support/raw_ostream.h"
22#include "llvm/Transforms/Vectorize/SandboxVectorizer/InstrMaps.h"
23#include "llvm/Transforms/Vectorize/SandboxVectorizer/Scheduler.h"
24
25namespace llvm::sandboxir {
26
27class LegalityAnalysis;
28class Value;
29class InstrMaps;
30
31class ShuffleMask {
32public:
33 using IndicesVecT = SmallVector<int, 8>;
34
35private:
36 IndicesVecT Indices;
37
38public:
39 ShuffleMask(SmallVectorImpl<int> &&Indices) : Indices(std::move(Indices)) {}
40 ShuffleMask(std::initializer_list<int> Indices) : Indices(Indices) {}
41 explicit ShuffleMask(ArrayRef<int> Indices) : Indices(Indices) {}
42 operator ArrayRef<int>() const { return Indices; }
43 /// Creates and returns an identity shuffle mask of size \p Sz.
44 /// For example if Sz == 4 the returned mask is {0, 1, 2, 3}.
45 static ShuffleMask getIdentity(unsigned Sz) {
46 IndicesVecT Indices;
47 Indices.reserve(N: Sz);
48 llvm::append_range(C&: Indices, R: seq<int>(Begin: 0, End: (int)Sz));
49 return ShuffleMask(std::move(Indices));
50 }
51 /// \Returns true if the mask is a perfect identity mask with consecutive
52 /// indices, i.e., performs no lane shuffling, like 0,1,2,3...
53 bool isIdentity() const {
54 for (auto [Idx, Elm] : enumerate(First: Indices)) {
55 if ((int)Idx != Elm)
56 return false;
57 }
58 return true;
59 }
60 bool operator==(const ShuffleMask &Other) const {
61 return Indices == Other.Indices;
62 }
63 bool operator!=(const ShuffleMask &Other) const { return !(*this == Other); }
64 size_t size() const { return Indices.size(); }
65 int operator[](int Idx) const { return Indices[Idx]; }
66 using const_iterator = IndicesVecT::const_iterator;
67 const_iterator begin() const { return Indices.begin(); }
68 const_iterator end() const { return Indices.end(); }
69#ifndef NDEBUG
70 friend raw_ostream &operator<<(raw_ostream &OS, const ShuffleMask &Mask) {
71 Mask.print(OS);
72 return OS;
73 }
74 void print(raw_ostream &OS) const {
75 interleave(Indices, OS, [&OS](auto Elm) { OS << Elm; }, ",");
76 }
77 LLVM_DUMP_METHOD void dump() const;
78#endif
79};
80
81enum class LegalityResultID {
82 Pack, ///> Collect scalar values.
83 Widen, ///> Vectorize by combining scalars to a vector.
84 DiamondReuse, ///> Don't generate new code, reuse existing vector.
85 DiamondReuseWithShuffle, ///> Reuse the existing vector but add a shuffle.
86 DiamondReuseMultiInput, ///> Reuse more than one vector and/or scalars.
87};
88
89/// The reason for vectorizing or not vectorizing.
90enum class ResultReason {
91 NotInstructions,
92 DiffOpcodes,
93 DiffTypes,
94 DiffMathFlags,
95 DiffWrapFlags,
96 DiffBBs,
97 RepeatedInstrs,
98 NotConsecutive,
99 AlignmentNotSupported,
100 CantSchedule,
101 Unimplemented,
102 Infeasible,
103 ForcePackForDebugging,
104};
105
106#ifndef NDEBUG
107struct ToStr {
108 static const char *getLegalityResultID(LegalityResultID ID) {
109 switch (ID) {
110 case LegalityResultID::Pack:
111 return "Pack";
112 case LegalityResultID::Widen:
113 return "Widen";
114 case LegalityResultID::DiamondReuse:
115 return "DiamondReuse";
116 case LegalityResultID::DiamondReuseWithShuffle:
117 return "DiamondReuseWithShuffle";
118 case LegalityResultID::DiamondReuseMultiInput:
119 return "DiamondReuseMultiInput";
120 }
121 llvm_unreachable("Unknown LegalityResultID enum");
122 }
123
124 static const char *getVecReason(ResultReason Reason) {
125 switch (Reason) {
126 case ResultReason::NotInstructions:
127 return "NotInstructions";
128 case ResultReason::DiffOpcodes:
129 return "DiffOpcodes";
130 case ResultReason::DiffTypes:
131 return "DiffTypes";
132 case ResultReason::DiffMathFlags:
133 return "DiffMathFlags";
134 case ResultReason::DiffWrapFlags:
135 return "DiffWrapFlags";
136 case ResultReason::DiffBBs:
137 return "DiffBBs";
138 case ResultReason::RepeatedInstrs:
139 return "RepeatedInstrs";
140 case ResultReason::NotConsecutive:
141 return "NotConsecutive";
142 case ResultReason::AlignmentNotSupported:
143 return "AlignmentNotSupported";
144 case ResultReason::CantSchedule:
145 return "CantSchedule";
146 case ResultReason::Unimplemented:
147 return "Unimplemented";
148 case ResultReason::Infeasible:
149 return "Infeasible";
150 case ResultReason::ForcePackForDebugging:
151 return "ForcePackForDebugging";
152 }
153 llvm_unreachable("Unknown ResultReason enum");
154 }
155};
156#endif // NDEBUG
157
158/// The legality outcome is represented by a class rather than an enum class
159/// because in some cases the legality checks are expensive and look for a
160/// particular instruction that can be passed along to the vectorizer to avoid
161/// repeating the same expensive computation.
162class LegalityResult {
163protected:
164 LegalityResultID ID;
165 /// Only Legality can create LegalityResults.
166 LegalityResult(LegalityResultID ID) : ID(ID) {}
167 friend class LegalityAnalysis;
168
169 /// We shouldn't need copies.
170 LegalityResult(const LegalityResult &) = delete;
171 LegalityResult &operator=(const LegalityResult &) = delete;
172
173public:
174 virtual ~LegalityResult() = default;
175 LegalityResultID getSubclassID() const { return ID; }
176#ifndef NDEBUG
177 virtual void print(raw_ostream &OS) const {
178 OS << ToStr::getLegalityResultID(ID);
179 }
180 LLVM_DUMP_METHOD void dump() const;
181 friend raw_ostream &operator<<(raw_ostream &OS, const LegalityResult &LR) {
182 LR.print(OS);
183 return OS;
184 }
185#endif // NDEBUG
186};
187
188/// Base class for results with reason.
189class LegalityResultWithReason : public LegalityResult {
190 ResultReason Reason;
191 LegalityResultWithReason(LegalityResultID ID, ResultReason Reason)
192 : LegalityResult(ID), Reason(Reason) {}
193 friend class Pack; // For constructor.
194
195public:
196 ResultReason getReason() const { return Reason; }
197#ifndef NDEBUG
198 void print(raw_ostream &OS) const override {
199 LegalityResult::print(OS);
200 OS << " Reason: " << ToStr::getVecReason(Reason);
201 }
202#endif
203};
204
205class Widen final : public LegalityResult {
206 friend class LegalityAnalysis;
207 Widen() : LegalityResult(LegalityResultID::Widen) {}
208
209public:
210 static bool classof(const LegalityResult *From) {
211 return From->getSubclassID() == LegalityResultID::Widen;
212 }
213};
214
215class DiamondReuse final : public LegalityResult {
216 friend class LegalityAnalysis;
217 Action *Vec;
218 DiamondReuse(Action *Vec)
219 : LegalityResult(LegalityResultID::DiamondReuse), Vec(Vec) {}
220
221public:
222 static bool classof(const LegalityResult *From) {
223 return From->getSubclassID() == LegalityResultID::DiamondReuse;
224 }
225 Action *getVector() const { return Vec; }
226};
227
228class DiamondReuseWithShuffle final : public LegalityResult {
229 friend class LegalityAnalysis;
230 Action *Vec;
231 ShuffleMask Mask;
232 DiamondReuseWithShuffle(Action *Vec, const ShuffleMask &Mask)
233 : LegalityResult(LegalityResultID::DiamondReuseWithShuffle), Vec(Vec),
234 Mask(Mask) {}
235
236public:
237 static bool classof(const LegalityResult *From) {
238 return From->getSubclassID() == LegalityResultID::DiamondReuseWithShuffle;
239 }
240 Action *getVector() const { return Vec; }
241 const ShuffleMask &getMask() const { return Mask; }
242};
243
244class Pack final : public LegalityResultWithReason {
245 Pack(ResultReason Reason)
246 : LegalityResultWithReason(LegalityResultID::Pack, Reason) {}
247 friend class LegalityAnalysis; // For constructor.
248
249public:
250 static bool classof(const LegalityResult *From) {
251 return From->getSubclassID() == LegalityResultID::Pack;
252 }
253};
254
255/// Describes how to collect the values needed by each lane.
256class CollectDescr {
257public:
258 /// Describes how to get a value element. If the value is a vector then it
259 /// also provides the index to extract it from.
260 class ExtractElementDescr {
261 PointerUnion<Action *, Value *> V = nullptr;
262 /// The index in `V` that the value can be extracted from.
263 int ExtractIdx = 0;
264
265 public:
266 ExtractElementDescr(Action *V, int ExtractIdx)
267 : V(V), ExtractIdx(ExtractIdx) {}
268 ExtractElementDescr(Value *V) : V(V) {}
269 Action *getValue() const { return cast<Action *>(Val: V); }
270 Value *getScalar() const { return cast<Value *>(Val: V); }
271 bool needsExtract() const { return isa<Action *>(Val: V); }
272 int getExtractIdx() const { return ExtractIdx; }
273 };
274
275 using DescrVecT = SmallVector<ExtractElementDescr, 4>;
276 DescrVecT Descrs;
277
278public:
279 CollectDescr(SmallVectorImpl<ExtractElementDescr> &&Descrs)
280 : Descrs(std::move(Descrs)) {}
281 /// If all elements come from a single vector input, then return that vector
282 /// and also the shuffle mask required to get them in order.
283 std::optional<std::pair<Action *, ShuffleMask>> getSingleInput() const {
284 const auto &Descr0 = *Descrs.begin();
285 if (!Descr0.needsExtract())
286 return std::nullopt;
287 auto *V0 = Descr0.getValue();
288 ShuffleMask::IndicesVecT MaskIndices;
289 MaskIndices.push_back(Elt: Descr0.getExtractIdx());
290 for (const auto &Descr : drop_begin(RangeOrContainer: Descrs)) {
291 if (!Descr.needsExtract())
292 return std::nullopt;
293 if (Descr.getValue() != V0)
294 return std::nullopt;
295 MaskIndices.push_back(Elt: Descr.getExtractIdx());
296 }
297 return std::make_pair(x&: V0, y: ShuffleMask(std::move(MaskIndices)));
298 }
299 bool hasVectorInputs() const {
300 return any_of(Range: Descrs, P: [](const auto &D) { return D.needsExtract(); });
301 }
302 const SmallVector<ExtractElementDescr, 4> &getDescrs() const {
303 return Descrs;
304 }
305};
306
307class DiamondReuseMultiInput final : public LegalityResult {
308 friend class LegalityAnalysis;
309 CollectDescr Descr;
310 DiamondReuseMultiInput(CollectDescr &&Descr)
311 : LegalityResult(LegalityResultID::DiamondReuseMultiInput),
312 Descr(std::move(Descr)) {}
313
314public:
315 static bool classof(const LegalityResult *From) {
316 return From->getSubclassID() == LegalityResultID::DiamondReuseMultiInput;
317 }
318 const CollectDescr &getCollectDescr() const { return Descr; }
319};
320
321/// Performs the legality analysis and returns a LegalityResult object.
322class LegalityAnalysis {
323 Scheduler Sched;
324 /// Owns the legality result objects created by createLegalityResult().
325 SmallVector<std::unique_ptr<LegalityResult>> ResultPool;
326 /// Checks opcodes, types and other IR-specifics and returns a ResultReason
327 /// object if not vectorizable, or nullptr otherwise.
328 std::optional<ResultReason>
329 notVectorizableBasedOnOpcodesAndTypes(BndlRef<Value *> Bndl);
330
331 ScalarEvolution &SE;
332 const DataLayout &DL;
333 TargetTransformInfo &TTI;
334 InstrMaps &IMaps;
335
336 /// Finds how we can collect the values in \p Bndl from the vectorized or
337 /// non-vectorized code. It returns a map of the value we should extract from
338 /// and the corresponding shuffle mask we need to use.
339 CollectDescr getHowToCollectValues(BndlRef<Value *> Bndl) const;
340
341public:
342 LegalityAnalysis(AAResults &AA, ScalarEvolution &SE, const DataLayout &DL,
343 TargetTransformInfo &TTI, Context &Ctx, InstrMaps &IMaps,
344 SchedDirection Dir)
345 : Sched(AA, Ctx, Dir), SE(SE), DL(DL), TTI(TTI), IMaps(IMaps) {}
346 /// A LegalityResult factory.
347 template <typename ResultT, typename... ArgsT>
348 ResultT &createLegalityResult(ArgsT &&...Args) {
349 ResultPool.push_back(
350 std::unique_ptr<ResultT>(new ResultT(std::move(Args)...)));
351 return cast<ResultT>(*ResultPool.back());
352 }
353
354 /// \returns true if \p Instrs are in different blocks.
355 template <typename ValueT>
356 static bool differentBlock(BndlRef<ValueT *> Instrs) {
357 auto *BB0 = cast<Instruction>(Instrs[0])->getParent();
358 return any_of(drop_begin(Instrs), [BB0](auto *V) {
359 return cast<Instruction>(V)->getParent() != BB0;
360 });
361 }
362
363 /// \returns true if all values in \p Values are unique.
364 template <typename ValueT> static bool areUnique(BndlRef<ValueT *> Values) {
365 SmallPtrSet<Value *, 8> Unique(llvm::from_range, Values);
366 return Unique.size() == Values.size();
367 }
368
369 /// \returns true if the alignment of the vector composed of \p Values has
370 /// alignment that is supported by the target.
371 LLVM_ABI bool isAlignmentSupported(ArrayRef<Value *> Values) const;
372
373 /// Checks if it's legal to vectorize the instructions in \p Bndl.
374 /// \Returns a LegalityResult object owned by LegalityAnalysis.
375 /// \p SkipScheduling skips the scheduler check and is only meant for testing.
376 // TODO: Try to remove the SkipScheduling argument by refactoring the tests.
377 LLVM_ABI const LegalityResult &canVectorize(BndlRef<Value *> Bndl,
378 bool SkipScheduling = false);
379 /// \Returns a Pack with reason 'ForcePackForDebugging'.
380 const LegalityResult &getForcedPackForDebugging() {
381 return createLegalityResult<Pack>(Args: ResultReason::ForcePackForDebugging);
382 }
383 LLVM_ABI void clear();
384};
385
386} // namespace llvm::sandboxir
387
388#endif // LLVM_TRANSFORMS_VECTORIZE_SANDBOXVECTORIZER_LEGALITY_H
389