[llvm] [SLP] Cost reduce.add(mul(ext, ext)) as a dot-product reduction (PR #224066)
Alexey Bataev via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 17 03:58:15 PDT 2026
================
@@ -32305,6 +32305,85 @@ class HorizontalReduction {
cast<VectorType>(getWidenedType(RType, ReduxWidth)), FMF,
CostKind);
}
+ // reduce.add(mul(ext(A), ext(B))) lowers to a single dot-product
+ // reduction (e.g. UDOT/SDOT) where the target supports it. Prefer
+ // that fused cost when it is cheaper. The extends and the multiply
+ // are already counted in the tree cost, so subtract them here to
+ // avoid double counting (mirrors the FMA handling below).
+ if (RdxKind == RecurKind::Add && !ReducedVals.empty()) {
+ Type *SrcElemTy = nullptr;
+ bool IsZExt = true;
+ bool SameOperands = true;
+ // Match one reduced lane as mul(ext(a), ext(b)) where both
+ // factors use the same widening extend. Reports the extend
+ // signedness, the pre-extension scalar type, and whether both
+ // factors are the same extend value (a single extend column).
+ auto MatchMulAccLane = [](Value *V, bool &ZExt, Type *&SrcTy,
+ bool &SharedExt) {
+ Value *E0, *E1, *A, *B;
+ if (match(V,
+ m_Mul(m_CombineAnd(m_Value(E0), m_ZExt(m_Value(A))),
+ m_CombineAnd(m_Value(E1), m_ZExt(m_Value(B))))))
+ ZExt = true;
+ else if (match(V, m_Mul(m_CombineAnd(m_Value(E0),
+ m_SExt(m_Value(A))),
+ m_CombineAnd(m_Value(E1),
+ m_SExt(m_Value(B))))))
+ ZExt = false;
+ else
+ return false;
+ if (A->getType() != B->getType())
+ return false;
+ SrcTy = A->getType()->getScalarType();
+ SharedExt = E0 == E1;
+ return true;
+ };
+ bool IsMulAcc = all_of(ReducedVals, [&](Value *RdxVal) {
+ bool ThisZExt;
+ Type *ThisSrcTy;
+ bool SharedExt;
+ if (!MatchMulAccLane(RdxVal, ThisZExt, ThisSrcTy, SharedExt))
+ return false;
+ if (!SharedExt)
+ SameOperands = false;
+ if (!SrcElemTy) {
+ SrcElemTy = ThisSrcTy;
+ IsZExt = ThisZExt;
+ return true;
+ }
+ return SrcElemTy == ThisSrcTy && IsZExt == ThisZExt;
+ });
+ // Only fuse when the multiply factors are actually vectorized
+ // (a live non-gather tree entry); otherwise the dot-product does
+ // not form and the fused cost would not apply.
+ if (IsMulAcc && SrcElemTy &&
+ R.isVectorized(ReducedVals.front())) {
+ auto *SrcVecTy =
+ cast<VectorType>(getWidenedType(SrcElemTy, ReduxWidth));
+ InstructionCost RedCost = TTI->getMulAccReductionCost(
+ IsZExt, RdxOpcode, RedTy, SrcVecTy, CostKind);
+ if (RedCost.isValid()) {
+ // Derive operand info and cast context from the actual
+ // mul(ext, ext) so the subtracted costs match the tree.
+ auto *Mul = cast<Instruction>(ReducedVals.front());
+ auto *Ext0 = cast<Instruction>(Mul->getOperand(0));
+ auto *Ext1 = cast<Instruction>(Mul->getOperand(1));
----------------
alexey-bataev wrote:
Better to operate on TreeEntries here, not single instruction, we may see non-matching operand info here
https://github.com/llvm/llvm-project/pull/224066
More information about the llvm-commits
mailing list