1//===- ComplexDeinterleavingPass.cpp --------------------------------------===//
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// Identification:
10// This step is responsible for finding the patterns that can be lowered to
11// complex instructions, and building a graph to represent the complex
12// structures. Starting from the "Converging Shuffle" (a shuffle that
13// reinterleaves the complex components, with a mask of <0, 2, 1, 3>), the
14// operands are evaluated and identified as "Composite Nodes" (collections of
15// instructions that can potentially be lowered to a single complex
16// instruction). This is performed by checking the real and imaginary components
17// and tracking the data flow for each component while following the operand
18// pairs. Validity of each node is expected to be done upon creation, and any
19// validation errors should halt traversal and prevent further graph
20// construction.
21// Instead of relying on Shuffle operations, vector interleaving and
22// deinterleaving can be represented by vector.interleave2 and
23// vector.deinterleave2 intrinsics. Scalable vectors can be represented only by
24// these intrinsics, whereas, fixed-width vectors are recognized for both
25// shufflevector instruction and intrinsics.
26//
27// Replacement:
28// This step traverses the graph built up by identification, delegating to the
29// target to validate and generate the correct intrinsics, and plumbs them
30// together connecting each end of the new intrinsics graph to the existing
31// use-def chain. This step is assumed to finish successfully, as all
32// information is expected to be correct by this point.
33//
34//
35// Internal data structure:
36// ComplexDeinterleavingGraph:
37// Keeps references to all the valid CompositeNodes formed as part of the
38// transformation, and every Instruction contained within said nodes. It also
39// holds onto a reference to the root Instruction, and the root node that should
40// replace it.
41//
42// ComplexDeinterleavingCompositeNode:
43// A CompositeNode represents a single transformation point; each node should
44// transform into a single complex instruction (ignoring vector splitting, which
45// would generate more instructions per node). They are identified in a
46// depth-first manner, traversing and identifying the operands of each
47// instruction in the order they appear in the IR.
48// Each node maintains a reference to its Real and Imaginary instructions,
49// as well as any additional instructions that make up the identified operation
50// (Internal instructions should only have uses within their containing node).
51// A Node also contains the rotation and operation type that it represents.
52// Operands contains pointers to other CompositeNodes, acting as the edges in
53// the graph. ReplacementValue is the transformed Value* that has been emitted
54// to the IR.
55//
56// Note: If the operation of a Node is Shuffle, only the Real, Imaginary, and
57// ReplacementValue fields of that Node are relevant, where the ReplacementValue
58// should be pre-populated.
59//
60//===----------------------------------------------------------------------===//
61
62#include "llvm/CodeGen/ComplexDeinterleavingPass.h"
63#include "llvm/ADT/AllocatorList.h"
64#include "llvm/ADT/MapVector.h"
65#include "llvm/ADT/Statistic.h"
66#include "llvm/Analysis/TargetLibraryInfo.h"
67#include "llvm/Analysis/TargetTransformInfo.h"
68#include "llvm/CodeGen/TargetLowering.h"
69#include "llvm/CodeGen/TargetSubtargetInfo.h"
70#include "llvm/IR/IRBuilder.h"
71#include "llvm/IR/Intrinsics.h"
72#include "llvm/IR/PatternMatch.h"
73#include "llvm/InitializePasses.h"
74#include "llvm/Support/Allocator.h"
75#include "llvm/Target/TargetMachine.h"
76#include "llvm/Transforms/Utils/Local.h"
77#include <algorithm>
78
79using namespace llvm;
80using namespace PatternMatch;
81
82#define DEBUG_TYPE "complex-deinterleaving"
83
84STATISTIC(NumComplexTransformations, "Amount of complex patterns transformed");
85
86static cl::opt<bool> ComplexDeinterleavingEnabled(
87 "enable-complex-deinterleaving",
88 cl::desc("Enable generation of complex instructions"), cl::init(Val: true),
89 cl::Hidden);
90
91/// Checks the given mask, and determines whether said mask is interleaving.
92///
93/// To be interleaving, a mask must alternate between `i` and `i + (Length /
94/// 2)`, and must contain all numbers within the range of `[0..Length)` (e.g. a
95/// 4x vector interleaving mask would be <0, 2, 1, 3>).
96static bool isInterleavingMask(ArrayRef<int> Mask);
97
98/// Checks the given mask, and determines whether said mask is deinterleaving.
99///
100/// To be deinterleaving, a mask must increment in steps of 2, and either start
101/// with 0 or 1.
102/// (e.g. an 8x vector deinterleaving mask would be either <0, 2, 4, 6> or
103/// <1, 3, 5, 7>).
104static bool isDeinterleavingMask(ArrayRef<int> Mask);
105
106/// Returns true if the operation is a negation of V, and it works for both
107/// integers and floats.
108static bool isNeg(Value *V);
109
110/// Returns the operand for negation operation.
111static Value *getNegOperand(Value *V);
112
113namespace {
114struct ComplexValue {
115 Value *Real = nullptr;
116 Value *Imag = nullptr;
117
118 bool operator==(const ComplexValue &Other) const {
119 return Real == Other.Real && Imag == Other.Imag;
120 }
121};
122hash_code hash_value(const ComplexValue &Arg) {
123 return hash_combine(args: DenseMapInfo<Value *>::getHashValue(PtrVal: Arg.Real),
124 args: DenseMapInfo<Value *>::getHashValue(PtrVal: Arg.Imag));
125}
126} // end namespace
127typedef SmallVector<struct ComplexValue, 2> ComplexValues;
128
129template <> struct llvm::DenseMapInfo<ComplexValue> {
130 static unsigned getHashValue(const ComplexValue &Val) {
131 return hash_combine(args: DenseMapInfo<Value *>::getHashValue(PtrVal: Val.Real),
132 args: DenseMapInfo<Value *>::getHashValue(PtrVal: Val.Imag));
133 }
134 static bool isEqual(const ComplexValue &LHS, const ComplexValue &RHS) {
135 return LHS.Real == RHS.Real && LHS.Imag == RHS.Imag;
136 }
137};
138
139namespace {
140template <typename T, typename IterT>
141std::optional<T> findCommonBetweenCollections(IterT A, IterT B) {
142 auto Common = llvm::find_if(A, [B](T I) { return llvm::is_contained(B, I); });
143 if (Common != A.end())
144 return std::make_optional(*Common);
145 return std::nullopt;
146}
147
148class ComplexDeinterleavingLegacyPass : public FunctionPass {
149public:
150 static char ID;
151
152 ComplexDeinterleavingLegacyPass(const TargetMachine *TM = nullptr)
153 : FunctionPass(ID), TM(TM) {}
154
155 StringRef getPassName() const override {
156 return "Complex Deinterleaving Pass";
157 }
158
159 bool runOnFunction(Function &F) override;
160 void getAnalysisUsage(AnalysisUsage &AU) const override {
161 AU.addRequired<TargetLibraryInfoWrapperPass>();
162 AU.setPreservesCFG();
163 }
164
165private:
166 const TargetMachine *TM;
167};
168
169class ComplexDeinterleavingGraph;
170struct ComplexDeinterleavingCompositeNode {
171
172 ComplexDeinterleavingCompositeNode(ComplexDeinterleavingOperation Op,
173 Value *R, Value *I)
174 : Operation(Op) {
175 Vals.push_back(Elt: {.Real: R, .Imag: I});
176 }
177
178 ComplexDeinterleavingCompositeNode(ComplexDeinterleavingOperation Op,
179 ComplexValues &Other)
180 : Operation(Op), Vals(Other) {}
181
182private:
183 friend class ComplexDeinterleavingGraph;
184 using CompositeNode = ComplexDeinterleavingCompositeNode;
185 bool OperandsValid = true;
186
187public:
188 ComplexDeinterleavingOperation Operation;
189 ComplexValues Vals;
190
191 // This two members are required exclusively for generating
192 // ComplexDeinterleavingOperation::Symmetric operations.
193 unsigned Opcode;
194 std::optional<FastMathFlags> Flags;
195
196 ComplexDeinterleavingRotation Rotation =
197 ComplexDeinterleavingRotation::Rotation_0;
198 SmallVector<CompositeNode *> Operands;
199 Value *ReplacementNode = nullptr;
200
201 void addOperand(CompositeNode *Node) {
202 if (!Node)
203 OperandsValid = false;
204 Operands.push_back(Elt: Node);
205 }
206
207 void dump() { dump(OS&: dbgs()); }
208 void dump(raw_ostream &OS) {
209 auto PrintValue = [&](Value *V) {
210 if (V) {
211 OS << "\"";
212 V->print(O&: OS, IsForDebug: true);
213 OS << "\"\n";
214 } else
215 OS << "nullptr\n";
216 };
217 auto PrintNodeRef = [&](CompositeNode *Ptr) {
218 if (Ptr)
219 OS << Ptr << "\n";
220 else
221 OS << "nullptr\n";
222 };
223
224 OS << "- CompositeNode: " << this << "\n";
225 for (unsigned I = 0; I < Vals.size(); I++) {
226 OS << " Real(" << I << ") : ";
227 PrintValue(Vals[I].Real);
228 OS << " Imag(" << I << ") : ";
229 PrintValue(Vals[I].Imag);
230 }
231 OS << " ReplacementNode: ";
232 PrintValue(ReplacementNode);
233 OS << " Operation: " << (int)Operation << "\n";
234 OS << " Rotation: " << ((int)Rotation * 90) << "\n";
235 OS << " Operands: \n";
236 for (const auto &Op : Operands) {
237 OS << " - ";
238 PrintNodeRef(Op);
239 }
240 }
241
242 bool areOperandsValid() { return OperandsValid; }
243};
244
245class ComplexDeinterleavingGraph {
246public:
247 struct Product {
248 Value *Multiplier;
249 Value *Multiplicand;
250 bool IsPositive;
251 };
252
253 using Addend = std::pair<Value *, bool>;
254 using AddendList = BumpPtrList<Addend>;
255 using CompositeNode = ComplexDeinterleavingCompositeNode::CompositeNode;
256
257 // Helper struct for holding info about potential partial multiplication
258 // candidates
259 struct PartialMulCandidate {
260 Value *Common;
261 CompositeNode *Node;
262 unsigned RealIdx;
263 unsigned ImagIdx;
264 bool IsNodeInverted;
265 };
266
267 explicit ComplexDeinterleavingGraph(const TargetLowering *TL,
268 const TargetLibraryInfo *TLI,
269 unsigned Factor)
270 : TL(TL), TLI(TLI), Factor(Factor) {}
271
272private:
273 const TargetLowering *TL = nullptr;
274 const TargetLibraryInfo *TLI = nullptr;
275 unsigned Factor;
276 SmallVector<CompositeNode *> CompositeNodes;
277 DenseMap<ComplexValues, CompositeNode *> CachedResult;
278 SpecificBumpPtrAllocator<ComplexDeinterleavingCompositeNode> Allocator;
279
280 SmallPtrSet<Instruction *, 16> FinalInstructions;
281
282 /// Root instructions are instructions from which complex computation starts
283 DenseMap<Instruction *, CompositeNode *> RootToNode;
284
285 /// Topologically sorted root instructions
286 SmallVector<Instruction *, 1> OrderedRoots;
287
288 /// When examining a basic block for complex deinterleaving, if it is a simple
289 /// one-block loop, then the only incoming block is 'Incoming' and the
290 /// 'BackEdge' block is the block itself."
291 BasicBlock *BackEdge = nullptr;
292 BasicBlock *Incoming = nullptr;
293
294 /// ReductionInfo maps from %ReductionOp to %PHInode and Instruction
295 /// %OutsideUser as it is shown in the IR:
296 ///
297 /// vector.body:
298 /// %PHInode = phi <vector type> [ zeroinitializer, %entry ],
299 /// [ %ReductionOp, %vector.body ]
300 /// ...
301 /// %ReductionOp = fadd i64 ...
302 /// ...
303 /// br i1 %condition, label %vector.body, %middle.block
304 ///
305 /// middle.block:
306 /// %OutsideUser = llvm.vector.reduce.fadd(..., %ReductionOp)
307 ///
308 /// %OutsideUser can be `llvm.vector.reduce.fadd` or `fadd` preceding
309 /// `llvm.vector.reduce.fadd` when unroll factor isn't one.
310 MapVector<Instruction *, std::pair<PHINode *, Instruction *>> ReductionInfo;
311
312 /// In the process of detecting a reduction, we consider a pair of
313 /// %ReductionOP, which we refer to as real and imag (or vice versa), and
314 /// traverse the use-tree to detect complex operations. As this is a reduction
315 /// operation, it will eventually reach RealPHI and ImagPHI, which corresponds
316 /// to the %ReductionOPs that we suspect to be complex.
317 /// RealPHI and ImagPHI are used by the identifyPHINode method.
318 PHINode *RealPHI = nullptr;
319 PHINode *ImagPHI = nullptr;
320
321 /// Set this flag to true if RealPHI and ImagPHI were reached during reduction
322 /// detection.
323 bool PHIsFound = false;
324
325 /// OldToNewPHI maps the original real PHINode to a new, double-sized PHINode.
326 /// The new PHINode corresponds to a vector of deinterleaved complex numbers.
327 /// This mapping is populated during
328 /// ComplexDeinterleavingOperation::ReductionPHI node replacement. It is then
329 /// used in the ComplexDeinterleavingOperation::ReductionOperation node
330 /// replacement process.
331 DenseMap<PHINode *, PHINode *> OldToNewPHI;
332
333 CompositeNode *prepareCompositeNode(ComplexDeinterleavingOperation Operation,
334 Value *R, Value *I) {
335 assert(((Operation != ComplexDeinterleavingOperation::ReductionPHI &&
336 Operation != ComplexDeinterleavingOperation::ReductionOperation) ||
337 (R && I)) &&
338 "Reduction related nodes must have Real and Imaginary parts");
339 return new (Allocator.Allocate())
340 ComplexDeinterleavingCompositeNode(Operation, R, I);
341 }
342
343 CompositeNode *prepareCompositeNode(ComplexDeinterleavingOperation Operation,
344 ComplexValues &Vals) {
345#ifndef NDEBUG
346 for (auto &V : Vals) {
347 assert(
348 ((Operation != ComplexDeinterleavingOperation::ReductionPHI &&
349 Operation != ComplexDeinterleavingOperation::ReductionOperation) ||
350 (V.Real && V.Imag)) &&
351 "Reduction related nodes must have Real and Imaginary parts");
352 }
353#endif
354 return new (Allocator.Allocate())
355 ComplexDeinterleavingCompositeNode(Operation, Vals);
356 }
357
358 CompositeNode *submitCompositeNode(CompositeNode *Node) {
359 CompositeNodes.push_back(Elt: Node);
360 if (Node->Vals[0].Real)
361 CachedResult[Node->Vals] = Node;
362 return Node;
363 }
364
365 /// Identifies a complex partial multiply pattern and its rotation, based on
366 /// the following patterns
367 ///
368 /// 0: r: cr + ar * br
369 /// i: ci + ar * bi
370 /// 90: r: cr - ai * bi
371 /// i: ci + ai * br
372 /// 180: r: cr - ar * br
373 /// i: ci - ar * bi
374 /// 270: r: cr + ai * bi
375 /// i: ci - ai * br
376 CompositeNode *identifyPartialMul(Instruction *Real, Instruction *Imag);
377
378 /// Identify the other branch of a Partial Mul, taking the CommonOperandI that
379 /// is partially known from identifyPartialMul, filling in the other half of
380 /// the complex pair.
381 CompositeNode *
382 identifyNodeWithImplicitAdd(Instruction *I, Instruction *J,
383 std::pair<Value *, Value *> &CommonOperandI);
384
385 /// Identifies a complex add pattern and its rotation, based on the following
386 /// patterns.
387 ///
388 /// 90: r: ar - bi
389 /// i: ai + br
390 /// 270: r: ar + bi
391 /// i: ai - br
392 CompositeNode *identifyAdd(Instruction *Real, Instruction *Imag);
393 CompositeNode *identifySymmetricOperation(ComplexValues &Vals);
394 CompositeNode *identifyPartialReduction(Value *R, Value *I);
395 CompositeNode *identifyDotProduct(Value *Inst);
396
397 CompositeNode *identifyNode(ComplexValues &Vals);
398
399 CompositeNode *identifyNode(Value *R, Value *I) {
400 ComplexValues Vals;
401 Vals.push_back(Elt: {.Real: R, .Imag: I});
402 return identifyNode(Vals);
403 }
404
405 /// Determine if a sum of complex numbers can be formed from \p RealAddends
406 /// and \p ImagAddens. If \p Accumulator is not null, add the result to it.
407 /// Return nullptr if it is not possible to construct a complex number.
408 /// \p Flags are needed to generate symmetric Add and Sub operations.
409 CompositeNode *identifyAdditions(AddendList &RealAddends,
410 AddendList &ImagAddends,
411 std::optional<FastMathFlags> Flags,
412 CompositeNode *Accumulator);
413
414 /// Extract one addend that have both real and imaginary parts positive.
415 CompositeNode *extractPositiveAddend(AddendList &RealAddends,
416 AddendList &ImagAddends);
417
418 /// Determine if sum of multiplications of complex numbers can be formed from
419 /// \p RealMuls and \p ImagMuls. If \p Accumulator is not null, add the result
420 /// to it. Return nullptr if it is not possible to construct a complex number.
421 CompositeNode *identifyMultiplications(SmallVectorImpl<Product> &RealMuls,
422 SmallVectorImpl<Product> &ImagMuls,
423 CompositeNode *Accumulator);
424
425 /// Go through pairs of multiplication (one Real and one Imag) and find all
426 /// possible candidates for partial multiplication and put them into \p
427 /// Candidates. Returns true if all Product has pair with common operand
428 bool collectPartialMuls(ArrayRef<Product> RealMuls,
429 ArrayRef<Product> ImagMuls,
430 SmallVectorImpl<PartialMulCandidate> &Candidates);
431
432 /// If the code is compiled with -Ofast or expressions have `reassoc` flag,
433 /// the order of complex computation operations may be significantly altered,
434 /// and the real and imaginary parts may not be executed in parallel. This
435 /// function takes this into consideration and employs a more general approach
436 /// to identify complex computations. Initially, it gathers all the addends
437 /// and multiplicands and then constructs a complex expression from them.
438 CompositeNode *identifyReassocNodes(Instruction *I, Instruction *J);
439
440 CompositeNode *identifyRoot(Instruction *I);
441
442 /// Identifies the Deinterleave operation applied to a vector containing
443 /// complex numbers. There are two ways to represent the Deinterleave
444 /// operation:
445 /// * Using two shufflevectors with even indices for /pReal instruction and
446 /// odd indices for /pImag instructions (only for fixed-width vectors)
447 /// * Using N extractvalue instructions applied to `vector.deinterleaveN`
448 /// intrinsics (for both fixed and scalable vectors) where N is a multiple of
449 /// 2.
450 CompositeNode *identifyDeinterleave(ComplexValues &Vals);
451
452 /// identifying the operation that represents a complex number repeated in a
453 /// Splat vector. There are two possible types of splats: ConstantExpr with
454 /// the opcode ShuffleVector and ShuffleVectorInstr. Both should have an
455 /// initialization mask with all values set to zero.
456 CompositeNode *identifySplat(ComplexValues &Vals);
457
458 CompositeNode *identifyPHINode(Instruction *Real, Instruction *Imag);
459
460 /// Identifies SelectInsts in a loop that has reduction with predication masks
461 /// and/or predicated tail folding
462 CompositeNode *identifySelectNode(Instruction *Real, Instruction *Imag);
463
464 Value *replaceNode(IRBuilderBase &Builder, CompositeNode *Node);
465
466 /// Complete IR modifications after producing new reduction operation:
467 /// * Populate the PHINode generated for
468 /// ComplexDeinterleavingOperation::ReductionPHI
469 /// * Deinterleave the final value outside of the loop and repurpose original
470 /// reduction users
471 void processReductionOperation(Value *OperationReplacement,
472 CompositeNode *Node);
473 void processReductionSingle(Value *OperationReplacement, CompositeNode *Node);
474
475public:
476 void dump() { dump(OS&: dbgs()); }
477 void dump(raw_ostream &OS) {
478 for (const auto &Node : CompositeNodes)
479 Node->dump(OS);
480 }
481
482 /// Returns false if the deinterleaving operation should be cancelled for the
483 /// current graph.
484 bool identifyNodes(Instruction *RootI);
485
486 /// In case \pB is one-block loop, this function seeks potential reductions
487 /// and populates ReductionInfo. Returns true if any reductions were
488 /// identified.
489 bool collectPotentialReductions(BasicBlock *B);
490
491 void identifyReductionNodes();
492
493 /// Check that every instruction, from the roots to the leaves, has internal
494 /// uses.
495 bool checkNodes();
496
497 /// Perform the actual replacement of the underlying instruction graph.
498 void replaceNodes();
499};
500
501class ComplexDeinterleaving {
502public:
503 ComplexDeinterleaving(const TargetLowering *tl, const TargetLibraryInfo *tli)
504 : TL(tl), TLI(tli) {}
505 bool runOnFunction(Function &F);
506
507private:
508 bool evaluateBasicBlock(BasicBlock *B, unsigned Factor);
509
510 const TargetLowering *TL = nullptr;
511 const TargetLibraryInfo *TLI = nullptr;
512};
513
514} // namespace
515
516char ComplexDeinterleavingLegacyPass::ID = 0;
517
518INITIALIZE_PASS_BEGIN(ComplexDeinterleavingLegacyPass, DEBUG_TYPE,
519 "Complex Deinterleaving", false, false)
520INITIALIZE_PASS_END(ComplexDeinterleavingLegacyPass, DEBUG_TYPE,
521 "Complex Deinterleaving", false, false)
522
523PreservedAnalyses ComplexDeinterleavingPass::run(Function &F,
524 FunctionAnalysisManager &AM) {
525 const TargetLowering *TL = TM->getSubtargetImpl(F)->getTargetLowering();
526 auto &TLI = AM.getResult<llvm::TargetLibraryAnalysis>(IR&: F);
527 if (!ComplexDeinterleaving(TL, &TLI).runOnFunction(F))
528 return PreservedAnalyses::all();
529
530 PreservedAnalyses PA;
531 PA.preserve<FunctionAnalysisManagerModuleProxy>();
532 return PA;
533}
534
535FunctionPass *llvm::createComplexDeinterleavingPass(const TargetMachine *TM) {
536 return new ComplexDeinterleavingLegacyPass(TM);
537}
538
539bool ComplexDeinterleavingLegacyPass::runOnFunction(Function &F) {
540 const auto *TL = TM->getSubtargetImpl(F)->getTargetLowering();
541 auto TLI = getAnalysis<TargetLibraryInfoWrapperPass>().getTLI(F);
542 return ComplexDeinterleaving(TL, &TLI).runOnFunction(F);
543}
544
545bool ComplexDeinterleaving::runOnFunction(Function &F) {
546 if (!ComplexDeinterleavingEnabled) {
547 LLVM_DEBUG(
548 dbgs() << "Complex deinterleaving has been explicitly disabled.\n");
549 return false;
550 }
551
552 if (!TL->isComplexDeinterleavingSupported()) {
553 LLVM_DEBUG(
554 dbgs() << "Complex deinterleaving has been disabled, target does "
555 "not support lowering of complex number operations.\n");
556 return false;
557 }
558
559 bool Changed = false;
560 for (auto &B : F)
561 Changed |= evaluateBasicBlock(B: &B, Factor: 2);
562
563 // TODO: Permit changes for both interleave factors in the same function.
564 if (!Changed) {
565 for (auto &B : F)
566 Changed |= evaluateBasicBlock(B: &B, Factor: 4);
567 }
568
569 // TODO: We can also support interleave factors of 6 and 8 if needed.
570
571 return Changed;
572}
573
574static bool isInterleavingMask(ArrayRef<int> Mask) {
575 // If the size is not even, it's not an interleaving mask
576 if ((Mask.size() & 1))
577 return false;
578
579 int HalfNumElements = Mask.size() / 2;
580 for (int Idx = 0; Idx < HalfNumElements; ++Idx) {
581 int MaskIdx = Idx * 2;
582 if (Mask[MaskIdx] != Idx || Mask[MaskIdx + 1] != (Idx + HalfNumElements))
583 return false;
584 }
585
586 return true;
587}
588
589static bool isDeinterleavingMask(ArrayRef<int> Mask) {
590 int Offset = Mask[0];
591 int HalfNumElements = Mask.size() / 2;
592
593 for (int Idx = 1; Idx < HalfNumElements; ++Idx) {
594 if (Mask[Idx] != (Idx * 2) + Offset)
595 return false;
596 }
597
598 return true;
599}
600
601bool isNeg(Value *V) {
602 return match(V, P: m_FNeg(X: m_Value())) || match(V, P: m_Neg(V: m_Value()));
603}
604
605Value *getNegOperand(Value *V) {
606 assert(isNeg(V));
607 auto *I = cast<Instruction>(Val: V);
608 if (I->getOpcode() == Instruction::FNeg)
609 return I->getOperand(i: 0);
610
611 return I->getOperand(i: 1);
612}
613
614bool ComplexDeinterleaving::evaluateBasicBlock(BasicBlock *B, unsigned Factor) {
615 ComplexDeinterleavingGraph Graph(TL, TLI, Factor);
616 if (Graph.collectPotentialReductions(B))
617 Graph.identifyReductionNodes();
618
619 for (auto &I : *B)
620 Graph.identifyNodes(RootI: &I);
621
622 if (Graph.checkNodes()) {
623 Graph.replaceNodes();
624 return true;
625 }
626
627 return false;
628}
629
630ComplexDeinterleavingGraph::CompositeNode *
631ComplexDeinterleavingGraph::identifyNodeWithImplicitAdd(
632 Instruction *Real, Instruction *Imag,
633 std::pair<Value *, Value *> &PartialMatch) {
634 LLVM_DEBUG(dbgs() << "identifyNodeWithImplicitAdd " << *Real << " / " << *Imag
635 << "\n");
636
637 if (!Real->hasOneUse() || !Imag->hasOneUse()) {
638 LLVM_DEBUG(dbgs() << " - Mul operand has multiple uses.\n");
639 return nullptr;
640 }
641
642 if ((Real->getOpcode() != Instruction::FMul &&
643 Real->getOpcode() != Instruction::Mul) ||
644 (Imag->getOpcode() != Instruction::FMul &&
645 Imag->getOpcode() != Instruction::Mul)) {
646 LLVM_DEBUG(
647 dbgs() << " - Real or imaginary instruction is not fmul or mul\n");
648 return nullptr;
649 }
650
651 Value *R0 = Real->getOperand(i: 0);
652 Value *R1 = Real->getOperand(i: 1);
653 Value *I0 = Imag->getOperand(i: 0);
654 Value *I1 = Imag->getOperand(i: 1);
655
656 // A +/+ has a rotation of 0. If any of the operands are fneg, we flip the
657 // rotations and use the operand.
658 unsigned Negs = 0;
659 if (isNeg(V: R0)) {
660 Negs |= 1;
661 R0 = getNegOperand(V: R0);
662 } else if (isNeg(V: R1)) {
663 Negs |= 1;
664 R1 = getNegOperand(V: R1);
665 }
666
667 if (isNeg(V: I0)) {
668 Negs |= 2;
669 Negs ^= 1;
670 I0 = getNegOperand(V: I0);
671 } else if (isNeg(V: I1)) {
672 Negs |= 2;
673 Negs ^= 1;
674 I1 = getNegOperand(V: I1);
675 }
676
677 ComplexDeinterleavingRotation Rotation = (ComplexDeinterleavingRotation)Negs;
678
679 Value *CommonOperand;
680 Value *UncommonRealOp;
681 Value *UncommonImagOp;
682
683 if (R0 == I0 || R0 == I1) {
684 CommonOperand = R0;
685 UncommonRealOp = R1;
686 } else if (R1 == I0 || R1 == I1) {
687 CommonOperand = R1;
688 UncommonRealOp = R0;
689 } else {
690 LLVM_DEBUG(dbgs() << " - No equal operand\n");
691 return nullptr;
692 }
693
694 UncommonImagOp = (CommonOperand == I0) ? I1 : I0;
695 if (Rotation == ComplexDeinterleavingRotation::Rotation_90 ||
696 Rotation == ComplexDeinterleavingRotation::Rotation_270)
697 std::swap(a&: UncommonRealOp, b&: UncommonImagOp);
698
699 // Between identifyPartialMul and here we need to have found a complete valid
700 // pair from the CommonOperand of each part.
701 if (Rotation == ComplexDeinterleavingRotation::Rotation_0 ||
702 Rotation == ComplexDeinterleavingRotation::Rotation_180)
703 PartialMatch.first = CommonOperand;
704 else
705 PartialMatch.second = CommonOperand;
706
707 if (!PartialMatch.first || !PartialMatch.second) {
708 LLVM_DEBUG(dbgs() << " - Incomplete partial match\n");
709 return nullptr;
710 }
711
712 CompositeNode *CommonNode =
713 identifyNode(R: PartialMatch.first, I: PartialMatch.second);
714 if (!CommonNode) {
715 LLVM_DEBUG(dbgs() << " - No CommonNode identified\n");
716 return nullptr;
717 }
718
719 CompositeNode *UncommonNode = identifyNode(R: UncommonRealOp, I: UncommonImagOp);
720 if (!UncommonNode) {
721 LLVM_DEBUG(dbgs() << " - No UncommonNode identified\n");
722 return nullptr;
723 }
724
725 CompositeNode *Node = prepareCompositeNode(
726 Operation: ComplexDeinterleavingOperation::CMulPartial, R: Real, I: Imag);
727 Node->Rotation = Rotation;
728 Node->addOperand(Node: CommonNode);
729 Node->addOperand(Node: UncommonNode);
730 return submitCompositeNode(Node);
731}
732
733ComplexDeinterleavingGraph::CompositeNode *
734ComplexDeinterleavingGraph::identifyPartialMul(Instruction *Real,
735 Instruction *Imag) {
736 LLVM_DEBUG(dbgs() << "identifyPartialMul " << *Real << " / " << *Imag
737 << "\n");
738
739 // Determine rotation
740 auto IsAdd = [](unsigned Op) {
741 return Op == Instruction::FAdd || Op == Instruction::Add;
742 };
743 auto IsSub = [](unsigned Op) {
744 return Op == Instruction::FSub || Op == Instruction::Sub;
745 };
746 ComplexDeinterleavingRotation Rotation;
747 if (IsAdd(Real->getOpcode()) && IsAdd(Imag->getOpcode()))
748 Rotation = ComplexDeinterleavingRotation::Rotation_0;
749 else if (IsSub(Real->getOpcode()) && IsAdd(Imag->getOpcode()))
750 Rotation = ComplexDeinterleavingRotation::Rotation_90;
751 else if (IsSub(Real->getOpcode()) && IsSub(Imag->getOpcode()))
752 Rotation = ComplexDeinterleavingRotation::Rotation_180;
753 else if (IsAdd(Real->getOpcode()) && IsSub(Imag->getOpcode()))
754 Rotation = ComplexDeinterleavingRotation::Rotation_270;
755 else {
756 LLVM_DEBUG(dbgs() << " - Unhandled rotation.\n");
757 return nullptr;
758 }
759
760 if (isa<FPMathOperator>(Val: Real) &&
761 (!Real->getFastMathFlags().allowContract() ||
762 !Imag->getFastMathFlags().allowContract())) {
763 LLVM_DEBUG(dbgs() << " - Contract is missing from the FastMath flags.\n");
764 return nullptr;
765 }
766
767 Value *CR = Real->getOperand(i: 0);
768 Instruction *RealMulI = dyn_cast<Instruction>(Val: Real->getOperand(i: 1));
769 if (!RealMulI)
770 return nullptr;
771 Value *CI = Imag->getOperand(i: 0);
772 Instruction *ImagMulI = dyn_cast<Instruction>(Val: Imag->getOperand(i: 1));
773 if (!ImagMulI)
774 return nullptr;
775
776 if (!RealMulI->hasOneUse() || !ImagMulI->hasOneUse()) {
777 LLVM_DEBUG(dbgs() << " - Mul instruction has multiple uses\n");
778 return nullptr;
779 }
780
781 Value *R0 = RealMulI->getOperand(i: 0);
782 Value *R1 = RealMulI->getOperand(i: 1);
783 Value *I0 = ImagMulI->getOperand(i: 0);
784 Value *I1 = ImagMulI->getOperand(i: 1);
785
786 Value *CommonOperand;
787 Value *UncommonRealOp;
788 Value *UncommonImagOp;
789
790 if (R0 == I0 || R0 == I1) {
791 CommonOperand = R0;
792 UncommonRealOp = R1;
793 } else if (R1 == I0 || R1 == I1) {
794 CommonOperand = R1;
795 UncommonRealOp = R0;
796 } else {
797 LLVM_DEBUG(dbgs() << " - No equal operand\n");
798 return nullptr;
799 }
800
801 UncommonImagOp = (CommonOperand == I0) ? I1 : I0;
802 if (Rotation == ComplexDeinterleavingRotation::Rotation_90 ||
803 Rotation == ComplexDeinterleavingRotation::Rotation_270)
804 std::swap(a&: UncommonRealOp, b&: UncommonImagOp);
805
806 std::pair<Value *, Value *> PartialMatch(
807 (Rotation == ComplexDeinterleavingRotation::Rotation_0 ||
808 Rotation == ComplexDeinterleavingRotation::Rotation_180)
809 ? CommonOperand
810 : nullptr,
811 (Rotation == ComplexDeinterleavingRotation::Rotation_90 ||
812 Rotation == ComplexDeinterleavingRotation::Rotation_270)
813 ? CommonOperand
814 : nullptr);
815
816 auto *CRInst = dyn_cast<Instruction>(Val: CR);
817 auto *CIInst = dyn_cast<Instruction>(Val: CI);
818
819 if (!CRInst || !CIInst) {
820 LLVM_DEBUG(dbgs() << " - Common operands are not instructions.\n");
821 return nullptr;
822 }
823
824 CompositeNode *CNode =
825 identifyNodeWithImplicitAdd(Real: CRInst, Imag: CIInst, PartialMatch);
826 if (!CNode) {
827 LLVM_DEBUG(dbgs() << " - No cnode identified\n");
828 return nullptr;
829 }
830
831 CompositeNode *UncommonRes = identifyNode(R: UncommonRealOp, I: UncommonImagOp);
832 if (!UncommonRes) {
833 LLVM_DEBUG(dbgs() << " - No UncommonRes identified\n");
834 return nullptr;
835 }
836
837 assert(PartialMatch.first && PartialMatch.second);
838 CompositeNode *CommonRes =
839 identifyNode(R: PartialMatch.first, I: PartialMatch.second);
840 if (!CommonRes) {
841 LLVM_DEBUG(dbgs() << " - No CommonRes identified\n");
842 return nullptr;
843 }
844
845 CompositeNode *Node = prepareCompositeNode(
846 Operation: ComplexDeinterleavingOperation::CMulPartial, R: Real, I: Imag);
847 Node->Rotation = Rotation;
848 Node->addOperand(Node: CommonRes);
849 Node->addOperand(Node: UncommonRes);
850 Node->addOperand(Node: CNode);
851 return submitCompositeNode(Node);
852}
853
854ComplexDeinterleavingGraph::CompositeNode *
855ComplexDeinterleavingGraph::identifyAdd(Instruction *Real, Instruction *Imag) {
856 LLVM_DEBUG(dbgs() << "identifyAdd " << *Real << " / " << *Imag << "\n");
857
858 // Determine rotation
859 ComplexDeinterleavingRotation Rotation;
860 if ((Real->getOpcode() == Instruction::FSub &&
861 Imag->getOpcode() == Instruction::FAdd) ||
862 (Real->getOpcode() == Instruction::Sub &&
863 Imag->getOpcode() == Instruction::Add))
864 Rotation = ComplexDeinterleavingRotation::Rotation_90;
865 else if ((Real->getOpcode() == Instruction::FAdd &&
866 Imag->getOpcode() == Instruction::FSub) ||
867 (Real->getOpcode() == Instruction::Add &&
868 Imag->getOpcode() == Instruction::Sub))
869 Rotation = ComplexDeinterleavingRotation::Rotation_270;
870 else {
871 LLVM_DEBUG(dbgs() << " - Unhandled case, rotation is not assigned.\n");
872 return nullptr;
873 }
874
875 auto *AR = dyn_cast<Instruction>(Val: Real->getOperand(i: 0));
876 auto *BI = dyn_cast<Instruction>(Val: Real->getOperand(i: 1));
877 auto *AI = dyn_cast<Instruction>(Val: Imag->getOperand(i: 0));
878 auto *BR = dyn_cast<Instruction>(Val: Imag->getOperand(i: 1));
879
880 if (!AR || !AI || !BR || !BI) {
881 LLVM_DEBUG(dbgs() << " - Not all operands are instructions.\n");
882 return nullptr;
883 }
884
885 CompositeNode *ResA = identifyNode(R: AR, I: AI);
886 if (!ResA) {
887 LLVM_DEBUG(dbgs() << " - AR/AI is not identified as a composite node.\n");
888 return nullptr;
889 }
890 CompositeNode *ResB = identifyNode(R: BR, I: BI);
891 if (!ResB) {
892 LLVM_DEBUG(dbgs() << " - BR/BI is not identified as a composite node.\n");
893 return nullptr;
894 }
895
896 CompositeNode *Node =
897 prepareCompositeNode(Operation: ComplexDeinterleavingOperation::CAdd, R: Real, I: Imag);
898 Node->Rotation = Rotation;
899 Node->addOperand(Node: ResA);
900 Node->addOperand(Node: ResB);
901 return submitCompositeNode(Node);
902}
903
904static bool isInstructionPairAdd(Instruction *A, Instruction *B) {
905 unsigned OpcA = A->getOpcode();
906 unsigned OpcB = B->getOpcode();
907
908 return (OpcA == Instruction::FSub && OpcB == Instruction::FAdd) ||
909 (OpcA == Instruction::FAdd && OpcB == Instruction::FSub) ||
910 (OpcA == Instruction::Sub && OpcB == Instruction::Add) ||
911 (OpcA == Instruction::Add && OpcB == Instruction::Sub);
912}
913
914static bool isInstructionPairMul(Instruction *A, Instruction *B) {
915 auto Pattern =
916 m_BinOp(L: m_FMul(L: m_Value(), R: m_Value()), R: m_FMul(L: m_Value(), R: m_Value()));
917
918 return match(V: A, P: Pattern) && match(V: B, P: Pattern);
919}
920
921static bool isInstructionPotentiallySymmetric(Instruction *I) {
922 switch (I->getOpcode()) {
923 case Instruction::FAdd:
924 case Instruction::FSub:
925 case Instruction::FMul:
926 case Instruction::FNeg:
927 case Instruction::Add:
928 case Instruction::Sub:
929 case Instruction::Mul:
930 return true;
931 default:
932 return false;
933 }
934}
935
936ComplexDeinterleavingGraph::CompositeNode *
937ComplexDeinterleavingGraph::identifySymmetricOperation(ComplexValues &Vals) {
938 auto *FirstReal = cast<Instruction>(Val: Vals[0].Real);
939 unsigned FirstOpc = FirstReal->getOpcode();
940 for (auto &V : Vals) {
941 auto *Real = cast<Instruction>(Val: V.Real);
942 auto *Imag = cast<Instruction>(Val: V.Imag);
943 if (Real->getOpcode() != FirstOpc || Imag->getOpcode() != FirstOpc)
944 return nullptr;
945
946 if (!isInstructionPotentiallySymmetric(I: Real) ||
947 !isInstructionPotentiallySymmetric(I: Imag))
948 return nullptr;
949
950 if (isa<FPMathOperator>(Val: FirstReal))
951 if (Real->getFastMathFlags() != FirstReal->getFastMathFlags() ||
952 Imag->getFastMathFlags() != FirstReal->getFastMathFlags())
953 return nullptr;
954 }
955
956 ComplexValues OpVals;
957 for (auto &V : Vals) {
958 auto *R0 = cast<Instruction>(Val: V.Real)->getOperand(i: 0);
959 auto *I0 = cast<Instruction>(Val: V.Imag)->getOperand(i: 0);
960 OpVals.push_back(Elt: {.Real: R0, .Imag: I0});
961 }
962
963 CompositeNode *Op0 = identifyNode(Vals&: OpVals);
964 CompositeNode *Op1 = nullptr;
965 if (Op0 == nullptr)
966 return nullptr;
967
968 if (FirstReal->isBinaryOp()) {
969 OpVals.clear();
970 for (auto &V : Vals) {
971 auto *R1 = cast<Instruction>(Val: V.Real)->getOperand(i: 1);
972 auto *I1 = cast<Instruction>(Val: V.Imag)->getOperand(i: 1);
973 OpVals.push_back(Elt: {.Real: R1, .Imag: I1});
974 }
975 Op1 = identifyNode(Vals&: OpVals);
976 if (Op1 == nullptr)
977 return nullptr;
978 }
979
980 auto Node =
981 prepareCompositeNode(Operation: ComplexDeinterleavingOperation::Symmetric, Vals);
982 Node->Opcode = FirstReal->getOpcode();
983 if (isa<FPMathOperator>(Val: FirstReal))
984 Node->Flags = FirstReal->getFastMathFlags();
985
986 Node->addOperand(Node: Op0);
987 if (FirstReal->isBinaryOp())
988 Node->addOperand(Node: Op1);
989
990 return submitCompositeNode(Node);
991}
992
993ComplexDeinterleavingGraph::CompositeNode *
994ComplexDeinterleavingGraph::identifyDotProduct(Value *V) {
995 if (!TL->isComplexDeinterleavingOperationSupported(
996 Operation: ComplexDeinterleavingOperation::CDot, Ty: V->getType())) {
997 LLVM_DEBUG(dbgs() << "Target doesn't support complex deinterleaving "
998 "operation CDot with the type "
999 << *V->getType() << "\n");
1000 return nullptr;
1001 }
1002
1003 auto *Inst = cast<Instruction>(Val: V);
1004 auto *RealUser = cast<Instruction>(Val: *Inst->user_begin());
1005
1006 CompositeNode *CN =
1007 prepareCompositeNode(Operation: ComplexDeinterleavingOperation::CDot, R: Inst, I: nullptr);
1008
1009 CompositeNode *ANode = nullptr;
1010
1011 const Intrinsic::ID PartialReduceInt = Intrinsic::vector_partial_reduce_add;
1012
1013 Value *AReal = nullptr;
1014 Value *AImag = nullptr;
1015 Value *BReal = nullptr;
1016 Value *BImag = nullptr;
1017 Value *Phi = nullptr;
1018
1019 auto UnwrapCast = [](Value *V) -> Value * {
1020 if (auto *CI = dyn_cast<CastInst>(Val: V))
1021 return CI->getOperand(i_nocapture: 0);
1022 return V;
1023 };
1024
1025 auto PatternRot0 = m_Intrinsic<PartialReduceInt>(
1026 Ops: m_Intrinsic<PartialReduceInt>(Ops: m_Value(V&: Phi),
1027 Ops: m_Mul(L: m_Value(V&: BReal), R: m_Value(V&: AReal))),
1028 Ops: m_Neg(V: m_Mul(L: m_Value(V&: BImag), R: m_Value(V&: AImag))));
1029
1030 auto PatternRot270 = m_Intrinsic<PartialReduceInt>(
1031 Ops: m_Intrinsic<PartialReduceInt>(
1032 Ops: m_Value(V&: Phi), Ops: m_Neg(V: m_Mul(L: m_Value(V&: BReal), R: m_Value(V&: AImag)))),
1033 Ops: m_Mul(L: m_Value(V&: BImag), R: m_Value(V&: AReal)));
1034
1035 if (match(V: Inst, P: PatternRot0)) {
1036 CN->Rotation = ComplexDeinterleavingRotation::Rotation_0;
1037 } else if (match(V: Inst, P: PatternRot270)) {
1038 CN->Rotation = ComplexDeinterleavingRotation::Rotation_270;
1039 } else {
1040 Value *A0, *A1;
1041 // The rotations 90 and 180 share the same operation pattern, so inspect the
1042 // order of the operands, identifying where the real and imaginary
1043 // components of A go, to discern between the aforementioned rotations.
1044 auto PatternRot90Rot180 = m_Intrinsic<PartialReduceInt>(
1045 Ops: m_Intrinsic<PartialReduceInt>(Ops: m_Value(V&: Phi),
1046 Ops: m_Mul(L: m_Value(V&: BReal), R: m_Value(V&: A0))),
1047 Ops: m_Mul(L: m_Value(V&: BImag), R: m_Value(V&: A1)));
1048
1049 if (!match(V: Inst, P: PatternRot90Rot180))
1050 return nullptr;
1051
1052 A0 = UnwrapCast(A0);
1053 A1 = UnwrapCast(A1);
1054
1055 // Test if A0 is real/A1 is imag
1056 ANode = identifyNode(R: A0, I: A1);
1057 if (!ANode) {
1058 // Test if A0 is imag/A1 is real
1059 ANode = identifyNode(R: A1, I: A0);
1060 // Unable to identify operand components, thus unable to identify rotation
1061 if (!ANode)
1062 return nullptr;
1063 CN->Rotation = ComplexDeinterleavingRotation::Rotation_90;
1064 AReal = A1;
1065 AImag = A0;
1066 } else {
1067 AReal = A0;
1068 AImag = A1;
1069 CN->Rotation = ComplexDeinterleavingRotation::Rotation_180;
1070 }
1071 }
1072
1073 AReal = UnwrapCast(AReal);
1074 AImag = UnwrapCast(AImag);
1075 BReal = UnwrapCast(BReal);
1076 BImag = UnwrapCast(BImag);
1077
1078 VectorType *VTy = cast<VectorType>(Val: V->getType());
1079 Type *ExpectedOperandTy = VectorType::getSubdividedVectorType(VTy, NumSubdivs: 2);
1080 if (AReal->getType() != ExpectedOperandTy)
1081 return nullptr;
1082 if (AImag->getType() != ExpectedOperandTy)
1083 return nullptr;
1084 if (BReal->getType() != ExpectedOperandTy)
1085 return nullptr;
1086 if (BImag->getType() != ExpectedOperandTy)
1087 return nullptr;
1088
1089 if (Phi->getType() != VTy && RealUser->getType() != VTy)
1090 return nullptr;
1091
1092 CompositeNode *Node = identifyNode(R: AReal, I: AImag);
1093
1094 // In the case that a node was identified to figure out the rotation, ensure
1095 // that trying to identify a node with AReal and AImag post-unwrap results in
1096 // the same node
1097 if (ANode && Node != ANode) {
1098 LLVM_DEBUG(
1099 dbgs()
1100 << "Identified node is different from previously identified node. "
1101 "Unable to confidently generate a complex operation node\n");
1102 return nullptr;
1103 }
1104
1105 CN->addOperand(Node);
1106 CN->addOperand(Node: identifyNode(R: BReal, I: BImag));
1107 CN->addOperand(Node: identifyNode(R: Phi, I: RealUser));
1108
1109 return submitCompositeNode(Node: CN);
1110}
1111
1112ComplexDeinterleavingGraph::CompositeNode *
1113ComplexDeinterleavingGraph::identifyPartialReduction(Value *R, Value *I) {
1114 // Partial reductions don't support non-vector types, so check these first
1115 if (!isa<VectorType>(Val: R->getType()) || !isa<VectorType>(Val: I->getType()))
1116 return nullptr;
1117
1118 if (!R->hasUseList() || !I->hasUseList())
1119 return nullptr;
1120
1121 auto CommonUser =
1122 findCommonBetweenCollections<Value *>(A: R->users(), B: I->users());
1123 if (!CommonUser)
1124 return nullptr;
1125
1126 auto *IInst = dyn_cast<IntrinsicInst>(Val: *CommonUser);
1127 if (!IInst || IInst->getIntrinsicID() != Intrinsic::vector_partial_reduce_add)
1128 return nullptr;
1129
1130 if (CompositeNode *CN = identifyDotProduct(V: IInst))
1131 return CN;
1132
1133 return nullptr;
1134}
1135
1136ComplexDeinterleavingGraph::CompositeNode *
1137ComplexDeinterleavingGraph::identifyNode(ComplexValues &Vals) {
1138 auto It = CachedResult.find(Val: Vals);
1139 if (It != CachedResult.end()) {
1140 LLVM_DEBUG(dbgs() << " - Folding to existing node\n");
1141 return It->second;
1142 }
1143
1144 if (Vals.size() == 1) {
1145 assert(Factor == 2 && "Can only handle interleave factors of 2");
1146 Value *R = Vals[0].Real;
1147 Value *I = Vals[0].Imag;
1148 if (CompositeNode *CN = identifyPartialReduction(R, I))
1149 return CN;
1150 bool IsReduction = RealPHI == R && (!ImagPHI || ImagPHI == I);
1151 if (!IsReduction && R->getType() != I->getType())
1152 return nullptr;
1153 }
1154
1155 if (CompositeNode *CN = identifySplat(Vals))
1156 return CN;
1157
1158 for (auto &V : Vals) {
1159 auto *Real = dyn_cast<Instruction>(Val: V.Real);
1160 auto *Imag = dyn_cast<Instruction>(Val: V.Imag);
1161 if (!Real || !Imag)
1162 return nullptr;
1163 }
1164
1165 if (CompositeNode *CN = identifyDeinterleave(Vals))
1166 return CN;
1167
1168 if (Vals.size() == 1) {
1169 assert(Factor == 2 && "Can only handle interleave factors of 2");
1170 auto *Real = dyn_cast<Instruction>(Val: Vals[0].Real);
1171 auto *Imag = dyn_cast<Instruction>(Val: Vals[0].Imag);
1172 if (CompositeNode *CN = identifyPHINode(Real, Imag))
1173 return CN;
1174
1175 if (CompositeNode *CN = identifySelectNode(Real, Imag))
1176 return CN;
1177
1178 auto *VTy = cast<VectorType>(Val: Real->getType());
1179 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
1180
1181 bool HasCMulSupport = TL->isComplexDeinterleavingOperationSupported(
1182 Operation: ComplexDeinterleavingOperation::CMulPartial, Ty: NewVTy);
1183 bool HasCAddSupport = TL->isComplexDeinterleavingOperationSupported(
1184 Operation: ComplexDeinterleavingOperation::CAdd, Ty: NewVTy);
1185
1186 if (HasCMulSupport && isInstructionPairMul(A: Real, B: Imag)) {
1187 if (CompositeNode *CN = identifyPartialMul(Real, Imag))
1188 return CN;
1189 }
1190
1191 if (HasCAddSupport && isInstructionPairAdd(A: Real, B: Imag)) {
1192 if (CompositeNode *CN = identifyAdd(Real, Imag))
1193 return CN;
1194 }
1195
1196 if (HasCMulSupport && HasCAddSupport) {
1197 if (CompositeNode *CN = identifyReassocNodes(I: Real, J: Imag)) {
1198 return CN;
1199 }
1200 }
1201 }
1202
1203 if (CompositeNode *CN = identifySymmetricOperation(Vals))
1204 return CN;
1205
1206 LLVM_DEBUG(dbgs() << " - Not recognised as a valid pattern.\n");
1207 CachedResult[Vals] = nullptr;
1208 return nullptr;
1209}
1210
1211ComplexDeinterleavingGraph::CompositeNode *
1212ComplexDeinterleavingGraph::identifyReassocNodes(Instruction *Real,
1213 Instruction *Imag) {
1214 auto IsOperationSupported = [](Instruction *I) -> bool {
1215 unsigned Opcode = I->getOpcode();
1216 return match(V: I, P: m_AnyIntrinsic<Intrinsic::fma, Intrinsic::fmuladd>()) ||
1217 Opcode == Instruction::FAdd || Opcode == Instruction::FSub ||
1218 Opcode == Instruction::FNeg || Opcode == Instruction::Add ||
1219 Opcode == Instruction::Sub;
1220 };
1221
1222 if (!IsOperationSupported(Real) || !IsOperationSupported(Imag))
1223 return nullptr;
1224
1225 std::optional<FastMathFlags> Flags;
1226 if (isa<FPMathOperator>(Val: Real)) {
1227 if (Real->getFastMathFlags() != Imag->getFastMathFlags()) {
1228 LLVM_DEBUG(dbgs() << "The flags in Real and Imaginary instructions are "
1229 "not identical\n");
1230 return nullptr;
1231 }
1232
1233 Flags = Real->getFastMathFlags();
1234 if (!Flags->allowReassoc()) {
1235 LLVM_DEBUG(
1236 dbgs()
1237 << "the 'Reassoc' attribute is missing in the FastMath flags\n");
1238 return nullptr;
1239 }
1240 }
1241
1242 // Collect multiplications and addend instructions from the given instruction
1243 // while traversing it operands. Additionally, verify that all instructions
1244 // have the same fast math flags.
1245 auto Collect = [&Flags](Instruction *Insn, SmallVectorImpl<Product> &Muls,
1246 AddendList &Addends) -> bool {
1247 SmallVector<PointerIntPair<Value *, 1, bool>> Worklist = {{Insn, true}};
1248 while (!Worklist.empty()) {
1249 auto [V, IsPositive] = Worklist.pop_back_val();
1250
1251 Instruction *I = dyn_cast<Instruction>(Val: V);
1252 if (!I) {
1253 Addends.emplace_back(Vs&: V, Vs&: IsPositive);
1254 continue;
1255 }
1256
1257 // If an instruction has more than one user, it indicates that it either
1258 // has an external user, which will be later checked by the checkNodes
1259 // function, or it is a subexpression utilized by multiple expressions. In
1260 // the latter case, we will attempt to separately identify the complex
1261 // operation from here in order to create a shared
1262 // ComplexDeinterleavingCompositeNode.
1263 if (I != Insn && I->hasNUsesOrMore(N: 2)) {
1264 LLVM_DEBUG(dbgs() << "Found potential sub-expression: " << *I << "\n");
1265 Addends.emplace_back(Vs&: I, Vs&: IsPositive);
1266 continue;
1267 }
1268 switch (I->getOpcode()) {
1269 case Instruction::FAdd:
1270 case Instruction::Add:
1271 Worklist.emplace_back(Args: I->getOperand(i: 1), Args&: IsPositive);
1272 Worklist.emplace_back(Args: I->getOperand(i: 0), Args&: IsPositive);
1273 break;
1274 case Instruction::FSub:
1275 Worklist.emplace_back(Args: I->getOperand(i: 1), Args: !IsPositive);
1276 Worklist.emplace_back(Args: I->getOperand(i: 0), Args&: IsPositive);
1277 break;
1278 case Instruction::Sub:
1279 if (isNeg(V: I)) {
1280 Worklist.emplace_back(Args: getNegOperand(V: I), Args: !IsPositive);
1281 } else {
1282 Worklist.emplace_back(Args: I->getOperand(i: 1), Args: !IsPositive);
1283 Worklist.emplace_back(Args: I->getOperand(i: 0), Args&: IsPositive);
1284 }
1285 break;
1286 case Instruction::FMul:
1287 case Instruction::Mul: {
1288 Value *A, *B;
1289 if (isNeg(V: I->getOperand(i: 0))) {
1290 A = getNegOperand(V: I->getOperand(i: 0));
1291 IsPositive = !IsPositive;
1292 } else {
1293 A = I->getOperand(i: 0);
1294 }
1295
1296 if (isNeg(V: I->getOperand(i: 1))) {
1297 B = getNegOperand(V: I->getOperand(i: 1));
1298 IsPositive = !IsPositive;
1299 } else {
1300 B = I->getOperand(i: 1);
1301 }
1302 Muls.push_back(Elt: Product{.Multiplier: A, .Multiplicand: B, .IsPositive: IsPositive});
1303 break;
1304 }
1305 case Instruction::FNeg:
1306 Worklist.emplace_back(Args: I->getOperand(i: 0), Args: !IsPositive);
1307 break;
1308 case Instruction::Call: {
1309 Value *A, *B, *C;
1310 if (!match(V: I, P: m_Intrinsic<Intrinsic::fma>(Ops: m_Value(V&: A), Ops: m_Value(V&: B),
1311 Ops: m_Value(V&: C))) &&
1312 !match(V: I, P: m_Intrinsic<Intrinsic::fmuladd>(Ops: m_Value(V&: A), Ops: m_Value(V&: B),
1313 Ops: m_Value(V&: C)))) {
1314 Addends.emplace_back(Vs&: I, Vs&: IsPositive);
1315 continue;
1316 }
1317
1318 if (isNeg(V: A)) {
1319 A = getNegOperand(V: A);
1320 IsPositive = !IsPositive;
1321 }
1322
1323 if (isNeg(V: B)) {
1324 B = getNegOperand(V: B);
1325 IsPositive = !IsPositive;
1326 }
1327
1328 Muls.push_back(Elt: Product{.Multiplier: A, .Multiplicand: B, .IsPositive: IsPositive});
1329 Worklist.emplace_back(Args&: C, Args&: IsPositive);
1330 break;
1331 }
1332 default:
1333 Addends.emplace_back(Vs&: I, Vs&: IsPositive);
1334 continue;
1335 }
1336
1337 if (Flags && I->getFastMathFlags() != *Flags) {
1338 LLVM_DEBUG(dbgs() << "The instruction's fast math flags are "
1339 "inconsistent with the root instructions' flags: "
1340 << *I << "\n");
1341 return false;
1342 }
1343 }
1344 return true;
1345 };
1346
1347 SmallVector<Product> RealMuls, ImagMuls;
1348 AddendList RealAddends, ImagAddends;
1349 if (!Collect(Real, RealMuls, RealAddends) ||
1350 !Collect(Imag, ImagMuls, ImagAddends))
1351 return nullptr;
1352
1353 if (RealAddends.size() != ImagAddends.size())
1354 return nullptr;
1355
1356 CompositeNode *FinalNode = nullptr;
1357 if (!RealMuls.empty() || !ImagMuls.empty()) {
1358 // If there are multiplicands, extract positive addend and use it as an
1359 // accumulator
1360 FinalNode = extractPositiveAddend(RealAddends, ImagAddends);
1361 FinalNode = identifyMultiplications(RealMuls, ImagMuls, Accumulator: FinalNode);
1362 if (!FinalNode)
1363 return nullptr;
1364 }
1365
1366 // Identify and process remaining additions
1367 if (!RealAddends.empty() || !ImagAddends.empty()) {
1368 FinalNode = identifyAdditions(RealAddends, ImagAddends, Flags, Accumulator: FinalNode);
1369 if (!FinalNode)
1370 return nullptr;
1371 }
1372 assert(FinalNode && "FinalNode can not be nullptr here");
1373 assert(FinalNode->Vals.size() == 1);
1374 // Set the Real and Imag fields of the final node and submit it
1375 FinalNode->Vals[0].Real = Real;
1376 FinalNode->Vals[0].Imag = Imag;
1377 submitCompositeNode(Node: FinalNode);
1378 return FinalNode;
1379}
1380
1381bool ComplexDeinterleavingGraph::collectPartialMuls(
1382 ArrayRef<Product> RealMuls, ArrayRef<Product> ImagMuls,
1383 SmallVectorImpl<PartialMulCandidate> &PartialMulCandidates) {
1384 // Helper function to extract a common operand from two products
1385 auto FindCommonInstruction = [](const Product &Real,
1386 const Product &Imag) -> Value * {
1387 if (Real.Multiplicand == Imag.Multiplicand ||
1388 Real.Multiplicand == Imag.Multiplier)
1389 return Real.Multiplicand;
1390
1391 if (Real.Multiplier == Imag.Multiplicand ||
1392 Real.Multiplier == Imag.Multiplier)
1393 return Real.Multiplier;
1394
1395 return nullptr;
1396 };
1397
1398 // Iterating over real and imaginary multiplications to find common operands
1399 // If a common operand is found, a partial multiplication candidate is created
1400 // and added to the candidates vector The function returns false if no common
1401 // operands are found for any product
1402 for (unsigned i = 0; i < RealMuls.size(); ++i) {
1403 bool FoundCommon = false;
1404 for (unsigned j = 0; j < ImagMuls.size(); ++j) {
1405 auto *Common = FindCommonInstruction(RealMuls[i], ImagMuls[j]);
1406 if (!Common)
1407 continue;
1408
1409 auto *A = RealMuls[i].Multiplicand == Common ? RealMuls[i].Multiplier
1410 : RealMuls[i].Multiplicand;
1411 auto *B = ImagMuls[j].Multiplicand == Common ? ImagMuls[j].Multiplier
1412 : ImagMuls[j].Multiplicand;
1413
1414 auto Node = identifyNode(R: A, I: B);
1415 if (Node) {
1416 FoundCommon = true;
1417 PartialMulCandidates.push_back(Elt: {.Common: Common, .Node: Node, .RealIdx: i, .ImagIdx: j, .IsNodeInverted: false});
1418 }
1419
1420 Node = identifyNode(R: B, I: A);
1421 if (Node) {
1422 FoundCommon = true;
1423 PartialMulCandidates.push_back(Elt: {.Common: Common, .Node: Node, .RealIdx: i, .ImagIdx: j, .IsNodeInverted: true});
1424 }
1425 }
1426 if (!FoundCommon)
1427 return false;
1428 }
1429 return true;
1430}
1431
1432ComplexDeinterleavingGraph::CompositeNode *
1433ComplexDeinterleavingGraph::identifyMultiplications(
1434 SmallVectorImpl<Product> &RealMuls, SmallVectorImpl<Product> &ImagMuls,
1435 CompositeNode *Accumulator = nullptr) {
1436 if (RealMuls.size() != ImagMuls.size())
1437 return nullptr;
1438
1439 SmallVector<PartialMulCandidate> Info;
1440 if (!collectPartialMuls(RealMuls, ImagMuls, PartialMulCandidates&: Info))
1441 return nullptr;
1442
1443 // Map to store common instruction to node pointers
1444 DenseMap<Value *, CompositeNode *> CommonToNode;
1445 SmallVector<bool> Processed(Info.size(), false);
1446 for (unsigned I = 0; I < Info.size(); ++I) {
1447 if (Processed[I])
1448 continue;
1449
1450 PartialMulCandidate &InfoA = Info[I];
1451 for (unsigned J = I + 1; J < Info.size(); ++J) {
1452 if (Processed[J])
1453 continue;
1454
1455 PartialMulCandidate &InfoB = Info[J];
1456 auto *InfoReal = &InfoA;
1457 auto *InfoImag = &InfoB;
1458
1459 auto NodeFromCommon = identifyNode(R: InfoReal->Common, I: InfoImag->Common);
1460 if (!NodeFromCommon) {
1461 std::swap(a&: InfoReal, b&: InfoImag);
1462 NodeFromCommon = identifyNode(R: InfoReal->Common, I: InfoImag->Common);
1463 }
1464 if (!NodeFromCommon)
1465 continue;
1466
1467 CommonToNode[InfoReal->Common] = NodeFromCommon;
1468 CommonToNode[InfoImag->Common] = NodeFromCommon;
1469 Processed[I] = true;
1470 Processed[J] = true;
1471 }
1472 }
1473
1474 SmallVector<bool> ProcessedReal(RealMuls.size(), false);
1475 SmallVector<bool> ProcessedImag(ImagMuls.size(), false);
1476 CompositeNode *Result = Accumulator;
1477 for (auto &PMI : Info) {
1478 if (ProcessedReal[PMI.RealIdx] || ProcessedImag[PMI.ImagIdx])
1479 continue;
1480
1481 auto It = CommonToNode.find(Val: PMI.Common);
1482 // TODO: Process independent complex multiplications. Cases like this:
1483 // A.real() * B where both A and B are complex numbers.
1484 if (It == CommonToNode.end()) {
1485 LLVM_DEBUG({
1486 dbgs() << "Unprocessed independent partial multiplication:\n";
1487 for (auto *Mul : {&RealMuls[PMI.RealIdx], &RealMuls[PMI.RealIdx]})
1488 dbgs().indent(4) << (Mul->IsPositive ? "+" : "-") << *Mul->Multiplier
1489 << " multiplied by " << *Mul->Multiplicand << "\n";
1490 });
1491 return nullptr;
1492 }
1493
1494 auto &RealMul = RealMuls[PMI.RealIdx];
1495 auto &ImagMul = ImagMuls[PMI.ImagIdx];
1496
1497 auto NodeA = It->second;
1498 auto NodeB = PMI.Node;
1499 auto IsMultiplicandReal = PMI.Common == NodeA->Vals[0].Real;
1500 // The following table illustrates the relationship between multiplications
1501 // and rotations. If we consider the multiplication (X + iY) * (U + iV), we
1502 // can see:
1503 //
1504 // Rotation | Real | Imag |
1505 // ---------+--------+--------+
1506 // 0 | x * u | x * v |
1507 // 90 | -y * v | y * u |
1508 // 180 | -x * u | -x * v |
1509 // 270 | y * v | -y * u |
1510 //
1511 // Check if the candidate can indeed be represented by partial
1512 // multiplication
1513 // TODO: Add support for multiplication by complex one
1514 if ((IsMultiplicandReal && PMI.IsNodeInverted) ||
1515 (!IsMultiplicandReal && !PMI.IsNodeInverted))
1516 continue;
1517
1518 // Determine the rotation based on the multiplications
1519 ComplexDeinterleavingRotation Rotation;
1520 if (IsMultiplicandReal) {
1521 // Detect 0 and 180 degrees rotation
1522 if (RealMul.IsPositive && ImagMul.IsPositive)
1523 Rotation = llvm::ComplexDeinterleavingRotation::Rotation_0;
1524 else if (!RealMul.IsPositive && !ImagMul.IsPositive)
1525 Rotation = llvm::ComplexDeinterleavingRotation::Rotation_180;
1526 else
1527 continue;
1528
1529 } else {
1530 // Detect 90 and 270 degrees rotation
1531 if (!RealMul.IsPositive && ImagMul.IsPositive)
1532 Rotation = llvm::ComplexDeinterleavingRotation::Rotation_90;
1533 else if (RealMul.IsPositive && !ImagMul.IsPositive)
1534 Rotation = llvm::ComplexDeinterleavingRotation::Rotation_270;
1535 else
1536 continue;
1537 }
1538
1539 LLVM_DEBUG({
1540 dbgs() << "Identified partial multiplication (X, Y) * (U, V):\n";
1541 dbgs().indent(4) << "X: " << *NodeA->Vals[0].Real << "\n";
1542 dbgs().indent(4) << "Y: " << *NodeA->Vals[0].Imag << "\n";
1543 dbgs().indent(4) << "U: " << *NodeB->Vals[0].Real << "\n";
1544 dbgs().indent(4) << "V: " << *NodeB->Vals[0].Imag << "\n";
1545 dbgs().indent(4) << "Rotation - " << (int)Rotation * 90 << "\n";
1546 });
1547
1548 CompositeNode *NodeMul = prepareCompositeNode(
1549 Operation: ComplexDeinterleavingOperation::CMulPartial, R: nullptr, I: nullptr);
1550 NodeMul->Rotation = Rotation;
1551 NodeMul->addOperand(Node: NodeA);
1552 NodeMul->addOperand(Node: NodeB);
1553 if (Result)
1554 NodeMul->addOperand(Node: Result);
1555 submitCompositeNode(Node: NodeMul);
1556 Result = NodeMul;
1557 ProcessedReal[PMI.RealIdx] = true;
1558 ProcessedImag[PMI.ImagIdx] = true;
1559 }
1560
1561 // Ensure all products have been processed, if not return nullptr.
1562 if (!all_of(Range&: ProcessedReal, P: [](bool V) { return V; }) ||
1563 !all_of(Range&: ProcessedImag, P: [](bool V) { return V; })) {
1564
1565 // Dump debug information about which partial multiplications are not
1566 // processed.
1567 LLVM_DEBUG({
1568 dbgs() << "Unprocessed products (Real):\n";
1569 for (size_t i = 0; i < ProcessedReal.size(); ++i) {
1570 if (!ProcessedReal[i])
1571 dbgs().indent(4) << (RealMuls[i].IsPositive ? "+" : "-")
1572 << *RealMuls[i].Multiplier << " multiplied by "
1573 << *RealMuls[i].Multiplicand << "\n";
1574 }
1575 dbgs() << "Unprocessed products (Imag):\n";
1576 for (size_t i = 0; i < ProcessedImag.size(); ++i) {
1577 if (!ProcessedImag[i])
1578 dbgs().indent(4) << (ImagMuls[i].IsPositive ? "+" : "-")
1579 << *ImagMuls[i].Multiplier << " multiplied by "
1580 << *ImagMuls[i].Multiplicand << "\n";
1581 }
1582 });
1583 return nullptr;
1584 }
1585
1586 return Result;
1587}
1588
1589ComplexDeinterleavingGraph::CompositeNode *
1590ComplexDeinterleavingGraph::identifyAdditions(
1591 AddendList &RealAddends, AddendList &ImagAddends,
1592 std::optional<FastMathFlags> Flags, CompositeNode *Accumulator = nullptr) {
1593 if (RealAddends.size() != ImagAddends.size())
1594 return nullptr;
1595
1596 CompositeNode *Result = nullptr;
1597 // If we have accumulator use it as first addend
1598 if (Accumulator)
1599 Result = Accumulator;
1600 // Otherwise find an element with both positive real and imaginary parts.
1601 else
1602 Result = extractPositiveAddend(RealAddends, ImagAddends);
1603
1604 if (!Result)
1605 return nullptr;
1606
1607 while (!RealAddends.empty()) {
1608 auto ItR = RealAddends.begin();
1609 auto [R, IsPositiveR] = *ItR;
1610
1611 bool FoundImag = false;
1612 for (auto ItI = ImagAddends.begin(); ItI != ImagAddends.end(); ++ItI) {
1613 auto [I, IsPositiveI] = *ItI;
1614 ComplexDeinterleavingRotation Rotation;
1615 if (IsPositiveR && IsPositiveI)
1616 Rotation = ComplexDeinterleavingRotation::Rotation_0;
1617 else if (!IsPositiveR && IsPositiveI)
1618 Rotation = ComplexDeinterleavingRotation::Rotation_90;
1619 else if (!IsPositiveR && !IsPositiveI)
1620 Rotation = ComplexDeinterleavingRotation::Rotation_180;
1621 else
1622 Rotation = ComplexDeinterleavingRotation::Rotation_270;
1623
1624 CompositeNode *AddNode = nullptr;
1625 if (Rotation == ComplexDeinterleavingRotation::Rotation_0 ||
1626 Rotation == ComplexDeinterleavingRotation::Rotation_180) {
1627 AddNode = identifyNode(R, I);
1628 } else {
1629 AddNode = identifyNode(R: I, I: R);
1630 }
1631 if (AddNode) {
1632 LLVM_DEBUG({
1633 dbgs() << "Identified addition:\n";
1634 dbgs().indent(4) << "X: " << *R << "\n";
1635 dbgs().indent(4) << "Y: " << *I << "\n";
1636 dbgs().indent(4) << "Rotation - " << (int)Rotation * 90 << "\n";
1637 });
1638
1639 CompositeNode *TmpNode = nullptr;
1640 if (Rotation == llvm::ComplexDeinterleavingRotation::Rotation_0) {
1641 TmpNode = prepareCompositeNode(
1642 Operation: ComplexDeinterleavingOperation::Symmetric, R: nullptr, I: nullptr);
1643 if (Flags) {
1644 TmpNode->Opcode = Instruction::FAdd;
1645 TmpNode->Flags = *Flags;
1646 } else {
1647 TmpNode->Opcode = Instruction::Add;
1648 }
1649 } else if (Rotation ==
1650 llvm::ComplexDeinterleavingRotation::Rotation_180) {
1651 TmpNode = prepareCompositeNode(
1652 Operation: ComplexDeinterleavingOperation::Symmetric, R: nullptr, I: nullptr);
1653 if (Flags) {
1654 TmpNode->Opcode = Instruction::FSub;
1655 TmpNode->Flags = *Flags;
1656 } else {
1657 TmpNode->Opcode = Instruction::Sub;
1658 }
1659 } else {
1660 TmpNode = prepareCompositeNode(Operation: ComplexDeinterleavingOperation::CAdd,
1661 R: nullptr, I: nullptr);
1662 TmpNode->Rotation = Rotation;
1663 }
1664
1665 TmpNode->addOperand(Node: Result);
1666 TmpNode->addOperand(Node: AddNode);
1667 submitCompositeNode(Node: TmpNode);
1668 Result = TmpNode;
1669 RealAddends.erase(I: ItR);
1670 ImagAddends.erase(I: ItI);
1671 FoundImag = true;
1672 break;
1673 }
1674 }
1675 if (!FoundImag)
1676 return nullptr;
1677 }
1678 return Result;
1679}
1680
1681ComplexDeinterleavingGraph::CompositeNode *
1682ComplexDeinterleavingGraph::extractPositiveAddend(AddendList &RealAddends,
1683 AddendList &ImagAddends) {
1684 for (auto ItR = RealAddends.begin(); ItR != RealAddends.end(); ++ItR) {
1685 for (auto ItI = ImagAddends.begin(); ItI != ImagAddends.end(); ++ItI) {
1686 auto [R, IsPositiveR] = *ItR;
1687 auto [I, IsPositiveI] = *ItI;
1688 if (IsPositiveR && IsPositiveI) {
1689 auto Result = identifyNode(R, I);
1690 if (Result) {
1691 RealAddends.erase(I: ItR);
1692 ImagAddends.erase(I: ItI);
1693 return Result;
1694 }
1695 }
1696 }
1697 }
1698 return nullptr;
1699}
1700
1701bool ComplexDeinterleavingGraph::identifyNodes(Instruction *RootI) {
1702 // This potential root instruction might already have been recognized as
1703 // reduction. Because RootToNode maps both Real and Imaginary parts to
1704 // CompositeNode we should choose only one either Real or Imag instruction to
1705 // use as an anchor for generating complex instruction.
1706 auto It = RootToNode.find(Val: RootI);
1707 if (It != RootToNode.end()) {
1708 auto RootNode = It->second;
1709 assert(RootNode->Operation ==
1710 ComplexDeinterleavingOperation::ReductionOperation ||
1711 RootNode->Operation ==
1712 ComplexDeinterleavingOperation::ReductionSingle);
1713 assert(RootNode->Vals.size() == 1 &&
1714 "Cannot handle reductions involving multiple complex values");
1715 // Find out which part, Real or Imag, comes later, and only if we come to
1716 // the latest part, add it to OrderedRoots.
1717 auto *R = cast<Instruction>(Val: RootNode->Vals[0].Real);
1718 auto *I = RootNode->Vals[0].Imag ? cast<Instruction>(Val: RootNode->Vals[0].Imag)
1719 : nullptr;
1720
1721 Instruction *ReplacementAnchor;
1722 if (I)
1723 ReplacementAnchor = R->comesBefore(Other: I) ? I : R;
1724 else
1725 ReplacementAnchor = R;
1726
1727 if (ReplacementAnchor != RootI)
1728 return false;
1729 OrderedRoots.push_back(Elt: RootI);
1730 return true;
1731 }
1732
1733 auto RootNode = identifyRoot(I: RootI);
1734 if (!RootNode)
1735 return false;
1736
1737 LLVM_DEBUG({
1738 Function *F = RootI->getFunction();
1739 BasicBlock *B = RootI->getParent();
1740 dbgs() << "Complex deinterleaving graph for " << F->getName()
1741 << "::" << B->getName() << ".\n";
1742 dump(dbgs());
1743 dbgs() << "\n";
1744 });
1745 RootToNode[RootI] = RootNode;
1746 OrderedRoots.push_back(Elt: RootI);
1747 return true;
1748}
1749
1750bool ComplexDeinterleavingGraph::collectPotentialReductions(BasicBlock *B) {
1751 bool FoundPotentialReduction = false;
1752 if (Factor != 2)
1753 return false;
1754
1755 auto *Br = dyn_cast<CondBrInst>(Val: B->getTerminator());
1756 if (!Br)
1757 return false;
1758
1759 // Identify simple one-block loop
1760 if (Br->getSuccessor(i: 0) != B && Br->getSuccessor(i: 1) != B)
1761 return false;
1762
1763 for (auto &PHI : B->phis()) {
1764 if (PHI.getNumIncomingValues() != 2)
1765 continue;
1766
1767 if (!PHI.getType()->isVectorTy())
1768 continue;
1769
1770 auto *ReductionOp = dyn_cast<Instruction>(Val: PHI.getIncomingValueForBlock(BB: B));
1771 if (!ReductionOp)
1772 continue;
1773
1774 // Check if final instruction is reduced outside of current block
1775 Instruction *FinalReduction = nullptr;
1776 auto NumUsers = 0u;
1777 for (auto *U : ReductionOp->users()) {
1778 ++NumUsers;
1779 if (U == &PHI)
1780 continue;
1781 FinalReduction = dyn_cast<Instruction>(Val: U);
1782 }
1783
1784 if (NumUsers != 2 || !FinalReduction || FinalReduction->getParent() == B ||
1785 isa<PHINode>(Val: FinalReduction))
1786 continue;
1787
1788 ReductionInfo[ReductionOp] = {&PHI, FinalReduction};
1789 BackEdge = B;
1790 auto BackEdgeIdx = PHI.getBasicBlockIndex(BB: B);
1791 auto IncomingIdx = BackEdgeIdx == 0 ? 1 : 0;
1792 Incoming = PHI.getIncomingBlock(i: IncomingIdx);
1793 FoundPotentialReduction = true;
1794
1795 // If the initial value of PHINode is an Instruction, consider it a leaf
1796 // value of a complex deinterleaving graph.
1797 if (auto *InitPHI =
1798 dyn_cast<Instruction>(Val: PHI.getIncomingValueForBlock(BB: Incoming)))
1799 FinalInstructions.insert(Ptr: InitPHI);
1800 }
1801 return FoundPotentialReduction;
1802}
1803
1804void ComplexDeinterleavingGraph::identifyReductionNodes() {
1805 assert(Factor == 2 && "Cannot handle multiple complex values");
1806
1807 SmallVector<bool> Processed(ReductionInfo.size(), false);
1808 SmallVector<Instruction *> OperationInstruction;
1809 for (auto &P : ReductionInfo)
1810 OperationInstruction.push_back(Elt: P.first);
1811
1812 // Identify a complex computation by evaluating two reduction operations that
1813 // potentially could be involved
1814 for (size_t i = 0; i < OperationInstruction.size(); ++i) {
1815 if (Processed[i])
1816 continue;
1817 for (size_t j = i + 1; j < OperationInstruction.size(); ++j) {
1818 if (Processed[j])
1819 continue;
1820 auto *Real = OperationInstruction[i];
1821 auto *Imag = OperationInstruction[j];
1822 if (Real->getType() != Imag->getType())
1823 continue;
1824
1825 RealPHI = ReductionInfo[Real].first;
1826 ImagPHI = ReductionInfo[Imag].first;
1827 PHIsFound = false;
1828 auto Node = identifyNode(R: Real, I: Imag);
1829 if (!Node) {
1830 std::swap(a&: Real, b&: Imag);
1831 std::swap(a&: RealPHI, b&: ImagPHI);
1832 Node = identifyNode(R: Real, I: Imag);
1833 }
1834
1835 // If a node is identified and reduction PHINode is used in the chain of
1836 // operations, mark its operation instructions as used to prevent
1837 // re-identification and attach the node to the real part
1838 if (Node && PHIsFound) {
1839 LLVM_DEBUG(dbgs() << "Identified reduction starting from instructions: "
1840 << *Real << " / " << *Imag << "\n");
1841 Processed[i] = true;
1842 Processed[j] = true;
1843 auto RootNode = prepareCompositeNode(
1844 Operation: ComplexDeinterleavingOperation::ReductionOperation, R: Real, I: Imag);
1845 RootNode->addOperand(Node);
1846 RootToNode[Real] = RootNode;
1847 RootToNode[Imag] = RootNode;
1848 submitCompositeNode(Node: RootNode);
1849 break;
1850 }
1851 }
1852
1853 auto *Real = OperationInstruction[i];
1854 // We want to check that we have 2 operands, but the function attributes
1855 // being counted as operands bloats this value.
1856 if (Processed[i] || Real->getNumOperands() < 2)
1857 continue;
1858
1859 // Can only combined integer reductions at the moment.
1860 if (!ReductionInfo[Real].second->getType()->isIntegerTy())
1861 continue;
1862
1863 RealPHI = ReductionInfo[Real].first;
1864 ImagPHI = nullptr;
1865 PHIsFound = false;
1866 auto Node = identifyNode(R: Real->getOperand(i: 0), I: Real->getOperand(i: 1));
1867 if (Node && PHIsFound) {
1868 LLVM_DEBUG(
1869 dbgs() << "Identified single reduction starting from instruction: "
1870 << *Real << "/" << *ReductionInfo[Real].second << "\n");
1871
1872 // Reducing to a single vector is not supported, only permit reducing down
1873 // to scalar values.
1874 // Doing this here will leave the prior node in the graph,
1875 // however with no uses the node will be unreachable by the replacement
1876 // process. That along with the usage outside the graph should prevent the
1877 // replacement process from kicking off at all for this graph.
1878 // TODO Add support for reducing to a single vector value
1879 if (ReductionInfo[Real].second->getType()->isVectorTy())
1880 continue;
1881
1882 Processed[i] = true;
1883 auto RootNode = prepareCompositeNode(
1884 Operation: ComplexDeinterleavingOperation::ReductionSingle, R: Real, I: nullptr);
1885 RootNode->addOperand(Node);
1886 RootToNode[Real] = RootNode;
1887 submitCompositeNode(Node: RootNode);
1888 }
1889 }
1890
1891 RealPHI = nullptr;
1892 ImagPHI = nullptr;
1893}
1894
1895bool ComplexDeinterleavingGraph::checkNodes() {
1896 bool FoundDeinterleaveNode = false;
1897 for (CompositeNode *N : CompositeNodes) {
1898 if (!N->areOperandsValid())
1899 return false;
1900
1901 if (N->Operation == ComplexDeinterleavingOperation::Deinterleave)
1902 FoundDeinterleaveNode = true;
1903 }
1904
1905 // We need a deinterleave node in order to guarantee that we're working with
1906 // complex numbers.
1907 if (!FoundDeinterleaveNode) {
1908 LLVM_DEBUG(
1909 dbgs() << "Couldn't find a deinterleave node within the graph, cannot "
1910 "guarantee safety during graph transformation.\n");
1911 return false;
1912 }
1913
1914 // Collect all instructions from roots to leaves
1915 SmallPtrSet<Instruction *, 16> AllInstructions;
1916 SmallVector<Instruction *, 8> Worklist;
1917 for (auto &Pair : RootToNode)
1918 Worklist.push_back(Elt: Pair.first);
1919
1920 // Extract all instructions that are used by all XCMLA/XCADD/ADD/SUB/NEG
1921 // chains
1922 while (!Worklist.empty()) {
1923 auto *I = Worklist.pop_back_val();
1924
1925 if (!AllInstructions.insert(Ptr: I).second)
1926 continue;
1927
1928 for (Value *Op : I->operands()) {
1929 if (auto *OpI = dyn_cast<Instruction>(Val: Op)) {
1930 if (!FinalInstructions.count(Ptr: I))
1931 Worklist.emplace_back(Args&: OpI);
1932 }
1933 }
1934 }
1935
1936 // Find instructions that have users outside of chain
1937 for (auto *I : AllInstructions) {
1938 // Skip root nodes
1939 if (RootToNode.count(Val: I))
1940 continue;
1941
1942 for (User *U : I->users()) {
1943 if (AllInstructions.count(Ptr: cast<Instruction>(Val: U)))
1944 continue;
1945
1946 // Found an instruction that is not used by XCMLA/XCADD chain
1947 Worklist.emplace_back(Args&: I);
1948 break;
1949 }
1950 }
1951
1952 // If any instructions are found to be used outside, find and remove roots
1953 // that somehow connect to those instructions.
1954 SmallPtrSet<Instruction *, 16> Visited;
1955 while (!Worklist.empty()) {
1956 auto *I = Worklist.pop_back_val();
1957 if (!Visited.insert(Ptr: I).second)
1958 continue;
1959
1960 // Found an impacted root node. Removing it from the nodes to be
1961 // deinterleaved
1962 if (RootToNode.count(Val: I)) {
1963 LLVM_DEBUG(dbgs() << "Instruction " << *I
1964 << " could be deinterleaved but its chain of complex "
1965 "operations have an outside user\n");
1966 RootToNode.erase(Val: I);
1967 }
1968
1969 if (!AllInstructions.count(Ptr: I) || FinalInstructions.count(Ptr: I))
1970 continue;
1971
1972 for (User *U : I->users())
1973 Worklist.emplace_back(Args: cast<Instruction>(Val: U));
1974
1975 for (Value *Op : I->operands()) {
1976 if (auto *OpI = dyn_cast<Instruction>(Val: Op))
1977 Worklist.emplace_back(Args&: OpI);
1978 }
1979 }
1980 return !RootToNode.empty();
1981}
1982
1983ComplexDeinterleavingGraph::CompositeNode *
1984ComplexDeinterleavingGraph::identifyRoot(Instruction *RootI) {
1985 if (auto *Intrinsic = dyn_cast<IntrinsicInst>(Val: RootI)) {
1986 if (Intrinsic::getInterleaveIntrinsicID(Factor) !=
1987 Intrinsic->getIntrinsicID())
1988 return nullptr;
1989
1990 ComplexValues Vals;
1991 for (unsigned I = 0; I < Factor; I += 2) {
1992 auto *Real = dyn_cast<Instruction>(Val: Intrinsic->getOperand(i_nocapture: I));
1993 auto *Imag = dyn_cast<Instruction>(Val: Intrinsic->getOperand(i_nocapture: I + 1));
1994 if (!Real || !Imag)
1995 return nullptr;
1996 Vals.push_back(Elt: {.Real: Real, .Imag: Imag});
1997 }
1998
1999 ComplexDeinterleavingGraph::CompositeNode *Node1 = identifyNode(Vals);
2000 if (!Node1)
2001 return nullptr;
2002 return Node1;
2003 }
2004
2005 // TODO: We could also add support for fixed-width interleave factors of 4
2006 // and above, but currently for symmetric operations the interleaves and
2007 // deinterleaves are already removed by VectorCombine. If we extend this to
2008 // permit complex multiplications, reductions, etc. then we should also add
2009 // support for fixed-width here.
2010 if (Factor != 2)
2011 return nullptr;
2012
2013 auto *SVI = dyn_cast<ShuffleVectorInst>(Val: RootI);
2014 if (!SVI)
2015 return nullptr;
2016
2017 // Look for a shufflevector that takes separate vectors of the real and
2018 // imaginary components and recombines them into a single vector.
2019 if (!isInterleavingMask(Mask: SVI->getShuffleMask()))
2020 return nullptr;
2021
2022 Instruction *Real;
2023 Instruction *Imag;
2024 if (!match(V: RootI, P: m_Shuffle(v1: m_Instruction(I&: Real), v2: m_Instruction(I&: Imag))))
2025 return nullptr;
2026
2027 return identifyNode(R: Real, I: Imag);
2028}
2029
2030ComplexDeinterleavingGraph::CompositeNode *
2031ComplexDeinterleavingGraph::identifyDeinterleave(ComplexValues &Vals) {
2032 Instruction *II = nullptr;
2033
2034 // Must be at least one complex value.
2035 auto CheckExtract = [&](Value *V, unsigned ExpectedIdx,
2036 Instruction *ExpectedInsn) -> ExtractValueInst * {
2037 auto *EVI = dyn_cast<ExtractValueInst>(Val: V);
2038 if (!EVI || EVI->getNumIndices() != 1 ||
2039 EVI->getIndices()[0] != ExpectedIdx ||
2040 !isa<Instruction>(Val: EVI->getAggregateOperand()) ||
2041 (ExpectedInsn && ExpectedInsn != EVI->getAggregateOperand()))
2042 return nullptr;
2043 return EVI;
2044 };
2045
2046 for (unsigned Idx = 0; Idx < Vals.size(); Idx++) {
2047 ExtractValueInst *RealEVI = CheckExtract(Vals[Idx].Real, Idx * 2, II);
2048 if (RealEVI && Idx == 0)
2049 II = cast<Instruction>(Val: RealEVI->getAggregateOperand());
2050 if (!RealEVI || !CheckExtract(Vals[Idx].Imag, (Idx * 2) + 1, II)) {
2051 II = nullptr;
2052 break;
2053 }
2054 }
2055
2056 if (auto *IntrinsicII = dyn_cast_or_null<IntrinsicInst>(Val: II)) {
2057 if (IntrinsicII->getIntrinsicID() !=
2058 Intrinsic::getDeinterleaveIntrinsicID(Factor: 2 * Vals.size()))
2059 return nullptr;
2060
2061 // The remaining should match too.
2062 CompositeNode *PlaceholderNode = prepareCompositeNode(
2063 Operation: llvm::ComplexDeinterleavingOperation::Deinterleave, Vals);
2064 PlaceholderNode->ReplacementNode = II->getOperand(i: 0);
2065 for (auto &V : Vals) {
2066 FinalInstructions.insert(Ptr: cast<Instruction>(Val: V.Real));
2067 FinalInstructions.insert(Ptr: cast<Instruction>(Val: V.Imag));
2068 }
2069 return submitCompositeNode(Node: PlaceholderNode);
2070 }
2071
2072 if (Vals.size() != 1)
2073 return nullptr;
2074
2075 Value *Real = Vals[0].Real;
2076 Value *Imag = Vals[0].Imag;
2077 auto *RealShuffle = dyn_cast<ShuffleVectorInst>(Val: Real);
2078 auto *ImagShuffle = dyn_cast<ShuffleVectorInst>(Val: Imag);
2079 if (!RealShuffle || !ImagShuffle) {
2080 if (RealShuffle || ImagShuffle)
2081 LLVM_DEBUG(dbgs() << " - There's a shuffle where there shouldn't be.\n");
2082 return nullptr;
2083 }
2084
2085 Value *RealOp1 = RealShuffle->getOperand(i_nocapture: 1);
2086 if (!isa<UndefValue>(Val: RealOp1) && !match(V: RealOp1, P: m_Zero())) {
2087 LLVM_DEBUG(dbgs() << " - RealOp1 is not undef or zero.\n");
2088 return nullptr;
2089 }
2090 Value *ImagOp1 = ImagShuffle->getOperand(i_nocapture: 1);
2091 if (!isa<UndefValue>(Val: ImagOp1) && !match(V: ImagOp1, P: m_Zero())) {
2092 LLVM_DEBUG(dbgs() << " - ImagOp1 is not undef or zero.\n");
2093 return nullptr;
2094 }
2095
2096 Value *RealOp0 = RealShuffle->getOperand(i_nocapture: 0);
2097 Value *ImagOp0 = ImagShuffle->getOperand(i_nocapture: 0);
2098
2099 if (RealOp0 != ImagOp0) {
2100 LLVM_DEBUG(dbgs() << " - Shuffle operands are not equal.\n");
2101 return nullptr;
2102 }
2103
2104 ArrayRef<int> RealMask = RealShuffle->getShuffleMask();
2105 ArrayRef<int> ImagMask = ImagShuffle->getShuffleMask();
2106 if (!isDeinterleavingMask(Mask: RealMask) || !isDeinterleavingMask(Mask: ImagMask)) {
2107 LLVM_DEBUG(dbgs() << " - Masks are not deinterleaving.\n");
2108 return nullptr;
2109 }
2110
2111 if (RealMask[0] != 0 || ImagMask[0] != 1) {
2112 LLVM_DEBUG(dbgs() << " - Masks do not have the correct initial value.\n");
2113 return nullptr;
2114 }
2115
2116 // Type checking, the shuffle type should be a vector type of the same
2117 // scalar type, but half the size
2118 auto CheckType = [&](ShuffleVectorInst *Shuffle) {
2119 Value *Op = Shuffle->getOperand(i_nocapture: 0);
2120 auto *ShuffleTy = cast<FixedVectorType>(Val: Shuffle->getType());
2121 auto *OpTy = cast<FixedVectorType>(Val: Op->getType());
2122
2123 if (OpTy->getScalarType() != ShuffleTy->getScalarType())
2124 return false;
2125 if ((ShuffleTy->getNumElements() * 2) != OpTy->getNumElements())
2126 return false;
2127
2128 return true;
2129 };
2130
2131 auto CheckDeinterleavingShuffle = [&](ShuffleVectorInst *Shuffle) -> bool {
2132 if (!CheckType(Shuffle))
2133 return false;
2134
2135 ArrayRef<int> Mask = Shuffle->getShuffleMask();
2136 int Last = *Mask.rbegin();
2137
2138 Value *Op = Shuffle->getOperand(i_nocapture: 0);
2139 auto *OpTy = cast<FixedVectorType>(Val: Op->getType());
2140 int NumElements = OpTy->getNumElements();
2141
2142 // Ensure that the deinterleaving shuffle only pulls from the first
2143 // shuffle operand.
2144 return Last < NumElements;
2145 };
2146
2147 if (RealShuffle->getType() != ImagShuffle->getType()) {
2148 LLVM_DEBUG(dbgs() << " - Shuffle types aren't equal.\n");
2149 return nullptr;
2150 }
2151 if (!CheckDeinterleavingShuffle(RealShuffle)) {
2152 LLVM_DEBUG(dbgs() << " - RealShuffle is invalid type.\n");
2153 return nullptr;
2154 }
2155 if (!CheckDeinterleavingShuffle(ImagShuffle)) {
2156 LLVM_DEBUG(dbgs() << " - ImagShuffle is invalid type.\n");
2157 return nullptr;
2158 }
2159
2160 CompositeNode *PlaceholderNode =
2161 prepareCompositeNode(Operation: llvm::ComplexDeinterleavingOperation::Deinterleave,
2162 R: RealShuffle, I: ImagShuffle);
2163 PlaceholderNode->ReplacementNode = RealShuffle->getOperand(i_nocapture: 0);
2164 FinalInstructions.insert(Ptr: RealShuffle);
2165 FinalInstructions.insert(Ptr: ImagShuffle);
2166 return submitCompositeNode(Node: PlaceholderNode);
2167}
2168
2169ComplexDeinterleavingGraph::CompositeNode *
2170ComplexDeinterleavingGraph::identifySplat(ComplexValues &Vals) {
2171 auto IsSplat = [](Value *V) -> bool {
2172 // Fixed-width vector with constants
2173 if (isa<ConstantDataVector>(Val: V))
2174 return true;
2175
2176 if (isa<ConstantInt>(Val: V) || isa<ConstantFP>(Val: V))
2177 return isa<VectorType>(Val: V->getType());
2178
2179 VectorType *VTy;
2180 ArrayRef<int> Mask;
2181 // Splats are represented differently depending on whether the repeated
2182 // value is a constant or an Instruction
2183 if (auto *Const = dyn_cast<ConstantExpr>(Val: V)) {
2184 if (Const->getOpcode() != Instruction::ShuffleVector)
2185 return false;
2186 VTy = cast<VectorType>(Val: Const->getType());
2187 Mask = Const->getShuffleMask();
2188 } else if (auto *Shuf = dyn_cast<ShuffleVectorInst>(Val: V)) {
2189 VTy = Shuf->getType();
2190 Mask = Shuf->getShuffleMask();
2191 } else {
2192 return false;
2193 }
2194
2195 // When the data type is <1 x Type>, it's not possible to differentiate
2196 // between the ComplexDeinterleaving::Deinterleave and
2197 // ComplexDeinterleaving::Splat operations.
2198 if (!VTy->isScalableTy() && VTy->getElementCount().getKnownMinValue() == 1)
2199 return false;
2200
2201 return all_equal(Range&: Mask) && Mask[0] == 0;
2202 };
2203
2204 // The splats must meet the following requirements:
2205 // 1. Must either be all instructions or all values.
2206 // 2. Non-constant splats must live in the same block.
2207 if (auto *FirstValAsInstruction = dyn_cast<Instruction>(Val: Vals[0].Real)) {
2208 BasicBlock *FirstBB = FirstValAsInstruction->getParent();
2209 for (auto &V : Vals) {
2210 if (!IsSplat(V.Real) || !IsSplat(V.Imag))
2211 return nullptr;
2212
2213 auto *Real = dyn_cast<Instruction>(Val: V.Real);
2214 auto *Imag = dyn_cast<Instruction>(Val: V.Imag);
2215 if (!Real || !Imag || Real->getParent() != FirstBB ||
2216 Imag->getParent() != FirstBB)
2217 return nullptr;
2218 }
2219 } else {
2220 for (auto &V : Vals) {
2221 if (!IsSplat(V.Real) || !IsSplat(V.Imag) || isa<Instruction>(Val: V.Real) ||
2222 isa<Instruction>(Val: V.Imag))
2223 return nullptr;
2224 }
2225 }
2226
2227 for (auto &V : Vals) {
2228 auto *Real = dyn_cast<Instruction>(Val: V.Real);
2229 auto *Imag = dyn_cast<Instruction>(Val: V.Imag);
2230 if (Real && Imag) {
2231 FinalInstructions.insert(Ptr: Real);
2232 FinalInstructions.insert(Ptr: Imag);
2233 }
2234 }
2235 CompositeNode *PlaceholderNode =
2236 prepareCompositeNode(Operation: ComplexDeinterleavingOperation::Splat, Vals);
2237 return submitCompositeNode(Node: PlaceholderNode);
2238}
2239
2240ComplexDeinterleavingGraph::CompositeNode *
2241ComplexDeinterleavingGraph::identifyPHINode(Instruction *Real,
2242 Instruction *Imag) {
2243 if (Real != RealPHI || (ImagPHI && Imag != ImagPHI))
2244 return nullptr;
2245
2246 PHIsFound = true;
2247 CompositeNode *PlaceholderNode = prepareCompositeNode(
2248 Operation: ComplexDeinterleavingOperation::ReductionPHI, R: Real, I: Imag);
2249 return submitCompositeNode(Node: PlaceholderNode);
2250}
2251
2252ComplexDeinterleavingGraph::CompositeNode *
2253ComplexDeinterleavingGraph::identifySelectNode(Instruction *Real,
2254 Instruction *Imag) {
2255 auto *SelectReal = dyn_cast<SelectInst>(Val: Real);
2256 auto *SelectImag = dyn_cast<SelectInst>(Val: Imag);
2257 if (!SelectReal || !SelectImag)
2258 return nullptr;
2259
2260 Instruction *MaskA, *MaskB;
2261 Instruction *AR, *AI, *RA, *BI;
2262 if (!match(V: Real, P: m_Select(C: m_Instruction(I&: MaskA), L: m_Instruction(I&: AR),
2263 R: m_Instruction(I&: RA))) ||
2264 !match(V: Imag, P: m_Select(C: m_Instruction(I&: MaskB), L: m_Instruction(I&: AI),
2265 R: m_Instruction(I&: BI))))
2266 return nullptr;
2267
2268 if (MaskA != MaskB && !MaskA->isIdenticalTo(I: MaskB))
2269 return nullptr;
2270
2271 if (!MaskA->getType()->isVectorTy())
2272 return nullptr;
2273
2274 auto NodeA = identifyNode(R: AR, I: AI);
2275 if (!NodeA)
2276 return nullptr;
2277
2278 auto NodeB = identifyNode(R: RA, I: BI);
2279 if (!NodeB)
2280 return nullptr;
2281
2282 CompositeNode *PlaceholderNode = prepareCompositeNode(
2283 Operation: ComplexDeinterleavingOperation::ReductionSelect, R: Real, I: Imag);
2284 PlaceholderNode->addOperand(Node: NodeA);
2285 PlaceholderNode->addOperand(Node: NodeB);
2286 FinalInstructions.insert(Ptr: MaskA);
2287 FinalInstructions.insert(Ptr: MaskB);
2288 return submitCompositeNode(Node: PlaceholderNode);
2289}
2290
2291static Value *replaceSymmetricNode(IRBuilderBase &B, unsigned Opcode,
2292 std::optional<FastMathFlags> Flags,
2293 Value *InputA, Value *InputB) {
2294 Value *I;
2295 switch (Opcode) {
2296 case Instruction::FNeg:
2297 I = B.CreateFNeg(V: InputA);
2298 break;
2299 case Instruction::FAdd:
2300 I = B.CreateFAdd(L: InputA, R: InputB);
2301 break;
2302 case Instruction::Add:
2303 I = B.CreateAdd(LHS: InputA, RHS: InputB);
2304 break;
2305 case Instruction::FSub:
2306 I = B.CreateFSub(L: InputA, R: InputB);
2307 break;
2308 case Instruction::Sub:
2309 I = B.CreateSub(LHS: InputA, RHS: InputB);
2310 break;
2311 case Instruction::FMul:
2312 I = B.CreateFMul(L: InputA, R: InputB);
2313 break;
2314 case Instruction::Mul:
2315 I = B.CreateMul(LHS: InputA, RHS: InputB);
2316 break;
2317 default:
2318 llvm_unreachable("Incorrect symmetric opcode");
2319 }
2320 if (Flags)
2321 cast<Instruction>(Val: I)->setFastMathFlags(*Flags);
2322 return I;
2323}
2324
2325Value *ComplexDeinterleavingGraph::replaceNode(IRBuilderBase &Builder,
2326 CompositeNode *Node) {
2327 if (Node->ReplacementNode)
2328 return Node->ReplacementNode;
2329
2330 auto ReplaceOperandIfExist = [&](CompositeNode *Node,
2331 unsigned Idx) -> Value * {
2332 return Node->Operands.size() > Idx
2333 ? replaceNode(Builder, Node: Node->Operands[Idx])
2334 : nullptr;
2335 };
2336
2337 Value *ReplacementNode = nullptr;
2338 switch (Node->Operation) {
2339 case ComplexDeinterleavingOperation::CDot: {
2340 Value *Input0 = ReplaceOperandIfExist(Node, 0);
2341 Value *Input1 = ReplaceOperandIfExist(Node, 1);
2342 Value *Accumulator = ReplaceOperandIfExist(Node, 2);
2343 assert(!Input1 || (Input0->getType() == Input1->getType() &&
2344 "Node inputs need to be of the same type"));
2345 ReplacementNode = TL->createComplexDeinterleavingIR(
2346 B&: Builder, OperationType: Node->Operation, Rotation: Node->Rotation, InputA: Input0, InputB: Input1, Accumulator);
2347 break;
2348 }
2349 case ComplexDeinterleavingOperation::CAdd:
2350 case ComplexDeinterleavingOperation::CMulPartial:
2351 case ComplexDeinterleavingOperation::Symmetric: {
2352 Value *Input0 = ReplaceOperandIfExist(Node, 0);
2353 Value *Input1 = ReplaceOperandIfExist(Node, 1);
2354 Value *Accumulator = ReplaceOperandIfExist(Node, 2);
2355 assert(!Input1 || (Input0->getType() == Input1->getType() &&
2356 "Node inputs need to be of the same type"));
2357 assert(!Accumulator ||
2358 (Input0->getType() == Accumulator->getType() &&
2359 "Accumulator and input need to be of the same type"));
2360 if (Node->Operation == ComplexDeinterleavingOperation::Symmetric)
2361 ReplacementNode = replaceSymmetricNode(B&: Builder, Opcode: Node->Opcode, Flags: Node->Flags,
2362 InputA: Input0, InputB: Input1);
2363 else
2364 ReplacementNode = TL->createComplexDeinterleavingIR(
2365 B&: Builder, OperationType: Node->Operation, Rotation: Node->Rotation, InputA: Input0, InputB: Input1,
2366 Accumulator);
2367 break;
2368 }
2369 case ComplexDeinterleavingOperation::Deinterleave:
2370 llvm_unreachable("Deinterleave node should already have ReplacementNode");
2371 break;
2372 case ComplexDeinterleavingOperation::Splat: {
2373 SmallVector<Value *> Ops;
2374 for (auto &V : Node->Vals) {
2375 Ops.push_back(Elt: V.Real);
2376 Ops.push_back(Elt: V.Imag);
2377 }
2378 auto *R = dyn_cast<Instruction>(Val: Node->Vals[0].Real);
2379 auto *I = dyn_cast<Instruction>(Val: Node->Vals[0].Imag);
2380 if (R && I) {
2381 // Splats that are not constant are interleaved where they are located
2382 Instruction *InsertPoint = R;
2383 for (auto V : Node->Vals) {
2384 if (InsertPoint->comesBefore(Other: cast<Instruction>(Val: V.Real)))
2385 InsertPoint = cast<Instruction>(Val: V.Real);
2386 if (InsertPoint->comesBefore(Other: cast<Instruction>(Val: V.Imag)))
2387 InsertPoint = cast<Instruction>(Val: V.Imag);
2388 }
2389 InsertPoint = InsertPoint->getNextNode();
2390 IRBuilder<> IRB(InsertPoint);
2391 ReplacementNode = IRB.CreateVectorInterleave(Ops);
2392 } else {
2393 ReplacementNode = Builder.CreateVectorInterleave(Ops);
2394 }
2395 break;
2396 }
2397 case ComplexDeinterleavingOperation::ReductionPHI: {
2398 // If Operation is ReductionPHI, a new empty PHINode is created.
2399 // It is filled later when the ReductionOperation is processed.
2400 auto *OldPHI = cast<PHINode>(Val: Node->Vals[0].Real);
2401 auto *VTy = cast<VectorType>(Val: Node->Vals[0].Real->getType());
2402 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
2403 auto *NewPHI = PHINode::Create(Ty: NewVTy, NumReservedValues: 0, NameStr: "", InsertBefore: BackEdge->getFirstNonPHIIt());
2404 OldToNewPHI[OldPHI] = NewPHI;
2405 ReplacementNode = NewPHI;
2406 break;
2407 }
2408 case ComplexDeinterleavingOperation::ReductionSingle:
2409 ReplacementNode = replaceNode(Builder, Node: Node->Operands[0]);
2410 processReductionSingle(OperationReplacement: ReplacementNode, Node);
2411 break;
2412 case ComplexDeinterleavingOperation::ReductionOperation:
2413 ReplacementNode = replaceNode(Builder, Node: Node->Operands[0]);
2414 processReductionOperation(OperationReplacement: ReplacementNode, Node);
2415 break;
2416 case ComplexDeinterleavingOperation::ReductionSelect: {
2417 auto *MaskReal = cast<Instruction>(Val: Node->Vals[0].Real)->getOperand(i: 0);
2418 auto *MaskImag = cast<Instruction>(Val: Node->Vals[0].Imag)->getOperand(i: 0);
2419 auto *A = replaceNode(Builder, Node: Node->Operands[0]);
2420 auto *B = replaceNode(Builder, Node: Node->Operands[1]);
2421 auto *NewMask = Builder.CreateVectorInterleave(Ops: {MaskReal, MaskImag});
2422 ReplacementNode = Builder.CreateSelect(C: NewMask, True: A, False: B);
2423 break;
2424 }
2425 }
2426
2427 assert(ReplacementNode && "Target failed to create Intrinsic call.");
2428 NumComplexTransformations += 1;
2429 Node->ReplacementNode = ReplacementNode;
2430 return ReplacementNode;
2431}
2432
2433void ComplexDeinterleavingGraph::processReductionSingle(
2434 Value *OperationReplacement, CompositeNode *Node) {
2435 auto *Real = cast<Instruction>(Val: Node->Vals[0].Real);
2436 auto *OldPHI = ReductionInfo[Real].first;
2437 auto *NewPHI = OldToNewPHI[OldPHI];
2438 auto *VTy = cast<VectorType>(Val: Real->getType());
2439 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
2440
2441 Value *Init = OldPHI->getIncomingValueForBlock(BB: Incoming);
2442
2443 IRBuilder<> Builder(Incoming->getTerminator());
2444
2445 Value *NewInit = nullptr;
2446 if (auto *C = dyn_cast<Constant>(Val: Init)) {
2447 if (C->isNullValue())
2448 NewInit = Constant::getNullValue(Ty: NewVTy);
2449 }
2450
2451 if (!NewInit)
2452 NewInit =
2453 Builder.CreateVectorInterleave(Ops: {Init, Constant::getNullValue(Ty: VTy)});
2454
2455 NewPHI->addIncoming(V: NewInit, BB: Incoming);
2456 NewPHI->addIncoming(V: OperationReplacement, BB: BackEdge);
2457
2458 auto *FinalReduction = ReductionInfo[Real].second;
2459 Builder.SetInsertPoint(&*FinalReduction->getParent()->getFirstInsertionPt());
2460
2461 auto *AddReduce = Builder.CreateAddReduce(Src: OperationReplacement);
2462 FinalReduction->replaceAllUsesWith(V: AddReduce);
2463}
2464
2465void ComplexDeinterleavingGraph::processReductionOperation(
2466 Value *OperationReplacement, CompositeNode *Node) {
2467 auto *Real = cast<Instruction>(Val: Node->Vals[0].Real);
2468 auto *Imag = cast<Instruction>(Val: Node->Vals[0].Imag);
2469 auto *OldPHIReal = ReductionInfo[Real].first;
2470 auto *OldPHIImag = ReductionInfo[Imag].first;
2471 auto *NewPHI = OldToNewPHI[OldPHIReal];
2472
2473 // We have to interleave initial origin values coming from IncomingBlock
2474 Value *InitReal = OldPHIReal->getIncomingValueForBlock(BB: Incoming);
2475 Value *InitImag = OldPHIImag->getIncomingValueForBlock(BB: Incoming);
2476
2477 IRBuilder<> Builder(Incoming->getTerminator());
2478 auto *NewInit = Builder.CreateVectorInterleave(Ops: {InitReal, InitImag});
2479
2480 NewPHI->addIncoming(V: NewInit, BB: Incoming);
2481 NewPHI->addIncoming(V: OperationReplacement, BB: BackEdge);
2482
2483 // Deinterleave complex vector outside of loop so that it can be finally
2484 // reduced
2485 auto *FinalReductionReal = ReductionInfo[Real].second;
2486 auto *FinalReductionImag = ReductionInfo[Imag].second;
2487
2488 auto *Br = cast<CondBrInst>(Val: BackEdge->getTerminator());
2489 BasicBlock *ExitBB = Br->getSuccessor(i: Br->getSuccessor(i: 0) == BackEdge);
2490 Builder.SetInsertPoint(&*ExitBB->getFirstInsertionPt());
2491
2492 auto *Deinterleave = Builder.CreateIntrinsic(ID: Intrinsic::vector_deinterleave2,
2493 OverloadTypes: OperationReplacement->getType(),
2494 Args: OperationReplacement);
2495
2496 auto *NewReal = Builder.CreateExtractValue(Agg: Deinterleave, Idxs: (uint64_t)0);
2497 FinalReductionReal->replaceUsesOfWith(From: Real, To: NewReal);
2498
2499 Builder.SetInsertPoint(FinalReductionImag);
2500 auto *NewImag = Builder.CreateExtractValue(Agg: Deinterleave, Idxs: 1);
2501 FinalReductionImag->replaceUsesOfWith(From: Imag, To: NewImag);
2502}
2503
2504void ComplexDeinterleavingGraph::replaceNodes() {
2505 SmallVector<Instruction *, 16> DeadInstrRoots;
2506 for (auto *RootInstruction : OrderedRoots) {
2507 // Check if this potential root went through check process and we can
2508 // deinterleave it
2509 if (!RootToNode.count(Val: RootInstruction))
2510 continue;
2511
2512 IRBuilder<> Builder(RootInstruction);
2513 auto RootNode = RootToNode[RootInstruction];
2514 Value *R = replaceNode(Builder, Node: RootNode);
2515
2516 if (RootNode->Operation ==
2517 ComplexDeinterleavingOperation::ReductionOperation) {
2518 auto *RootReal = cast<Instruction>(Val: RootNode->Vals[0].Real);
2519 auto *RootImag = cast<Instruction>(Val: RootNode->Vals[0].Imag);
2520 ReductionInfo[RootReal].first->removeIncomingValue(BB: BackEdge);
2521 ReductionInfo[RootImag].first->removeIncomingValue(BB: BackEdge);
2522 DeadInstrRoots.push_back(Elt: RootReal);
2523 DeadInstrRoots.push_back(Elt: RootImag);
2524 } else if (RootNode->Operation ==
2525 ComplexDeinterleavingOperation::ReductionSingle) {
2526 auto *RootInst = cast<Instruction>(Val: RootNode->Vals[0].Real);
2527 auto &Info = ReductionInfo[RootInst];
2528 Info.first->removeIncomingValue(BB: BackEdge);
2529 DeadInstrRoots.push_back(Elt: Info.second);
2530 } else {
2531 assert(R && "Unable to find replacement for RootInstruction");
2532 DeadInstrRoots.push_back(Elt: RootInstruction);
2533 RootInstruction->replaceAllUsesWith(V: R);
2534 }
2535 }
2536
2537 for (auto *I : DeadInstrRoots)
2538 RecursivelyDeleteTriviallyDeadInstructions(V: I, TLI);
2539}
2540