[llvm] [AArch64] Reflect cost of integer sub-reductions. (PR #194594)
via llvm-commits
llvm-commits at lists.llvm.org
Tue Apr 28 03:56:32 PDT 2026
llvmbot wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-transforms
Author: Sander de Smalen (sdesmalen-arm)
<details>
<summary>Changes</summary>
The cost of sub-reductions is either the cost of *mlslb + *mlslt, or the cost of a dot operation with 2 negations:
```
partial_reduce_umls acc, lhs, rhs
<=> -partial_reduce_umla -acc, lhs, rhs
```
(codegen for this was added by #<!-- -->186809)
The cost-model was previously a bit of a hack, since sub-reductions were expanded and therefore expensive, although we made the expansion cost artifically cheaper so that it would still be a candidate for cdot instructions.
---
Full diff: https://github.com/llvm/llvm-project/pull/194594.diff
3 Files Affected:
- (modified) llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp (+16-20)
- (modified) llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-chained.ll (+8-6)
- (modified) llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-sub-sdot.ll (+2-2)
``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
index aff89e00523c0..755321c65881c 100644
--- a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
+++ b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
@@ -6047,26 +6047,27 @@ InstructionCost AArch64TTIImpl::getPartialReductionCost(
bool IsSub = Opcode == Instruction::Sub;
InstructionCost Cost = InputLT.first * TTI::TCC_Basic;
+ InstructionCost INegCost = IsSub ? 2 * InputLT.first * TTI::TCC_Basic : 0;
if (AccumLT.second.getScalarType() == MVT::i32 &&
- InputLT.second.getScalarType() == MVT::i8 && !IsSub) {
+ InputLT.second.getScalarType() == MVT::i8) {
// i8 -> i32 is natively supported with udot/sdot for both NEON and SVE.
if (!IsUSDot && IsSupported(true, ST->hasDotProd()))
- return Cost;
+ return Cost + INegCost;
// i8 -> i32 usdot requires +i8mm
if (IsUSDot && IsSupported(ST->hasMatMulInt8(), ST->hasMatMulInt8()))
- return Cost;
+ return Cost + INegCost;
}
- if (ST->isSVEorStreamingSVEAvailable() && !IsUSDot && !IsSub) {
+ if (ST->isSVEorStreamingSVEAvailable() && !IsUSDot) {
// i16 -> i64 is natively supported for udot/sdot
if (AccumLT.second.getScalarType() == MVT::i64 &&
InputLT.second.getScalarType() == MVT::i16)
- return Cost;
+ return Cost + INegCost;
// i16 -> i32 is natively supported with SVE2p1
if (AccumLT.second.getScalarType() == MVT::i32 &&
InputLT.second.getScalarType() == MVT::i16 &&
- (ST->hasSVE2p1() || ST->hasSME2()))
+ (ST->hasSVE2p1() || ST->hasSME2()) && !IsSub)
return Cost;
// i8 -> i64 is supported with an extra level of extends
if (AccumLT.second.getScalarType() == MVT::i64 &&
@@ -6076,12 +6077,12 @@ InstructionCost AArch64TTIImpl::getPartialReductionCost(
// that now, a regular reduction would be cheaper because the costs of
// the extends in the IR are still counted. This can be fixed
// after https://github.com/llvm/llvm-project/pull/147302 has landed.
- return Cost;
+ return Cost + INegCost;
// i8 -> i16 is natively supported with SVE2p3
if (AccumLT.second.getScalarType() == MVT::i16 &&
InputLT.second.getScalarType() == MVT::i8 &&
- (ST->hasSVE2p3() || ST->hasSME2p3()))
- return Cost;
+ (ST->hasSVE2p3() || ST->hasSME2p3()) && !IsSub)
+ return Cost + INegCost;
}
// f16 -> f32 is natively supported for fdot using either
@@ -6092,11 +6093,11 @@ InstructionCost AArch64TTIImpl::getPartialReductionCost(
InputLT.second.getScalarType() == MVT::f16)
return Cost;
- // For a ratio of 2, we can use *mlal top/bottom instructions.
- if (Ratio == 2 && !IsSub) {
+ // For a ratio of 2, we can use *mlal and *mlsl top/bottom instructions.
+ if (Ratio == 2) {
MVT InVT = InputLT.second.getScalarType();
- // SVE2 [us]mlalb/t and NEON [us]mlal(2)
+ // SVE2 [us]ml[as]lb/t and NEON [us]ml[as]l(2)
if (IsSupported(ST->hasSVE2(), true) &&
llvm::is_contained({MVT::i8, MVT::i16, MVT::i32}, InVT.SimpleTy))
return Cost * 2;
@@ -6110,14 +6111,9 @@ InstructionCost AArch64TTIImpl::getPartialReductionCost(
return Cost * 2;
}
- 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 ExpandCost.isValid() && IsSub ? ((8 * ExpandCost) / 10) : ExpandCost;
+ return BaseT::getPartialReductionCost(Opcode, InputTypeA, InputTypeB,
+ AccumType, VF, OpAExtend, OpBExtend,
+ BinOp, CostKind, FMF);
}
InstructionCost
diff --git a/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-chained.ll b/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-chained.ll
index e54d1e9ddb1a3..b054d34fe597a 100644
--- a/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-chained.ll
+++ b/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-chained.ll
@@ -892,7 +892,7 @@ define i32 @chained_partial_reduce_sub_add_sub(ptr %a, ptr %b, ptr %c, i32 %N) #
; CHECK-SVE-MAXBW-NEXT: br label [[VECTOR_BODY:%.*]]
; CHECK-SVE-MAXBW: vector.body:
; CHECK-SVE-MAXBW-NEXT: [[INDEX:%.*]] = phi i64 [ 0, [[VECTOR_PH]] ], [ [[INDEX_NEXT:%.*]], [[VECTOR_BODY]] ]
-; CHECK-SVE-MAXBW-NEXT: [[VEC_PHI:%.*]] = phi <vscale x 8 x i32> [ zeroinitializer, [[VECTOR_PH]] ], [ [[TMP15:%.*]], [[VECTOR_BODY]] ]
+; CHECK-SVE-MAXBW-NEXT: [[VEC_PHI:%.*]] = phi <vscale x 2 x i32> [ zeroinitializer, [[VECTOR_PH]] ], [ [[PARTIAL_REDUCE4:%.*]], [[VECTOR_BODY]] ]
; CHECK-SVE-MAXBW-NEXT: [[TMP7:%.*]] = getelementptr inbounds nuw i8, ptr [[A]], i64 [[INDEX]]
; CHECK-SVE-MAXBW-NEXT: [[TMP8:%.*]] = getelementptr inbounds nuw i8, ptr [[B]], i64 [[INDEX]]
; CHECK-SVE-MAXBW-NEXT: [[TMP9:%.*]] = getelementptr inbounds nuw i8, ptr [[C]], i64 [[INDEX]]
@@ -901,18 +901,20 @@ define i32 @chained_partial_reduce_sub_add_sub(ptr %a, ptr %b, ptr %c, i32 %N) #
; CHECK-SVE-MAXBW-NEXT: [[WIDE_LOAD2:%.*]] = load <vscale x 8 x i8>, ptr [[TMP9]], align 1
; CHECK-SVE-MAXBW-NEXT: [[TMP13:%.*]] = sext <vscale x 8 x i8> [[WIDE_LOAD]] to <vscale x 8 x i32>
; CHECK-SVE-MAXBW-NEXT: [[TMP14:%.*]] = sext <vscale x 8 x i8> [[WIDE_LOAD1]] to <vscale x 8 x i32>
-; CHECK-SVE-MAXBW-NEXT: [[TMP12:%.*]] = sext <vscale x 8 x i8> [[WIDE_LOAD2]] to <vscale x 8 x i32>
; CHECK-SVE-MAXBW-NEXT: [[TMP10:%.*]] = mul nsw <vscale x 8 x i32> [[TMP13]], [[TMP14]]
-; CHECK-SVE-MAXBW-NEXT: [[TMP11:%.*]] = sub <vscale x 8 x i32> [[VEC_PHI]], [[TMP10]]
+; CHECK-SVE-MAXBW-NEXT: [[TMP11:%.*]] = sub <vscale x 8 x i32> zeroinitializer, [[TMP10]]
+; CHECK-SVE-MAXBW-NEXT: [[PARTIAL_REDUCE:%.*]] = call <vscale x 2 x i32> @llvm.vector.partial.reduce.add.nxv2i32.nxv8i32(<vscale x 2 x i32> [[VEC_PHI]], <vscale x 8 x i32> [[TMP11]])
+; CHECK-SVE-MAXBW-NEXT: [[TMP12:%.*]] = sext <vscale x 8 x i8> [[WIDE_LOAD2]] to <vscale x 8 x i32>
; CHECK-SVE-MAXBW-NEXT: [[TMP18:%.*]] = mul nsw <vscale x 8 x i32> [[TMP13]], [[TMP12]]
-; CHECK-SVE-MAXBW-NEXT: [[TMP16:%.*]] = add <vscale x 8 x i32> [[TMP11]], [[TMP18]]
+; CHECK-SVE-MAXBW-NEXT: [[PARTIAL_REDUCE3:%.*]] = call <vscale x 2 x i32> @llvm.vector.partial.reduce.add.nxv2i32.nxv8i32(<vscale x 2 x i32> [[PARTIAL_REDUCE]], <vscale x 8 x i32> [[TMP18]])
; CHECK-SVE-MAXBW-NEXT: [[TMP19:%.*]] = mul nsw <vscale x 8 x i32> [[TMP14]], [[TMP12]]
-; CHECK-SVE-MAXBW-NEXT: [[TMP15]] = sub <vscale x 8 x i32> [[TMP16]], [[TMP19]]
+; CHECK-SVE-MAXBW-NEXT: [[TMP15:%.*]] = sub <vscale x 8 x i32> zeroinitializer, [[TMP19]]
+; CHECK-SVE-MAXBW-NEXT: [[PARTIAL_REDUCE4]] = call <vscale x 2 x i32> @llvm.vector.partial.reduce.add.nxv2i32.nxv8i32(<vscale x 2 x i32> [[PARTIAL_REDUCE3]], <vscale x 8 x i32> [[TMP15]])
; CHECK-SVE-MAXBW-NEXT: [[INDEX_NEXT]] = add nuw i64 [[INDEX]], [[TMP3]]
; CHECK-SVE-MAXBW-NEXT: [[TMP22:%.*]] = icmp eq i64 [[INDEX_NEXT]], [[N_VEC]]
; CHECK-SVE-MAXBW-NEXT: br i1 [[TMP22]], label [[MIDDLE_BLOCK:%.*]], label [[VECTOR_BODY]], !llvm.loop [[LOOP12:![0-9]+]]
; CHECK-SVE-MAXBW: middle.block:
-; CHECK-SVE-MAXBW-NEXT: [[TMP17:%.*]] = call i32 @llvm.vector.reduce.add.nxv8i32(<vscale x 8 x i32> [[TMP15]])
+; CHECK-SVE-MAXBW-NEXT: [[TMP16:%.*]] = call i32 @llvm.vector.reduce.add.nxv2i32(<vscale x 2 x i32> [[PARTIAL_REDUCE4]])
; CHECK-SVE-MAXBW-NEXT: [[CMP_N:%.*]] = icmp eq i64 [[WIDE_TRIP_COUNT]], [[N_VEC]]
; CHECK-SVE-MAXBW-NEXT: br i1 [[CMP_N]], label [[FOR_COND_CLEANUP:%.*]], label [[SCALAR_PH]]
; CHECK-SVE-MAXBW: scalar.ph:
diff --git a/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-sub-sdot.ll b/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-sub-sdot.ll
index ff3881f2c7fda..d40b984618d45 100644
--- a/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-sub-sdot.ll
+++ b/llvm/test/Transforms/LoopVectorize/AArch64/partial-reduce-sub-sdot.ll
@@ -15,9 +15,9 @@
; COMMON: LV: Checking a loop in 'add_sub_chained_reduction'
; SVE: Cost of 1 for VF vscale x 16: EXPRESSION vp<{{.*}}> = ir<%acc> + partial.reduce.add (mul (ir<%load1> sext to i32), (ir<%load2> sext to i32))
-; SVE: Cost of 16 for VF vscale x 16: EXPRESSION vp<{{.*}}> = vp<%9> + partial.reduce.add (sub (0, mul (ir<%load2> sext to i32), (ir<%load3> sext to i32)))
+; SVE: Cost of 3 for VF vscale x 16: EXPRESSION vp<{{.*}}> = vp<%9> + partial.reduce.add (sub (0, mul (ir<%load2> sext to i32), (ir<%load3> sext to i32)))
; NEON: Cost of 1 for VF 16: EXPRESSION vp<{{.*}}> = ir<%acc> + partial.reduce.add (mul (ir<%load1> sext to i32), (ir<%load2> sext to i32))
-; NEON: Cost of 16 for VF 16: EXPRESSION vp<{{.*}}> = vp<%9> + partial.reduce.add (sub (0, mul (ir<%load2> sext to i32), (ir<%load3> sext to i32)))
+; NEON: Cost of 3 for VF 16: EXPRESSION vp<{{.*}}> = vp<%9> + partial.reduce.add (sub (0, mul (ir<%load2> sext to i32), (ir<%load3> sext to i32)))
target triple = "aarch64"
``````````
</details>
https://github.com/llvm/llvm-project/pull/194594
More information about the llvm-commits
mailing list