[llvm] [LV] Handle complex multiply reductions (PR #204349)

via llvm-commits llvm-commits at lists.llvm.org
Wed Aug 5 13:56:06 PDT 2026


https://github.com/jon-gibney updated https://github.com/llvm/llvm-project/pull/204349

>From 2dc78baf33949638c5bc0da39a9b2ba1eb360dee Mon Sep 17 00:00:00 2001
From: Jon Gibney <jonathon.gibney at hpe.com>
Date: Tue, 16 Jun 2026 14:30:38 -0500
Subject: [PATCH 1/4] [LV] Handle complex multiply reductions Reductions of
 some types of operations on complex numbers are already being handled because
 the operations on the real and imaginary parts can be handled as separate
 reductions. However, with multiplication both parts contribute to the new
 values for each, so it's more complicated. With this patch, we look for a
 combination of 2 phi nodes that are used for calculations that amount to a
 complex multiply and then treat that as 2 linked reductions. This lets us
 vectorize something like this (with Flang):

  subroutine f(n,a,t)
    complex(kind=8) a(n),t
    do i = 1,n
      t = t * a(i)
    end do
  end

Assisted-by: Claude Opus
---
 llvm/include/llvm/Analysis/IVDescriptors.h    |  36 +++-
 .../include/llvm/Transforms/Utils/LoopUtils.h |   5 +
 llvm/lib/Analysis/IVDescriptors.cpp           |  79 +++++++++
 llvm/lib/Transforms/Utils/LoopUtils.cpp       |  40 +++++
 .../Vectorize/LoopVectorizationLegality.cpp   |  30 ++++
 .../Transforms/Vectorize/LoopVectorize.cpp    |  42 ++++-
 .../Transforms/Vectorize/SLPVectorizer.cpp    |   3 +
 llvm/lib/Transforms/Vectorize/VPlan.h         |  21 ++-
 .../lib/Transforms/Vectorize/VPlanRecipes.cpp |  35 +++-
 llvm/lib/Transforms/Vectorize/VPlanUnroll.cpp |  13 ++
 llvm/lib/Transforms/Vectorize/VPlanUtils.cpp  |   8 +-
 .../LoopVectorize/reduction-complex.ll        | 165 ++++++++++++++++++
 12 files changed, 462 insertions(+), 15 deletions(-)
 create mode 100644 llvm/test/Transforms/LoopVectorize/reduction-complex.ll

diff --git a/llvm/include/llvm/Analysis/IVDescriptors.h b/llvm/include/llvm/Analysis/IVDescriptors.h
index bad372421dfae..cb154b43884a6 100644
--- a/llvm/include/llvm/Analysis/IVDescriptors.h
+++ b/llvm/include/llvm/Analysis/IVDescriptors.h
@@ -59,6 +59,9 @@ enum class RecurKind {
   FMinimumNum, ///< FP min with llvm.minimumnum semantics
   FMaximumNum, ///< FP max with llvm.maximumnum semantics
   FMulAdd,  ///< Sum of float products with llvm.fmuladd(a * b + sum).
+  ComplexFMul, ///< Complex float multiply reduction. Represents one half
+               ///< (real or imaginary) of a paired complex multiplication
+               ///< reduction across two coupled PHI nodes.
   AnyOf,    ///< AnyOf reduction with select(cmp(),x,y) where one of (x,y) is
             ///< loop invariant, and both x and y are integer type.
   FindIV,   ///< FindIV reduction with select(icmp(),x,y) where one of (x,y) is
@@ -94,16 +97,20 @@ class RecurrenceDescriptor {
                        Type *RT, bool Signed, bool Ordered,
                        SmallPtrSetImpl<Instruction *> &CI,
                        unsigned MinWidthCastToRecurTy,
-                       bool PhiHasUsesOutsideReductionChain = false)
+                       bool PhiHasUsesOutsideReductionChain = false,
+                       PHINode *PartnerPhi = nullptr,
+                       bool IsComplexRealPart = false)
       : IntermediateStore(Store), StartValue(Start), LoopExitInstr(Exit),
         Kind(K), FMF(FMF), ExactFPMathInst(ExactFP), RecurrenceType(RT),
         IsSigned(Signed), IsOrdered(Ordered),
         PhiHasUsesOutsideReductionChain(PhiHasUsesOutsideReductionChain),
-        MinWidthCastToRecurrenceType(MinWidthCastToRecurTy) {
+        MinWidthCastToRecurrenceType(MinWidthCastToRecurTy),
+        PartnerPhi(PartnerPhi), IsComplexRealPart(IsComplexRealPart) {
     CastInsts.insert_range(CI);
     assert(
-        (!PhiHasUsesOutsideReductionChain || isMinMaxRecurrenceKind(K)) &&
-        "Only min/max recurrences are allowed to have multiple uses currently");
+        (!PhiHasUsesOutsideReductionChain || isMinMaxRecurrenceKind(K) ||
+         K == RecurKind::ComplexFMul) &&
+        "Only min/max/complex recurrences are allowed to have multiple uses");
   }
 
   /// Simpler constructor for min/max recurrences that don't track cast
@@ -299,6 +306,23 @@ class RecurrenceDescriptor {
     return isFindLastRecurrenceKind(Kind) || isFindIVRecurrenceKind(Kind);
   }
 
+  /// Returns true if the recurrence kind is a complex multiply reduction.
+  static bool isComplexRecurrenceKind(RecurKind Kind) {
+    return Kind == RecurKind::ComplexFMul;
+  }
+
+  /// Returns the partner PHI node for complex multiply reductions.
+  PHINode *getPartnerPhi() const { return PartnerPhi; }
+
+  /// Returns true if this is the real part of a complex multiply reduction.
+  bool isComplexRealPart() const { return IsComplexRealPart; }
+
+  /// Check if two PHI nodes form a complex multiply reduction pair.
+  LLVM_ABI static bool
+  isComplexMultiplyReduction(PHINode *PhiA, PHINode *PhiB, Loop *TheLoop,
+                             RecurrenceDescriptor &RdxDescA,
+                             RecurrenceDescriptor &RdxDescB);
+
   /// Returns the type of the recurrence. This type can be narrower than the
   /// actual type of the Phi if the recurrence has been type-promoted.
   Type *getRecurrenceType() const { return RecurrenceType; }
@@ -369,6 +393,10 @@ class RecurrenceDescriptor {
   SmallPtrSet<Instruction *, 8> CastInsts;
   // The minimum width used by the recurrence.
   unsigned MinWidthCastToRecurrenceType;
+  // For ComplexFMul: the partner PHI node.
+  PHINode *PartnerPhi = nullptr;
+  // For ComplexFMul: true if this is the real part.
+  bool IsComplexRealPart = false;
 };
 
 /// A struct for saving information about induction variables.
diff --git a/llvm/include/llvm/Transforms/Utils/LoopUtils.h b/llvm/include/llvm/Transforms/Utils/LoopUtils.h
index 0bf2d866b72bf..747010ca29eeb 100644
--- a/llvm/include/llvm/Transforms/Utils/LoopUtils.h
+++ b/llvm/include/llvm/Transforms/Utils/LoopUtils.h
@@ -540,6 +540,11 @@ LLVM_ABI Value *getShuffleReduction(IRBuilderBase &Builder, Value *Src,
 /// Fast-math-flags are propagated using the IRBuilder's setting.
 LLVM_ABI Value *createSimpleReduction(IRBuilderBase &B, Value *Src,
                                       RecurKind RdxKind);
+
+/// Create a horizontal complex multiply reduction from two vectors
+/// (real and imaginary parts). Returns {scalar_re, scalar_im}.
+LLVM_ABI std::pair<Value *, Value *>
+createComplexReduction(IRBuilderBase &Builder, Value *ReVec, Value *ImVec);
 /// Overloaded function to generate vector-predication intrinsics for
 /// reduction.
 LLVM_ABI Value *createSimpleReduction(IRBuilderBase &B, Value *Src,
diff --git a/llvm/lib/Analysis/IVDescriptors.cpp b/llvm/lib/Analysis/IVDescriptors.cpp
index 800b64e9a29af..576566c7a6764 100644
--- a/llvm/lib/Analysis/IVDescriptors.cpp
+++ b/llvm/lib/Analysis/IVDescriptors.cpp
@@ -1248,6 +1248,8 @@ unsigned RecurrenceDescriptor::getOpcode(RecurKind Kind) {
     return Instruction::FAdd;
   case RecurKind::FSub:
     return Instruction::FSub;
+  case RecurKind::ComplexFMul:
+    return Instruction::FMul;
   case RecurKind::SMax:
   case RecurKind::SMin:
   case RecurKind::UMax:
@@ -1381,6 +1383,83 @@ RecurrenceDescriptor::getReductionOpChain(PHINode *Phi, Loop *L) const {
   return ReductionOperations;
 }
 
+// A complex multiply recurrence is going to have 2 incoming phi nodes, one for
+// the real part and one for the imaginary part, and the calculation is going to
+// look something like
+//   next.real = (X.real * incoming.real) - (X.imag * incoming.imag)
+//   next.imag = (X.real * incoming.imag) + (X.imag * incoming.real)
+// If we see a pattern like this, we create 2 descriptors linked to each other.
+bool RecurrenceDescriptor::isComplexMultiplyReduction(
+    PHINode *PhiA, PHINode *PhiB, Loop *TheLoop, RecurrenceDescriptor &RdxDescA,
+    RecurrenceDescriptor &RdxDescB) {
+  if (!PhiA->getType()->isFloatingPointTy() ||
+      !PhiB->getType()->isFloatingPointTy())
+    return false;
+  if (PhiA->getType() != PhiB->getType())
+    return false;
+
+  BasicBlock *Header = TheLoop->getHeader();
+  if (PhiA->getParent() != Header || PhiB->getParent() != Header)
+    return false;
+  if (PhiA->getNumIncomingValues() != 2 || PhiB->getNumIncomingValues() != 2)
+    return false;
+
+  BasicBlock *Latch = TheLoop->getLoopLatch();
+  if (!Latch)
+    return false;
+
+  Value *BackA = PhiA->getIncomingValueForBlock(Latch);
+  Value *BackB = PhiB->getIncomingValueForBlock(Latch);
+  if (!BackA || !BackB)
+    return false;
+
+  auto *OpA = dyn_cast<BinaryOperator>(BackA);
+  auto *OpB = dyn_cast<BinaryOperator>(BackB);
+  if (!OpA || !OpB)
+    return false;
+  PHINode *PhiReal;
+  if (OpA->getOpcode() == Instruction::FSub &&
+      OpB->getOpcode() == Instruction::FAdd)
+    PhiReal = PhiA;
+  else if (OpA->getOpcode() == Instruction::FAdd &&
+           OpB->getOpcode() == Instruction::FSub)
+    PhiReal = PhiB;
+  else
+    return false;
+  if (!OpA->hasAllowReassoc() || !OpB->hasAllowReassoc())
+    return false;
+
+  Value *ABExt, *BAExt, *AAExt, *BBExt;
+  if (!match(OpA,
+             m_c_BinOp(
+                 m_AllowReassoc(m_c_FMul(m_Value(ABExt), m_Specific(PhiB))),
+                 m_AllowReassoc(m_c_FMul(m_Value(AAExt), m_Specific(PhiA))))))
+    return false;
+  if (!match(OpB,
+             m_c_BinOp(
+                 m_AllowReassoc(m_c_FMul(m_Value(BAExt), m_Specific(PhiA))),
+                 m_AllowReassoc(m_c_FMul(m_Value(BBExt), m_Specific(PhiB))))))
+    return false;
+  if (ABExt != BAExt || AAExt != BBExt)
+    return false;
+
+  Type *Ty = PhiA->getType();
+  FastMathFlags FMF = OpA->getFastMathFlags() & OpB->getFastMathFlags();
+
+  SmallPtrSet<Instruction *, 8> CastInsts;
+
+  RdxDescA = RecurrenceDescriptor(
+      PhiA->getIncomingValueForBlock(TheLoop->getLoopPredecessor()),
+      cast<Instruction>(BackA), nullptr, RecurKind::ComplexFMul, FMF, nullptr,
+      Ty, false, false, CastInsts, 0, false, PhiB, PhiReal == PhiA);
+  RdxDescB = RecurrenceDescriptor(
+      PhiB->getIncomingValueForBlock(TheLoop->getLoopPredecessor()),
+      cast<Instruction>(BackB), nullptr, RecurKind::ComplexFMul, FMF, nullptr,
+      Ty, false, false, CastInsts, 0, false, PhiA, PhiReal == PhiB);
+
+  return true;
+}
+
 InductionDescriptor::InductionDescriptor(
     Value *Start, InductionKind K, const SCEV *Step, BinaryOperator *BOp,
     SmallVectorImpl<Instruction *> *Casts,
diff --git a/llvm/lib/Transforms/Utils/LoopUtils.cpp b/llvm/lib/Transforms/Utils/LoopUtils.cpp
index f5dcdb4ace162..168b8eaf39d39 100644
--- a/llvm/lib/Transforms/Utils/LoopUtils.cpp
+++ b/llvm/lib/Transforms/Utils/LoopUtils.cpp
@@ -1569,6 +1569,9 @@ Value *llvm::getReductionIdentity(Intrinsic::ID RdxID, Type *Ty,
 }
 
 Value *llvm::getRecurrenceIdentity(RecurKind K, Type *Tp, FastMathFlags FMF) {
+  // ComplexFMul: real part identity is 1.0 (caller handles imag = 0.0)
+  if (K == RecurKind::ComplexFMul)
+    return ConstantFP::get(Tp, 1.0);
   assert((!(K == RecurKind::FMin || K == RecurKind::FMax) ||
           (FMF.noNaNs() && FMF.noSignedZeros())) &&
          "nnan, nsz is expected to be set for FP min/max reduction.");
@@ -1578,6 +1581,8 @@ Value *llvm::getRecurrenceIdentity(RecurKind K, Type *Tp, FastMathFlags FMF) {
 
 Value *llvm::createSimpleReduction(IRBuilderBase &Builder, Value *Src,
                                    RecurKind RdxKind) {
+  assert(RdxKind != RecurKind::ComplexFMul &&
+         "ComplexFMul uses createComplexReduction instead");
   auto *SrcVecEltTy = cast<VectorType>(Src->getType())->getElementType();
   auto getIdentity = [&]() {
     return getRecurrenceIdentity(RdxKind, SrcVecEltTy,
@@ -1631,6 +1636,41 @@ Value *llvm::createSimpleReduction(IRBuilderBase &Builder, Value *Src,
   return Builder.CreateIntrinsic(EltTy, VPID, Ops);
 }
 
+std::pair<Value *, Value *> llvm::createComplexReduction(IRBuilderBase &Builder,
+                                                         Value *ReVec,
+                                                         Value *ImVec) {
+  auto *VecTy = cast<FixedVectorType>(ReVec->getType());
+  unsigned VF = VecTy->getNumElements();
+  assert(VF > 0 && (VF & (VF - 1)) == 0 && "VF must be a power of 2");
+
+  Value *Re = ReVec;
+  Value *Im = ImVec;
+
+  for (unsigned Width = VF; Width > 1; Width >>= 1) {
+    unsigned Half = Width / 2;
+    SmallVector<int, 16> LoMask(Half), HiMask(Half);
+    for (unsigned i = 0; i < Half; ++i) {
+      LoMask[i] = i;
+      HiMask[i] = i + Half;
+    }
+    Value *LoRe = Builder.CreateShuffleVector(Re, LoMask, "lo.re");
+    Value *HiRe = Builder.CreateShuffleVector(Re, HiMask, "hi.re");
+    Value *LoIm = Builder.CreateShuffleVector(Im, LoMask, "lo.im");
+    Value *HiIm = Builder.CreateShuffleVector(Im, HiMask, "hi.im");
+
+    Value *ReRe = Builder.CreateFMul(LoRe, HiRe);
+    Value *ImIm = Builder.CreateFMul(LoIm, HiIm);
+    Value *ReIm = Builder.CreateFMul(LoRe, HiIm);
+    Value *ImRe = Builder.CreateFMul(LoIm, HiRe);
+    Re = Builder.CreateFSub(ReRe, ImIm, "red.re");
+    Im = Builder.CreateFAdd(ReIm, ImRe, "red.im");
+  }
+
+  Value *ScalarRe = Builder.CreateExtractElement(Re, uint64_t(0), "final.re");
+  Value *ScalarIm = Builder.CreateExtractElement(Im, uint64_t(0), "final.im");
+  return {ScalarRe, ScalarIm};
+}
+
 Value *llvm::createOrderedReduction(IRBuilderBase &B, RecurKind Kind,
                                     Value *Src, Value *Start) {
   assert((Kind == RecurKind::FAdd || Kind == RecurKind::FMulAdd) &&
diff --git a/llvm/lib/Transforms/Vectorize/LoopVectorizationLegality.cpp b/llvm/lib/Transforms/Vectorize/LoopVectorizationLegality.cpp
index 9086880599231..7a6228a5617f5 100644
--- a/llvm/lib/Transforms/Vectorize/LoopVectorizationLegality.cpp
+++ b/llvm/lib/Transforms/Vectorize/LoopVectorizationLegality.cpp
@@ -896,6 +896,36 @@ bool LoopVectorizationLegality::canVectorizeInstr(Instruction &I) {
       return true;
     }
 
+    // Check if this PHI is part of a complex multiply reduction pair.
+    if (Phi->getType()->isFloatingPointTy()) {
+      bool Found = false;
+      for (PHINode &OtherPhi : Phi->getParent()->phis()) {
+        if (&OtherPhi == Phi || !OtherPhi.getType()->isFloatingPointTy())
+          continue;
+        if (Inductions.count(&OtherPhi) ||
+            FixedOrderRecurrences.count(&OtherPhi))
+          continue;
+        if (Reductions.count(&OtherPhi) &&
+            Reductions[&OtherPhi].getRecurrenceKind() ==
+                RecurKind::ComplexFMul &&
+            Reductions[&OtherPhi].getPartnerPhi() == Phi) {
+          Found = true;
+          break;
+        }
+        RecurrenceDescriptor RdxDescA, RdxDescB;
+        if (RecurrenceDescriptor::isComplexMultiplyReduction(
+                Phi, &OtherPhi, TheLoop, RdxDescA, RdxDescB)) {
+          Reductions[Phi] = RdxDescA;
+          if (!Reductions.count(&OtherPhi))
+            Reductions[&OtherPhi] = RdxDescB;
+          Found = true;
+          break;
+        }
+      }
+      if (Found)
+        return true;
+    }
+
     reportVectorizationFailure("Found an unidentified PHI",
                                "value that could not be identified as "
                                "reduction is used outside the loop",
diff --git a/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp b/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp
index 3a7a4b9aecc9e..14cc26d26f086 100644
--- a/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp
+++ b/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp
@@ -7040,6 +7040,29 @@ void LoopVectorizationPlanner::addReductionResultComputation(
       Builder.setInsertPoint(MiddleVPBB, IP);
       FinalReductionResult =
           Builder.createAnyOfReduction(NewExiting, NewVal, Start, ExitDL);
+    } else if (RecurrenceDescriptor::isComplexRecurrenceKind(RecurrenceKind)) {
+      // Find the recipe corresponding to the partner PHI.
+      PHINode *PartnerPhi = RdxDesc.getPartnerPhi();
+      assert(PartnerPhi && "ComplexFMul must have a partner PHI");
+      VPReductionPHIRecipe *PartnerPhiR = nullptr;
+      for (VPRecipeBase &OtherR :
+           Plan->getVectorLoopRegion()->getEntryBasicBlock()->phis()) {
+        auto *OtherPhiR = dyn_cast<VPReductionPHIRecipe>(&OtherR);
+        if (OtherPhiR && OtherPhiR->getUnderlyingInstr() == PartnerPhi) {
+          PartnerPhiR = OtherPhiR;
+          break;
+        }
+      }
+      assert(PartnerPhiR && "Must find partner reduction PHI recipe");
+      auto *PartnerExitingVPV = PartnerPhiR->getBackedgeValue();
+
+      FastMathFlags FMFs = RdxDesc.getFastMathFlags();
+      bool IsReal = RdxDesc.isComplexRealPart();
+      VPIRFlags Flags(RecurrenceKind, /*IsOrdered=*/false, PhiR->isInLoop(),
+                      FMFs, /*IsComplexRealPart=*/IsReal);
+      FinalReductionResult = Builder.createNaryOp(
+          VPInstruction::ComputeComplexReductionResult,
+          {NewExitingVPV, PartnerExitingVPV}, Flags, ExitDL);
     } else {
       // If the vector reduction can be performed in a smaller type, we
       // truncate then extend the loop exit value to enable InstCombine to
@@ -7084,6 +7107,8 @@ void LoopVectorizationPlanner::addReductionResultComputation(
       // Skip ComputeReductionResult and FindIV reductions when they are not the
       // final result.
       if (match(U, m_VPInstruction<VPInstruction::ComputeReductionResult>()) ||
+          match(U, m_VPInstruction<
+                       VPInstruction::ComputeComplexReductionResult>()) ||
           (RecurrenceDescriptor::isFindIVRecurrenceKind(RecurrenceKind) &&
            match(U, m_VPInstruction<Instruction::ICmp>())))
         continue;
@@ -7104,8 +7129,16 @@ void LoopVectorizationPlanner::addReductionResultComputation(
          !RecurrenceDescriptor::isMinMaxRecurrenceKind(RK) &&
          !RecurrenceDescriptor::isFindLastRecurrenceKind(RK))) {
       VPBuilder PHBuilder(Plan->getVectorPreheader());
-      VPValue *Iden = Plan->getOrAddLiveIn(
-          getRecurrenceIdentity(RK, PhiTy, PhiR->getFastMathFlagsOrNone()));
+      Value *IdenVal;
+      // For ComplexFMul, the identity for the imaginary part is 0.0,
+      // while the real part uses 1.0.
+      if (RecurrenceDescriptor::isComplexRecurrenceKind(RK) &&
+          !RdxDesc.isComplexRealPart())
+        IdenVal = ConstantFP::get(PhiTy, 0.0);
+      else
+        IdenVal =
+            getRecurrenceIdentity(RK, PhiTy, PhiR->getFastMathFlagsOrNone());
+      VPValue *Iden = Plan->getOrAddLiveIn(IdenVal);
       auto *ScaleFactorVPV = Plan->getConstantInt(32, 1);
       VPValue *StartV = PHBuilder.createNaryOp(
           VPInstruction::ReductionStartVector,
@@ -7628,7 +7661,10 @@ static SmallVector<Instruction *> preparePlanForEpilogueVectorLoop(
       // value.
       auto IsReductionResult = [](VPRecipeBase *R) {
         auto *VPI = dyn_cast<VPInstruction>(R);
-        return VPI && VPI->getOpcode() == VPInstruction::ComputeReductionResult;
+        return VPI &&
+               (VPI->getOpcode() == VPInstruction::ComputeReductionResult ||
+                VPI->getOpcode() ==
+                    VPInstruction::ComputeComplexReductionResult);
       };
       auto *RdxResult = cast<VPInstruction>(
           vputils::findRecipe(ReductionPhi->getBackedgeValue(), IsReductionResult));
diff --git a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
index 3c14f2a0b595b..23eb6c0e8ba9b 100644
--- a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
+++ b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
@@ -31807,6 +31807,7 @@ class HorizontalReduction {
         case RecurKind::FMaximumNum:
         case RecurKind::FMinimumNum:
         case RecurKind::None:
+        case RecurKind::ComplexFMul:
           llvm_unreachable("Unexpected reduction kind for repeated scalar.");
         }
       }
@@ -31965,6 +31966,7 @@ class HorizontalReduction {
     case RecurKind::FMaximumNum:
     case RecurKind::FMinimumNum:
     case RecurKind::None:
+    case RecurKind::ComplexFMul:
       llvm_unreachable("Unexpected reduction kind for repeated scalar.");
     }
     return nullptr;
@@ -32069,6 +32071,7 @@ class HorizontalReduction {
     case RecurKind::FMinNum:
     case RecurKind::FMaximumNum:
     case RecurKind::FMinimumNum:
+    case RecurKind::ComplexFMul:
     case RecurKind::None:
       llvm_unreachable("Unexpected reduction kind for reused scalars.");
     }
diff --git a/llvm/lib/Transforms/Vectorize/VPlan.h b/llvm/lib/Transforms/Vectorize/VPlan.h
index 814b77a96e825..9134a8835eb0e 100644
--- a/llvm/lib/Transforms/Vectorize/VPlan.h
+++ b/llvm/lib/Transforms/Vectorize/VPlan.h
@@ -772,12 +772,14 @@ class VPIRFlags {
     // TODO: Derive order/in-loop from plan and remove here.
     unsigned char IsOrdered : 1;
     unsigned char IsInLoop : 1;
+    unsigned char IsComplexRealPart : 1;
     FastMathFlagsTy FMFs;
 
     ReductionFlagsTy(RecurKind Kind, bool IsOrdered, bool IsInLoop,
-                     FastMathFlags FMFs)
+                     FastMathFlags FMFs, bool IsComplexRealPart = false)
         : Kind(static_cast<unsigned char>(Kind)), IsOrdered(IsOrdered),
-          IsInLoop(IsInLoop), FMFs(FMFs) {}
+          IsInLoop(IsInLoop), IsComplexRealPart(IsComplexRealPart), FMFs(FMFs) {
+    }
   };
 
   OperationType OpType;
@@ -883,9 +885,11 @@ class VPIRFlags {
     GEPFlagsStorage = GEPFlags.getRaw();
   }
 
-  VPIRFlags(RecurKind Kind, bool IsOrdered, bool IsInLoop, FastMathFlags FMFs)
+  VPIRFlags(RecurKind Kind, bool IsOrdered, bool IsInLoop, FastMathFlags FMFs,
+            bool IsComplexRealPart = false)
       : OpType(OperationType::ReductionOp), AllFlags() {
-    ReductionFlags = ReductionFlagsTy(Kind, IsOrdered, IsInLoop, FMFs);
+    ReductionFlags =
+        ReductionFlagsTy(Kind, IsOrdered, IsInLoop, FMFs, IsComplexRealPart);
   }
 
   void transferFlags(VPIRFlags &Other) {
@@ -1080,6 +1084,12 @@ class VPIRFlags {
     return ReductionFlags.IsInLoop;
   }
 
+  bool isReductionRealPart() const {
+    assert(OpType == OperationType::ReductionOp &&
+           "recipe doesn't have reduction flags");
+    return ReductionFlags.IsComplexRealPart;
+  }
+
 private:
   /// Get a reference to the fast-math flags for FPMathOp, FCmp or ReductionOp.
   FastMathFlagsTy &getFMFsRef() {
@@ -1117,7 +1127,7 @@ class VPIRFlags {
 };
 LLVM_PACKED_END
 
-static_assert(sizeof(VPIRFlags) <= 3, "VPIRFlags should not grow");
+static_assert(sizeof(VPIRFlags) <= 4, "VPIRFlags should not grow");
 
 /// A pure-virtual common base class for recipes defining a single VPValue and
 /// using IR flags.
@@ -1278,6 +1288,7 @@ class LLVM_ABI_FOR_TEST VPInstruction : public VPRecipeWithIRFlags,
     /// Reduce the operands to the final reduction result using the operation
     /// specified via the operation's VPIRFlags.
     ComputeReductionResult,
+    ComputeComplexReductionResult,
     // Extracts the last part of its operand. Removed during unrolling.
     ExtractLastPart,
     // Extracts the last lane of its vector operand, per part.
diff --git a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
index ca63d1498316b..6a208a1957a6a 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
@@ -683,6 +683,7 @@ unsigned VPInstruction::getNumOperandsForOpcode() const {
   case VPInstruction::Intrinsic:
   case VPInstruction::CanonicalIVIncrementForPart:
   case VPInstruction::ComputeReductionResult:
+  case VPInstruction::ComputeComplexReductionResult:
   case VPInstruction::FirstActiveLane:
   case VPInstruction::LastActiveLane:
   case VPInstruction::ExtractLane:
@@ -977,6 +978,25 @@ Value *VPInstruction::generate(VPTransformState &State) {
 
     return ReducedPartRdx;
   }
+  case VPInstruction::ComputeComplexReductionResult: {
+    bool IsRealPart = isReductionRealPart();
+    bool IsInLoop = isReductionInLoop();
+
+    Value *OwnVec = State.get(getOperand(0), IsInLoop);
+    Value *PartnerVec = State.get(getOperand(1), IsInLoop);
+
+    IRBuilderBase::FastMathFlagGuard FMFG(Builder);
+    if (hasFastMathFlags())
+      Builder.setFastMathFlags(getFastMathFlagsOrNone());
+
+    if (State.VF.isVector() && !IsInLoop) {
+      Value *ReVec = IsRealPart ? OwnVec : PartnerVec;
+      Value *ImVec = IsRealPart ? PartnerVec : OwnVec;
+      auto [ScalarRe, ScalarIm] = createComplexReduction(Builder, ReVec, ImVec);
+      return IsRealPart ? ScalarRe : ScalarIm;
+    }
+    return OwnVec;
+  }
   case VPInstruction::ExtractLastLane:
   case VPInstruction::ExtractPenultimateElement: {
     unsigned Offset =
@@ -1508,6 +1528,7 @@ bool VPInstruction::isVectorToScalar() const {
          getOpcode() == VPInstruction::LastActiveLane ||
          getOpcode() == VPInstruction::ExtractLastActive ||
          getOpcode() == VPInstruction::ComputeReductionResult ||
+         getOpcode() == VPInstruction::ComputeComplexReductionResult ||
          getOpcode() == VPInstruction::AnyOf ||
          getOpcode() == VPInstruction::NumActiveLanes;
 }
@@ -1536,6 +1557,7 @@ void VPInstruction::addOperand(VPValue *Op) {
            "types of operand 0 and new operand must match");
     break;
   case VPInstruction::ComputeReductionResult:
+  case VPInstruction::ComputeComplexReductionResult:
   case VPInstruction::BuildVector:
   case VPInstruction::BuildStructVector:
     assert(Ty == getOperand(0)->getScalarType() &&
@@ -1810,6 +1832,9 @@ void VPInstruction::printRecipe(raw_ostream &O, const Twine &Indent,
   case VPInstruction::ComputeReductionResult:
     O << "compute-reduction-result";
     break;
+  case VPInstruction::ComputeComplexReductionResult:
+    O << "compute-complex-reduction-result";
+    break;
   case VPInstruction::LogicalAnd:
     O << "logical-and";
     break;
@@ -2535,6 +2560,7 @@ VPIRFlags VPIRFlags::getDefaultFlags(unsigned Opcode, Type *ResultTy) {
   case Instruction::ICmp:
   case Instruction::FCmp:
   case VPInstruction::ComputeReductionResult:
+  case VPInstruction::ComputeComplexReductionResult:
     llvm_unreachable("opcode requires explicit flags");
   default:
     return VPIRFlags();
@@ -2576,7 +2602,8 @@ bool VPIRFlags::flagsValidForOpcode(unsigned Opcode) const {
   case OperationType::Cmp:
     return Opcode == Instruction::FCmp || Opcode == Instruction::ICmp;
   case OperationType::ReductionOp:
-    return Opcode == VPInstruction::ComputeReductionResult;
+    return Opcode == VPInstruction::ComputeReductionResult ||
+           Opcode == VPInstruction::ComputeComplexReductionResult;
   case OperationType::Other:
     return true;
   }
@@ -2589,7 +2616,8 @@ bool VPIRFlags::hasRequiredFlagsForOpcode(unsigned Opcode) const {
     return OpType == OperationType::Cmp;
   if (Opcode == Instruction::FCmp)
     return OpType == OperationType::FCmp;
-  if (Opcode == VPInstruction::ComputeReductionResult)
+  if (Opcode == VPInstruction::ComputeReductionResult ||
+      Opcode == VPInstruction::ComputeComplexReductionResult)
     return OpType == OperationType::ReductionOp;
 
   OperationType Required = getDefaultFlags(Opcode).OpType;
@@ -2684,6 +2712,9 @@ static void printRecurrenceKind(raw_ostream &OS, const RecurKind &Kind) {
   case RecurKind::FindLast:
     OS << "find-last";
     break;
+  case RecurKind::ComplexFMul:
+    OS << "complex-fmul";
+    break;
   }
 }
 
diff --git a/llvm/lib/Transforms/Vectorize/VPlanUnroll.cpp b/llvm/lib/Transforms/Vectorize/VPlanUnroll.cpp
index a430de94cca14..9f59c6e52b159 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanUnroll.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanUnroll.cpp
@@ -430,6 +430,19 @@ void UnrollState::unrollBlock(VPBlockBase *VPB) {
         VPI->addOperand(getValueForPart(Op1, Part));
       continue;
     }
+    // ComputeComplexReductionResult has 2 operands that need expansion.
+    if (auto *VPI = dyn_cast<VPInstruction>(&R);
+        VPI &&
+        VPI->getOpcode() == VPInstruction::ComputeComplexReductionResult) {
+      addUniformForAllParts(VPI);
+      VPValue *OwnOp = VPI->getOperand(0);
+      VPValue *PartnerOp = VPI->getOperand(1);
+      for (unsigned Part = 1; Part != UF; ++Part) {
+        VPI->addOperand(getValueForPart(OwnOp, Part));
+        VPI->addOperand(getValueForPart(PartnerOp, Part));
+      }
+      continue;
+    }
     VPValue *Op0;
     if (match(&R, m_ExtractLane(m_VPValue(Op0), m_VPValue(Op1)))) {
       auto *VPI = cast<VPInstruction>(&R);
diff --git a/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp b/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
index 93b18b31e9e7d..e9e7cd846b7de 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
@@ -766,13 +766,19 @@ VPInstruction *vputils::findComputeReductionResult(VPReductionPHIRecipe *PhiR) {
   if (auto *Res =
           findUserOf<VPInstruction::ComputeReductionResult>(BackedgeVal))
     return Res;
+  if (auto *Res =
+          findUserOf<VPInstruction::ComputeComplexReductionResult>(BackedgeVal))
+    return Res;
 
   // Look through selects inserted for tail folding or predicated reductions.
   VPRecipeBase *SelR =
       findUserOf(BackedgeVal, m_Select(m_VPValue(), m_VPValue(), m_VPValue()));
   if (!SelR)
     return nullptr;
-  return findUserOf<VPInstruction::ComputeReductionResult>(
+  if (auto *Res = findUserOf<VPInstruction::ComputeReductionResult>(
+          cast<VPSingleDefRecipe>(SelR)))
+    return Res;
+  return findUserOf<VPInstruction::ComputeComplexReductionResult>(
       cast<VPSingleDefRecipe>(SelR));
 }
 
diff --git a/llvm/test/Transforms/LoopVectorize/reduction-complex.ll b/llvm/test/Transforms/LoopVectorize/reduction-complex.ll
new file mode 100644
index 0000000000000..86eb2a17123ad
--- /dev/null
+++ b/llvm/test/Transforms/LoopVectorize/reduction-complex.ll
@@ -0,0 +1,165 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt < %s -passes=loop-vectorize -force-vector-interleave=1 -force-vector-width=4 -S | FileCheck %s
+
+define void @reduction_complex_prod_(i64 %n, ptr %A, ptr %R) {
+; CHECK-LABEL: define void @reduction_complex_prod_(
+; CHECK-SAME: i64 [[N:%.*]], ptr [[A:%.*]], ptr [[R:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*]]:
+; CHECK-NEXT:    [[R_GEP_IM:%.*]] = getelementptr inbounds nuw i8, ptr [[R]], i64 4
+; CHECK-NEXT:    [[R_RE:%.*]] = load float, ptr [[R]], align 4
+; CHECK-NEXT:    [[R_IM:%.*]] = load float, ptr [[R_GEP_IM]], align 4
+; CHECK-NEXT:    [[MIN_ITERS_CHECK:%.*]] = icmp ult i64 [[N]], 4
+; CHECK-NEXT:    br i1 [[MIN_ITERS_CHECK]], label %[[SCALAR_PH:.*]], label %[[VECTOR_PH:.*]]
+; CHECK:       [[VECTOR_PH]]:
+; CHECK-NEXT:    [[N_MOD_VF:%.*]] = urem i64 [[N]], 4
+; CHECK-NEXT:    [[N_VEC:%.*]] = sub i64 [[N]], [[N_MOD_VF]]
+; CHECK-NEXT:    [[TMP0:%.*]] = insertelement <4 x float> zeroinitializer, float [[R_IM]], i32 0
+; CHECK-NEXT:    [[TMP1:%.*]] = insertelement <4 x float> splat (float 1.000000e+00), float [[R_RE]], i32 0
+; CHECK-NEXT:    br label %[[VECTOR_BODY:.*]]
+; CHECK:       [[VECTOR_BODY]]:
+; CHECK-NEXT:    [[INDEX:%.*]] = phi i64 [ 0, %[[VECTOR_PH]] ], [ [[INDEX_NEXT:%.*]], %[[VECTOR_BODY]] ]
+; CHECK-NEXT:    [[VEC_PHI:%.*]] = phi <4 x float> [ [[TMP0]], %[[VECTOR_PH]] ], [ [[TMP34:%.*]], %[[VECTOR_BODY]] ]
+; CHECK-NEXT:    [[VEC_PHI1:%.*]] = phi <4 x float> [ [[TMP1]], %[[VECTOR_PH]] ], [ [[TMP31:%.*]], %[[VECTOR_BODY]] ]
+; CHECK-NEXT:    [[TMP2:%.*]] = add i64 [[INDEX]], 1
+; CHECK-NEXT:    [[TMP3:%.*]] = add i64 [[INDEX]], 2
+; CHECK-NEXT:    [[TMP4:%.*]] = add i64 [[INDEX]], 3
+; CHECK-NEXT:    [[TMP5:%.*]] = getelementptr [8 x i8], ptr [[A]], i64 [[INDEX]]
+; CHECK-NEXT:    [[TMP6:%.*]] = getelementptr [8 x i8], ptr [[A]], i64 [[TMP2]]
+; CHECK-NEXT:    [[TMP7:%.*]] = getelementptr [8 x i8], ptr [[A]], i64 [[TMP3]]
+; CHECK-NEXT:    [[TMP8:%.*]] = getelementptr [8 x i8], ptr [[A]], i64 [[TMP4]]
+; CHECK-NEXT:    [[TMP9:%.*]] = load float, ptr [[TMP5]], align 4
+; CHECK-NEXT:    [[TMP10:%.*]] = load float, ptr [[TMP6]], align 4
+; CHECK-NEXT:    [[TMP11:%.*]] = load float, ptr [[TMP7]], align 4
+; CHECK-NEXT:    [[TMP12:%.*]] = load float, ptr [[TMP8]], align 4
+; CHECK-NEXT:    [[TMP13:%.*]] = insertelement <4 x float> poison, float [[TMP9]], i32 0
+; CHECK-NEXT:    [[TMP14:%.*]] = insertelement <4 x float> [[TMP13]], float [[TMP10]], i32 1
+; CHECK-NEXT:    [[TMP15:%.*]] = insertelement <4 x float> [[TMP14]], float [[TMP11]], i32 2
+; CHECK-NEXT:    [[TMP16:%.*]] = insertelement <4 x float> [[TMP15]], float [[TMP12]], i32 3
+; CHECK-NEXT:    [[TMP17:%.*]] = getelementptr i8, ptr [[TMP5]], i64 4
+; CHECK-NEXT:    [[TMP18:%.*]] = getelementptr i8, ptr [[TMP6]], i64 4
+; CHECK-NEXT:    [[TMP19:%.*]] = getelementptr i8, ptr [[TMP7]], i64 4
+; CHECK-NEXT:    [[TMP20:%.*]] = getelementptr i8, ptr [[TMP8]], i64 4
+; CHECK-NEXT:    [[TMP21:%.*]] = load float, ptr [[TMP17]], align 4
+; CHECK-NEXT:    [[TMP22:%.*]] = load float, ptr [[TMP18]], align 4
+; CHECK-NEXT:    [[TMP23:%.*]] = load float, ptr [[TMP19]], align 4
+; CHECK-NEXT:    [[TMP24:%.*]] = load float, ptr [[TMP20]], align 4
+; CHECK-NEXT:    [[TMP25:%.*]] = insertelement <4 x float> poison, float [[TMP21]], i32 0
+; CHECK-NEXT:    [[TMP26:%.*]] = insertelement <4 x float> [[TMP25]], float [[TMP22]], i32 1
+; CHECK-NEXT:    [[TMP27:%.*]] = insertelement <4 x float> [[TMP26]], float [[TMP23]], i32 2
+; CHECK-NEXT:    [[TMP28:%.*]] = insertelement <4 x float> [[TMP27]], float [[TMP24]], i32 3
+; CHECK-NEXT:    [[TMP29:%.*]] = fmul fast <4 x float> [[TMP16]], [[VEC_PHI1]]
+; CHECK-NEXT:    [[TMP30:%.*]] = fmul fast <4 x float> [[TMP28]], [[VEC_PHI]]
+; CHECK-NEXT:    [[TMP31]] = fsub fast <4 x float> [[TMP29]], [[TMP30]]
+; CHECK-NEXT:    [[TMP32:%.*]] = fmul fast <4 x float> [[TMP16]], [[VEC_PHI]]
+; CHECK-NEXT:    [[TMP33:%.*]] = fmul fast <4 x float> [[TMP28]], [[VEC_PHI1]]
+; CHECK-NEXT:    [[TMP34]] = fadd fast <4 x float> [[TMP33]], [[TMP32]]
+; CHECK-NEXT:    [[INDEX_NEXT]] = add nuw i64 [[INDEX]], 4
+; CHECK-NEXT:    [[TMP35:%.*]] = icmp eq i64 [[INDEX_NEXT]], [[N_VEC]]
+; CHECK-NEXT:    br i1 [[TMP35]], label %[[MIDDLE_BLOCK:.*]], label %[[VECTOR_BODY]], !llvm.loop [[LOOP0:![0-9]+]]
+; CHECK:       [[MIDDLE_BLOCK]]:
+; CHECK-NEXT:    [[LO_RE:%.*]] = shufflevector <4 x float> [[TMP31]], <4 x float> poison, <2 x i32> <i32 0, i32 1>
+; CHECK-NEXT:    [[HI_RE:%.*]] = shufflevector <4 x float> [[TMP31]], <4 x float> poison, <2 x i32> <i32 2, i32 3>
+; CHECK-NEXT:    [[LO_IM:%.*]] = shufflevector <4 x float> [[TMP34]], <4 x float> poison, <2 x i32> <i32 0, i32 1>
+; CHECK-NEXT:    [[HI_IM:%.*]] = shufflevector <4 x float> [[TMP34]], <4 x float> poison, <2 x i32> <i32 2, i32 3>
+; CHECK-NEXT:    [[TMP36:%.*]] = fmul fast <2 x float> [[LO_RE]], [[HI_RE]]
+; CHECK-NEXT:    [[TMP37:%.*]] = fmul fast <2 x float> [[LO_IM]], [[HI_IM]]
+; CHECK-NEXT:    [[TMP38:%.*]] = fmul fast <2 x float> [[LO_RE]], [[HI_IM]]
+; CHECK-NEXT:    [[TMP39:%.*]] = fmul fast <2 x float> [[LO_IM]], [[HI_RE]]
+; CHECK-NEXT:    [[RED_RE:%.*]] = fsub fast <2 x float> [[TMP36]], [[TMP37]]
+; CHECK-NEXT:    [[RED_IM:%.*]] = fadd fast <2 x float> [[TMP38]], [[TMP39]]
+; CHECK-NEXT:    [[LO_RE2:%.*]] = shufflevector <2 x float> [[RED_RE]], <2 x float> poison, <1 x i32> zeroinitializer
+; CHECK-NEXT:    [[HI_RE3:%.*]] = shufflevector <2 x float> [[RED_RE]], <2 x float> poison, <1 x i32> <i32 1>
+; CHECK-NEXT:    [[LO_IM4:%.*]] = shufflevector <2 x float> [[RED_IM]], <2 x float> poison, <1 x i32> zeroinitializer
+; CHECK-NEXT:    [[HI_IM5:%.*]] = shufflevector <2 x float> [[RED_IM]], <2 x float> poison, <1 x i32> <i32 1>
+; CHECK-NEXT:    [[TMP40:%.*]] = fmul fast <1 x float> [[LO_RE2]], [[HI_RE3]]
+; CHECK-NEXT:    [[TMP41:%.*]] = fmul fast <1 x float> [[LO_IM4]], [[HI_IM5]]
+; CHECK-NEXT:    [[TMP42:%.*]] = fmul fast <1 x float> [[LO_RE2]], [[HI_IM5]]
+; CHECK-NEXT:    [[TMP43:%.*]] = fmul fast <1 x float> [[LO_IM4]], [[HI_RE3]]
+; CHECK-NEXT:    [[RED_RE6:%.*]] = fsub fast <1 x float> [[TMP40]], [[TMP41]]
+; CHECK-NEXT:    [[RED_IM7:%.*]] = fadd fast <1 x float> [[TMP42]], [[TMP43]]
+; CHECK-NEXT:    [[FINAL_RE:%.*]] = extractelement <1 x float> [[RED_RE6]], i64 0
+; CHECK-NEXT:    [[FINAL_IM:%.*]] = extractelement <1 x float> [[RED_IM7]], i64 0
+; CHECK-NEXT:    [[LO_RE8:%.*]] = shufflevector <4 x float> [[TMP31]], <4 x float> poison, <2 x i32> <i32 0, i32 1>
+; CHECK-NEXT:    [[HI_RE9:%.*]] = shufflevector <4 x float> [[TMP31]], <4 x float> poison, <2 x i32> <i32 2, i32 3>
+; CHECK-NEXT:    [[LO_IM10:%.*]] = shufflevector <4 x float> [[TMP34]], <4 x float> poison, <2 x i32> <i32 0, i32 1>
+; CHECK-NEXT:    [[HI_IM11:%.*]] = shufflevector <4 x float> [[TMP34]], <4 x float> poison, <2 x i32> <i32 2, i32 3>
+; CHECK-NEXT:    [[TMP44:%.*]] = fmul fast <2 x float> [[LO_RE8]], [[HI_RE9]]
+; CHECK-NEXT:    [[TMP45:%.*]] = fmul fast <2 x float> [[LO_IM10]], [[HI_IM11]]
+; CHECK-NEXT:    [[TMP46:%.*]] = fmul fast <2 x float> [[LO_RE8]], [[HI_IM11]]
+; CHECK-NEXT:    [[TMP47:%.*]] = fmul fast <2 x float> [[LO_IM10]], [[HI_RE9]]
+; CHECK-NEXT:    [[RED_RE12:%.*]] = fsub fast <2 x float> [[TMP44]], [[TMP45]]
+; CHECK-NEXT:    [[RED_IM13:%.*]] = fadd fast <2 x float> [[TMP46]], [[TMP47]]
+; CHECK-NEXT:    [[LO_RE14:%.*]] = shufflevector <2 x float> [[RED_RE12]], <2 x float> poison, <1 x i32> zeroinitializer
+; CHECK-NEXT:    [[HI_RE15:%.*]] = shufflevector <2 x float> [[RED_RE12]], <2 x float> poison, <1 x i32> <i32 1>
+; CHECK-NEXT:    [[LO_IM16:%.*]] = shufflevector <2 x float> [[RED_IM13]], <2 x float> poison, <1 x i32> zeroinitializer
+; CHECK-NEXT:    [[HI_IM17:%.*]] = shufflevector <2 x float> [[RED_IM13]], <2 x float> poison, <1 x i32> <i32 1>
+; CHECK-NEXT:    [[TMP48:%.*]] = fmul fast <1 x float> [[LO_RE14]], [[HI_RE15]]
+; CHECK-NEXT:    [[TMP49:%.*]] = fmul fast <1 x float> [[LO_IM16]], [[HI_IM17]]
+; CHECK-NEXT:    [[TMP50:%.*]] = fmul fast <1 x float> [[LO_RE14]], [[HI_IM17]]
+; CHECK-NEXT:    [[TMP51:%.*]] = fmul fast <1 x float> [[LO_IM16]], [[HI_RE15]]
+; CHECK-NEXT:    [[RED_RE18:%.*]] = fsub fast <1 x float> [[TMP48]], [[TMP49]]
+; CHECK-NEXT:    [[RED_IM19:%.*]] = fadd fast <1 x float> [[TMP50]], [[TMP51]]
+; CHECK-NEXT:    [[FINAL_RE20:%.*]] = extractelement <1 x float> [[RED_RE18]], i64 0
+; CHECK-NEXT:    [[FINAL_IM21:%.*]] = extractelement <1 x float> [[RED_IM19]], i64 0
+; CHECK-NEXT:    [[CMP_N:%.*]] = icmp eq i64 [[N]], [[N_VEC]]
+; CHECK-NEXT:    br i1 [[CMP_N]], label %[[EXIT:.*]], label %[[SCALAR_PH]]
+; CHECK:       [[SCALAR_PH]]:
+; CHECK-NEXT:    [[BC_RESUME_VAL:%.*]] = phi i64 [ [[N_VEC]], %[[MIDDLE_BLOCK]] ], [ 0, %[[ENTRY]] ]
+; CHECK-NEXT:    [[BC_MERGE_RDX:%.*]] = phi float [ [[FINAL_IM]], %[[MIDDLE_BLOCK]] ], [ [[R_IM]], %[[ENTRY]] ]
+; CHECK-NEXT:    [[BC_MERGE_RDX22:%.*]] = phi float [ [[FINAL_RE20]], %[[MIDDLE_BLOCK]] ], [ [[R_RE]], %[[ENTRY]] ]
+; CHECK-NEXT:    br label %[[BODY:.*]]
+; CHECK:       [[BODY]]:
+; CHECK-NEXT:    [[IV:%.*]] = phi i64 [ [[BC_RESUME_VAL]], %[[SCALAR_PH]] ], [ [[IV_NEXT:%.*]], %[[BODY]] ]
+; CHECK-NEXT:    [[PROD_IM:%.*]] = phi float [ [[BC_MERGE_RDX]], %[[SCALAR_PH]] ], [ [[PROD_IM_NEXT:%.*]], %[[BODY]] ]
+; CHECK-NEXT:    [[PROD_RE:%.*]] = phi float [ [[BC_MERGE_RDX22]], %[[SCALAR_PH]] ], [ [[PROD_RE_NEXT:%.*]], %[[BODY]] ]
+; CHECK-NEXT:    [[ELT_GEP:%.*]] = getelementptr [8 x i8], ptr [[A]], i64 [[IV]]
+; CHECK-NEXT:    [[ELT_RE:%.*]] = load float, ptr [[ELT_GEP]], align 4
+; CHECK-NEXT:    [[ELT_GEP_IM:%.*]] = getelementptr i8, ptr [[ELT_GEP]], i64 4
+; CHECK-NEXT:    [[ELT_IM:%.*]] = load float, ptr [[ELT_GEP_IM]], align 4
+; CHECK-NEXT:    [[RERE:%.*]] = fmul fast float [[ELT_RE]], [[PROD_RE]]
+; CHECK-NEXT:    [[IMIM:%.*]] = fmul fast float [[ELT_IM]], [[PROD_IM]]
+; CHECK-NEXT:    [[PROD_RE_NEXT]] = fsub fast float [[RERE]], [[IMIM]]
+; CHECK-NEXT:    [[REIM:%.*]] = fmul fast float [[ELT_RE]], [[PROD_IM]]
+; CHECK-NEXT:    [[IMRE:%.*]] = fmul fast float [[ELT_IM]], [[PROD_RE]]
+; CHECK-NEXT:    [[PROD_IM_NEXT]] = fadd fast float [[IMRE]], [[REIM]]
+; CHECK-NEXT:    [[IV_NEXT]] = add nuw nsw i64 [[IV]], 1
+; CHECK-NEXT:    [[EXITCOND:%.*]] = icmp eq i64 [[IV_NEXT]], [[N]]
+; CHECK-NEXT:    br i1 [[EXITCOND]], label %[[EXIT]], label %[[BODY]], !llvm.loop [[LOOP3:![0-9]+]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    [[R_OUT_IM:%.*]] = phi float [ [[PROD_IM_NEXT]], %[[BODY]] ], [ [[FINAL_IM]], %[[MIDDLE_BLOCK]] ]
+; CHECK-NEXT:    [[R_OUT_RE:%.*]] = phi float [ [[PROD_RE_NEXT]], %[[BODY]] ], [ [[FINAL_RE20]], %[[MIDDLE_BLOCK]] ]
+; CHECK-NEXT:    store float [[R_OUT_RE]], ptr [[R]], align 4
+; CHECK-NEXT:    store float [[R_OUT_IM]], ptr [[R_GEP_IM]], align 4
+; CHECK-NEXT:    ret void
+;
+entry:
+  %R.gep.im = getelementptr inbounds nuw i8, ptr %R, i64 4
+  %R.re = load float, ptr %R, align 4
+  %R.im = load float, ptr %R.gep.im, align 4
+  br label %body
+
+body:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %body ]
+  %prod.im = phi float [ %R.im, %entry ], [ %prod.im.next, %body ]
+  %prod.re = phi float [ %R.re, %entry ], [ %prod.re.next, %body ]
+  %elt.gep = getelementptr [8 x i8], ptr %A, i64 %iv
+  %elt.re = load float, ptr %elt.gep, align 4
+  %elt.gep.im = getelementptr i8, ptr %elt.gep, i64 4
+  %elt.im = load float, ptr %elt.gep.im, align 4
+  %rere = fmul fast float %elt.re, %prod.re
+  %imim = fmul fast float %elt.im, %prod.im
+  %prod.re.next = fsub fast float %rere, %imim
+  %reim = fmul fast float %elt.re, %prod.im
+  %imre = fmul fast float %elt.im, %prod.re
+  %prod.im.next = fadd fast float %imre, %reim
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, %n
+  br i1 %exitcond, label %exit, label %body
+
+exit:
+  %R.out.im = phi float [ %prod.im.next, %body ]
+  %R.out.re = phi float [ %prod.re.next, %body ]
+  store float %R.out.re, ptr %R, align 4
+  store float %R.out.im, ptr %R.gep.im, align 4
+  ret void
+}

>From 169160cffe06a4212e2bdf6c4e3707dbae0db907 Mon Sep 17 00:00:00 2001
From: Jon Gibney <jonathon.gibney at hpe.com>
Date: Wed, 5 Aug 2026 12:06:45 -0500
Subject: [PATCH 2/4] Update based on review feedback

- Rework logic in isComplexMultiplyReduction to avoid false positive
  when the real update has its sign flipped

- Update generate() for ComputeComplexReductionResult to handle more
  than 2 operands.

- Update release notes.
---
 llvm/docs/ReleaseNotes.md                     |  2 +
 llvm/lib/Analysis/IVDescriptors.cpp           | 66 +++++++++++--------
 .../lib/Transforms/Vectorize/VPlanRecipes.cpp | 27 ++++++++
 3 files changed, 66 insertions(+), 29 deletions(-)

diff --git a/llvm/docs/ReleaseNotes.md b/llvm/docs/ReleaseNotes.md
index f8fce847b62f3..3931152529522 100644
--- a/llvm/docs/ReleaseNotes.md
+++ b/llvm/docs/ReleaseNotes.md
@@ -69,6 +69,8 @@ Makes programs 10x faster by doing Special New Thing.
 
 ### Changes to Vectorizers
 
+* LoopVectorizer can now recognize multiply reductions of complex numbers.
+
 ### Changes to the AArch64 Backend
 
 ### Changes to the AMDGPU Backend
diff --git a/llvm/lib/Analysis/IVDescriptors.cpp b/llvm/lib/Analysis/IVDescriptors.cpp
index 576566c7a6764..15b1103ec2842 100644
--- a/llvm/lib/Analysis/IVDescriptors.cpp
+++ b/llvm/lib/Analysis/IVDescriptors.cpp
@@ -1417,45 +1417,53 @@ bool RecurrenceDescriptor::isComplexMultiplyReduction(
   auto *OpB = dyn_cast<BinaryOperator>(BackB);
   if (!OpA || !OpB)
     return false;
-  PHINode *PhiReal;
+  PHINode *PhiRe, *PhiIm;
+  BinaryOperator *NextRe, *NextIm;
   if (OpA->getOpcode() == Instruction::FSub &&
-      OpB->getOpcode() == Instruction::FAdd)
-    PhiReal = PhiA;
-  else if (OpA->getOpcode() == Instruction::FAdd &&
-           OpB->getOpcode() == Instruction::FSub)
-    PhiReal = PhiB;
-  else
-    return false;
-  if (!OpA->hasAllowReassoc() || !OpB->hasAllowReassoc())
+      OpB->getOpcode() == Instruction::FAdd) {
+    PhiRe = PhiA;
+    PhiIm = PhiB;
+    NextRe = OpA;
+    NextIm = OpB;
+  } else if (OpA->getOpcode() == Instruction::FAdd &&
+             OpB->getOpcode() == Instruction::FSub) {
+    PhiRe = PhiB;
+    PhiIm = PhiA;
+    NextRe = OpB;
+    NextIm = OpA;
+  } else
     return false;
 
-  Value *ABExt, *BAExt, *AAExt, *BBExt;
-  if (!match(OpA,
-             m_c_BinOp(
-                 m_AllowReassoc(m_c_FMul(m_Value(ABExt), m_Specific(PhiB))),
-                 m_AllowReassoc(m_c_FMul(m_Value(AAExt), m_Specific(PhiA))))))
+  Value *NthReForRe, *NthImForRe, *NthReForIm, *NthImForIm;
+  if (!match(NextRe->getOperand(0),
+             m_AllowReassoc(m_c_FMul(m_Value(NthReForRe), m_Specific(PhiRe)))))
+    return false;
+  if (!match(NextRe->getOperand(1),
+             m_AllowReassoc(m_c_FMul(m_Value(NthImForRe), m_Specific(PhiIm)))))
     return false;
-  if (!match(OpB,
-             m_c_BinOp(
-                 m_AllowReassoc(m_c_FMul(m_Value(BAExt), m_Specific(PhiA))),
-                 m_AllowReassoc(m_c_FMul(m_Value(BBExt), m_Specific(PhiB))))))
+  if (!match(NextIm, m_c_FAdd(m_AllowReassoc(m_c_FMul(m_Value(NthImForIm),
+                                                      m_Specific(PhiRe))),
+                              m_AllowReassoc(m_c_FMul(m_Value(NthReForIm),
+                                                      m_Specific(PhiIm))))))
     return false;
-  if (ABExt != BAExt || AAExt != BBExt)
+  if (NthImForRe != NthImForIm || NthReForRe != NthReForIm)
     return false;
 
-  Type *Ty = PhiA->getType();
-  FastMathFlags FMF = OpA->getFastMathFlags() & OpB->getFastMathFlags();
+  Type *Ty = PhiRe->getType();
+  FastMathFlags FMF = NextRe->getFastMathFlags() & NextIm->getFastMathFlags();
 
   SmallPtrSet<Instruction *, 8> CastInsts;
 
-  RdxDescA = RecurrenceDescriptor(
-      PhiA->getIncomingValueForBlock(TheLoop->getLoopPredecessor()),
-      cast<Instruction>(BackA), nullptr, RecurKind::ComplexFMul, FMF, nullptr,
-      Ty, false, false, CastInsts, 0, false, PhiB, PhiReal == PhiA);
-  RdxDescB = RecurrenceDescriptor(
-      PhiB->getIncomingValueForBlock(TheLoop->getLoopPredecessor()),
-      cast<Instruction>(BackB), nullptr, RecurKind::ComplexFMul, FMF, nullptr,
-      Ty, false, false, CastInsts, 0, false, PhiA, PhiReal == PhiB);
+  auto RdxDescRe = RecurrenceDescriptor(
+      PhiRe->getIncomingValueForBlock(TheLoop->getLoopPredecessor()),
+      cast<Instruction>(NextRe), nullptr, RecurKind::ComplexFMul, FMF, nullptr,
+      Ty, false, false, CastInsts, 0, false, PhiIm, true);
+  auto RdxDescIm = RecurrenceDescriptor(
+      PhiIm->getIncomingValueForBlock(TheLoop->getLoopPredecessor()),
+      cast<Instruction>(NextIm), nullptr, RecurKind::ComplexFMul, FMF, nullptr,
+      Ty, false, false, CastInsts, 0, false, PhiRe, false);
+  RdxDescA = PhiA == PhiRe ? RdxDescRe : RdxDescIm;
+  RdxDescB = PhiB == PhiRe ? RdxDescRe : RdxDescIm;
 
   return true;
 }
diff --git a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
index 6a208a1957a6a..41718394733c0 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
@@ -982,6 +982,10 @@ Value *VPInstruction::generate(VPTransformState &State) {
     bool IsRealPart = isReductionRealPart();
     bool IsInLoop = isReductionInLoop();
 
+    unsigned NumOperands = getNumOperands();
+    assert(NumOperands >= 2 && NumOperands % 2 == 0 &&
+           "Expected pairs of own/partner operands");
+
     Value *OwnVec = State.get(getOperand(0), IsInLoop);
     Value *PartnerVec = State.get(getOperand(1), IsInLoop);
 
@@ -989,6 +993,29 @@ Value *VPInstruction::generate(VPTransformState &State) {
     if (hasFastMathFlags())
       Builder.setFastMathFlags(getFastMathFlagsOrNone());
 
+    // Reduce across unroll parts with element-wise complex multiplication.
+    // Operands come in (own, partner) pairs from VPlanUnroll.
+    for (unsigned I = 2; I < NumOperands; I += 2) {
+      Value *OwnPart = State.get(getOperand(I), IsInLoop);
+      Value *PartnerPart = State.get(getOperand(I + 1), IsInLoop);
+
+      Value *Re1 = IsRealPart ? OwnVec : PartnerVec;
+      Value *Im1 = IsRealPart ? PartnerVec : OwnVec;
+      Value *Re2 = IsRealPart ? OwnPart : PartnerPart;
+      Value *Im2 = IsRealPart ? PartnerPart : OwnPart;
+
+      // (Re1 + i*Im1) * (Re2 + i*Im2)
+      Value *NewRe = Builder.CreateFSub(
+          Builder.CreateFMul(Re1, Re2), Builder.CreateFMul(Im1, Im2),
+          "rdx.re");
+      Value *NewIm = Builder.CreateFAdd(
+          Builder.CreateFMul(Re1, Im2), Builder.CreateFMul(Im1, Re2),
+          "rdx.im");
+
+      OwnVec = IsRealPart ? NewRe : NewIm;
+      PartnerVec = IsRealPart ? NewIm : NewRe;
+    }
+
     if (State.VF.isVector() && !IsInLoop) {
       Value *ReVec = IsRealPart ? OwnVec : PartnerVec;
       Value *ImVec = IsRealPart ? PartnerVec : OwnVec;

>From debde29f344e8de0947e1ff8f1b9f040c45a5d3f Mon Sep 17 00:00:00 2001
From: Jon Gibney <jonathon.gibney at hpe.com>
Date: Wed, 5 Aug 2026 15:10:51 -0500
Subject: [PATCH 3/4] Fix formatting

---
 llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp | 10 ++++------
 1 file changed, 4 insertions(+), 6 deletions(-)

diff --git a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
index 41718394733c0..7fa1fd475abed 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
@@ -1005,12 +1005,10 @@ Value *VPInstruction::generate(VPTransformState &State) {
       Value *Im2 = IsRealPart ? PartnerPart : OwnPart;
 
       // (Re1 + i*Im1) * (Re2 + i*Im2)
-      Value *NewRe = Builder.CreateFSub(
-          Builder.CreateFMul(Re1, Re2), Builder.CreateFMul(Im1, Im2),
-          "rdx.re");
-      Value *NewIm = Builder.CreateFAdd(
-          Builder.CreateFMul(Re1, Im2), Builder.CreateFMul(Im1, Re2),
-          "rdx.im");
+      Value *NewRe = Builder.CreateFSub(Builder.CreateFMul(Re1, Re2),
+                                        Builder.CreateFMul(Im1, Im2), "rdx.re");
+      Value *NewIm = Builder.CreateFAdd(Builder.CreateFMul(Re1, Im2),
+                                        Builder.CreateFMul(Im1, Re2), "rdx.im");
 
       OwnVec = IsRealPart ? NewRe : NewIm;
       PartnerVec = IsRealPart ? NewIm : NewRe;

>From c69e7ce7dfbc66739d24f0ed43bd6808f14bf659 Mon Sep 17 00:00:00 2001
From: Jon Gibney <jonathon.gibney at hpe.com>
Date: Wed, 5 Aug 2026 15:55:16 -0500
Subject: [PATCH 4/4] Fix RISCV switch

---
 llvm/lib/Target/RISCV/RISCVTargetTransformInfo.h | 1 +
 1 file changed, 1 insertion(+)

diff --git a/llvm/lib/Target/RISCV/RISCVTargetTransformInfo.h b/llvm/lib/Target/RISCV/RISCVTargetTransformInfo.h
index 40c6204ee0380..731e0e358f29c 100644
--- a/llvm/lib/Target/RISCV/RISCVTargetTransformInfo.h
+++ b/llvm/lib/Target/RISCV/RISCVTargetTransformInfo.h
@@ -456,6 +456,7 @@ class RISCVTTIImpl final : public BasicTTIImplBase<RISCVTTIImpl> {
     case RecurKind::FMinimumNum:
     case RecurKind::FMaximumNum:
     case RecurKind::FAddChainWithSubs:
+    case RecurKind::ComplexFMul:
       return false;
     case RecurKind::None:
       llvm_unreachable("Unknown reduction kind.");



More information about the llvm-commits mailing list