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 if (isNeg(A)) {
1319 A = getNegOperand(A);
1320 IsPositive = !IsPositive;
1321 }
1322
1323 if (isNeg(B)) {
1324 B = getNegOperand(B);
1325 IsPositive = !IsPositive;
1326 }
1327
1328 Muls.push_back(Product{A, B, IsPositive});
1329 Worklist.emplace_back(C, IsPositive);
1330 break;
1331 }
1332 default:
1333 Addends.emplace_back(I, IsPositive);
1334 continue;
1335 }
1336
1337 if (Flags && I->getFastMathFlags() != *Flags) {
1338 LLVM_DEBUG(dbgs() << "The instruction's fast math flags are "
1339 "inconsistent with the root instructions' flags: "
1340 << *I << "\n");
1341 return false;
1342 }
1343 }
1344 return true;
1345 };
1346
1347 SmallVector<Product> RealMuls, ImagMuls;
1348 AddendList RealAddends, ImagAddends;
1349 if (!Collect(Real, RealMuls, RealAddends) ||
1350 !Collect(Imag, ImagMuls, ImagAddends))
1351 return nullptr;
1352
1353 if (RealAddends.size() != ImagAddends.size())
1354 return nullptr;
1355
1356 CompositeNode *FinalNode = nullptr;
1357 if (!RealMuls.empty() || !ImagMuls.empty()) {
1358 // If there are multiplicands, extract positive addend and use it as an
1359 // accumulator
1360 FinalNode = extractPositiveAddend(RealAddends, ImagAddends);
1361 FinalNode = identifyMultiplications(RealMuls, ImagMuls, FinalNode);
1362 if (!FinalNode)
1363 return nullptr;
1364 }
1365
1366 // Identify and process remaining additions
1367 if (!RealAddends.empty() || !ImagAddends.empty()) {
1368 FinalNode = identifyAdditions(RealAddends, ImagAddends, Flags, FinalNode);
1369 if (!FinalNode)
1370 return nullptr;
1371 }
1372 assert(FinalNode && "FinalNode can not be nullptr here");
1373 assert(FinalNode->Vals.size() == 1);
1374 // Set the Real and Imag fields of the final node and submit it
1375 FinalNode->Vals[0].Real = Real;
1376 FinalNode->Vals[0].Imag = Imag;
1377 submitCompositeNode(FinalNode);
1378 return FinalNode;
1379}
1380
1381bool ComplexDeinterleavingGraph::collectPartialMuls(
1382 ArrayRef<Product> RealMuls, ArrayRef<Product> ImagMuls,
1383 SmallVectorImpl<PartialMulCandidate> &PartialMulCandidates) {
1384 // Helper function to extract a common operand from two products
1385 auto FindCommonInstruction = [](const Product &Real,
1386 const Product &Imag) -> Value * {
1387 if (Real.Multiplicand == Imag.Multiplicand ||
1388 Real.Multiplicand == Imag.Multiplier)
1389 return Real.Multiplicand;
1390
1391 if (Real.Multiplier == Imag.Multiplicand ||
1392 Real.Multiplier == Imag.Multiplier)
1393 return Real.Multiplier;
1394
1395 return nullptr;
1396 };
1397
1398 // Iterating over real and imaginary multiplications to find common operands
1399 // If a common operand is found, a partial multiplication candidate is created
1400 // and added to the candidates vector The function returns false if no common
1401 // operands are found for any product
1402 for (unsigned i = 0; i < RealMuls.size(); ++i) {
1403 bool FoundCommon = false;
1404 for (unsigned j = 0; j < ImagMuls.size(); ++j) {
1405 auto *Common = FindCommonInstruction(RealMuls[i], ImagMuls[j]);
1406 if (!Common)
1407 continue;
1408
1409 auto *A = RealMuls[i].Multiplicand == Common ? RealMuls[i].Multiplier
1410 : RealMuls[i].Multiplicand;
1411 auto *B = ImagMuls[j].Multiplicand == Common ? ImagMuls[j].Multiplier
1412 : ImagMuls[j].Multiplicand;
1413
1414 auto Node = identifyNode(A, B);
1415 if (Node) {
1416 FoundCommon = true;
1417 PartialMulCandidates.push_back({Common, Node, i, j, false});
1418 }
1419
1420 Node = identifyNode(B, A);
1421 if (Node) {
1422 FoundCommon = true;
1423 PartialMulCandidates.push_back({Common, Node, i, j, true});
1424 }
1425 }
1426 if (!FoundCommon)
1427 return false;
1428 }
1429 return true;
1430}
1431
1432ComplexDeinterleavingGraph::CompositeNode *
1433ComplexDeinterleavingGraph::identifyMultiplications(
1434 SmallVectorImpl<Product> &RealMuls, SmallVectorImpl<Product> &ImagMuls,
1435 CompositeNode *Accumulator = nullptr) {
1436 if (RealMuls.size() != ImagMuls.size())
1437 return nullptr;
1438
1440 if (!collectPartialMuls(RealMuls, ImagMuls, Info))
1441 return nullptr;
1442
1443 // Map to store common instruction to node pointers
1444 DenseMap<Value *, CompositeNode *> CommonToNode;
1445 SmallVector<bool> Processed(Info.size(), false);
1446 for (unsigned I = 0; I < Info.size(); ++I) {
1447 if (Processed[I])
1448 continue;
1449
1450 PartialMulCandidate &InfoA = Info[I];
1451 for (unsigned J = I + 1; J < Info.size(); ++J) {
1452 if (Processed[J])
1453 continue;
1454
1455 PartialMulCandidate &InfoB = Info[J];
1456 auto *InfoReal = &InfoA;
1457 auto *InfoImag = &InfoB;
1458
1459 auto NodeFromCommon = identifyNode(InfoReal->Common, InfoImag->Common);
1460 if (!NodeFromCommon) {
1461 std::swap(InfoReal, InfoImag);
1462 NodeFromCommon = identifyNode(InfoReal->Common, InfoImag->Common);
1463 }
1464 if (!NodeFromCommon)
1465 continue;
1466
1467 CommonToNode[InfoReal->Common] = NodeFromCommon;
1468 CommonToNode[InfoImag->Common] = NodeFromCommon;
1469 Processed[I] = true;
1470 Processed[J] = true;
1471 }
1472 }
1473
1474 SmallVector<bool> ProcessedReal(RealMuls.size(), false);
1475 SmallVector<bool> ProcessedImag(ImagMuls.size(), false);
1476 CompositeNode *Result = Accumulator;
1477 for (auto &PMI : Info) {
1478 if (ProcessedReal[PMI.RealIdx] || ProcessedImag[PMI.ImagIdx])
1479 continue;
1480
1481 auto It = CommonToNode.find(PMI.Common);
1482 // TODO: Process independent complex multiplications. Cases like this:
1483 // A.real() * B where both A and B are complex numbers.
1484 if (It == CommonToNode.end()) {
1485 LLVM_DEBUG({
1486 dbgs() << "Unprocessed independent partial multiplication:\n";
1487 for (auto *Mul : {&RealMuls[PMI.RealIdx], &RealMuls[PMI.RealIdx]})
1488 dbgs().indent(4) << (Mul->IsPositive ? "+" : "-") << *Mul->Multiplier
1489 << " multiplied by " << *Mul->Multiplicand << "\n";
1490 });
1491 return nullptr;
1492 }
1493
1494 auto &RealMul = RealMuls[PMI.RealIdx];
1495 auto &ImagMul = ImagMuls[PMI.ImagIdx];
1496
1497 auto NodeA = It->second;
1498 auto NodeB = PMI.Node;
1499 auto IsMultiplicandReal = PMI.Common == NodeA->Vals[0].Real;
1500 // The following table illustrates the relationship between multiplications
1501 // and rotations. If we consider the multiplication (X + iY) * (U + iV), we
1502 // can see:
1503 //
1504 // Rotation | Real | Imag |
1505 // ---------+--------+--------+
1506 // 0 | x * u | x * v |
1507 // 90 | -y * v | y * u |
1508 // 180 | -x * u | -x * v |
1509 // 270 | y * v | -y * u |
1510 //
1511 // Check if the candidate can indeed be represented by partial
1512 // multiplication
1513 // TODO: Add support for multiplication by complex one
1514 if ((IsMultiplicandReal && PMI.IsNodeInverted) ||
1515 (!IsMultiplicandReal && !PMI.IsNodeInverted))
1516 continue;
1517
1518 // Determine the rotation based on the multiplications
1520 if (IsMultiplicandReal) {
1521 // Detect 0 and 180 degrees rotation
1522 if (RealMul.IsPositive && ImagMul.IsPositive)
1524 else if (!RealMul.IsPositive && !ImagMul.IsPositive)
1526 else
1527 continue;
1528
1529 } else {
1530 // Detect 90 and 270 degrees rotation
1531 if (!RealMul.IsPositive && ImagMul.IsPositive)
1533 else if (RealMul.IsPositive && !ImagMul.IsPositive)
1535 else
1536 continue;
1537 }
1538
1539 LLVM_DEBUG({
1540 dbgs() << "Identified partial multiplication (X, Y) * (U, V):\n";
1541 dbgs().indent(4) << "X: " << *NodeA->Vals[0].Real << "\n";
1542 dbgs().indent(4) << "Y: " << *NodeA->Vals[0].Imag << "\n";
1543 dbgs().indent(4) << "U: " << *NodeB->Vals[0].Real << "\n";
1544 dbgs().indent(4) << "V: " << *NodeB->Vals[0].Imag << "\n";
1545 dbgs().indent(4) << "Rotation - " << (int)Rotation * 90 << "\n";
1546 });
1547
1548 CompositeNode *NodeMul = prepareCompositeNode(
1549 ComplexDeinterleavingOperation::CMulPartial, nullptr, nullptr);
1550 NodeMul->Rotation = Rotation;
1551 NodeMul->addOperand(NodeA);
1552 NodeMul->addOperand(NodeB);
1553 if (Result)
1554 NodeMul->addOperand(Result);
1555 submitCompositeNode(NodeMul);
1556 Result = NodeMul;
1557 ProcessedReal[PMI.RealIdx] = true;
1558 ProcessedImag[PMI.ImagIdx] = true;
1559 }
1560
1561 // Ensure all products have been processed, if not return nullptr.
1562 if (!all_of(ProcessedReal, [](bool V) { return V; }) ||
1563 !all_of(ProcessedImag, [](bool V) { return V; })) {
1564
1565 // Dump debug information about which partial multiplications are not
1566 // processed.
1567 LLVM_DEBUG({
1568 dbgs() << "Unprocessed products (Real):\n";
1569 for (size_t i = 0; i < ProcessedReal.size(); ++i) {
1570 if (!ProcessedReal[i])
1571 dbgs().indent(4) << (RealMuls[i].IsPositive ? "+" : "-")
1572 << *RealMuls[i].Multiplier << " multiplied by "
1573 << *RealMuls[i].Multiplicand << "\n";
1574 }
1575 dbgs() << "Unprocessed products (Imag):\n";
1576 for (size_t i = 0; i < ProcessedImag.size(); ++i) {
1577 if (!ProcessedImag[i])
1578 dbgs().indent(4) << (ImagMuls[i].IsPositive ? "+" : "-")
1579 << *ImagMuls[i].Multiplier << " multiplied by "
1580 << *ImagMuls[i].Multiplicand << "\n";
1581 }
1582 });
1583 return nullptr;
1584 }
1585
1586 return Result;
1587}
1588
1589ComplexDeinterleavingGraph::CompositeNode *
1590ComplexDeinterleavingGraph::identifyAdditions(
1591 AddendList &RealAddends, AddendList &ImagAddends,
1592 std::optional<FastMathFlags> Flags, CompositeNode *Accumulator = nullptr) {
1593 if (RealAddends.size() != ImagAddends.size())
1594 return nullptr;
1595
1596 CompositeNode *Result = nullptr;
1597 // If we have accumulator use it as first addend
1598 if (Accumulator)
1600 // Otherwise find an element with both positive real and imaginary parts.
1601 else
1602 Result = extractPositiveAddend(RealAddends, ImagAddends);
1603
1604 if (!Result)
1605 return nullptr;
1606
1607 while (!RealAddends.empty()) {
1608 auto ItR = RealAddends.begin();
1609 auto [R, IsPositiveR] = *ItR;
1610
1611 bool FoundImag = false;
1612 for (auto ItI = ImagAddends.begin(); ItI != ImagAddends.end(); ++ItI) {
1613 auto [I, IsPositiveI] = *ItI;
1615 if (IsPositiveR && IsPositiveI)
1616 Rotation = ComplexDeinterleavingRotation::Rotation_0;
1617 else if (!IsPositiveR && IsPositiveI)
1618 Rotation = ComplexDeinterleavingRotation::Rotation_90;
1619 else if (!IsPositiveR && !IsPositiveI)
1620 Rotation = ComplexDeinterleavingRotation::Rotation_180;
1621 else
1622 Rotation = ComplexDeinterleavingRotation::Rotation_270;
1623
1624 CompositeNode *AddNode = nullptr;
1625 if (Rotation == ComplexDeinterleavingRotation::Rotation_0 ||
1626 Rotation == ComplexDeinterleavingRotation::Rotation_180) {
1627 AddNode = identifyNode(R, I);
1628 } else {
1629 AddNode = identifyNode(I, R);
1630 }
1631 if (AddNode) {
1632 LLVM_DEBUG({
1633 dbgs() << "Identified addition:\n";
1634 dbgs().indent(4) << "X: " << *R << "\n";
1635 dbgs().indent(4) << "Y: " << *I << "\n";
1636 dbgs().indent(4) << "Rotation - " << (int)Rotation * 90 << "\n";
1637 });
1638
1639 CompositeNode *TmpNode = nullptr;
1641 TmpNode = prepareCompositeNode(
1642 ComplexDeinterleavingOperation::Symmetric, nullptr, nullptr);
1643 if (Flags) {
1644 TmpNode->Opcode = Instruction::FAdd;
1645 TmpNode->Flags = *Flags;
1646 } else {
1647 TmpNode->Opcode = Instruction::Add;
1648 }
1649 } else if (Rotation ==
1651 TmpNode = prepareCompositeNode(
1652 ComplexDeinterleavingOperation::Symmetric, nullptr, nullptr);
1653 if (Flags) {
1654 TmpNode->Opcode = Instruction::FSub;
1655 TmpNode->Flags = *Flags;
1656 } else {
1657 TmpNode->Opcode = Instruction::Sub;
1658 }
1659 } else {
1660 TmpNode = prepareCompositeNode(ComplexDeinterleavingOperation::CAdd,
1661 nullptr, nullptr);
1662 TmpNode->Rotation = Rotation;
1663 }
1664
1665 TmpNode->addOperand(Result);
1666 TmpNode->addOperand(AddNode);
1667 submitCompositeNode(TmpNode);
1668 Result = TmpNode;
1669 RealAddends.erase(ItR);
1670 ImagAddends.erase(ItI);
1671 FoundImag = true;
1672 break;
1673 }
1674 }
1675 if (!FoundImag)
1676 return nullptr;
1677 }
1678 return Result;
1679}
1680
1681ComplexDeinterleavingGraph::CompositeNode *
1682ComplexDeinterleavingGraph::extractPositiveAddend(AddendList &RealAddends,
1683 AddendList &ImagAddends) {
1684 for (auto ItR = RealAddends.begin(); ItR != RealAddends.end(); ++ItR) {
1685 for (auto ItI = ImagAddends.begin(); ItI != ImagAddends.end(); ++ItI) {
1686 auto [R, IsPositiveR] = *ItR;
1687 auto [I, IsPositiveI] = *ItI;
1688 if (IsPositiveR && IsPositiveI) {
1689 auto Result = identifyNode(R, I);
1690 if (Result) {
1691 RealAddends.erase(ItR);
1692 ImagAddends.erase(ItI);
1693 return Result;
1694 }
1695 }
1696 }
1697 }
1698 return nullptr;
1699}
1700
1701bool ComplexDeinterleavingGraph::identifyNodes(Instruction *RootI) {
1702 // This potential root instruction might already have been recognized as
1703 // reduction. Because RootToNode maps both Real and Imaginary parts to
1704 // CompositeNode we should choose only one either Real or Imag instruction to
1705 // use as an anchor for generating complex instruction.
1706 auto It = RootToNode.find(RootI);
1707 if (It != RootToNode.end()) {
1708 auto RootNode = It->second;
1709 assert(RootNode->Operation ==
1710 ComplexDeinterleavingOperation::ReductionOperation ||
1711 RootNode->Operation ==
1712 ComplexDeinterleavingOperation::ReductionSingle);
1713 assert(RootNode->Vals.size() == 1 &&
1714 "Cannot handle reductions involving multiple complex values");
1715 // Find out which part, Real or Imag, comes later, and only if we come to
1716 // the latest part, add it to OrderedRoots.
1717 auto *R = cast<Instruction>(RootNode->Vals[0].Real);
1718 auto *I = RootNode->Vals[0].Imag ? cast<Instruction>(RootNode->Vals[0].Imag)
1719 : nullptr;
1720
1721 Instruction *ReplacementAnchor;
1722 if (I)
1723 ReplacementAnchor = R->comesBefore(I) ? I : R;
1724 else
1725 ReplacementAnchor = R;
1726
1727 if (ReplacementAnchor != RootI)
1728 return false;
1729 OrderedRoots.push_back(RootI);
1730 return true;
1731 }
1732
1733 auto RootNode = identifyRoot(RootI);
1734 if (!RootNode)
1735 return false;
1736
1737 LLVM_DEBUG({
1738 Function *F = RootI->getFunction();
1739 BasicBlock *B = RootI->getParent();
1740 dbgs() << "Complex deinterleaving graph for " << F->getName()
1741 << "::" << B->getName() << ".\n";
1742 dump(dbgs());
1743 dbgs() << "\n";
1744 });
1745 RootToNode[RootI] = RootNode;
1746 OrderedRoots.push_back(RootI);
1747 return true;
1748}
1749
1750bool ComplexDeinterleavingGraph::collectPotentialReductions(BasicBlock *B) {
1751 bool FoundPotentialReduction = false;
1752 if (Factor != 2)
1753 return false;
1754
1755 auto *Br = dyn_cast<CondBrInst>(B->getTerminator());
1756 if (!Br)
1757 return false;
1758
1759 // Identify simple one-block loop
1760 if (Br->getSuccessor(0) != B && Br->getSuccessor(1) != B)
1761 return false;
1762
1763 for (auto &PHI : B->phis()) {
1764 if (PHI.getNumIncomingValues() != 2)
1765 continue;
1766
1767 if (!PHI.getType()->isVectorTy())
1768 continue;
1769
1770 auto *ReductionOp = dyn_cast<Instruction>(PHI.getIncomingValueForBlock(B));
1771 if (!ReductionOp)
1772 continue;
1773
1774 // Check if final instruction is reduced outside of current block
1775 Instruction *FinalReduction = nullptr;
1776 auto NumUsers = 0u;
1777 for (auto *U : ReductionOp->users()) {
1778 ++NumUsers;
1779 if (U == &PHI)
1780 continue;
1781 FinalReduction = dyn_cast<Instruction>(U);
1782 }
1783
1784 if (NumUsers != 2 || !FinalReduction || FinalReduction->getParent() == B ||
1785 isa<PHINode>(FinalReduction))
1786 continue;
1787
1788 ReductionInfo[ReductionOp] = {&PHI, FinalReduction};
1789 BackEdge = B;
1790 auto BackEdgeIdx = PHI.getBasicBlockIndex(B);
1791 auto IncomingIdx = BackEdgeIdx == 0 ? 1 : 0;
1792 Incoming = PHI.getIncomingBlock(IncomingIdx);
1793 FoundPotentialReduction = true;
1794
1795 // If the initial value of PHINode is an Instruction, consider it a leaf
1796 // value of a complex deinterleaving graph.
1797 if (auto *InitPHI =
1798 dyn_cast<Instruction>(PHI.getIncomingValueForBlock(Incoming)))
1799 FinalInstructions.insert(InitPHI);
1800 }
1801 return FoundPotentialReduction;
1802}
1803
1804void ComplexDeinterleavingGraph::identifyReductionNodes() {
1805 assert(Factor == 2 && "Cannot handle multiple complex values");
1806
1807 SmallVector<bool> Processed(ReductionInfo.size(), false);
1808 SmallVector<Instruction *> OperationInstruction;
1809 for (auto &P : ReductionInfo)
1810 OperationInstruction.push_back(P.first);
1811
1812 // Identify a complex computation by evaluating two reduction operations that
1813 // potentially could be involved
1814 for (size_t i = 0; i < OperationInstruction.size(); ++i) {
1815 if (Processed[i])
1816 continue;
1817 for (size_t j = i + 1; j < OperationInstruction.size(); ++j) {
1818 if (Processed[j])
1819 continue;
1820 auto *Real = OperationInstruction[i];
1821 auto *Imag = OperationInstruction[j];
1822 if (Real->getType() != Imag->getType())
1823 continue;
1824
1825 RealPHI = ReductionInfo[Real].first;
1826 ImagPHI = ReductionInfo[Imag].first;
1827 PHIsFound = false;
1828 auto Node = identifyNode(Real, Imag);
1829 if (!Node) {
1830 std::swap(Real, Imag);
1831 std::swap(RealPHI, ImagPHI);
1832 Node = identifyNode(Real, Imag);
1833 }
1834
1835 // If a node is identified and reduction PHINode is used in the chain of
1836 // operations, mark its operation instructions as used to prevent
1837 // re-identification and attach the node to the real part
1838 if (Node && PHIsFound) {
1839 LLVM_DEBUG(dbgs() << "Identified reduction starting from instructions: "
1840 << *Real << " / " << *Imag << "\n");
1841 Processed[i] = true;
1842 Processed[j] = true;
1843 auto RootNode = prepareCompositeNode(
1844 ComplexDeinterleavingOperation::ReductionOperation, Real, Imag);
1845 RootNode->addOperand(Node);
1846 RootToNode[Real] = RootNode;
1847 RootToNode[Imag] = RootNode;
1848 submitCompositeNode(RootNode);
1849 break;
1850 }
1851 }
1852
1853 auto *Real = OperationInstruction[i];
1854 // We want to check that we have 2 operands, but the function attributes
1855 // being counted as operands bloats this value.
1856 if (Processed[i] || Real->getNumOperands() < 2)
1857 continue;
1858
1859 // Can only combined integer reductions at the moment.
1860 if (!ReductionInfo[Real].second->getType()->isIntegerTy())
1861 continue;
1862
1863 RealPHI = ReductionInfo[Real].first;
1864 ImagPHI = nullptr;
1865 PHIsFound = false;
1866 auto Node = identifyNode(Real->getOperand(0), Real->getOperand(1));
1867 if (Node && PHIsFound) {
1868 LLVM_DEBUG(
1869 dbgs() << "Identified single reduction starting from instruction: "
1870 << *Real << "/" << *ReductionInfo[Real].second << "\n");
1871
1872 // Reducing to a single vector is not supported, only permit reducing down
1873 // to scalar values.
1874 // Doing this here will leave the prior node in the graph,
1875 // however with no uses the node will be unreachable by the replacement
1876 // process. That along with the usage outside the graph should prevent the
1877 // replacement process from kicking off at all for this graph.
1878 // TODO Add support for reducing to a single vector value
1879 if (ReductionInfo[Real].second->getType()->isVectorTy())
1880 continue;
1881
1882 Processed[i] = true;
1883 auto RootNode = prepareCompositeNode(
1884 ComplexDeinterleavingOperation::ReductionSingle, Real, nullptr);
1885 RootNode->addOperand(Node);
1886 RootToNode[Real] = RootNode;
1887 submitCompositeNode(RootNode);
1888 }
1889 }
1890
1891 RealPHI = nullptr;
1892 ImagPHI = nullptr;
1893}
1894
1895bool ComplexDeinterleavingGraph::checkNodes() {
1896 bool FoundDeinterleaveNode = false;
1897 for (CompositeNode *N : CompositeNodes) {
1898 if (!N->areOperandsValid())
1899 return false;
1900
1901 if (N->Operation == ComplexDeinterleavingOperation::Deinterleave)
1902 FoundDeinterleaveNode = true;
1903 }
1904
1905 // We need a deinterleave node in order to guarantee that we're working with
1906 // complex numbers.
1907 if (!FoundDeinterleaveNode) {
1908 LLVM_DEBUG(
1909 dbgs() << "Couldn't find a deinterleave node within the graph, cannot "
1910 "guarantee safety during graph transformation.\n");
1911 return false;
1912 }
1913
1914 // Collect all instructions from roots to leaves
1915 SmallPtrSet<Instruction *, 16> AllInstructions;
1916 SmallVector<Instruction *, 8> Worklist;
1917 for (auto &Pair : RootToNode)
1918 Worklist.push_back(Pair.first);
1919
1920 // Extract all instructions that are used by all XCMLA/XCADD/ADD/SUB/NEG
1921 // chains
1922 while (!Worklist.empty()) {
1923 auto *I = Worklist.pop_back_val();
1924
1925 if (!AllInstructions.insert(I).second)
1926 continue;
1927
1928 for (Value *Op : I->operands()) {
1929 if (auto *OpI = dyn_cast<Instruction>(Op)) {
1930 if (!FinalInstructions.count(I))
1931 Worklist.emplace_back(OpI);
1932 }
1933 }
1934 }
1935
1936 // Find instructions that have users outside of chain
1937 for (auto *I : AllInstructions) {
1938 // Skip root nodes
1939 if (RootToNode.count(I))
1940 continue;
1941
1942 for (User *U : I->users()) {
1943 if (AllInstructions.count(cast<Instruction>(U)))
1944 continue;
1945
1946 // Found an instruction that is not used by XCMLA/XCADD chain
1947 Worklist.emplace_back(I);
1948 break;
1949 }
1950 }
1951
1952 // If any instructions are found to be used outside, find and remove roots
1953 // that somehow connect to those instructions.
1954 SmallPtrSet<Instruction *, 16> Visited;
1955 while (!Worklist.empty()) {
1956 auto *I = Worklist.pop_back_val();
1957 if (!Visited.insert(I).second)
1958 continue;
1959
1960 // Found an impacted root node. Removing it from the nodes to be
1961 // deinterleaved
1962 if (RootToNode.count(I)) {
1963 LLVM_DEBUG(dbgs() << "Instruction " << *I
1964 << " could be deinterleaved but its chain of complex "
1965 "operations have an outside user\n");
1966 RootToNode.erase(I);
1967 }
1968
1969 if (!AllInstructions.count(I) || FinalInstructions.count(I))
1970 continue;
1971
1972 for (User *U : I->users())
1973 Worklist.emplace_back(cast<Instruction>(U));
1974
1975 for (Value *Op : I->operands()) {
1976 if (auto *OpI = dyn_cast<Instruction>(Op))
1977 Worklist.emplace_back(OpI);
1978 }
1979 }
1980 return !RootToNode.empty();
1981}
1982
1983ComplexDeinterleavingGraph::CompositeNode *
1984ComplexDeinterleavingGraph::identifyRoot(Instruction *RootI) {
1985 if (auto *Intrinsic = dyn_cast<IntrinsicInst>(RootI)) {
1987 Intrinsic->getIntrinsicID())
1988 return nullptr;
1989
1990 ComplexValues Vals;
1991 for (unsigned I = 0; I < Factor; I += 2) {
1992 auto *Real = dyn_cast<Instruction>(Intrinsic->getOperand(I));
1993 auto *Imag = dyn_cast<Instruction>(Intrinsic->getOperand(I + 1));
1994 if (!Real || !Imag)
1995 return nullptr;
1996 Vals.push_back({Real, Imag});
1997 }
1998
1999 ComplexDeinterleavingGraph::CompositeNode *Node1 = identifyNode(Vals);
2000 if (!Node1)
2001 return nullptr;
2002 return Node1;
2003 }
2004
2005 // TODO: We could also add support for fixed-width interleave factors of 4
2006 // and above, but currently for symmetric operations the interleaves and
2007 // deinterleaves are already removed by VectorCombine. If we extend this to
2008 // permit complex multiplications, reductions, etc. then we should also add
2009 // support for fixed-width here.
2010 if (Factor != 2)
2011 return nullptr;
2012
2013 auto *SVI = dyn_cast<ShuffleVectorInst>(RootI);
2014 if (!SVI)
2015 return nullptr;
2016
2017 // Look for a shufflevector that takes separate vectors of the real and
2018 // imaginary components and recombines them into a single vector.
2019 if (!isInterleavingMask(SVI->getShuffleMask()))
2020 return nullptr;
2021
2022 Instruction *Real;
2023 Instruction *Imag;
2024 if (!match(RootI, m_Shuffle(m_Instruction(Real), m_Instruction(Imag))))
2025 return nullptr;
2026
2027 return identifyNode(Real, Imag);
2028}
2029
2030ComplexDeinterleavingGraph::CompositeNode *
2031ComplexDeinterleavingGraph::identifyDeinterleave(ComplexValues &Vals) {
2032 Instruction *II = nullptr;
2033
2034 // Must be at least one complex value.
2035 auto CheckExtract = [&](Value *V, unsigned ExpectedIdx,
2036 Instruction *ExpectedInsn) -> ExtractValueInst * {
2037 auto *EVI = dyn_cast<ExtractValueInst>(V);
2038 if (!EVI || EVI->getNumIndices() != 1 ||
2039 EVI->getIndices()[0] != ExpectedIdx ||
2040 !isa<Instruction>(EVI->getAggregateOperand()) ||
2041 (ExpectedInsn && ExpectedInsn != EVI->getAggregateOperand()))
2042 return nullptr;
2043 return EVI;
2044 };
2045
2046 for (unsigned Idx = 0; Idx < Vals.size(); Idx++) {
2047 ExtractValueInst *RealEVI = CheckExtract(Vals[Idx].Real, Idx * 2, II);
2048 if (RealEVI && Idx == 0)
2050 if (!RealEVI || !CheckExtract(Vals[Idx].Imag, (Idx * 2) + 1, II)) {
2051 II = nullptr;
2052 break;
2053 }
2054 }
2055
2056 if (auto *IntrinsicII = dyn_cast_or_null<IntrinsicInst>(II)) {
2057 if (IntrinsicII->getIntrinsicID() !=
2059 return nullptr;
2060
2061 // The remaining should match too.
2062 CompositeNode *PlaceholderNode = prepareCompositeNode(
2064 PlaceholderNode->ReplacementNode = II->getOperand(0);
2065 for (auto &V : Vals) {
2066 FinalInstructions.insert(cast<Instruction>(V.Real));
2067 FinalInstructions.insert(cast<Instruction>(V.Imag));
2068 }
2069 return submitCompositeNode(PlaceholderNode);
2070 }
2071
2072 if (Vals.size() != 1)
2073 return nullptr;
2074
2075 Value *Real = Vals[0].Real;
2076 Value *Imag = Vals[0].Imag;
2077 auto *RealShuffle = dyn_cast<ShuffleVectorInst>(Real);
2078 auto *ImagShuffle = dyn_cast<ShuffleVectorInst>(Imag);
2079 if (!RealShuffle || !ImagShuffle) {
2080 if (RealShuffle || ImagShuffle)
2081 LLVM_DEBUG(dbgs() << " - There's a shuffle where there shouldn't be.\n");
2082 return nullptr;
2083 }
2084
2085 Value *RealOp1 = RealShuffle->getOperand(1);
2086 if (!isa<UndefValue>(RealOp1) && !match(RealOp1, m_Zero())) {
2087 LLVM_DEBUG(dbgs() << " - RealOp1 is not undef or zero.\n");
2088 return nullptr;
2089 }
2090 Value *ImagOp1 = ImagShuffle->getOperand(1);
2091 if (!isa<UndefValue>(ImagOp1) && !match(ImagOp1, m_Zero())) {
2092 LLVM_DEBUG(dbgs() << " - ImagOp1 is not undef or zero.\n");
2093 return nullptr;
2094 }
2095
2096 Value *RealOp0 = RealShuffle->getOperand(0);
2097 Value *ImagOp0 = ImagShuffle->getOperand(0);
2098
2099 if (RealOp0 != ImagOp0) {
2100 LLVM_DEBUG(dbgs() << " - Shuffle operands are not equal.\n");
2101 return nullptr;
2102 }
2103
2104 ArrayRef<int> RealMask = RealShuffle->getShuffleMask();
2105 ArrayRef<int> ImagMask = ImagShuffle->getShuffleMask();
2106 if (!isDeinterleavingMask(RealMask) || !isDeinterleavingMask(ImagMask)) {
2107 LLVM_DEBUG(dbgs() << " - Masks are not deinterleaving.\n");
2108 return nullptr;
2109 }
2110
2111 if (RealMask[0] != 0 || ImagMask[0] != 1) {
2112 LLVM_DEBUG(dbgs() << " - Masks do not have the correct initial value.\n");
2113 return nullptr;
2114 }
2115
2116 // Type checking, the shuffle type should be a vector type of the same
2117 // scalar type, but half the size
2118 auto CheckType = [&](ShuffleVectorInst *Shuffle) {
2119 Value *Op = Shuffle->getOperand(0);
2120 auto *ShuffleTy = cast<FixedVectorType>(Shuffle->getType());
2121 auto *OpTy = cast<FixedVectorType>(Op->getType());
2122
2123 if (OpTy->getScalarType() != ShuffleTy->getScalarType())
2124 return false;
2125 if ((ShuffleTy->getNumElements() * 2) != OpTy->getNumElements())
2126 return false;
2127
2128 return true;
2129 };
2130
2131 auto CheckDeinterleavingShuffle = [&](ShuffleVectorInst *Shuffle) -> bool {
2132 if (!CheckType(Shuffle))
2133 return false;
2134
2135 ArrayRef<int> Mask = Shuffle->getShuffleMask();
2136 int Last = *Mask.rbegin();
2137
2138 Value *Op = Shuffle->getOperand(0);
2139 auto *OpTy = cast<FixedVectorType>(Op->getType());
2140 int NumElements = OpTy->getNumElements();
2141
2142 // Ensure that the deinterleaving shuffle only pulls from the first
2143 // shuffle operand.
2144 return Last < NumElements;
2145 };
2146
2147 if (RealShuffle->getType() != ImagShuffle->getType()) {
2148 LLVM_DEBUG(dbgs() << " - Shuffle types aren't equal.\n");
2149 return nullptr;
2150 }
2151 if (!CheckDeinterleavingShuffle(RealShuffle)) {
2152 LLVM_DEBUG(dbgs() << " - RealShuffle is invalid type.\n");
2153 return nullptr;
2154 }
2155 if (!CheckDeinterleavingShuffle(ImagShuffle)) {
2156 LLVM_DEBUG(dbgs() << " - ImagShuffle is invalid type.\n");
2157 return nullptr;
2158 }
2159
2160 CompositeNode *PlaceholderNode =
2162 RealShuffle, ImagShuffle);
2163 PlaceholderNode->ReplacementNode = RealShuffle->getOperand(0);
2164 FinalInstructions.insert(RealShuffle);
2165 FinalInstructions.insert(ImagShuffle);
2166 return submitCompositeNode(PlaceholderNode);
2167}
2168
2169ComplexDeinterleavingGraph::CompositeNode *
2170ComplexDeinterleavingGraph::identifySplat(ComplexValues &Vals) {
2171 auto IsSplat = [](Value *V) -> bool {
2172 // Fixed-width vector with constants
2174 return true;
2175
2176 if (isa<ConstantInt>(V) || isa<ConstantFP>(V))
2177 return isa<VectorType>(V->getType());
2178
2179 VectorType *VTy;
2180 ArrayRef<int> Mask;
2181 // Splats are represented differently depending on whether the repeated
2182 // value is a constant or an Instruction
2183 if (auto *Const = dyn_cast<ConstantExpr>(V)) {
2184 if (Const->getOpcode() != Instruction::ShuffleVector)
2185 return false;
2186 VTy = cast<VectorType>(Const->getType());
2187 Mask = Const->getShuffleMask();
2188 } else if (auto *Shuf = dyn_cast<ShuffleVectorInst>(V)) {
2189 VTy = Shuf->getType();
2190 Mask = Shuf->getShuffleMask();
2191 } else {
2192 return false;
2193 }
2194
2195 // When the data type is <1 x Type>, it's not possible to differentiate
2196 // between the ComplexDeinterleaving::Deinterleave and
2197 // ComplexDeinterleaving::Splat operations.
2198 if (!VTy->isScalableTy() && VTy->getElementCount().getKnownMinValue() == 1)
2199 return false;
2200
2201 return all_equal(Mask) && Mask[0] == 0;
2202 };
2203
2204 // The splats must meet the following requirements:
2205 // 1. Must either be all instructions or all values.
2206 // 2. Non-constant splats must live in the same block.
2207 if (auto *FirstValAsInstruction = dyn_cast<Instruction>(Vals[0].Real)) {
2208 BasicBlock *FirstBB = FirstValAsInstruction->getParent();
2209 for (auto &V : Vals) {
2210 if (!IsSplat(V.Real) || !IsSplat(V.Imag))
2211 return nullptr;
2212
2213 auto *Real = dyn_cast<Instruction>(V.Real);
2214 auto *Imag = dyn_cast<Instruction>(V.Imag);
2215 if (!Real || !Imag || Real->getParent() != FirstBB ||
2216 Imag->getParent() != FirstBB)
2217 return nullptr;
2218 }
2219 } else {
2220 for (auto &V : Vals) {
2221 if (!IsSplat(V.Real) || !IsSplat(V.Imag) || isa<Instruction>(V.Real) ||
2222 isa<Instruction>(V.Imag))
2223 return nullptr;
2224 }
2225 }
2226
2227 for (auto &V : Vals) {
2228 auto *Real = dyn_cast<Instruction>(V.Real);
2229 auto *Imag = dyn_cast<Instruction>(V.Imag);
2230 if (Real && Imag) {
2231 FinalInstructions.insert(Real);
2232 FinalInstructions.insert(Imag);
2233 }
2234 }
2235 CompositeNode *PlaceholderNode =
2236 prepareCompositeNode(ComplexDeinterleavingOperation::Splat, Vals);
2237 return submitCompositeNode(PlaceholderNode);
2238}
2239
2240ComplexDeinterleavingGraph::CompositeNode *
2241ComplexDeinterleavingGraph::identifyPHINode(Instruction *Real,
2242 Instruction *Imag) {
2243 if (Real != RealPHI || (ImagPHI && Imag != ImagPHI))
2244 return nullptr;
2245
2246 PHIsFound = true;
2247 CompositeNode *PlaceholderNode = prepareCompositeNode(
2248 ComplexDeinterleavingOperation::ReductionPHI, Real, Imag);
2249 return submitCompositeNode(PlaceholderNode);
2250}
2251
2252ComplexDeinterleavingGraph::CompositeNode *
2253ComplexDeinterleavingGraph::identifySelectNode(Instruction *Real,
2254 Instruction *Imag) {
2255 auto *SelectReal = dyn_cast<SelectInst>(Real);
2256 auto *SelectImag = dyn_cast<SelectInst>(Imag);
2257 if (!SelectReal || !SelectImag)
2258 return nullptr;
2259
2260 Instruction *MaskA, *MaskB;
2261 Instruction *AR, *AI, *RA, *BI;
2262 if (!match(Real, m_Select(m_Instruction(MaskA), m_Instruction(AR),
2263 m_Instruction(RA))) ||
2264 !match(Imag, m_Select(m_Instruction(MaskB), m_Instruction(AI),
2265 m_Instruction(BI))))
2266 return nullptr;
2267
2268 if (MaskA != MaskB && !MaskA->isIdenticalTo(MaskB))
2269 return nullptr;
2270
2271 if (!MaskA->getType()->isVectorTy())
2272 return nullptr;
2273
2274 auto NodeA = identifyNode(AR, AI);
2275 if (!NodeA)
2276 return nullptr;
2277
2278 auto NodeB = identifyNode(RA, BI);
2279 if (!NodeB)
2280 return nullptr;
2281
2282 CompositeNode *PlaceholderNode = prepareCompositeNode(
2283 ComplexDeinterleavingOperation::ReductionSelect, Real, Imag);
2284 PlaceholderNode->addOperand(NodeA);
2285 PlaceholderNode->addOperand(NodeB);
2286 FinalInstructions.insert(MaskA);
2287 FinalInstructions.insert(MaskB);
2288 return submitCompositeNode(PlaceholderNode);
2289}
2290
2291static Value *replaceSymmetricNode(IRBuilderBase &B, unsigned Opcode,
2292 std::optional<FastMathFlags> Flags,
2293 Value *InputA, Value *InputB) {
2294 Value *I;
2295 switch (Opcode) {
2296 case Instruction::FNeg:
2297 I = B.CreateFNeg(InputA);
2298 break;
2299 case Instruction::FAdd:
2300 I = B.CreateFAdd(InputA, InputB);
2301 break;
2302 case Instruction::Add:
2303 I = B.CreateAdd(InputA, InputB);
2304 break;
2305 case Instruction::FSub:
2306 I = B.CreateFSub(InputA, InputB);
2307 break;
2308 case Instruction::Sub:
2309 I = B.CreateSub(InputA, InputB);
2310 break;
2311 case Instruction::FMul:
2312 I = B.CreateFMul(InputA, InputB);
2313 break;
2314 case Instruction::Mul:
2315 I = B.CreateMul(InputA, InputB);
2316 break;
2317 default:
2318 llvm_unreachable("Incorrect symmetric opcode");
2319 }
2320 if (Flags)
2321 cast<Instruction>(I)->setFastMathFlags(*Flags);
2322 return I;
2323}
2324
2325Value *ComplexDeinterleavingGraph::replaceNode(IRBuilderBase &Builder,
2326 CompositeNode *Node) {
2327 if (Node->ReplacementNode)
2328 return Node->ReplacementNode;
2329
2330 auto ReplaceOperandIfExist = [&](CompositeNode *Node,
2331 unsigned Idx) -> Value * {
2332 return Node->Operands.size() > Idx
2333 ? replaceNode(Builder, Node->Operands[Idx])
2334 : nullptr;
2335 };
2336
2337 Value *ReplacementNode = nullptr;
2338 switch (Node->Operation) {
2339 case ComplexDeinterleavingOperation::CDot: {
2340 Value *Input0 = ReplaceOperandIfExist(Node, 0);
2341 Value *Input1 = ReplaceOperandIfExist(Node, 1);
2342 Value *Accumulator = ReplaceOperandIfExist(Node, 2);
2343 assert(!Input1 || (Input0->getType() == Input1->getType() &&
2344 "Node inputs need to be of the same type"));
2345 ReplacementNode = TL->createComplexDeinterleavingIR(
2346 Builder, Node->Operation, Node->Rotation, Input0, Input1, Accumulator);
2347 break;
2348 }
2349 case ComplexDeinterleavingOperation::CAdd:
2350 case ComplexDeinterleavingOperation::CMulPartial:
2351 case ComplexDeinterleavingOperation::Symmetric: {
2352 Value *Input0 = ReplaceOperandIfExist(Node, 0);
2353 Value *Input1 = ReplaceOperandIfExist(Node, 1);
2354 Value *Accumulator = ReplaceOperandIfExist(Node, 2);
2355 assert(!Input1 || (Input0->getType() == Input1->getType() &&
2356 "Node inputs need to be of the same type"));
2358 (Input0->getType() == Accumulator->getType() &&
2359 "Accumulator and input need to be of the same type"));
2360 if (Node->Operation == ComplexDeinterleavingOperation::Symmetric)
2361 ReplacementNode = replaceSymmetricNode(Builder, Node->Opcode, Node->Flags,
2362 Input0, Input1);
2363 else
2364 ReplacementNode = TL->createComplexDeinterleavingIR(
2365 Builder, Node->Operation, Node->Rotation, Input0, Input1,
2366 Accumulator);
2367 break;
2368 }
2369 case ComplexDeinterleavingOperation::Deinterleave:
2370 llvm_unreachable("Deinterleave node should already have ReplacementNode");
2371 break;
2372 case ComplexDeinterleavingOperation::Splat: {
2374 for (auto &V : Node->Vals) {
2375 Ops.push_back(V.Real);
2376 Ops.push_back(V.Imag);
2377 }
2378 auto *R = dyn_cast<Instruction>(Node->Vals[0].Real);
2379 auto *I = dyn_cast<Instruction>(Node->Vals[0].Imag);
2380 if (R && I) {
2381 // Splats that are not constant are interleaved where they are located
2382 Instruction *InsertPoint = R;
2383 for (auto V : Node->Vals) {
2384 if (InsertPoint->comesBefore(cast<Instruction>(V.Real)))
2385 InsertPoint = cast<Instruction>(V.Real);
2386 if (InsertPoint->comesBefore(cast<Instruction>(V.Imag)))
2387 InsertPoint = cast<Instruction>(V.Imag);
2388 }
2389 InsertPoint = InsertPoint->getNextNode();
2390 IRBuilder<> IRB(InsertPoint);
2391 ReplacementNode = IRB.CreateVectorInterleave(Ops);
2392 } else {
2393 ReplacementNode = Builder.CreateVectorInterleave(Ops);
2394 }
2395 break;
2396 }
2397 case ComplexDeinterleavingOperation::ReductionPHI: {
2398 // If Operation is ReductionPHI, a new empty PHINode is created.
2399 // It is filled later when the ReductionOperation is processed.
2400 auto *OldPHI = cast<PHINode>(Node->Vals[0].Real);
2401 auto *VTy = cast<VectorType>(Node->Vals[0].Real->getType());
2402 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
2403 auto *NewPHI = PHINode::Create(NewVTy, 0, "", BackEdge->getFirstNonPHIIt());
2404 OldToNewPHI[OldPHI] = NewPHI;
2405 ReplacementNode = NewPHI;
2406 break;
2407 }
2408 case ComplexDeinterleavingOperation::ReductionSingle:
2409 ReplacementNode = replaceNode(Builder, Node->Operands[0]);
2410 processReductionSingle(ReplacementNode, Node);
2411 break;
2412 case ComplexDeinterleavingOperation::ReductionOperation:
2413 ReplacementNode = replaceNode(Builder, Node->Operands[0]);
2414 processReductionOperation(ReplacementNode, Node);
2415 break;
2416 case ComplexDeinterleavingOperation::ReductionSelect: {
2417 auto *MaskReal = cast<Instruction>(Node->Vals[0].Real)->getOperand(0);
2418 auto *MaskImag = cast<Instruction>(Node->Vals[0].Imag)->getOperand(0);
2419 auto *A = replaceNode(Builder, Node->Operands[0]);
2420 auto *B = replaceNode(Builder, Node->Operands[1]);
2421 auto *NewMask = Builder.CreateVectorInterleave({MaskReal, MaskImag});
2422 ReplacementNode = Builder.CreateSelect(NewMask, A, B);
2423 break;
2424 }
2425 }
2426
2427 assert(ReplacementNode && "Target failed to create Intrinsic call.");
2428 NumComplexTransformations += 1;
2429 Node->ReplacementNode = ReplacementNode;
2430 return ReplacementNode;
2431}
2432
2433void ComplexDeinterleavingGraph::processReductionSingle(
2434 Value *OperationReplacement, CompositeNode *Node) {
2435 auto *Real = cast<Instruction>(Node->Vals[0].Real);
2436 auto *OldPHI = ReductionInfo[Real].first;
2437 auto *NewPHI = OldToNewPHI[OldPHI];
2438 auto *VTy = cast<VectorType>(Real->getType());
2439 auto *NewVTy = VectorType::getDoubleElementsVectorType(VTy);
2440
2441 Value *Init = OldPHI->getIncomingValueForBlock(Incoming);
2442
2443 IRBuilder<> Builder(Incoming->getTerminator());
2444
2445 Value *NewInit = nullptr;
2446 if (auto *C = dyn_cast<Constant>(Init)) {
2447 if (C->isNullValue())
2448 NewInit = Constant::getNullValue(NewVTy);
2449 }
2450
2451 if (!NewInit)
2452 NewInit =
2453 Builder.CreateVectorInterleave({Init, Constant::getNullValue(VTy)});
2454
2455 NewPHI->addIncoming(NewInit, Incoming);
2456 NewPHI->addIncoming(OperationReplacement, BackEdge);
2457
2458 auto *FinalReduction = ReductionInfo[Real].second;
2459 Builder.SetInsertPoint(&*FinalReduction->getParent()->getFirstInsertionPt());
2460
2461 auto *AddReduce = Builder.CreateAddReduce(OperationReplacement);
2462 FinalReduction->replaceAllUsesWith(AddReduce);
2463}
2464
2465void ComplexDeinterleavingGraph::processReductionOperation(
2466 Value *OperationReplacement, CompositeNode *Node) {
2467 auto *Real = cast<Instruction>(Node->Vals[0].Real);
2468 auto *Imag = cast<Instruction>(Node->Vals[0].Imag);
2469 auto *OldPHIReal = ReductionInfo[Real].first;
2470 auto *OldPHIImag = ReductionInfo[Imag].first;
2471 auto *NewPHI = OldToNewPHI[OldPHIReal];
2472
2473 // We have to interleave initial origin values coming from IncomingBlock
2474 Value *InitReal = OldPHIReal->getIncomingValueForBlock(Incoming);
2475 Value *InitImag = OldPHIImag->getIncomingValueForBlock(Incoming);
2476
2477 IRBuilder<> Builder(Incoming->getTerminator());
2478 auto *NewInit = Builder.CreateVectorInterleave({InitReal, InitImag});
2479
2480 NewPHI->addIncoming(NewInit, Incoming);
2481 NewPHI->addIncoming(OperationReplacement, BackEdge);
2482
2483 // Deinterleave complex vector outside of loop so that it can be finally
2484 // reduced
2485 auto *FinalReductionReal = ReductionInfo[Real].second;
2486 auto *FinalReductionImag = ReductionInfo[Imag].second;
2487
2488 auto *Br = cast<CondBrInst>(BackEdge->getTerminator());
2489 BasicBlock *ExitBB = Br->getSuccessor(Br->getSuccessor(0) == BackEdge);
2490 Builder.SetInsertPoint(&*ExitBB->getFirstInsertionPt());
2491
2492 auto *Deinterleave = Builder.CreateIntrinsic(Intrinsic::vector_deinterleave2,
2493 OperationReplacement->getType(),
2494 OperationReplacement);
2495
2496 auto *NewReal = Builder.CreateExtractValue(Deinterleave, (uint64_t)0);
2497 FinalReductionReal->replaceUsesOfWith(Real, NewReal);
2498
2499 Builder.SetInsertPoint(FinalReductionImag);
2500 auto *NewImag = Builder.CreateExtractValue(Deinterleave, 1);
2501 FinalReductionImag->replaceUsesOfWith(Imag, NewImag);
2502}
2503
2504void ComplexDeinterleavingGraph::replaceNodes() {
2505 SmallVector<Instruction *, 16> DeadInstrRoots;
2506 for (auto *RootInstruction : OrderedRoots) {
2507 // Check if this potential root went through check process and we can
2508 // deinterleave it
2509 if (!RootToNode.count(RootInstruction))
2510 continue;
2511
2512 IRBuilder<> Builder(RootInstruction);
2513 auto RootNode = RootToNode[RootInstruction];
2514 Value *R = replaceNode(Builder, RootNode);
2515
2516 if (RootNode->Operation ==
2517 ComplexDeinterleavingOperation::ReductionOperation) {
2518 auto *RootReal = cast<Instruction>(RootNode->Vals[0].Real);
2519 auto *RootImag = cast<Instruction>(RootNode->Vals[0].Imag);
2520 ReductionInfo[RootReal].first->removeIncomingValue(BackEdge);
2521 ReductionInfo[RootImag].first->removeIncomingValue(BackEdge);
2522 DeadInstrRoots.push_back(RootReal);
2523 DeadInstrRoots.push_back(RootImag);
2524 } else if (RootNode->Operation ==
2525 ComplexDeinterleavingOperation::ReductionSingle) {
2526 auto *RootInst = cast<Instruction>(RootNode->Vals[0].Real);
2527 auto &Info = ReductionInfo[RootInst];
2528 Info.first->removeIncomingValue(BackEdge);
2529 DeadInstrRoots.push_back(Info.second);
2530 } else {
2531 assert(R && "Unable to find replacement for RootInstruction");
2532 DeadInstrRoots.push_back(RootInstruction);
2533 RootInstruction->replaceAllUsesWith(R);
2534 }
2535 }
2536
2537 for (auto *I : DeadInstrRoots)
2539}
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:275
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:2709
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:578
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:535
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...