[llvm] Reapply "[X86] EltsFromConsecutiveLoads - handle trunc(wideload()) patterns" (#199371) (PR #208999)
Simon Pilgrim via llvm-commits
llvm-commits at lists.llvm.org
Mon Aug 10 03:05:52 PDT 2026
================
@@ -8071,6 +8071,95 @@ static SDValue EltsFromConsecutiveLoads(EVT VT, ArrayRef<SDValue> Elts,
}
}
+ // STRIDED - element loads at a uniform byte stride larger than the element
+ // size are folded into wide load(s) + vector truncation.
+ // Depth 1 lets the REVERSE block below recurse into us to catch reverse
+ // strides.
+ if (Depth <= 1 && Subtarget.hasAVX2() && !IsConsecutiveLoad &&
+ LoadMask.isAllOnes() && isPowerOf2_32(NumElems) && BaseSizeInBits <= 32 &&
+ VT.getSizeInBits() >= 128) {
+ unsigned WideEltBits = 0;
+ for (unsigned Trial : {16u, 32u, 64u}) {
+ if (Trial <= BaseSizeInBits || Trial % BaseSizeInBits != 0)
+ continue;
+ unsigned LaneStride = Trial / BaseSizeInBits;
+ bool AllMatch = true;
+ for (unsigned K = 1; K < NumElems && AllMatch; ++K) {
+ AllMatch = AllMatch && ByteOffsets[K] == 0 &&
+ DAG.areNonVolatileConsecutiveLoads(
+ Loads[K], LDBase, BaseSizeInBytes, K * LaneStride);
+ }
+ if (AllMatch) {
+ WideEltBits = Trial;
+ break;
+ }
+ }
+ if (WideEltBits != 0) {
+ MVT WideEltVT = MVT::getIntegerVT(WideEltBits);
+ MVT SrcEltVT = MVT::getIntegerVT(BaseSizeInBits);
+ // The wide loads read (stride - element) bytes past the last element.
+ // A stride-aligned base keeps every stride block inside a page that an
+ // accessed element touches.
+ unsigned StrideBytes = WideEltBits / 8;
+ bool RangeSafe =
+ LDBase->getBaseAlign() >= Align(StrideBytes) ||
+ LDBase->getPointerInfo().isDereferenceable(
+ NumElems * StrideBytes, *DAG.getContext(), DAG.getDataLayout());
+ for (unsigned WideRegBits : {512u, 256u, 128u}) {
+ unsigned LanesPerWideLoad = WideRegBits / WideEltBits;
+ if (LanesPerWideLoad < 2 || NumElems % LanesPerWideLoad != 0)
+ continue;
+ // Truncate to a legal piece at least 128-bit and no narrower
+ // than the source element.
+ unsigned PieceEltBits =
+ std::max(BaseSizeInBits, 128 / LanesPerWideLoad);
+ MVT PieceEltVT = MVT::getIntegerVT(PieceEltBits);
+ MVT WideVT = MVT::getVectorVT(WideEltVT, LanesPerWideLoad);
+ MVT PieceVT = MVT::getVectorVT(PieceEltVT, LanesPerWideLoad);
+ MVT TruncVT = MVT::getVectorVT(SrcEltVT, NumElems);
+ MVT ConcatVT = MVT::getVectorVT(PieceEltVT, NumElems);
+ if (!TLI.isTypeLegal(WideVT) || !TLI.isTypeLegal(PieceVT) ||
+ !TLI.isTypeLegal(ConcatVT) || !TLI.isTypeLegal(TruncVT))
+ continue;
+ unsigned NumWideLoads = NumElems / LanesPerWideLoad;
+ unsigned BytesPerWideLoad = WideRegBits / 8;
+ // Without a provable range the last load is shifted back to end at
+ // the last element. It overlaps the previous load so two are needed.
+ bool CanShift = NumWideLoads >= 2;
+ if (!RangeSafe && !CanShift)
+ continue;
+ unsigned TailBits = WideEltBits - BaseSizeInBits;
+ auto MMOFlags = LDBase->getMemOperand()->getFlags();
+ SDValue BasePtr = LDBase->getBasePtr();
+ SmallVector<SDValue, 8> Pieces;
+ for (unsigned K = 0; K != NumWideLoads; ++K) {
+ bool ShiftLast = !RangeSafe && K == NumWideLoads - 1;
+ unsigned Offset = K * BytesPerWideLoad;
+ if (ShiftLast)
+ Offset -= TailBits / 8;
+ SDValue Ptr =
+ DAG.getMemBasePlusOffset(BasePtr, TypeSize::getFixed(Offset), DL);
+ Align LdAlign = commonAlignment(LDBase->getBaseAlign(), Offset);
+ SDValue Ld =
+ DAG.getLoad(WideVT, DL, LDBase->getChain(), Ptr,
+ LDBase->getPointerInfo().getWithOffset(Offset),
+ LdAlign, MMOFlags);
+ for (auto *LD : Loads)
+ if (LD)
+ DAG.makeEquivalentMemoryOrdering(LD, Ld);
+ // The shifted load leaves each element in its lane's high bits.
+ if (ShiftLast)
+ Ld = DAG.getNode(ISD::SRL, DL, WideVT, Ld,
+ DAG.getConstant(TailBits, DL, WideVT));
----------------
RKSimon wrote:
getShiftAmountConstant
https://github.com/llvm/llvm-project/pull/208999
More information about the llvm-commits
mailing list