LLVM 24.0.0git
ComplexDeinterleavingPass.cpp
Go to the documentation of this file.
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
64#include "llvm/ADT/MapVector.h"
65#include "llvm/ADT/Statistic.h"
70#include "llvm/IR/IRBuilder.h"
71#include "llvm/IR/Intrinsics.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
87 "enable-complex-deinterleaving",
88 cl::desc("Enable generation of complex instructions"), cl::init(true),
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) {
125}
126} // end namespace
128
129template <> struct llvm::DenseMapInfo<ComplexValue> {
130 static unsigned getHashValue(const ComplexValue &Val) {
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({R, I});
176 }
177
178 ComplexDeinterleavingCompositeNode(ComplexDeinterleavingOperation Op,
180 : Operation(Op), Vals(Other) {}
181
182private:
183 friend class ComplexDeinterleavingGraph;
184 using CompositeNode = ComplexDeinterleavingCompositeNode;
185 bool OperandsValid = true;
186
187public:
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
197 ComplexDeinterleavingRotation::Rotation_0;
199 Value *ReplacementNode = nullptr;
200
201 void addOperand(CompositeNode *Node) {
202 if (!Node)
203 OperandsValid = false;
204 Operands.push_back(Node);
205 }
206
207 void dump() { dump(dbgs()); }
208 void dump(raw_ostream &OS) {
209 auto PrintValue = [&](Value *V) {
210 if (V) {
211 OS << "\"";
212 V->print(OS, 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
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(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({R, 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(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
525 const TargetLowering *TL = TM->getSubtargetImpl(F)->getTargetLowering();
526 auto &TLI = AM.getResult<llvm::TargetLibraryAnalysis>(F);
527 if (!ComplexDeinterleaving(TL, &TLI).runOnFunction(F))
528 return PreservedAnalyses::all();
529
532 return PA;
533}
534
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) {
548 dbgs() << "Complex deinterleaving has been explicitly disabled.\n");
549 return false;
550 }
551
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, 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, 4);
567 }
568
569 // TODO: We can also support interleave factors of 6 and 8 if needed.
570
571 return Changed;
572}
573
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
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, m_FNeg(m_Value())) || match(V, m_Neg(m_Value()));
603}
604
606 assert(isNeg(V));
607 auto *I = cast<Instruction>(V);
608 if (I->getOpcode() == Instruction::FNeg)
609 return I->getOperand(0);
610
611 return I->getOperand(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(&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)) {
647 dbgs() << " - Real or imaginary instruction is not fmul or mul\n");
648 return nullptr;
649 }
650
651 Value *R0 = Real->getOperand(0);
652 Value *R1 = Real->getOperand(1);
653 Value *I0 = Imag->getOperand(0);
654 Value *I1 = Imag->getOperand(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(R0)) {
660 Negs |= 1;
661 R0 = getNegOperand(R0);
662 } else if (isNeg(R1)) {
663 Negs |= 1;
664 R1 = getNegOperand(R1);
665 }
666
667 if (isNeg(I0)) {
668 Negs |= 2;
669 Negs ^= 1;
670 I0 = getNegOperand(I0);
671 } else if (isNeg(I1)) {
672 Negs |= 2;
673 Negs ^= 1;
674 I1 = getNegOperand(I1);
675 }
676
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(UncommonRealOp, 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(PartialMatch.first, PartialMatch.second);
714 if (!CommonNode) {
715 LLVM_DEBUG(dbgs() << " - No CommonNode identified\n");
716 return nullptr;
717 }
718
719 CompositeNode *UncommonNode = identifyNode(UncommonRealOp, UncommonImagOp);
720 if (!UncommonNode) {
721 LLVM_DEBUG(dbgs() << " - No UncommonNode identified\n");
722 return nullptr;
723 }
724
725 CompositeNode *Node = prepareCompositeNode(
726 ComplexDeinterleavingOperation::CMulPartial, Real, Imag);
727 Node->Rotation = Rotation;
728 Node->addOperand(CommonNode);
729 Node->addOperand(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 };
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>(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(0);
768 Instruction *RealMulI = dyn_cast<Instruction>(Real->getOperand(1));
769 if (!RealMulI)
770 return nullptr;
771 Value *CI = Imag->getOperand(0);
772 Instruction *ImagMulI = dyn_cast<Instruction>(Imag->getOperand(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(0);
782 Value *R1 = RealMulI->getOperand(1);
783 Value *I0 = ImagMulI->getOperand(0);
784 Value *I1 = ImagMulI->getOperand(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(UncommonRealOp, 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>(CR);
817 auto *CIInst = dyn_cast<Instruction>(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(CRInst, CIInst, PartialMatch);
826 if (!CNode) {
827 LLVM_DEBUG(dbgs() << " - No cnode identified\n");
828 return nullptr;
829 }
830
831 CompositeNode *UncommonRes = identifyNode(UncommonRealOp, 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(PartialMatch.first, PartialMatch.second);
840 if (!CommonRes) {
841 LLVM_DEBUG(dbgs() << " - No CommonRes identified\n");
842 return nullptr;
843 }
844
845 CompositeNode *Node = prepareCompositeNode(
846 ComplexDeinterleavingOperation::CMulPartial, Real, Imag);
847 Node->Rotation = Rotation;
848 Node->addOperand(CommonRes);
849 Node->addOperand(UncommonRes);
850 Node->addOperand(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
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>(Real->getOperand(0));
876 auto *BI = dyn_cast<Instruction>(Real->getOperand(1));
877 auto *AI = dyn_cast<Instruction>(Imag->getOperand(0));
878 auto *BR = dyn_cast<Instruction>(Imag->getOperand(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(AR, 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(BR, 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(ComplexDeinterleavingOperation::CAdd, Real, Imag);
898 Node->Rotation = Rotation;
899 Node->addOperand(ResA);
900 Node->addOperand(ResB);
901 return submitCompositeNode(Node);
902}
903
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
915 auto Pattern =
917
918 return match(A, Pattern) && match(B, Pattern);
919}
920
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>(Vals[0].Real);
939 unsigned FirstOpc = FirstReal->getOpcode();
940 for (auto &V : Vals) {
941 auto *Real = cast<Instruction>(V.Real);
942 auto *Imag = cast<Instruction>(V.Imag);
943 if (Real->getOpcode() != FirstOpc || Imag->getOpcode() != FirstOpc)
944 return nullptr;
945
948 return nullptr;
949
950 if (isa<FPMathOperator>(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>(V.Real)->getOperand(0);
959 auto *I0 = cast<Instruction>(V.Imag)->getOperand(0);
960 OpVals.push_back({R0, I0});
961 }
962
963 CompositeNode *Op0 = identifyNode(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>(V.Real)->getOperand(1);
972 auto *I1 = cast<Instruction>(V.Imag)->getOperand(1);
973 OpVals.push_back({R1, I1});
974 }
975 Op1 = identifyNode(OpVals);
976 if (Op1 == nullptr)
977 return nullptr;
978 }
979
980 auto Node =
981 prepareCompositeNode(ComplexDeinterleavingOperation::Symmetric, Vals);
982 Node->Opcode = FirstReal->getOpcode();
983 if (isa<FPMathOperator>(FirstReal))
984 Node->Flags = FirstReal->getFastMathFlags();
985
986 Node->addOperand(Op0);
987 if (FirstReal->isBinaryOp())
988 Node->addOperand(Op1);
989
990 return submitCompositeNode(Node);
991}
992
993ComplexDeinterleavingGraph::CompositeNode *
994ComplexDeinterleavingGraph::identifyDotProduct(Value *V) {
996 ComplexDeinterleavingOperation::CDot, 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>(V);
1004 auto *RealUser = cast<Instruction>(*Inst->user_begin());
1005
1006 CompositeNode *CN =
1007 prepareCompositeNode(ComplexDeinterleavingOperation::CDot, Inst, 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>(V))
1021 return CI->getOperand(0);
1022 return V;
1023 };
1024
1025 auto PatternRot0 = m_Intrinsic<PartialReduceInt>(
1027 m_Mul(m_Value(BReal), m_Value(AReal))),
1028 m_Neg(m_Mul(m_Value(BImag), m_Value(AImag))));
1029
1030 auto PatternRot270 = m_Intrinsic<PartialReduceInt>(
1032 m_Value(Phi), m_Neg(m_Mul(m_Value(BReal), m_Value(AImag)))),
1033 m_Mul(m_Value(BImag), m_Value(AReal)));
1034
1035 if (match(Inst, PatternRot0)) {
1036 CN->Rotation = ComplexDeinterleavingRotation::Rotation_0;
1037 } else if (match(Inst, 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>(
1046 m_Mul(m_Value(BReal), m_Value(A0))),
1047 m_Mul(m_Value(BImag), m_Value(A1)));
1048
1049 if (!match(Inst, 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(A0, A1);
1057 if (!ANode) {
1058 // Test if A0 is imag/A1 is real
1059 ANode = identifyNode(A1, 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>(V->getType());
1079 Type *ExpectedOperandTy = VectorType::getSubdividedVectorType(VTy, 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(AReal, 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(identifyNode(BReal, BImag));
1107 CN->addOperand(identifyNode(Phi, RealUser));
1108
1109 return submitCompositeNode(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>(R->getType()) || !isa<VectorType>(I->getType()))
1116 return nullptr;
1117
1118 if (!R->hasUseList() || !I->hasUseList())
1119 return nullptr;
1120
1121 auto CommonUser =
1122 findCommonBetweenCollections<Value *>(R->users(), I->users());
1123 if (!CommonUser)
1124 return nullptr;
1125
1126 auto *IInst = dyn_cast<IntrinsicInst>(*CommonUser);
1127 if (!IInst || IInst->getIntrinsicID() != Intrinsic::vector_partial_reduce_add)
1128 return nullptr;
1129
1130 if (CompositeNode *CN = identifyDotProduct(IInst))
1131 return CN;
1132
1133 return nullptr;
1134}
1135
1136ComplexDeinterleavingGraph::CompositeNode *
1137ComplexDeinterleavingGraph::identifyNode(ComplexValues &Vals) {
1138 auto It = CachedResult.find(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>(V.Real);
1160 auto *Imag = dyn_cast<Instruction>(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>(Vals[0].Real);
1171 auto *Imag = dyn_cast<Instruction>(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>(Real->getType());
1179 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
1180
1181 bool HasCMulSupport = TL->isComplexDeinterleavingOperationSupported(
1182 ComplexDeinterleavingOperation::CMulPartial, NewVTy);
1183 bool HasCAddSupport = TL->isComplexDeinterleavingOperationSupported(
1184 ComplexDeinterleavingOperation::CAdd, NewVTy);
1185
1186 if (HasCMulSupport && isInstructionPairMul(Real, Imag)) {
1187 if (CompositeNode *CN = identifyPartialMul(Real, Imag))
1188 return CN;
1189 }
1190
1191 if (HasCAddSupport && isInstructionPairAdd(Real, Imag)) {
1192 if (CompositeNode *CN = identifyAdd(Real, Imag))
1193 return CN;
1194 }
1195
1196 if (HasCMulSupport && HasCAddSupport) {
1197 if (CompositeNode *CN = identifyReassocNodes(Real, 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();
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>(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
1252 if (!I) {
1253 Addends.emplace_back(V, 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(2)) {
1264 LLVM_DEBUG(dbgs() << "Found potential sub-expression: " << *I << "\n");
1265 Addends.emplace_back(I, IsPositive);
1266 continue;
1267 }
1268 switch (I->getOpcode()) {
1269 case Instruction::FAdd:
1270 case Instruction::Add:
1271 Worklist.emplace_back(I->getOperand(1), IsPositive);
1272 Worklist.emplace_back(I->getOperand(0), IsPositive);
1273 break;
1274 case Instruction::FSub:
1275 Worklist.emplace_back(I->getOperand(1), !IsPositive);
1276 Worklist.emplace_back(I->getOperand(0), IsPositive);
1277 break;
1278 case Instruction::Sub:
1279 if (isNeg(I)) {
1280 Worklist.emplace_back(getNegOperand(I), !IsPositive);
1281 } else {
1282 Worklist.emplace_back(I->getOperand(1), !IsPositive);
1283 Worklist.emplace_back(I->getOperand(0), IsPositive);
1284 }
1285 break;
1286 case Instruction::FMul:
1287 case Instruction::Mul: {
1288 Value *A, *B;
1289 if (isNeg(I->getOperand(0))) {
1290 A = getNegOperand(I->getOperand(0));
1291 IsPositive = !IsPositive;
1292 } else {
1293 A = I->getOperand(0);
1294 }
1295
1296 if (isNeg(I->getOperand(1))) {
1297 B = getNegOperand(I->getOperand(1));
1298 IsPositive = !IsPositive;
1299 } else {
1300 B = I->getOperand(1);
1301 }
1302 Muls.push_back(Product{A, B, IsPositive});
1303 break;
1304 }
1305 case Instruction::FNeg:
1306 Worklist.emplace_back(I->getOperand(0), !IsPositive);
1307 break;
1308 case Instruction::Call: {
1309 Value *A, *B, *C;
1311 m_Value(C))) &&
1313 m_Value(C)))) {
1314 Addends.emplace_back(I, IsPositive);
1315 continue;
1316 }
1317
1318 bool IsProductPositive = IsPositive;
1319 if (isNeg(A)) {
1320 A = getNegOperand(A);
1321 IsProductPositive = !IsProductPositive;
1322 }
1323
1324 if (isNeg(B)) {
1325 B = getNegOperand(B);
1326 IsProductPositive = !IsProductPositive;
1327 }
1328
1329 Muls.push_back(Product{A, B, IsProductPositive});
1330 Worklist.emplace_back(C, IsPositive);
1331 break;
1332 }
1333 default:
1334 Addends.emplace_back(I, IsPositive);
1335 continue;
1336 }
1337
1338 if (Flags && I->getFastMathFlags() != *Flags) {
1339 LLVM_DEBUG(dbgs() << "The instruction's fast math flags are "
1340 "inconsistent with the root instructions' flags: "
1341 << *I << "\n");
1342 return false;
1343 }
1344 }
1345 return true;
1346 };
1347
1348 SmallVector<Product> RealMuls, ImagMuls;
1349 AddendList RealAddends, ImagAddends;
1350 if (!Collect(Real, RealMuls, RealAddends) ||
1351 !Collect(Imag, ImagMuls, ImagAddends))
1352 return nullptr;
1353
1354 if (RealAddends.size() != ImagAddends.size())
1355 return nullptr;
1356
1357 CompositeNode *FinalNode = nullptr;
1358 if (!RealMuls.empty() || !ImagMuls.empty()) {
1359 // If there are multiplicands, extract positive addend and use it as an
1360 // accumulator
1361 FinalNode = extractPositiveAddend(RealAddends, ImagAddends);
1362 FinalNode = identifyMultiplications(RealMuls, ImagMuls, FinalNode);
1363 if (!FinalNode)
1364 return nullptr;
1365 }
1366
1367 // Identify and process remaining additions
1368 if (!RealAddends.empty() || !ImagAddends.empty()) {
1369 FinalNode = identifyAdditions(RealAddends, ImagAddends, Flags, FinalNode);
1370 if (!FinalNode)
1371 return nullptr;
1372 }
1373 assert(FinalNode && "FinalNode can not be nullptr here");
1374 assert(FinalNode->Vals.size() == 1);
1375 // Set the Real and Imag fields of the final node and submit it
1376 FinalNode->Vals[0].Real = Real;
1377 FinalNode->Vals[0].Imag = Imag;
1378 submitCompositeNode(FinalNode);
1379 return FinalNode;
1380}
1381
1382bool ComplexDeinterleavingGraph::collectPartialMuls(
1383 ArrayRef<Product> RealMuls, ArrayRef<Product> ImagMuls,
1384 SmallVectorImpl<PartialMulCandidate> &PartialMulCandidates) {
1385 // Helper function to extract a common operand from two products
1386 auto FindCommonInstruction = [](const Product &Real,
1387 const Product &Imag) -> Value * {
1388 if (Real.Multiplicand == Imag.Multiplicand ||
1389 Real.Multiplicand == Imag.Multiplier)
1390 return Real.Multiplicand;
1391
1392 if (Real.Multiplier == Imag.Multiplicand ||
1393 Real.Multiplier == Imag.Multiplier)
1394 return Real.Multiplier;
1395
1396 return nullptr;
1397 };
1398
1399 // Iterating over real and imaginary multiplications to find common operands
1400 // If a common operand is found, a partial multiplication candidate is created
1401 // and added to the candidates vector The function returns false if no common
1402 // operands are found for any product
1403 for (unsigned i = 0; i < RealMuls.size(); ++i) {
1404 bool FoundCommon = false;
1405 for (unsigned j = 0; j < ImagMuls.size(); ++j) {
1406 auto *Common = FindCommonInstruction(RealMuls[i], ImagMuls[j]);
1407 if (!Common)
1408 continue;
1409
1410 auto *A = RealMuls[i].Multiplicand == Common ? RealMuls[i].Multiplier
1411 : RealMuls[i].Multiplicand;
1412 auto *B = ImagMuls[j].Multiplicand == Common ? ImagMuls[j].Multiplier
1413 : ImagMuls[j].Multiplicand;
1414
1415 auto Node = identifyNode(A, B);
1416 if (Node) {
1417 FoundCommon = true;
1418 PartialMulCandidates.push_back({Common, Node, i, j, false});
1419 }
1420
1421 Node = identifyNode(B, A);
1422 if (Node) {
1423 FoundCommon = true;
1424 PartialMulCandidates.push_back({Common, Node, i, j, true});
1425 }
1426 }
1427 if (!FoundCommon)
1428 return false;
1429 }
1430 return true;
1431}
1432
1433ComplexDeinterleavingGraph::CompositeNode *
1434ComplexDeinterleavingGraph::identifyMultiplications(
1435 SmallVectorImpl<Product> &RealMuls, SmallVectorImpl<Product> &ImagMuls,
1436 CompositeNode *Accumulator = nullptr) {
1437 if (RealMuls.size() != ImagMuls.size())
1438 return nullptr;
1439
1441 if (!collectPartialMuls(RealMuls, ImagMuls, Info))
1442 return nullptr;
1443
1444 // Map to store common instruction to node pointers
1445 DenseMap<Value *, CompositeNode *> CommonToNode;
1446 SmallVector<bool> Processed(Info.size(), false);
1447 for (unsigned I = 0; I < Info.size(); ++I) {
1448 if (Processed[I])
1449 continue;
1450
1451 PartialMulCandidate &InfoA = Info[I];
1452 for (unsigned J = I + 1; J < Info.size(); ++J) {
1453 if (Processed[J])
1454 continue;
1455
1456 PartialMulCandidate &InfoB = Info[J];
1457 auto *InfoReal = &InfoA;
1458 auto *InfoImag = &InfoB;
1459
1460 auto NodeFromCommon = identifyNode(InfoReal->Common, InfoImag->Common);
1461 if (!NodeFromCommon) {
1462 std::swap(InfoReal, InfoImag);
1463 NodeFromCommon = identifyNode(InfoReal->Common, InfoImag->Common);
1464 }
1465 if (!NodeFromCommon)
1466 continue;
1467
1468 CommonToNode[InfoReal->Common] = NodeFromCommon;
1469 CommonToNode[InfoImag->Common] = NodeFromCommon;
1470 Processed[I] = true;
1471 Processed[J] = true;
1472 }
1473 }
1474
1475 SmallVector<bool> ProcessedReal(RealMuls.size(), false);
1476 SmallVector<bool> ProcessedImag(ImagMuls.size(), false);
1477 CompositeNode *Result = Accumulator;
1478 for (auto &PMI : Info) {
1479 if (ProcessedReal[PMI.RealIdx] || ProcessedImag[PMI.ImagIdx])
1480 continue;
1481
1482 auto It = CommonToNode.find(PMI.Common);
1483 // TODO: Process independent complex multiplications. Cases like this:
1484 // A.real() * B where both A and B are complex numbers.
1485 if (It == CommonToNode.end()) {
1486 LLVM_DEBUG({
1487 dbgs() << "Unprocessed independent partial multiplication:\n";
1488 for (auto *Mul : {&RealMuls[PMI.RealIdx], &RealMuls[PMI.RealIdx]})
1489 dbgs().indent(4) << (Mul->IsPositive ? "+" : "-") << *Mul->Multiplier
1490 << " multiplied by " << *Mul->Multiplicand << "\n";
1491 });
1492 return nullptr;
1493 }
1494
1495 auto &RealMul = RealMuls[PMI.RealIdx];
1496 auto &ImagMul = ImagMuls[PMI.ImagIdx];
1497
1498 auto NodeA = It->second;
1499 auto NodeB = PMI.Node;
1500 auto IsMultiplicandReal = PMI.Common == NodeA->Vals[0].Real;
1501 // The following table illustrates the relationship between multiplications
1502 // and rotations. If we consider the multiplication (X + iY) * (U + iV), we
1503 // can see:
1504 //
1505 // Rotation | Real | Imag |
1506 // ---------+--------+--------+
1507 // 0 | x * u | x * v |
1508 // 90 | -y * v | y * u |
1509 // 180 | -x * u | -x * v |
1510 // 270 | y * v | -y * u |
1511 //
1512 // Check if the candidate can indeed be represented by partial
1513 // multiplication
1514 // TODO: Add support for multiplication by complex one
1515 if ((IsMultiplicandReal && PMI.IsNodeInverted) ||
1516 (!IsMultiplicandReal && !PMI.IsNodeInverted))
1517 continue;
1518
1519 // Determine the rotation based on the multiplications
1521 if (IsMultiplicandReal) {
1522 // Detect 0 and 180 degrees rotation
1523 if (RealMul.IsPositive && ImagMul.IsPositive)
1525 else if (!RealMul.IsPositive && !ImagMul.IsPositive)
1527 else
1528 continue;
1529
1530 } else {
1531 // Detect 90 and 270 degrees rotation
1532 if (!RealMul.IsPositive && ImagMul.IsPositive)
1534 else if (RealMul.IsPositive && !ImagMul.IsPositive)
1536 else
1537 continue;
1538 }
1539
1540 LLVM_DEBUG({
1541 dbgs() << "Identified partial multiplication (X, Y) * (U, V):\n";
1542 dbgs().indent(4) << "X: " << *NodeA->Vals[0].Real << "\n";
1543 dbgs().indent(4) << "Y: " << *NodeA->Vals[0].Imag << "\n";
1544 dbgs().indent(4) << "U: " << *NodeB->Vals[0].Real << "\n";
1545 dbgs().indent(4) << "V: " << *NodeB->Vals[0].Imag << "\n";
1546 dbgs().indent(4) << "Rotation - " << (int)Rotation * 90 << "\n";
1547 });
1548
1549 CompositeNode *NodeMul = prepareCompositeNode(
1550 ComplexDeinterleavingOperation::CMulPartial, nullptr, nullptr);
1551 NodeMul->Rotation = Rotation;
1552 NodeMul->addOperand(NodeA);
1553 NodeMul->addOperand(NodeB);
1554 if (Result)
1555 NodeMul->addOperand(Result);
1556 submitCompositeNode(NodeMul);
1557 Result = NodeMul;
1558 ProcessedReal[PMI.RealIdx] = true;
1559 ProcessedImag[PMI.ImagIdx] = true;
1560 }
1561
1562 // Ensure all products have been processed, if not return nullptr.
1563 if (!all_of(ProcessedReal, [](bool V) { return V; }) ||
1564 !all_of(ProcessedImag, [](bool V) { return V; })) {
1565
1566 // Dump debug information about which partial multiplications are not
1567 // processed.
1568 LLVM_DEBUG({
1569 dbgs() << "Unprocessed products (Real):\n";
1570 for (size_t i = 0; i < ProcessedReal.size(); ++i) {
1571 if (!ProcessedReal[i])
1572 dbgs().indent(4) << (RealMuls[i].IsPositive ? "+" : "-")
1573 << *RealMuls[i].Multiplier << " multiplied by "
1574 << *RealMuls[i].Multiplicand << "\n";
1575 }
1576 dbgs() << "Unprocessed products (Imag):\n";
1577 for (size_t i = 0; i < ProcessedImag.size(); ++i) {
1578 if (!ProcessedImag[i])
1579 dbgs().indent(4) << (ImagMuls[i].IsPositive ? "+" : "-")
1580 << *ImagMuls[i].Multiplier << " multiplied by "
1581 << *ImagMuls[i].Multiplicand << "\n";
1582 }
1583 });
1584 return nullptr;
1585 }
1586
1587 return Result;
1588}
1589
1590ComplexDeinterleavingGraph::CompositeNode *
1591ComplexDeinterleavingGraph::identifyAdditions(
1592 AddendList &RealAddends, AddendList &ImagAddends,
1593 std::optional<FastMathFlags> Flags, CompositeNode *Accumulator = nullptr) {
1594 if (RealAddends.size() != ImagAddends.size())
1595 return nullptr;
1596
1597 CompositeNode *Result = nullptr;
1598 // If we have accumulator use it as first addend
1599 if (Accumulator)
1601 // Otherwise find an element with both positive real and imaginary parts.
1602 else
1603 Result = extractPositiveAddend(RealAddends, ImagAddends);
1604
1605 if (!Result)
1606 return nullptr;
1607
1608 while (!RealAddends.empty()) {
1609 auto ItR = RealAddends.begin();
1610 auto [R, IsPositiveR] = *ItR;
1611
1612 bool FoundImag = false;
1613 for (auto ItI = ImagAddends.begin(); ItI != ImagAddends.end(); ++ItI) {
1614 auto [I, IsPositiveI] = *ItI;
1616 if (IsPositiveR && IsPositiveI)
1617 Rotation = ComplexDeinterleavingRotation::Rotation_0;
1618 else if (!IsPositiveR && IsPositiveI)
1619 Rotation = ComplexDeinterleavingRotation::Rotation_90;
1620 else if (!IsPositiveR && !IsPositiveI)
1621 Rotation = ComplexDeinterleavingRotation::Rotation_180;
1622 else
1623 Rotation = ComplexDeinterleavingRotation::Rotation_270;
1624
1625 CompositeNode *AddNode = nullptr;
1626 if (Rotation == ComplexDeinterleavingRotation::Rotation_0 ||
1627 Rotation == ComplexDeinterleavingRotation::Rotation_180) {
1628 AddNode = identifyNode(R, I);
1629 } else {
1630 AddNode = identifyNode(I, R);
1631 }
1632 if (AddNode) {
1633 LLVM_DEBUG({
1634 dbgs() << "Identified addition:\n";
1635 dbgs().indent(4) << "X: " << *R << "\n";
1636 dbgs().indent(4) << "Y: " << *I << "\n";
1637 dbgs().indent(4) << "Rotation - " << (int)Rotation * 90 << "\n";
1638 });
1639
1640 CompositeNode *TmpNode = nullptr;
1642 TmpNode = prepareCompositeNode(
1643 ComplexDeinterleavingOperation::Symmetric, nullptr, nullptr);
1644 if (Flags) {
1645 TmpNode->Opcode = Instruction::FAdd;
1646 TmpNode->Flags = *Flags;
1647 } else {
1648 TmpNode->Opcode = Instruction::Add;
1649 }
1650 } else if (Rotation ==
1652 TmpNode = prepareCompositeNode(
1653 ComplexDeinterleavingOperation::Symmetric, nullptr, nullptr);
1654 if (Flags) {
1655 TmpNode->Opcode = Instruction::FSub;
1656 TmpNode->Flags = *Flags;
1657 } else {
1658 TmpNode->Opcode = Instruction::Sub;
1659 }
1660 } else {
1661 TmpNode = prepareCompositeNode(ComplexDeinterleavingOperation::CAdd,
1662 nullptr, nullptr);
1663 TmpNode->Rotation = Rotation;
1664 }
1665
1666 TmpNode->addOperand(Result);
1667 TmpNode->addOperand(AddNode);
1668 submitCompositeNode(TmpNode);
1669 Result = TmpNode;
1670 RealAddends.erase(ItR);
1671 ImagAddends.erase(ItI);
1672 FoundImag = true;
1673 break;
1674 }
1675 }
1676 if (!FoundImag)
1677 return nullptr;
1678 }
1679 return Result;
1680}
1681
1682ComplexDeinterleavingGraph::CompositeNode *
1683ComplexDeinterleavingGraph::extractPositiveAddend(AddendList &RealAddends,
1684 AddendList &ImagAddends) {
1685 for (auto ItR = RealAddends.begin(); ItR != RealAddends.end(); ++ItR) {
1686 for (auto ItI = ImagAddends.begin(); ItI != ImagAddends.end(); ++ItI) {
1687 auto [R, IsPositiveR] = *ItR;
1688 auto [I, IsPositiveI] = *ItI;
1689 if (IsPositiveR && IsPositiveI) {
1690 auto Result = identifyNode(R, I);
1691 if (Result) {
1692 RealAddends.erase(ItR);
1693 ImagAddends.erase(ItI);
1694 return Result;
1695 }
1696 }
1697 }
1698 }
1699 return nullptr;
1700}
1701
1702bool ComplexDeinterleavingGraph::identifyNodes(Instruction *RootI) {
1703 // This potential root instruction might already have been recognized as
1704 // reduction. Because RootToNode maps both Real and Imaginary parts to
1705 // CompositeNode we should choose only one either Real or Imag instruction to
1706 // use as an anchor for generating complex instruction.
1707 auto It = RootToNode.find(RootI);
1708 if (It != RootToNode.end()) {
1709 auto RootNode = It->second;
1710 assert(RootNode->Operation ==
1711 ComplexDeinterleavingOperation::ReductionOperation ||
1712 RootNode->Operation ==
1713 ComplexDeinterleavingOperation::ReductionSingle);
1714 assert(RootNode->Vals.size() == 1 &&
1715 "Cannot handle reductions involving multiple complex values");
1716 // Find out which part, Real or Imag, comes later, and only if we come to
1717 // the latest part, add it to OrderedRoots.
1718 auto *R = cast<Instruction>(RootNode->Vals[0].Real);
1719 auto *I = RootNode->Vals[0].Imag ? cast<Instruction>(RootNode->Vals[0].Imag)
1720 : nullptr;
1721
1722 Instruction *ReplacementAnchor;
1723 if (I)
1724 ReplacementAnchor = R->comesBefore(I) ? I : R;
1725 else
1726 ReplacementAnchor = R;
1727
1728 if (ReplacementAnchor != RootI)
1729 return false;
1730 OrderedRoots.push_back(RootI);
1731 return true;
1732 }
1733
1734 auto RootNode = identifyRoot(RootI);
1735 if (!RootNode)
1736 return false;
1737
1738 LLVM_DEBUG({
1739 Function *F = RootI->getFunction();
1740 BasicBlock *B = RootI->getParent();
1741 dbgs() << "Complex deinterleaving graph for " << F->getName()
1742 << "::" << B->getName() << ".\n";
1743 dump(dbgs());
1744 dbgs() << "\n";
1745 });
1746 RootToNode[RootI] = RootNode;
1747 OrderedRoots.push_back(RootI);
1748 return true;
1749}
1750
1751bool ComplexDeinterleavingGraph::collectPotentialReductions(BasicBlock *B) {
1752 bool FoundPotentialReduction = false;
1753 if (Factor != 2)
1754 return false;
1755
1756 auto *Br = dyn_cast<CondBrInst>(B->getTerminator());
1757 if (!Br)
1758 return false;
1759
1760 // Identify simple one-block loop
1761 if (Br->getSuccessor(0) != B && Br->getSuccessor(1) != B)
1762 return false;
1763
1764 for (auto &PHI : B->phis()) {
1765 if (PHI.getNumIncomingValues() != 2)
1766 continue;
1767
1768 if (!PHI.getType()->isVectorTy())
1769 continue;
1770
1771 auto *ReductionOp = dyn_cast<Instruction>(PHI.getIncomingValueForBlock(B));
1772 if (!ReductionOp)
1773 continue;
1774
1775 // Check if final instruction is reduced outside of current block
1776 Instruction *FinalReduction = nullptr;
1777 auto NumUsers = 0u;
1778 for (auto *U : ReductionOp->users()) {
1779 ++NumUsers;
1780 if (U == &PHI)
1781 continue;
1782 FinalReduction = dyn_cast<Instruction>(U);
1783 }
1784
1785 if (NumUsers != 2 || !FinalReduction || FinalReduction->getParent() == B ||
1786 isa<PHINode>(FinalReduction))
1787 continue;
1788
1789 ReductionInfo[ReductionOp] = {&PHI, FinalReduction};
1790 BackEdge = B;
1791 auto BackEdgeIdx = PHI.getBasicBlockIndex(B);
1792 auto IncomingIdx = BackEdgeIdx == 0 ? 1 : 0;
1793 Incoming = PHI.getIncomingBlock(IncomingIdx);
1794 FoundPotentialReduction = true;
1795
1796 // If the initial value of PHINode is an Instruction, consider it a leaf
1797 // value of a complex deinterleaving graph.
1798 if (auto *InitPHI =
1799 dyn_cast<Instruction>(PHI.getIncomingValueForBlock(Incoming)))
1800 FinalInstructions.insert(InitPHI);
1801 }
1802 return FoundPotentialReduction;
1803}
1804
1805void ComplexDeinterleavingGraph::identifyReductionNodes() {
1806 assert(Factor == 2 && "Cannot handle multiple complex values");
1807
1808 SmallVector<bool> Processed(ReductionInfo.size(), false);
1809 SmallVector<Instruction *> OperationInstruction;
1810 for (auto &P : ReductionInfo)
1811 OperationInstruction.push_back(P.first);
1812
1813 // Identify a complex computation by evaluating two reduction operations that
1814 // potentially could be involved
1815 for (size_t i = 0; i < OperationInstruction.size(); ++i) {
1816 if (Processed[i])
1817 continue;
1818 for (size_t j = i + 1; j < OperationInstruction.size(); ++j) {
1819 if (Processed[j])
1820 continue;
1821 auto *Real = OperationInstruction[i];
1822 auto *Imag = OperationInstruction[j];
1823 if (Real->getType() != Imag->getType())
1824 continue;
1825
1826 RealPHI = ReductionInfo[Real].first;
1827 ImagPHI = ReductionInfo[Imag].first;
1828 PHIsFound = false;
1829 auto Node = identifyNode(Real, Imag);
1830 if (!Node) {
1831 std::swap(Real, Imag);
1832 std::swap(RealPHI, ImagPHI);
1833 Node = identifyNode(Real, Imag);
1834 }
1835
1836 // If a node is identified and reduction PHINode is used in the chain of
1837 // operations, mark its operation instructions as used to prevent
1838 // re-identification and attach the node to the real part
1839 if (Node && PHIsFound) {
1840 LLVM_DEBUG(dbgs() << "Identified reduction starting from instructions: "
1841 << *Real << " / " << *Imag << "\n");
1842 Processed[i] = true;
1843 Processed[j] = true;
1844 auto RootNode = prepareCompositeNode(
1845 ComplexDeinterleavingOperation::ReductionOperation, Real, Imag);
1846 RootNode->addOperand(Node);
1847 RootToNode[Real] = RootNode;
1848 RootToNode[Imag] = RootNode;
1849 submitCompositeNode(RootNode);
1850 break;
1851 }
1852 }
1853
1854 auto *Real = OperationInstruction[i];
1855 // We want to check that we have 2 operands, but the function attributes
1856 // being counted as operands bloats this value.
1857 if (Processed[i] || Real->getNumOperands() < 2)
1858 continue;
1859
1860 // Can only combined integer reductions at the moment.
1861 if (!ReductionInfo[Real].second->getType()->isIntegerTy())
1862 continue;
1863
1864 RealPHI = ReductionInfo[Real].first;
1865 ImagPHI = nullptr;
1866 PHIsFound = false;
1867 auto Node = identifyNode(Real->getOperand(0), Real->getOperand(1));
1868 if (Node && PHIsFound) {
1869 LLVM_DEBUG(
1870 dbgs() << "Identified single reduction starting from instruction: "
1871 << *Real << "/" << *ReductionInfo[Real].second << "\n");
1872
1873 // Reducing to a single vector is not supported, only permit reducing down
1874 // to scalar values.
1875 // Doing this here will leave the prior node in the graph,
1876 // however with no uses the node will be unreachable by the replacement
1877 // process. That along with the usage outside the graph should prevent the
1878 // replacement process from kicking off at all for this graph.
1879 // TODO Add support for reducing to a single vector value
1880 if (ReductionInfo[Real].second->getType()->isVectorTy())
1881 continue;
1882
1883 Processed[i] = true;
1884 auto RootNode = prepareCompositeNode(
1885 ComplexDeinterleavingOperation::ReductionSingle, Real, nullptr);
1886 RootNode->addOperand(Node);
1887 RootToNode[Real] = RootNode;
1888 submitCompositeNode(RootNode);
1889 }
1890 }
1891
1892 RealPHI = nullptr;
1893 ImagPHI = nullptr;
1894}
1895
1896bool ComplexDeinterleavingGraph::checkNodes() {
1897 bool FoundDeinterleaveNode = false;
1898 for (CompositeNode *N : CompositeNodes) {
1899 if (!N->areOperandsValid())
1900 return false;
1901
1902 if (N->Operation == ComplexDeinterleavingOperation::Deinterleave)
1903 FoundDeinterleaveNode = true;
1904 }
1905
1906 // We need a deinterleave node in order to guarantee that we're working with
1907 // complex numbers.
1908 if (!FoundDeinterleaveNode) {
1909 LLVM_DEBUG(
1910 dbgs() << "Couldn't find a deinterleave node within the graph, cannot "
1911 "guarantee safety during graph transformation.\n");
1912 return false;
1913 }
1914
1915 // Collect all instructions from roots to leaves
1916 SmallPtrSet<Instruction *, 16> AllInstructions;
1917 SmallVector<Instruction *, 8> Worklist;
1918 for (auto &Pair : RootToNode)
1919 Worklist.push_back(Pair.first);
1920
1921 // Extract all instructions that are used by all XCMLA/XCADD/ADD/SUB/NEG
1922 // chains
1923 while (!Worklist.empty()) {
1924 auto *I = Worklist.pop_back_val();
1925
1926 if (!AllInstructions.insert(I).second)
1927 continue;
1928
1929 for (Value *Op : I->operands()) {
1930 if (auto *OpI = dyn_cast<Instruction>(Op)) {
1931 if (!FinalInstructions.count(I))
1932 Worklist.emplace_back(OpI);
1933 }
1934 }
1935 }
1936
1937 // Find instructions that have users outside of chain
1938 for (auto *I : AllInstructions) {
1939 // Skip root nodes
1940 if (RootToNode.count(I))
1941 continue;
1942
1943 for (User *U : I->users()) {
1944 if (AllInstructions.count(cast<Instruction>(U)))
1945 continue;
1946
1947 // Found an instruction that is not used by XCMLA/XCADD chain
1948 Worklist.emplace_back(I);
1949 break;
1950 }
1951 }
1952
1953 // If any instructions are found to be used outside, find and remove roots
1954 // that somehow connect to those instructions.
1955 SmallPtrSet<Instruction *, 16> Visited;
1956 while (!Worklist.empty()) {
1957 auto *I = Worklist.pop_back_val();
1958 if (!Visited.insert(I).second)
1959 continue;
1960
1961 // Found an impacted root node. Removing it from the nodes to be
1962 // deinterleaved
1963 if (RootToNode.count(I)) {
1964 LLVM_DEBUG(dbgs() << "Instruction " << *I
1965 << " could be deinterleaved but its chain of complex "
1966 "operations have an outside user\n");
1967 RootToNode.erase(I);
1968 }
1969
1970 if (!AllInstructions.count(I) || FinalInstructions.count(I))
1971 continue;
1972
1973 for (User *U : I->users())
1974 Worklist.emplace_back(cast<Instruction>(U));
1975
1976 for (Value *Op : I->operands()) {
1977 if (auto *OpI = dyn_cast<Instruction>(Op))
1978 Worklist.emplace_back(OpI);
1979 }
1980 }
1981 return !RootToNode.empty();
1982}
1983
1984ComplexDeinterleavingGraph::CompositeNode *
1985ComplexDeinterleavingGraph::identifyRoot(Instruction *RootI) {
1986 if (auto *Intrinsic = dyn_cast<IntrinsicInst>(RootI)) {
1988 Intrinsic->getIntrinsicID())
1989 return nullptr;
1990
1991 ComplexValues Vals;
1992 for (unsigned I = 0; I < Factor; I += 2) {
1993 auto *Real = dyn_cast<Instruction>(Intrinsic->getOperand(I));
1994 auto *Imag = dyn_cast<Instruction>(Intrinsic->getOperand(I + 1));
1995 if (!Real || !Imag)
1996 return nullptr;
1997 Vals.push_back({Real, Imag});
1998 }
1999
2000 ComplexDeinterleavingGraph::CompositeNode *Node1 = identifyNode(Vals);
2001 if (!Node1)
2002 return nullptr;
2003 return Node1;
2004 }
2005
2006 // TODO: We could also add support for fixed-width interleave factors of 4
2007 // and above, but currently for symmetric operations the interleaves and
2008 // deinterleaves are already removed by VectorCombine. If we extend this to
2009 // permit complex multiplications, reductions, etc. then we should also add
2010 // support for fixed-width here.
2011 if (Factor != 2)
2012 return nullptr;
2013
2014 auto *SVI = dyn_cast<ShuffleVectorInst>(RootI);
2015 if (!SVI)
2016 return nullptr;
2017
2018 // Look for a shufflevector that takes separate vectors of the real and
2019 // imaginary components and recombines them into a single vector.
2020 if (!isInterleavingMask(SVI->getShuffleMask()))
2021 return nullptr;
2022
2023 Instruction *Real;
2024 Instruction *Imag;
2025 if (!match(RootI, m_Shuffle(m_Instruction(Real), m_Instruction(Imag))))
2026 return nullptr;
2027
2028 return identifyNode(Real, Imag);
2029}
2030
2031ComplexDeinterleavingGraph::CompositeNode *
2032ComplexDeinterleavingGraph::identifyDeinterleave(ComplexValues &Vals) {
2033 Instruction *II = nullptr;
2034
2035 // Must be at least one complex value.
2036 auto CheckExtract = [&](Value *V, unsigned ExpectedIdx,
2037 Instruction *ExpectedInsn) -> ExtractValueInst * {
2038 auto *EVI = dyn_cast<ExtractValueInst>(V);
2039 if (!EVI || EVI->getNumIndices() != 1 ||
2040 EVI->getIndices()[0] != ExpectedIdx ||
2041 !isa<Instruction>(EVI->getAggregateOperand()) ||
2042 (ExpectedInsn && ExpectedInsn != EVI->getAggregateOperand()))
2043 return nullptr;
2044 return EVI;
2045 };
2046
2047 for (unsigned Idx = 0; Idx < Vals.size(); Idx++) {
2048 ExtractValueInst *RealEVI = CheckExtract(Vals[Idx].Real, Idx * 2, II);
2049 if (RealEVI && Idx == 0)
2051 if (!RealEVI || !CheckExtract(Vals[Idx].Imag, (Idx * 2) + 1, II)) {
2052 II = nullptr;
2053 break;
2054 }
2055 }
2056
2057 if (auto *IntrinsicII = dyn_cast_or_null<IntrinsicInst>(II)) {
2058 if (IntrinsicII->getIntrinsicID() !=
2060 return nullptr;
2061
2062 // The remaining should match too.
2063 CompositeNode *PlaceholderNode = prepareCompositeNode(
2065 PlaceholderNode->ReplacementNode = II->getOperand(0);
2066 for (auto &V : Vals) {
2067 FinalInstructions.insert(cast<Instruction>(V.Real));
2068 FinalInstructions.insert(cast<Instruction>(V.Imag));
2069 }
2070 return submitCompositeNode(PlaceholderNode);
2071 }
2072
2073 if (Vals.size() != 1)
2074 return nullptr;
2075
2076 Value *Real = Vals[0].Real;
2077 Value *Imag = Vals[0].Imag;
2078 auto *RealShuffle = dyn_cast<ShuffleVectorInst>(Real);
2079 auto *ImagShuffle = dyn_cast<ShuffleVectorInst>(Imag);
2080 if (!RealShuffle || !ImagShuffle) {
2081 if (RealShuffle || ImagShuffle)
2082 LLVM_DEBUG(dbgs() << " - There's a shuffle where there shouldn't be.\n");
2083 return nullptr;
2084 }
2085
2086 Value *RealOp1 = RealShuffle->getOperand(1);
2087 if (!isa<UndefValue>(RealOp1) && !match(RealOp1, m_Zero())) {
2088 LLVM_DEBUG(dbgs() << " - RealOp1 is not undef or zero.\n");
2089 return nullptr;
2090 }
2091 Value *ImagOp1 = ImagShuffle->getOperand(1);
2092 if (!isa<UndefValue>(ImagOp1) && !match(ImagOp1, m_Zero())) {
2093 LLVM_DEBUG(dbgs() << " - ImagOp1 is not undef or zero.\n");
2094 return nullptr;
2095 }
2096
2097 Value *RealOp0 = RealShuffle->getOperand(0);
2098 Value *ImagOp0 = ImagShuffle->getOperand(0);
2099
2100 if (RealOp0 != ImagOp0) {
2101 LLVM_DEBUG(dbgs() << " - Shuffle operands are not equal.\n");
2102 return nullptr;
2103 }
2104
2105 ArrayRef<int> RealMask = RealShuffle->getShuffleMask();
2106 ArrayRef<int> ImagMask = ImagShuffle->getShuffleMask();
2107 if (!isDeinterleavingMask(RealMask) || !isDeinterleavingMask(ImagMask)) {
2108 LLVM_DEBUG(dbgs() << " - Masks are not deinterleaving.\n");
2109 return nullptr;
2110 }
2111
2112 if (RealMask[0] != 0 || ImagMask[0] != 1) {
2113 LLVM_DEBUG(dbgs() << " - Masks do not have the correct initial value.\n");
2114 return nullptr;
2115 }
2116
2117 // Type checking, the shuffle type should be a vector type of the same
2118 // scalar type, but half the size
2119 auto CheckType = [&](ShuffleVectorInst *Shuffle) {
2120 Value *Op = Shuffle->getOperand(0);
2121 auto *ShuffleTy = cast<FixedVectorType>(Shuffle->getType());
2122 auto *OpTy = cast<FixedVectorType>(Op->getType());
2123
2124 if (OpTy->getScalarType() != ShuffleTy->getScalarType())
2125 return false;
2126 if ((ShuffleTy->getNumElements() * 2) != OpTy->getNumElements())
2127 return false;
2128
2129 return true;
2130 };
2131
2132 auto CheckDeinterleavingShuffle = [&](ShuffleVectorInst *Shuffle) -> bool {
2133 if (!CheckType(Shuffle))
2134 return false;
2135
2136 ArrayRef<int> Mask = Shuffle->getShuffleMask();
2137 int Last = *Mask.rbegin();
2138
2139 Value *Op = Shuffle->getOperand(0);
2140 auto *OpTy = cast<FixedVectorType>(Op->getType());
2141 int NumElements = OpTy->getNumElements();
2142
2143 // Ensure that the deinterleaving shuffle only pulls from the first
2144 // shuffle operand.
2145 return Last < NumElements;
2146 };
2147
2148 if (RealShuffle->getType() != ImagShuffle->getType()) {
2149 LLVM_DEBUG(dbgs() << " - Shuffle types aren't equal.\n");
2150 return nullptr;
2151 }
2152 if (!CheckDeinterleavingShuffle(RealShuffle)) {
2153 LLVM_DEBUG(dbgs() << " - RealShuffle is invalid type.\n");
2154 return nullptr;
2155 }
2156 if (!CheckDeinterleavingShuffle(ImagShuffle)) {
2157 LLVM_DEBUG(dbgs() << " - ImagShuffle is invalid type.\n");
2158 return nullptr;
2159 }
2160
2161 CompositeNode *PlaceholderNode =
2163 RealShuffle, ImagShuffle);
2164 PlaceholderNode->ReplacementNode = RealShuffle->getOperand(0);
2165 FinalInstructions.insert(RealShuffle);
2166 FinalInstructions.insert(ImagShuffle);
2167 return submitCompositeNode(PlaceholderNode);
2168}
2169
2170ComplexDeinterleavingGraph::CompositeNode *
2171ComplexDeinterleavingGraph::identifySplat(ComplexValues &Vals) {
2172 auto IsSplat = [](Value *V) -> bool {
2173 // Fixed-width vector with constants
2175 return true;
2176
2177 if (isa<ConstantInt>(V) || isa<ConstantFP>(V))
2178 return isa<VectorType>(V->getType());
2179
2180 VectorType *VTy;
2181 ArrayRef<int> Mask;
2182 // Splats are represented differently depending on whether the repeated
2183 // value is a constant or an Instruction
2184 if (auto *Const = dyn_cast<ConstantExpr>(V)) {
2185 if (Const->getOpcode() != Instruction::ShuffleVector)
2186 return false;
2187 VTy = cast<VectorType>(Const->getType());
2188 Mask = Const->getShuffleMask();
2189 } else if (auto *Shuf = dyn_cast<ShuffleVectorInst>(V)) {
2190 VTy = Shuf->getType();
2191 Mask = Shuf->getShuffleMask();
2192 } else {
2193 return false;
2194 }
2195
2196 // When the data type is <1 x Type>, it's not possible to differentiate
2197 // between the ComplexDeinterleaving::Deinterleave and
2198 // ComplexDeinterleaving::Splat operations.
2199 if (!VTy->isScalableTy() && VTy->getElementCount().getKnownMinValue() == 1)
2200 return false;
2201
2202 return all_equal(Mask) && Mask[0] == 0;
2203 };
2204
2205 // The splats must meet the following requirements:
2206 // 1. Must either be all instructions or all values.
2207 // 2. Non-constant splats must live in the same block.
2208 if (auto *FirstValAsInstruction = dyn_cast<Instruction>(Vals[0].Real)) {
2209 BasicBlock *FirstBB = FirstValAsInstruction->getParent();
2210 for (auto &V : Vals) {
2211 if (!IsSplat(V.Real) || !IsSplat(V.Imag))
2212 return nullptr;
2213
2214 auto *Real = dyn_cast<Instruction>(V.Real);
2215 auto *Imag = dyn_cast<Instruction>(V.Imag);
2216 if (!Real || !Imag || Real->getParent() != FirstBB ||
2217 Imag->getParent() != FirstBB)
2218 return nullptr;
2219 }
2220 } else {
2221 for (auto &V : Vals) {
2222 if (!IsSplat(V.Real) || !IsSplat(V.Imag) || isa<Instruction>(V.Real) ||
2223 isa<Instruction>(V.Imag))
2224 return nullptr;
2225 }
2226 }
2227
2228 for (auto &V : Vals) {
2229 auto *Real = dyn_cast<Instruction>(V.Real);
2230 auto *Imag = dyn_cast<Instruction>(V.Imag);
2231 if (Real && Imag) {
2232 FinalInstructions.insert(Real);
2233 FinalInstructions.insert(Imag);
2234 }
2235 }
2236 CompositeNode *PlaceholderNode =
2237 prepareCompositeNode(ComplexDeinterleavingOperation::Splat, Vals);
2238 return submitCompositeNode(PlaceholderNode);
2239}
2240
2241ComplexDeinterleavingGraph::CompositeNode *
2242ComplexDeinterleavingGraph::identifyPHINode(Instruction *Real,
2243 Instruction *Imag) {
2244 if (Real != RealPHI || (ImagPHI && Imag != ImagPHI))
2245 return nullptr;
2246
2247 PHIsFound = true;
2248 CompositeNode *PlaceholderNode = prepareCompositeNode(
2249 ComplexDeinterleavingOperation::ReductionPHI, Real, Imag);
2250 return submitCompositeNode(PlaceholderNode);
2251}
2252
2253ComplexDeinterleavingGraph::CompositeNode *
2254ComplexDeinterleavingGraph::identifySelectNode(Instruction *Real,
2255 Instruction *Imag) {
2256 auto *SelectReal = dyn_cast<SelectInst>(Real);
2257 auto *SelectImag = dyn_cast<SelectInst>(Imag);
2258 if (!SelectReal || !SelectImag)
2259 return nullptr;
2260
2261 Instruction *MaskA, *MaskB;
2262 Instruction *AR, *AI, *RA, *BI;
2263 if (!match(Real, m_Select(m_Instruction(MaskA), m_Instruction(AR),
2264 m_Instruction(RA))) ||
2265 !match(Imag, m_Select(m_Instruction(MaskB), m_Instruction(AI),
2266 m_Instruction(BI))))
2267 return nullptr;
2268
2269 if (MaskA != MaskB && !MaskA->isIdenticalTo(MaskB))
2270 return nullptr;
2271
2272 if (!MaskA->getType()->isVectorTy())
2273 return nullptr;
2274
2275 auto NodeA = identifyNode(AR, AI);
2276 if (!NodeA)
2277 return nullptr;
2278
2279 auto NodeB = identifyNode(RA, BI);
2280 if (!NodeB)
2281 return nullptr;
2282
2283 CompositeNode *PlaceholderNode = prepareCompositeNode(
2284 ComplexDeinterleavingOperation::ReductionSelect, Real, Imag);
2285 PlaceholderNode->addOperand(NodeA);
2286 PlaceholderNode->addOperand(NodeB);
2287 FinalInstructions.insert(MaskA);
2288 FinalInstructions.insert(MaskB);
2289 return submitCompositeNode(PlaceholderNode);
2290}
2291
2292static Value *replaceSymmetricNode(IRBuilderBase &B, unsigned Opcode,
2293 std::optional<FastMathFlags> Flags,
2294 Value *InputA, Value *InputB) {
2295 Value *I;
2296 switch (Opcode) {
2297 case Instruction::FNeg:
2298 I = B.CreateFNeg(InputA);
2299 break;
2300 case Instruction::FAdd:
2301 I = B.CreateFAdd(InputA, InputB);
2302 break;
2303 case Instruction::Add:
2304 I = B.CreateAdd(InputA, InputB);
2305 break;
2306 case Instruction::FSub:
2307 I = B.CreateFSub(InputA, InputB);
2308 break;
2309 case Instruction::Sub:
2310 I = B.CreateSub(InputA, InputB);
2311 break;
2312 case Instruction::FMul:
2313 I = B.CreateFMul(InputA, InputB);
2314 break;
2315 case Instruction::Mul:
2316 I = B.CreateMul(InputA, InputB);
2317 break;
2318 default:
2319 llvm_unreachable("Incorrect symmetric opcode");
2320 }
2321 if (Flags)
2322 cast<Instruction>(I)->setFastMathFlags(*Flags);
2323 return I;
2324}
2325
2326Value *ComplexDeinterleavingGraph::replaceNode(IRBuilderBase &Builder,
2327 CompositeNode *Node) {
2328 if (Node->ReplacementNode)
2329 return Node->ReplacementNode;
2330
2331 auto ReplaceOperandIfExist = [&](CompositeNode *Node,
2332 unsigned Idx) -> Value * {
2333 return Node->Operands.size() > Idx
2334 ? replaceNode(Builder, Node->Operands[Idx])
2335 : nullptr;
2336 };
2337
2338 Value *ReplacementNode = nullptr;
2339 switch (Node->Operation) {
2340 case ComplexDeinterleavingOperation::CDot: {
2341 Value *Input0 = ReplaceOperandIfExist(Node, 0);
2342 Value *Input1 = ReplaceOperandIfExist(Node, 1);
2343 Value *Accumulator = ReplaceOperandIfExist(Node, 2);
2344 assert(!Input1 || (Input0->getType() == Input1->getType() &&
2345 "Node inputs need to be of the same type"));
2346 ReplacementNode = TL->createComplexDeinterleavingIR(
2347 Builder, Node->Operation, Node->Rotation, Input0, Input1, Accumulator);
2348 break;
2349 }
2350 case ComplexDeinterleavingOperation::CAdd:
2351 case ComplexDeinterleavingOperation::CMulPartial:
2352 case ComplexDeinterleavingOperation::Symmetric: {
2353 Value *Input0 = ReplaceOperandIfExist(Node, 0);
2354 Value *Input1 = ReplaceOperandIfExist(Node, 1);
2355 Value *Accumulator = ReplaceOperandIfExist(Node, 2);
2356 assert(!Input1 || (Input0->getType() == Input1->getType() &&
2357 "Node inputs need to be of the same type"));
2359 (Input0->getType() == Accumulator->getType() &&
2360 "Accumulator and input need to be of the same type"));
2361 if (Node->Operation == ComplexDeinterleavingOperation::Symmetric)
2362 ReplacementNode = replaceSymmetricNode(Builder, Node->Opcode, Node->Flags,
2363 Input0, Input1);
2364 else
2365 ReplacementNode = TL->createComplexDeinterleavingIR(
2366 Builder, Node->Operation, Node->Rotation, Input0, Input1,
2367 Accumulator);
2368 break;
2369 }
2370 case ComplexDeinterleavingOperation::Deinterleave:
2371 llvm_unreachable("Deinterleave node should already have ReplacementNode");
2372 break;
2373 case ComplexDeinterleavingOperation::Splat: {
2375 for (auto &V : Node->Vals) {
2376 Ops.push_back(V.Real);
2377 Ops.push_back(V.Imag);
2378 }
2379 auto *R = dyn_cast<Instruction>(Node->Vals[0].Real);
2380 auto *I = dyn_cast<Instruction>(Node->Vals[0].Imag);
2381 if (R && I) {
2382 // Splats that are not constant are interleaved where they are located
2383 Instruction *InsertPoint = R;
2384 for (auto V : Node->Vals) {
2385 if (InsertPoint->comesBefore(cast<Instruction>(V.Real)))
2386 InsertPoint = cast<Instruction>(V.Real);
2387 if (InsertPoint->comesBefore(cast<Instruction>(V.Imag)))
2388 InsertPoint = cast<Instruction>(V.Imag);
2389 }
2390 InsertPoint = InsertPoint->getNextNode();
2391 IRBuilder<> IRB(InsertPoint);
2392 ReplacementNode = IRB.CreateVectorInterleave(Ops);
2393 } else {
2394 ReplacementNode = Builder.CreateVectorInterleave(Ops);
2395 }
2396 break;
2397 }
2398 case ComplexDeinterleavingOperation::ReductionPHI: {
2399 // If Operation is ReductionPHI, a new empty PHINode is created.
2400 // It is filled later when the ReductionOperation is processed.
2401 auto *OldPHI = cast<PHINode>(Node->Vals[0].Real);
2402 auto *VTy = cast<VectorType>(Node->Vals[0].Real->getType());
2403 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
2404 auto *NewPHI = PHINode::Create(NewVTy, 0, "", BackEdge->getFirstNonPHIIt());
2405 OldToNewPHI[OldPHI] = NewPHI;
2406 ReplacementNode = NewPHI;
2407 break;
2408 }
2409 case ComplexDeinterleavingOperation::ReductionSingle:
2410 ReplacementNode = replaceNode(Builder, Node->Operands[0]);
2411 processReductionSingle(ReplacementNode, Node);
2412 break;
2413 case ComplexDeinterleavingOperation::ReductionOperation:
2414 ReplacementNode = replaceNode(Builder, Node->Operands[0]);
2415 processReductionOperation(ReplacementNode, Node);
2416 break;
2417 case ComplexDeinterleavingOperation::ReductionSelect: {
2418 auto *MaskReal = cast<Instruction>(Node->Vals[0].Real)->getOperand(0);
2419 auto *MaskImag = cast<Instruction>(Node->Vals[0].Imag)->getOperand(0);
2420 auto *A = replaceNode(Builder, Node->Operands[0]);
2421 auto *B = replaceNode(Builder, Node->Operands[1]);
2422 auto *NewMask = Builder.CreateVectorInterleave({MaskReal, MaskImag});
2423 ReplacementNode = Builder.CreateSelect(NewMask, A, B);
2424 break;
2425 }
2426 }
2427
2428 assert(ReplacementNode && "Target failed to create Intrinsic call.");
2429 NumComplexTransformations += 1;
2430 Node->ReplacementNode = ReplacementNode;
2431 return ReplacementNode;
2432}
2433
2434void ComplexDeinterleavingGraph::processReductionSingle(
2435 Value *OperationReplacement, CompositeNode *Node) {
2436 auto *Real = cast<Instruction>(Node->Vals[0].Real);
2437 auto *OldPHI = ReductionInfo[Real].first;
2438 auto *NewPHI = OldToNewPHI[OldPHI];
2439 auto *VTy = cast<VectorType>(Real->getType());
2440 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
2441
2442 Value *Init = OldPHI->getIncomingValueForBlock(Incoming);
2443
2444 IRBuilder<> Builder(Incoming->getTerminator());
2445
2446 Value *NewInit = nullptr;
2447 if (auto *C = dyn_cast<Constant>(Init)) {
2448 if (C->isNullValue())
2449 NewInit = Constant::getNullValue(NewVTy);
2450 }
2451
2452 if (!NewInit)
2453 NewInit =
2454 Builder.CreateVectorInterleave({Init, Constant::getNullValue(VTy)});
2455
2456 NewPHI->addIncoming(NewInit, Incoming);
2457 NewPHI->addIncoming(OperationReplacement, BackEdge);
2458
2459 auto *FinalReduction = ReductionInfo[Real].second;
2460 Builder.SetInsertPoint(&*FinalReduction->getParent()->getFirstInsertionPt());
2461
2462 auto *AddReduce = Builder.CreateAddReduce(OperationReplacement);
2463 FinalReduction->replaceAllUsesWith(AddReduce);
2464}
2465
2466void ComplexDeinterleavingGraph::processReductionOperation(
2467 Value *OperationReplacement, CompositeNode *Node) {
2468 auto *Real = cast<Instruction>(Node->Vals[0].Real);
2469 auto *Imag = cast<Instruction>(Node->Vals[0].Imag);
2470 auto *OldPHIReal = ReductionInfo[Real].first;
2471 auto *OldPHIImag = ReductionInfo[Imag].first;
2472 auto *NewPHI = OldToNewPHI[OldPHIReal];
2473
2474 // We have to interleave initial origin values coming from IncomingBlock
2475 Value *InitReal = OldPHIReal->getIncomingValueForBlock(Incoming);
2476 Value *InitImag = OldPHIImag->getIncomingValueForBlock(Incoming);
2477
2478 IRBuilder<> Builder(Incoming->getTerminator());
2479 auto *NewInit = Builder.CreateVectorInterleave({InitReal, InitImag});
2480
2481 NewPHI->addIncoming(NewInit, Incoming);
2482 NewPHI->addIncoming(OperationReplacement, BackEdge);
2483
2484 // Deinterleave complex vector outside of loop so that it can be finally
2485 // reduced
2486 auto *FinalReductionReal = ReductionInfo[Real].second;
2487 auto *FinalReductionImag = ReductionInfo[Imag].second;
2488
2489 auto *Br = cast<CondBrInst>(BackEdge->getTerminator());
2490 BasicBlock *ExitBB = Br->getSuccessor(Br->getSuccessor(0) == BackEdge);
2491 Builder.SetInsertPoint(&*ExitBB->getFirstInsertionPt());
2492
2493 auto *Deinterleave = Builder.CreateIntrinsic(Intrinsic::vector_deinterleave2,
2494 OperationReplacement->getType(),
2495 OperationReplacement);
2496
2497 auto *NewReal = Builder.CreateExtractValue(Deinterleave, (uint64_t)0);
2498 FinalReductionReal->replaceUsesOfWith(Real, NewReal);
2499
2500 Builder.SetInsertPoint(FinalReductionImag);
2501 auto *NewImag = Builder.CreateExtractValue(Deinterleave, 1);
2502 FinalReductionImag->replaceUsesOfWith(Imag, NewImag);
2503}
2504
2505void ComplexDeinterleavingGraph::replaceNodes() {
2506 SmallVector<Instruction *, 16> DeadInstrRoots;
2507 for (auto *RootInstruction : OrderedRoots) {
2508 // Check if this potential root went through check process and we can
2509 // deinterleave it
2510 if (!RootToNode.count(RootInstruction))
2511 continue;
2512
2513 IRBuilder<> Builder(RootInstruction);
2514 auto RootNode = RootToNode[RootInstruction];
2515 Value *R = replaceNode(Builder, RootNode);
2516
2517 if (RootNode->Operation ==
2518 ComplexDeinterleavingOperation::ReductionOperation) {
2519 auto *RootReal = cast<Instruction>(RootNode->Vals[0].Real);
2520 auto *RootImag = cast<Instruction>(RootNode->Vals[0].Imag);
2521 ReductionInfo[RootReal].first->removeIncomingValue(BackEdge);
2522 ReductionInfo[RootImag].first->removeIncomingValue(BackEdge);
2523 DeadInstrRoots.push_back(RootReal);
2524 DeadInstrRoots.push_back(RootImag);
2525 } else if (RootNode->Operation ==
2526 ComplexDeinterleavingOperation::ReductionSingle) {
2527 auto *RootInst = cast<Instruction>(RootNode->Vals[0].Real);
2528 auto &Info = ReductionInfo[RootInst];
2529 Info.first->removeIncomingValue(BackEdge);
2530 DeadInstrRoots.push_back(Info.second);
2531 } else {
2532 assert(R && "Unable to find replacement for RootInstruction");
2533 DeadInstrRoots.push_back(RootInstruction);
2534 RootInstruction->replaceAllUsesWith(R);
2535 }
2536 }
2537
2538 for (auto *I : DeadInstrRoots)
2540}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
unsigned uint64_t
static MCDisassembler::DecodeStatus addOperand(MCInst &Inst, const MCOperand &Opnd)
Rewrite undef for PHI
This file defines the BumpPtrAllocator interface.
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< ErlangGC > A("erlang", "erlang-compatible garbage collector")
static GCRegistry::Add< OcamlGC > B("ocaml", "ocaml 3.10-compatible GC")
static bool isInstructionPotentiallySymmetric(Instruction *I)
static Value * getNegOperand(Value *V)
Returns the operand for negation operation.
static bool isNeg(Value *V)
Returns true if the operation is a negation of V, and it works for both integers and floats.
static cl::opt< bool > ComplexDeinterleavingEnabled("enable-complex-deinterleaving", cl::desc("Enable generation of complex instructions"), cl::init(true), cl::Hidden)
static bool isInstructionPairAdd(Instruction *A, Instruction *B)
static Value * replaceSymmetricNode(IRBuilderBase &B, unsigned Opcode, std::optional< FastMathFlags > Flags, Value *InputA, Value *InputB)
static bool isInterleavingMask(ArrayRef< int > Mask)
Checks the given mask, and determines whether said mask is interleaving.
static bool isDeinterleavingMask(ArrayRef< int > Mask)
Checks the given mask, and determines whether said mask is deinterleaving.
SmallVector< struct ComplexValue, 2 > ComplexValues
static bool isInstructionPairMul(Instruction *A, Instruction *B)
static bool runOnFunction(Function &F, bool PostInlining)
#define DEBUG_TYPE
const AbstractManglingParser< Derived, Alloc >::OperatorInfo AbstractManglingParser< Derived, Alloc >::Ops[]
#define F(x, y, z)
Definition MD5.cpp:54
#define I(x, y, z)
Definition MD5.cpp:57
This file implements a map that provides insertion order iteration.
#define T
uint64_t IntrinsicInst * II
#define P(N)
PowerPC Reduce CR logical Operation
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
Definition PassSupport.h:44
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
Definition PassSupport.h:39
SI optimize exec mask operations pre RA
static LLVM_ATTRIBUTE_ALWAYS_INLINE bool CheckType(MVT::SimpleValueType VT, SDValue N, const TargetLowering *TLI, const DataLayout &DL)
This file defines the 'Statistic' class, which is designed to be an easy way to expose various metric...
#define STATISTIC(VARNAME, DESC)
Definition Statistic.h:171
#define LLVM_DEBUG(...)
Definition Debug.h:119
This file describes how to lower LLVM code to machine code.
This pass exposes codegen information to IR-level passes.
BinaryOperator * Mul
AnalysisUsage & addRequired()
LLVM_ABI void setPreservesCFG()
This function should be called by the pass, iff they do not:
Definition Pass.cpp:278
Represent a constant reference to an array (0 or more elements consecutively in memory),...
Definition ArrayRef.h:40
size_t size() const
Get the array size.
Definition ArrayRef.h:141
LLVM_ABI const_iterator getFirstInsertionPt() const
Returns an iterator to the first instruction in this block that is suitable for inserting a non-PHI i...
LLVM_ABI InstListType::const_iterator getFirstNonPHIIt() const
Returns an iterator to the first instruction in this block that is not a PHINode instruction.
const Instruction * getTerminator() const LLVM_READONLY
Returns the terminator instruction; assumes that the block is well-formed.
Definition BasicBlock.h:237
static LLVM_ABI Constant * getNullValue(Type *Ty)
Constructor to create a '0' constant of arbitrary type.
iterator find(const_arg_type_t< KeyT > Val)
Definition DenseMap.h:223
iterator end()
Definition DenseMap.h:141
bool allowContract() const
Definition FMF.h:69
FunctionPass class - This class is used to implement most global optimizations.
Definition Pass.h:314
Common base class shared among various IRBuilders.
Definition IRBuilder.h:114
Value * CreateExtractValue(Value *Agg, ArrayRef< unsigned > Idxs, const Twine &Name="")
Definition IRBuilder.h:2726
LLVM_ABI Value * CreateSelect(Value *C, Value *True, Value *False, const Twine &Name="", Instruction *MDFrom=nullptr)
LLVM_ABI Value * CreateAddReduce(Value *Src)
Create a vector int add reduction intrinsic of the source vector.
LLVM_ABI Value * CreateIntrinsic(Intrinsic::ID ID, ArrayRef< Type * > OverloadTypes, ArrayRef< Value * > Args, FMFSource FMFSource={}, const Twine &Name="", ArrayRef< OperandBundleDef > OpBundles={}, function_ref< void(CallInst *)> SetFn=[](CallInst *) {})
Variant to create a possibly constant-folded intrinsic.
void SetInsertPoint(BasicBlock *TheBB)
This specifies that created instructions should be appended to the end of the specified block.
Definition IRBuilder.h:181
LLVM_ABI Value * CreateVectorInterleave(ArrayRef< Value * > Ops, const Twine &Name="")
LLVM_ABI const Function * getFunction() const
Return the function this instruction belongs to.
LLVM_ABI bool comesBefore(const Instruction *Other) const
Given an instruction Other in the same basic block as this instruction, return true if this instructi...
LLVM_ABI FastMathFlags getFastMathFlags() const LLVM_READONLY
Convenience function for getting all the fast-math flags, which must be an operator which supports th...
unsigned getOpcode() const
Returns a member of one of the enums like Instruction::Add.
LLVM_ABI bool isIdenticalTo(const Instruction *I) const LLVM_READONLY
Return true if the specified instruction is exactly identical to the current one.
size_type size() const
Definition MapVector.h:58
static PHINode * Create(Type *Ty, unsigned NumReservedValues, const Twine &NameStr="", InsertPosition InsertBefore=nullptr)
Constructors - NumReservedValues is a hint for the number of incoming edges that this phi node will h...
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
PreservedAnalyses & preserve()
Mark an analysis as preserved.
Definition Analysis.h:132
size_type count(ConstPtrType Ptr) const
count - Return 1 if the specified pointer is in the set, 0 otherwise.
std::pair< iterator, bool > insert(PtrType Ptr)
Inserts Ptr if and only if there is no element in the container equal to Ptr.
reference emplace_back(ArgTypes &&... Args)
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
Analysis pass providing the TargetLibraryInfo.
virtual bool isComplexDeinterleavingOperationSupported(ComplexDeinterleavingOperation Operation, Type *Ty) const
Does this target support complex deinterleaving with the given operation and type.
virtual Value * createComplexDeinterleavingIR(IRBuilderBase &B, ComplexDeinterleavingOperation OperationType, ComplexDeinterleavingRotation Rotation, Value *InputA, Value *InputB, Value *Accumulator=nullptr) const
Create the IR node for the given complex deinterleaving operation.
virtual bool isComplexDeinterleavingSupported() const
Does this target support complex deinterleaving.
This class defines information used to lower LLVM code to legal SelectionDAG operators that the targe...
Primary interface to the complete machine description for the target machine.
virtual const TargetSubtargetInfo * getSubtargetImpl(const Function &) const
Virtual method implemented by subclasses that returns a reference to that target's TargetSubtargetInf...
virtual const TargetLowering * getTargetLowering() const
bool isVectorTy() const
True if this is an instance of VectorType.
Definition Type.h:288
Value * getOperand(unsigned i) const
Definition User.h:207
LLVM Value Representation.
Definition Value.h:75
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:255
bool hasOneUse() const
Return true if there is exactly one use of this value.
Definition Value.h:439
LLVM_ABI void replaceAllUsesWith(Value *V)
Change all uses of this to point to a new Value.
Definition Value.cpp:553
An opaque object representing a hash code.
Definition Hashing.h:77
const ParentTy * getParent() const
Definition ilist_node.h:34
NodeTy * getNextNode()
Get the next node, or nullptr for the list tail.
Definition ilist_node.h:348
raw_ostream & indent(unsigned NumSpaces)
indent - Insert 'NumSpaces' spaces.
Changed
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
constexpr std::underlying_type_t< E > Mask()
Get a bitmask with 1s in all places up to the high-order bit of E's largest value.
@ BR
Control flow instructions. These all have token chains.
@ BasicBlock
Various leaf nodes.
Definition ISDOpcodes.h:81
LLVM_ABI Intrinsic::ID getDeinterleaveIntrinsicID(unsigned Factor)
Returns the corresponding llvm.vector.deinterleaveN intrinsic for factor N.
LLVM_ABI Intrinsic::ID getInterleaveIntrinsicID(unsigned Factor)
Returns the corresponding llvm.vector.interleaveN intrinsic for factor N.
BinaryOp_match< SpecificConstantMatch, SrcTy, TargetOpcode::G_SUB > m_Neg(const SrcTy &&Src)
Matches a register negated by a G_SUB.
BinaryOp_match< LHS, RHS, Instruction::FMul > m_FMul(const LHS &L, const RHS &R)
bool match(Val *V, const Pattern &P)
match_bind< Instruction > m_Instruction(Instruction *&I)
Match an instruction, capturing it if we match.
ThreeOps_match< Cond, LHS, RHS, Instruction::Select > m_Select(const Cond &C, const LHS &L, const RHS &R)
Matches SelectInst.
auto m_BinOp()
Match an arbitrary binary operation and ignore it.
auto m_Value()
Match an arbitrary value and ignore it.
BinaryOp_match< LHS, RHS, Instruction::Mul > m_Mul(const LHS &L, const RHS &R)
TwoOps_match< V1_t, V2_t, Instruction::ShuffleVector > m_Shuffle(const V1_t &v1, const V2_t &v2)
Matches ShuffleVectorInst independently of mask value.
auto m_AnyIntrinsic()
Matches any intrinsic call and ignore it.
auto m_Intrinsic(const Ts &...Ops)
Match intrinsic calls like this: m_Intrinsic<Intrinsic::fabs>(m_Value(X))
FNeg_match< OpTy > m_FNeg(const OpTy &X)
Match 'fneg X' as 'fsub -0.0, X'.
is_zero m_Zero()
Match any null constant or a vector with all elements equal to 0.
initializer< Ty > init(const Ty &Val)
NodeAddr< PhiNode * > Phi
Definition RDFGraph.h:390
NodeAddr< NodeBase * > Node
Definition RDFGraph.h:381
friend class Instruction
Iterator for Instructions in a `BasicBlock.
Definition BasicBlock.h:73
This is an optimization pass for GlobalISel generic memory operations.
void dump(const SparseBitVector< ElementSize > &LHS, raw_ostream &out)
@ Offset
Definition DWP.cpp:577
bool all_of(R &&range, UnaryPredicate P)
Provide wrappers to std::all_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1739
hash_code hash_value(const FixedPointSemantics &Val)
LLVM_ABI bool RecursivelyDeleteTriviallyDeadInstructions(Value *V, const TargetLibraryInfo *TLI=nullptr, MemorySSAUpdater *MSSAU=nullptr, std::function< void(Value *)> AboutToDeleteCallback=std::function< void(Value *)>())
If the specified value is a trivially dead instruction, delete it.
Definition Local.cpp:522
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
InnerAnalysisManagerProxy< FunctionAnalysisManager, Module > FunctionAnalysisManagerModuleProxy
Provide the FunctionAnalysisManager to Module proxy.
bool operator==(const AddressRangeValuePair &LHS, const AddressRangeValuePair &RHS)
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Value
Definition InstrProf.h:143
auto dyn_cast_or_null(const Y &Val)
Definition Casting.h:753
LLVM_ABI FunctionPass * createComplexDeinterleavingPass(const TargetMachine *TM)
This pass implements generation of target-specific intrinsics to support handling of complex number a...
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
Definition Debug.cpp:209
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
Definition Casting.h:547
@ Other
Any other memory.
Definition ModRef.h:68
IRBuilder(LLVMContext &, FolderTy, InserterTy, MDNode *, ArrayRef< OperandBundleDef >) -> IRBuilder< FolderTy, InserterTy >
DWARFExpression::Operation Op
ArrayRef(const T &OneElt) -> ArrayRef< T >
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
auto find_if(R &&Range, UnaryPredicate P)
Provide wrappers to std::find_if which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1772
bool is_contained(R &&Range, const E &Element)
Returns true if Element is found in Range.
Definition STLExtras.h:1947
bool all_equal(std::initializer_list< T > Values)
Returns true if all Values in the initializer lists are equal or the list.
Definition STLExtras.h:2166
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
hash_code hash_combine(const Ts &...args)
Combine values into a single hash_code.
Definition Hashing.h:305
AllocatorList< T, BumpPtrAllocator > BumpPtrList
void swap(llvm::BitVector &LHS, llvm::BitVector &RHS)
Implement std::swap in terms of BitVector swap.
Definition BitVector.h:880
#define N
ComplexDeinterleavingPass(const TargetMachine &TM)
LLVM_ABI PreservedAnalyses run(Function &F, FunctionAnalysisManager &AM)
static bool isEqual(const ComplexValue &LHS, const ComplexValue &RHS)
static unsigned getHashValue(const ComplexValue &Val)
An information struct used to provide DenseMap with the various necessary components for a given valu...