================
@@ -7484,6 +7495,224 @@ void VPlanTransforms::makeMemOpWideningDecisions(VPlan 
&Plan, VFRange &Range,
                            });
 }
 
+void VPlanTransforms::multiversionForUnitStridedMemOps(
+    VPlan &Plan, VPCostContext &CostCtx, VFSelectionContext &Config,
+    VPRecipeBuilder &RecipeBuilder, VFRange &Range,
+    SmallVectorImpl<VPInstruction *> &MemOps) {
+  SmallVector<VPInstruction *> RemainingOps;
+
+  ScalarEvolution *SE = CostCtx.PSE.getSE();
+
+  PredicatedScalarEvolution StrideMVPSE(*SE, const_cast<Loop &>(*CostCtx.L));
+
+  SCEVUnionPredicate StridePredicates({}, *SE);
+
+  // Use `for_each` so that we could do `return Skip();`.
+  for_each(MemOps, [&](VPInstruction *VPI) {
+    auto Skip = [&]() { RemainingOps.push_back(VPI); };
+    if (RecipeBuilder.isConsecutiveWithoutVPlanBasedStrideSpeculation(VPI))
+      return Skip();
+    auto *PtrOp = VPI->getOpcode() == Instruction::Load ? VPI->getOperand(0)
+                                                        : VPI->getOperand(1);
+
+    const SCEV *PtrSCEV =
+        vputils::getSCEVExprForVPValue(PtrOp, CostCtx.PSE, CostCtx.L);
+    const SCEV *Start, *Stride;
+
+    if (!match(PtrSCEV, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Stride),
+                                            m_SpecificLoop(CostCtx.L))))
+      return Skip();
+
+    Type *ScalarTy = VPI->getOpcode() == Instruction::Load
+                         ? VPI->getScalarType()
+                         : VPI->getOperand(0)->getScalarType();
+
+    if (VPI->getMask()) {
+      Instruction *I = VPI->getUnderlyingInstr();
+      if (!LoopVectorizationPlanner::getDecisionAndClampRange(
+              [&](ElementCount VF) -> bool {
+                return Config.isLegalMaskedLoadOrStore(I, VF);
+              },
+              Range))
+        return Skip();
+    }
+
+    const auto *TypeSize = cast<SCEVConstant>(SE->getSizeOfExpr(
+        Stride->getType(), SE->getDataLayout().getTypeAllocSize(ScalarTy)));
+
+    auto ReplaceWithUnitStrided = [&]() {
+      VPBuilder Builder(VPI);
+      VPValue *StrideInBytes = Plan.getOrAddLiveIn(TypeSize->getValue());
+      auto *VecPtr =
+          new VPVectorPointerRecipe(PtrOp, ScalarTy, StrideInBytes,
+                                    GEPNoWrapFlags::none(), 
VPI->getDebugLoc());
+      Builder.insert(VecPtr);
+      if (VPI->getOpcode() == Instruction::Load) {
+        auto *WideLoad = new VPWidenLoadRecipe(
+            cast<LoadInst>(*VPI->getUnderlyingInstr()), VecPtr, VPI->getMask(),
+            true, *VPI, VPI->getDebugLoc());
+        Builder.insert(WideLoad);
+        VPI->replaceAllUsesWith(WideLoad);
+      } else {
+        auto *WideStore = new VPWidenStoreRecipe(
+            cast<StoreInst>(*VPI->getUnderlyingInstr()), VecPtr,
+            VPI->getOperand(0), VPI->getMask(), true, *VPI, 
VPI->getDebugLoc());
+        Builder.insert(WideStore);
+      }
+      VPI->eraseFromParent();
+    };
+
+    if (isa<SCEVConstant>(Stride)) {
+      if (Stride != TypeSize)
+        return Skip();
+
+      // Earlier MV helped with this memory operation too.
+      ReplaceWithUnitStrided();
+      return;
+    }
+
+    const SCEVConstant *StrideConstantMultiplier;
+    const SCEV *StrideNonConstantMultiplier;
+
+    const SCEV *ToMultiVersion = Stride;
+    const SCEV *MVConst = TypeSize;
+    if (match(Stride, m_scev_c_Mul(m_SCEVConstant(StrideConstantMultiplier),
+                                   m_SCEV(StrideNonConstantMultiplier)))) {
+      if (TypeSize != StrideConstantMultiplier) {
+        // TODO: Support `TypeSize = N * StrideConstantMultiplier`,
+        // including negative `N`. For now, only process when they're equal,
+        // which matches the useful part of the legacy behavior that
+        // multiversiones GEP index for stride one.
+        return Skip();
+      }
+      ToMultiVersion = StrideNonConstantMultiplier;
+      MVConst = SE->getOne(ToMultiVersion->getType());
+    } else if (!TypeSize->isOne()) {
+      // Likewise - try to match legacy behavior.
+      return Skip();
+    }
+
+    while (auto *C = dyn_cast<SCEVIntegralCastExpr>(ToMultiVersion)) {
+      ToMultiVersion = C->getOperand();
+      MVConst = SE->getTruncateOrSignExtend(MVConst, 
ToMultiVersion->getType());
+    }
+
+    if (match(ToMultiVersion, m_scev_UndefOrPoison()))
+      return Skip();
+
+    if (!isa<SCEVUnknown>(ToMultiVersion)) {
+      // Match legacy behavior.
+      // If/when changed, make sure that explicit poison/undef in the defining
+      // expression doesn't cause any issues.
+      return Skip();
+    }
+
+    Value *StrideVal = cast<SCEVUnknown>(ToMultiVersion)->getValue();
+
+    const SCEVPredicate *NewPred =
+        SE->getComparePredicate(CmpInst::ICMP_EQ, ToMultiVersion, MVConst);
+
+    // Check if new predicate implies that backedge is never taken. If so, 
there
+    // is no reason to multiversion for it.
+    auto *PredicatedMaxBTC = SE->rewriteUsingPredicate(
+        SE->getSymbolicMaxBackedgeTakenCount(CostCtx.L), CostCtx.L,
+        StridePredicates.getUnionWith(NewPred, *SE)
+            .getUnionWith(&CostCtx.PSE.getPredicate(), *SE));
+
+    if (LoopVectorizationPlanner::getDecisionAndClampRange(
+            [&](ElementCount VF) {
+              return SE->isKnownPredicate(
+                  ICmpInst::ICMP_SLT, PredicatedMaxBTC,
+                  SE->getConstant(PredicatedMaxBTC->getType(),
+                                  VF.isScalable() ? 1
+                                                  : VF.getFixedValue() - 1));
+            },
+            Range))
+      return Skip();
+
+    StridePredicates = StridePredicates.getUnionWith(NewPred, *SE);
+
+    auto ReplaceMVUses = [&](Value *V) {
+      VPValue *From = Plan.getLiveIn(V);
+      if (!From)
+        return;
+      VPValue *To = Plan.getConstantInt(
+          From->getScalarType(),
+          cast<SCEVConstant>(MVConst)->getAPInt().getLimitedValue());
+      // TODO: Why is "If" necessary?
+      From->replaceUsesWithIf(To, [&](VPUser &U, unsigned) {
+        auto *R = cast<VPRecipeBase>(&U);
+        return R->getRegion() ||
+               R->getParent() ==
+                   Plan.getVectorLoopRegion()->getSinglePredecessor();
+      });
+    };
+
+    ReplaceMVUses(StrideVal);
+    for (auto *U : StrideVal->users())
+      if (isa<SExtInst, ZExtInst>(U))
+        ReplaceMVUses(U);
+
+    ReplaceWithUnitStrided();
+  });
+
+  MemOps.swap(RemainingOps);
+
+  if (StridePredicates.isAlwaysTrue())
+    return;
+
+  VPBasicBlock *Entry = Plan.getEntry();
+  VPBuilder Builder(Entry);
+  DebugLoc DL = cast<VPIRBasicBlock>(Entry)
+                    ->getIRBasicBlock()
+                    ->getTerminator()
+                    ->getDebugLoc();
+  VPSCEVExpander Expander(Builder, *SE, DL);
+  VPValue *Pred = Expander.tryToExpandPredicate(&StridePredicates);
+  assert(Pred && "Must be expandable!");
+
+  auto *StridesCheckBB = Plan.createVPBasicBlock("strides.check");
+  VPBasicBlock *ScalarPH = Plan.getScalarPreheader();
+  VPBlockUtils::insertBlockBefore(StridesCheckBB, Plan.getVectorPreheader());
+  VPBlockUtils::connectBlocks(StridesCheckBB, ScalarPH);
+  // SCEVExpander/VPSCEVExpander::expandCodeForPredicate negate the condition,
+  // so scalar preheader should be the first successor.
+  std::swap(StridesCheckBB->getSuccessors()[0],
+            StridesCheckBB->getSuccessors()[1]);
+  Builder.setInsertPoint(StridesCheckBB);
+  Builder.createNaryOp(VPInstruction::BranchOnCond, Pred);
+
+  for (VPRecipeBase &R : ScalarPH->phis()) {
+    auto &Phi = cast<VPPhi>(R);
+    Phi.addIncoming(Phi.getIncomingValueForBlock(Entry));
+  }
+
+  auto RewriteVPExpandSCEV = [&](VPExpandSCEVRecipe *R) {
+    const SCEV *S = R->getSCEV();
+    Builder.setInsertPoint(R);
+    const SCEV *NewS =
+        SE->rewriteUsingPredicate(S, CostCtx.L, StridePredicates);
+    if (NewS == S)
+      return;
+    auto *NewR = Builder.createExpandSCEV(NewS);
+    R->replaceAllUsesWith(NewR);
+
+    // If this recipe is a trip count then we need to reset it explicitly.
+    if (R == Plan.getTripCount())
+      Plan.resetTripCount(NewR);
+
+    R->eraseFromParent();
+  };
+
+  if (auto *R = dyn_cast<VPExpandSCEVRecipe>(Plan.getTripCount())) {
+    RewriteVPExpandSCEV(R);
+  }
----------------
artagnon wrote:

```suggestion
  if (auto *R = dyn_cast<VPExpandSCEVRecipe>(Plan.getTripCount()))
    RewriteVPExpandSCEV(R);
```

Could move after the rewrite block that follows?

https://github.com/llvm/llvm-project/pull/182595
_______________________________________________
llvm-branch-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/llvm-branch-commits

Reply via email to