[llvm] [AArch64] SVE Shuffleopt: merge reduction reverse into tbl (PR #206047)
Paul Walker via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 25 05:32:52 PDT 2026
================
@@ -19380,6 +19380,194 @@ bool AArch64TargetLowering::optimizeExtendOrTruncateConversion(
return false;
}
+// Match a bitcasted tbl intrinsic, and bind the tbl along with the mask.
+static auto m_Tbl(Instruction *&Tbl, Value *&Mask) {
+ using namespace llvm::PatternMatch;
+ return m_OneUse(m_BitCast(
+ m_OneUse(m_Instruction(Tbl, m_Intrinsic<Intrinsic::aarch64_sve_tbl>(
+ m_Value(), m_Value(Mask))))));
+}
+
+// Match a tbl intrinsic whose result is converted to a floating point value.
+static auto m_UIToFPTbl(Instruction *&Tbl, Value *&Mask) {
+ using namespace llvm::PatternMatch;
+ return m_OneUse(m_UIToFP(m_Tbl(Tbl, Mask)));
+}
+
+// Match either of the above tbls, and recalculate the index of the
+// deinterleaved subvector. Bind the tbl and the index.
+struct deinterleaving_tbl_match {
+ Instruction *&Tbl;
+ unsigned &Idx;
+
+ deinterleaving_tbl_match(Instruction *&Tbl, unsigned &Idx)
+ : Tbl(Tbl), Idx(Idx) {}
+
+ template <typename ITy> bool match(ITy *V) const {
+ using namespace llvm::PatternMatch;
+
+ // Match the tbl.
+ Value *Mask;
+ if (!PatternMatch::match(
+ V, m_CombineOr(m_Tbl(Tbl, Mask), m_UIToFPTbl(Tbl, Mask))))
+ return false;
+
+ // For a deinterleaving+extending tbl, we will have a known constant values
+ // for the starting index and the step.
+ const APInt *Start;
+ if (!PatternMatch::match(
+ Mask, m_BitCast(m_Add(m_Mul(m_Intrinsic<Intrinsic::stepvector>(),
+ m_SpecificInt(4)),
+ m_APInt(Start)))))
+ return false;
+
+ unsigned SrcSize = Tbl->getType()->getScalarType()->getScalarSizeInBits();
+ unsigned ResSize = V->getType()->getScalarType()->getScalarSizeInBits();
+ // If the top bits are all ones, we know we're forcing an out-of-range
+ // index. With a deinterleave of 4, we should have 3 invalid indices for
+ // every valid one.
+ if (Start->countLeadingOnes() != ResSize - SrcSize)
+ return false;
+
+ // The start of the valid indices must be between 0 and 3, for the 4
+ // subvectors we're extracting.
+ Idx = Start->getZExtValue() & SrcSize - 1;
+ return Idx >= 0 && Idx < 4;
+ }
+};
+
+static auto m_DeinterleavingTbl(Instruction *&Tbl, unsigned &Idx) {
+ return deinterleaving_tbl_match(Tbl, Idx);
+}
+
+// We want to find a reverse used only by BinOps where the other term comes
+// from one of the deinterleave-and-extend tbls we created before. If this
+// BinOp is only used in a reduction operation in the loop (so the order of
+// elements within it do not matter), then we can potentially fold the reverse
+// into the tbl and remove the separate reverse operation.
+// Something like the following (possibly repeated multiple times):
+//
+// %acc.b.f64 = phi <vscale x 2 x double> [ splat(double 0.000000e+00),
+// %entry ], [ %fadd.b.f64, %loop ]
+// ...
+// %rev.load = load <vscale x 2 x double>, ptr %rev.ptr
+// %reversed = call <vscale x 2 x double> @llvm.vector.reverse.nxv2f64(
+// <vscale x 2 x double> %rev.load)
+// %bgra = load <vscale x 8 x i16>, ptr %src.gep
+// %stepvec = call <vscale x 2 x i64> @llvm.stepvector.nxv2i64()
+// %stride = mul nuw <vscale x 2 x i64> %stepvec, splat (i64 4)
+// %start = add nuw <vscale x 2 x i64> %stride, splat (i64 -65536)
+// %bc.to = bitcast <vscale x 2 x i64> %start to <vscale x 8 x i16>
+// %tbl = call <vscale x 8 x i16> @llvm.aarch64.sve.tbl.nxv8i16(
+// <vscale x 8 x i16> %bgra, <vscale x 8 x i16> %bc.to)
+// %bc.from = bitcast <vscale x 8 x i16> %tbl to <vscale x 2 x i64>
+// %b.f64 = uitofp <vscale x 2 x i64> %bc.from to <vscale x 2 x double>
+// %b.mul.f64 = fmul <vscale x 2 x double> %b.f64, %reversed
+// %fadd.b.f64 = call <vscale x 2 x double>
+// @llvm.vector.partial.reduce.fadd(%acc.b.f64, %b.mul.f64)
+static bool foldAdjacentReversesIntoTbls(Instruction *I) {
+ using namespace llvm::PatternMatch;
+ struct RevTblData {
+ Instruction *Tbl;
+ Use *RevUse;
+ unsigned Idx;
+ };
+
+ // Look for reverse intrinsics used with the results of tbl instructions.
+ SmallVector<RevTblData, 4> Tbls;
+ // Check all uses of the reverse; if there's anything which doesn't
+ // match our expected patterns, then give up on it. The intention is
+ // to remove the reverse, so any remaining users would prevent that.
+ for (Use &U : I->uses()) {
+ Instruction *UI = cast<Instruction>(U.getUser());
+
+ // Look for a tbl used to deinterleave and zero extend (and optionally
+ // convert to FP), the result of which is then used in a BinOp with
+ // the reverse.
+ Instruction *Tbl;
+ unsigned Idx;
+ if (!match(UI, m_OneUse(m_c_BinOp(m_DeinterleavingTbl(Tbl, Idx),
+ m_Specific(I)))))
+ return false;
+
+ // Check for a partial reduction user. Partial reductions permit reordering
+ // (and more) of the overall result within a vector, so it won't matter if
+ // we reverse the result.
+ if (!match(UI->getSingleUndroppableUse()->getUser(),
+ m_Intrinsic<Intrinsic::vector_partial_reduce_fadd>()))
+ return false;
+
+ Tbls.push_back({Tbl, &U, Idx});
+ }
+
+ // Convert each candidate we found to perform the reverse in the tbl instead,
+ // then remove the reverse.
+ for (auto [Tbl, RevUse, Idx] : Tbls) {
+ Instruction *Rev = cast<Instruction>(RevUse->get());
+ VectorType *SrcTy = cast<VectorType>(Tbl->getType());
+ VectorType *DstTy = cast<VectorType>(Rev->getType());
+ unsigned SrcBits = SrcTy->getScalarSizeInBits();
+ unsigned DstBits = DstTy->getScalarSizeInBits();
+ VectorType *StepVecTy = VectorType::getInteger(DstTy);
+
+ // We need to create a new mask for the tbl, so that we effectively reverse
+ // the elements with the tbl. This means the resulting vector will be
+ // backwards compared to the original, which is why we check for reductions
+ // where we don't care about the order.
+ //
+ // Similar to the normal deinterleaving+extending mask, we will use out-of
+ // range indices to perform the zero extension. However, we need a negative
+ // stride starting from the highest group of elements. Since we're grouping
+ // by 4 (for now), we need to subtract 4 from the total source element count
+ // for the vector to get the start, then add the index of the extraction.
+ //
+ // The result should be the following:
+ // <EltCnt - 4 + Idx, 0xFFFF, 0xFFFF, 0xFFFF, EltCnt - 8 + Idx, 0xFFFF...>.
+ IRBuilder<> Builder(Tbl);
+
+ // Create and splat the starting value.
+ APInt Invalid = APInt::getAllOnes(DstBits);
+ APInt StartIdx = Invalid << SrcBits;
+ StartIdx -= (4 - Idx);
+ Value *EltCnt = Builder.CreateVScale(StepVecTy->getScalarType());
+ EltCnt = Builder.CreateNUWMul(
+ EltCnt, ConstantInt::get(EltCnt->getType(),
+ AArch64::SVEBitsPerBlock / SrcBits));
+ Value *StartVal = Builder.CreateNUWAdd(
+ EltCnt, ConstantInt::get(EltCnt->getType(), StartIdx));
+ StartVal =
+ Builder.CreateVectorSplat(StepVecTy->getElementCount(), StartVal);
+
+ // Create the negative stride.
+ Value *StepVector = Builder.CreateStepVector(StepVecTy);
+ Value *ScaledSteps =
+ Builder.CreateMul(StepVector, ConstantInt::get(StepVecTy, -4));
+
+ // Add the start to the stride, replace the old mask.
+ ScaledSteps = Builder.CreateNUWAdd(ScaledSteps, StartVal);
----------------
paulwalker-arm wrote:
ScaledSteps is negative so the add is relying on wrapping behaviour and thus cannot be `nuw`?
https://github.com/llvm/llvm-project/pull/206047
More information about the llvm-commits
mailing list