[llvm] [CostModel] Move default expand cost for partial reductions to BasicTTIImpl (PR #189905)
Sander de Smalen via llvm-commits
llvm-commits at lists.llvm.org
Wed Apr 1 08:39:42 PDT 2026
https://github.com/sdesmalen-arm updated https://github.com/llvm/llvm-project/pull/189905
>From e62e10c72b628e9599748e6dbde142ac1e99928e Mon Sep 17 00:00:00 2001
From: Sander de Smalen <sander.desmalen at arm.com>
Date: Tue, 31 Mar 2026 14:31:15 +0000
Subject: [PATCH 1/3] [CostModel] Move default expand cost for partial
reductions to BasicTTIImpl.
This is a follow-up of the suggestion left here:
https://github.com/llvm/llvm-project/pull/181707#discussion_r2995733831
The override functions in AMDGPU/ARM/SystemZ/X86 are required to avoid
enabling partial reductions where they were previously disabled (I've added
this for all targets that implement getArithmeticReductionCost).
---
llvm/include/llvm/CodeGen/BasicTTIImpl.h | 41 +++++++++++++++++++
.../AArch64/AArch64TargetTransformInfo.cpp | 36 ++++------------
.../Target/AMDGPU/AMDGPUTargetTransformInfo.h | 9 ++++
llvm/lib/Target/ARM/ARMTargetTransformInfo.h | 9 ++++
.../SystemZ/SystemZTargetTransformInfo.h | 10 +++++
llvm/lib/Target/X86/X86TargetTransformInfo.h | 9 ++++
6 files changed, 86 insertions(+), 28 deletions(-)
diff --git a/llvm/include/llvm/CodeGen/BasicTTIImpl.h b/llvm/include/llvm/CodeGen/BasicTTIImpl.h
index 7812a301efbd7..02f054581529c 100644
--- a/llvm/include/llvm/CodeGen/BasicTTIImpl.h
+++ b/llvm/include/llvm/CodeGen/BasicTTIImpl.h
@@ -3435,6 +3435,47 @@ class BasicTTIImplBase : public TargetTransformInfoImplCRTPBase<T> {
return RedCost + MulCost + 2 * ExtCost;
}
+ InstructionCost getPartialReductionCost(
+ unsigned Opcode, Type *InputTypeA, Type *InputTypeB, Type *AccumType,
+ ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
+ TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
+ TTI::TargetCostKind CostKind,
+ std::optional<FastMathFlags> FMF) const override {
+ unsigned Ratio =
+ AccumType->getScalarSizeInBits() / InputTypeA->getScalarSizeInBits();
+ if (VF.getKnownMinValue() <= Ratio)
+ return InstructionCost::getInvalid();
+
+ Type *InputVectorType = VectorType::get(InputTypeA, VF);
+ Type *ExtInputVectorType = VectorType::get(AccumType, VF);
+ Type *AccumVectorType =
+ VectorType::get(AccumType, VF.divideCoefficientBy(Ratio));
+
+ auto ExtendCostA = InstructionCost(0);
+ if (OpAExtend != TTI::PartialReductionExtendKind::PR_None)
+ ExtendCostA = getCastInstrCost(
+ TTI::getOpcodeForPartialReductionExtendKind(OpAExtend),
+ ExtInputVectorType, InputVectorType, TTI::CastContextHint::None,
+ CostKind);
+
+ // TODO: add cost of extracting subvectors from the source vector that
+ // is to be partially reduced.
+ auto ReductionOpCost =
+ Ratio * getArithmeticInstrCost(Opcode, AccumVectorType, CostKind);
+
+ if (!BinOp)
+ return ExtendCostA + ReductionOpCost;
+
+ auto ExtendCostB = InstructionCost(0);
+ if (OpBExtend != TTI::PartialReductionExtendKind::PR_None)
+ ExtendCostB = getCastInstrCost(
+ TTI::getOpcodeForPartialReductionExtendKind(OpBExtend),
+ ExtInputVectorType, InputVectorType, TTI::CastContextHint::None,
+ CostKind);
+ return ExtendCostA + ExtendCostB + ReductionOpCost +
+ getArithmeticInstrCost(*BinOp, ExtInputVectorType, CostKind);
+ }
+
InstructionCost getVectorSplitCost() const { return 1; }
/// @}
diff --git a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
index de578ea29cbe9..e6eab7a3bdb87 100644
--- a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
+++ b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
@@ -6056,34 +6056,14 @@ InstructionCost AArch64TTIImpl::getPartialReductionCost(
return Cost * 2;
}
- // Returns cost of expanding the partial reduction in ISel.
- auto GetExpandCost = [&]() -> InstructionCost {
- Type *ExtVectorType =
- VectorType::get(AccumVectorType->getElementType(), VF);
- auto ExtendCostA = getCastInstrCost(
- TTI::getOpcodeForPartialReductionExtendKind(OpAExtend), ExtVectorType,
- InputVectorType, TTI::CastContextHint::None, CostKind);
- auto RedOpCost =
- Ratio * getArithmeticInstrCost(Opcode, AccumVectorType, CostKind);
- if (!BinOp)
- return ExtendCostA + RedOpCost;
-
- auto ExtendCostB = getCastInstrCost(
- TTI::getOpcodeForPartialReductionExtendKind(OpBExtend), ExtVectorType,
- InputVectorType, TTI::CastContextHint::None, CostKind);
- return ExtendCostA + ExtendCostB + RedOpCost +
- getArithmeticInstrCost(*BinOp, ExtVectorType, CostKind);
- };
-
- if (IsSub) {
- // Slightly lower the cost of a sub reduction so that it can be considered
- // as candidate for 'cdot' operations. This is a somewhat arbitrary number,
- // because we don't yet model these operations directly.
- return (8 * GetExpandCost()) / 10;
- }
-
- // By default, assume the operation is expanded.
- return GetExpandCost();
+ InstructionCost ExpandCost = BaseT::getPartialReductionCost(
+ Opcode, InputTypeA, InputTypeB, AccumType, VF, OpAExtend, OpBExtend,
+ BinOp, CostKind, FMF);
+
+ // Slightly lower the cost of a sub reduction so that it can be considered
+ // as candidate for 'cdot' operations. This is a somewhat arbitrary number,
+ // because we don't yet model these operations directly.
+ return IsSub ? ((8 * ExpandCost) / 10) : ExpandCost;
}
InstructionCost
diff --git a/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h b/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h
index ea2bf72836199..555c711a3b810 100644
--- a/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h
+++ b/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h
@@ -269,6 +269,15 @@ class GCNTTIImpl final : public BasicTTIImplBase<GCNTTIImpl> {
std::optional<FastMathFlags> FMF,
TTI::TargetCostKind CostKind) const override;
+ InstructionCost getPartialReductionCost(
+ unsigned Opcode, Type *InputTypeA, Type *InputTypeB, Type *AccumType,
+ ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
+ TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
+ TTI::TargetCostKind CostKind,
+ std::optional<FastMathFlags> FMF) const override {
+ return InstructionCost::getInvalid();
+ }
+
InstructionCost
getIntrinsicInstrCost(const IntrinsicCostAttributes &ICA,
TTI::TargetCostKind CostKind) const override;
diff --git a/llvm/lib/Target/ARM/ARMTargetTransformInfo.h b/llvm/lib/Target/ARM/ARMTargetTransformInfo.h
index f766deb884e0b..0d6d5d202bddf 100644
--- a/llvm/lib/Target/ARM/ARMTargetTransformInfo.h
+++ b/llvm/lib/Target/ARM/ARMTargetTransformInfo.h
@@ -413,6 +413,15 @@ class ARMTTIImpl final : public BasicTTIImplBase<ARMTTIImpl> {
getIntrinsicInstrCost(const IntrinsicCostAttributes &ICA,
TTI::TargetCostKind CostKind) const override;
+ InstructionCost getPartialReductionCost(
+ unsigned Opcode, Type *InputTypeA, Type *InputTypeB, Type *AccumType,
+ ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
+ TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
+ TTI::TargetCostKind CostKind,
+ std::optional<FastMathFlags> FMF) const override {
+ return InstructionCost::getInvalid();
+ }
+
/// getScalingFactorCost - Return the cost of the scaling used in
/// addressing mode represented by AM.
/// If the AM is supported, the return value must be >= 0.
diff --git a/llvm/lib/Target/SystemZ/SystemZTargetTransformInfo.h b/llvm/lib/Target/SystemZ/SystemZTargetTransformInfo.h
index d96036067c786..456604ef9f627 100644
--- a/llvm/lib/Target/SystemZ/SystemZTargetTransformInfo.h
+++ b/llvm/lib/Target/SystemZ/SystemZTargetTransformInfo.h
@@ -101,6 +101,16 @@ class SystemZTTIImpl final : public BasicTTIImplBase<SystemZTTIImpl> {
TTI::OperandValueInfo Op2Info = {TTI::OK_AnyValue, TTI::OP_None},
ArrayRef<const Value *> Args = {},
const Instruction *CxtI = nullptr) const override;
+
+ InstructionCost getPartialReductionCost(
+ unsigned Opcode, Type *InputTypeA, Type *InputTypeB, Type *AccumType,
+ ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
+ TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
+ TTI::TargetCostKind CostKind,
+ std::optional<FastMathFlags> FMF) const override {
+ return InstructionCost::getInvalid();
+ }
+
InstructionCost
getShuffleCost(TTI::ShuffleKind Kind, VectorType *DstTy, VectorType *SrcTy,
ArrayRef<int> Mask, TTI::TargetCostKind CostKind, int Index,
diff --git a/llvm/lib/Target/X86/X86TargetTransformInfo.h b/llvm/lib/Target/X86/X86TargetTransformInfo.h
index b3dde1555d0a0..b5124c3276896 100644
--- a/llvm/lib/Target/X86/X86TargetTransformInfo.h
+++ b/llvm/lib/Target/X86/X86TargetTransformInfo.h
@@ -224,6 +224,15 @@ class X86TTIImpl final : public BasicTTIImplBase<X86TTIImpl> {
std::optional<FastMathFlags> FMF,
TTI::TargetCostKind CostKind) const override;
+ InstructionCost getPartialReductionCost(
+ unsigned Opcode, Type *InputTypeA, Type *InputTypeB, Type *AccumType,
+ ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
+ TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
+ TTI::TargetCostKind CostKind,
+ std::optional<FastMathFlags> FMF) const override {
+ return InstructionCost::getInvalid();
+ }
+
InstructionCost getMinMaxCost(Intrinsic::ID IID, Type *Ty,
TTI::TargetCostKind CostKind,
FastMathFlags FMF) const;
>From 5f64abf238579a6e8a9ff566b81b7dbff847b70c Mon Sep 17 00:00:00 2001
From: Sander de Smalen <sander.desmalen at arm.com>
Date: Wed, 1 Apr 2026 08:17:28 +0000
Subject: [PATCH 2/3] Address CoPilot comments
---
llvm/include/llvm/CodeGen/BasicTTIImpl.h | 3 ++-
llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp | 2 +-
2 files changed, 3 insertions(+), 2 deletions(-)
diff --git a/llvm/include/llvm/CodeGen/BasicTTIImpl.h b/llvm/include/llvm/CodeGen/BasicTTIImpl.h
index 02f054581529c..cd94c5ff0e562 100644
--- a/llvm/include/llvm/CodeGen/BasicTTIImpl.h
+++ b/llvm/include/llvm/CodeGen/BasicTTIImpl.h
@@ -3443,7 +3443,8 @@ class BasicTTIImplBase : public TargetTransformInfoImplCRTPBase<T> {
std::optional<FastMathFlags> FMF) const override {
unsigned Ratio =
AccumType->getScalarSizeInBits() / InputTypeA->getScalarSizeInBits();
- if (VF.getKnownMinValue() <= Ratio)
+ if (VF.getKnownMinValue() <= Ratio || VF.getKnownMinValue() % Ratio != 0 ||
+ (BinOp && InputTypeA != InputTypeB))
return InstructionCost::getInvalid();
Type *InputVectorType = VectorType::get(InputTypeA, VF);
diff --git a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
index e6eab7a3bdb87..734339e5c7a05 100644
--- a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
+++ b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
@@ -6063,7 +6063,7 @@ InstructionCost AArch64TTIImpl::getPartialReductionCost(
// Slightly lower the cost of a sub reduction so that it can be considered
// as candidate for 'cdot' operations. This is a somewhat arbitrary number,
// because we don't yet model these operations directly.
- return IsSub ? ((8 * ExpandCost) / 10) : ExpandCost;
+ return ExpandCost.isValid() && IsSub ? ((8 * ExpandCost) / 10) : ExpandCost;
}
InstructionCost
>From 0bf032c44dc2a690ecdf13870c37a49c8318688e Mon Sep 17 00:00:00 2001
From: Sander de Smalen <sander.desmalen at arm.com>
Date: Wed, 1 Apr 2026 12:08:59 +0000
Subject: [PATCH 3/3] Address comments
---
llvm/include/llvm/CodeGen/BasicTTIImpl.h | 13 +++++++------
1 file changed, 7 insertions(+), 6 deletions(-)
diff --git a/llvm/include/llvm/CodeGen/BasicTTIImpl.h b/llvm/include/llvm/CodeGen/BasicTTIImpl.h
index cd94c5ff0e562..702437be7cf11 100644
--- a/llvm/include/llvm/CodeGen/BasicTTIImpl.h
+++ b/llvm/include/llvm/CodeGen/BasicTTIImpl.h
@@ -3441,10 +3441,11 @@ class BasicTTIImplBase : public TargetTransformInfoImplCRTPBase<T> {
TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
TTI::TargetCostKind CostKind,
std::optional<FastMathFlags> FMF) const override {
- unsigned Ratio =
- AccumType->getScalarSizeInBits() / InputTypeA->getScalarSizeInBits();
+ unsigned EltSizeAcc = AccumType->getScalarSizeInBits();
+ unsigned EltSizeInA = InputTypeA->getScalarSizeInBits();
+ unsigned Ratio = EltSizeAcc / EltSizeInA;
if (VF.getKnownMinValue() <= Ratio || VF.getKnownMinValue() % Ratio != 0 ||
- (BinOp && InputTypeA != InputTypeB))
+ EltSizeAcc % EltSizeInA != 0 || (BinOp && InputTypeA != InputTypeB))
return InstructionCost::getInvalid();
Type *InputVectorType = VectorType::get(InputTypeA, VF);
@@ -3452,7 +3453,7 @@ class BasicTTIImplBase : public TargetTransformInfoImplCRTPBase<T> {
Type *AccumVectorType =
VectorType::get(AccumType, VF.divideCoefficientBy(Ratio));
- auto ExtendCostA = InstructionCost(0);
+ InstructionCost ExtendCostA = 0;
if (OpAExtend != TTI::PartialReductionExtendKind::PR_None)
ExtendCostA = getCastInstrCost(
TTI::getOpcodeForPartialReductionExtendKind(OpAExtend),
@@ -3461,13 +3462,13 @@ class BasicTTIImplBase : public TargetTransformInfoImplCRTPBase<T> {
// TODO: add cost of extracting subvectors from the source vector that
// is to be partially reduced.
- auto ReductionOpCost =
+ InstructionCost ReductionOpCost =
Ratio * getArithmeticInstrCost(Opcode, AccumVectorType, CostKind);
if (!BinOp)
return ExtendCostA + ReductionOpCost;
- auto ExtendCostB = InstructionCost(0);
+ InstructionCost ExtendCostB = 0;
if (OpBExtend != TTI::PartialReductionExtendKind::PR_None)
ExtendCostB = getCastInstrCost(
TTI::getOpcodeForPartialReductionExtendKind(OpBExtend),
More information about the llvm-commits
mailing list