[llvm] [AMDGPU] Price the packed form of a vector of i1 (PR #217327)

via llvm-commits llvm-commits at lists.llvm.org
Tue Aug 25 06:02:17 PDT 2026


================
@@ -980,13 +980,57 @@ InstructionCost GCNTTIImpl::getCFInstrCost(unsigned Opcode,
   return BaseT::getCFInstrCost(Opcode, CostKind, I);
 }
 
+// A vector of i1 has no packed form on this target: every element lives in its
+// own mask. Measured instruction counts per element to pack the masks into an
+// integer, and to unpack them again.
+static constexpr unsigned MaskPackCostPerElt = 4;
+static constexpr unsigned MaskUnpackCostPerElt = 3;
+
+/// Returns the number of elements when \p Ty is a fixed vector of i1 with more
+/// than one element.
+static std::optional<unsigned> getPackedMaskElts(Type *Ty) {
+  auto *FVT = dyn_cast<FixedVectorType>(Ty);
+  if (FVT && FVT->getElementType()->isIntegerTy(1) && FVT->getNumElements() > 1)
+    return FVT->getNumElements();
+  return std::nullopt;
+}
+
+InstructionCost GCNTTIImpl::getCastInstrCost(unsigned Opcode, Type *Dst,
+                                             Type *Src,
+                                             TTI::CastContextHint CCH,
+                                             TTI::TargetCostKind CostKind,
+                                             const Instruction *I) const {
+  // A bitcast between a vector of i1 and an integer packs or unpacks a mask.
+  if (Opcode == Instruction::BitCast) {
+    if (std::optional<unsigned> Elts = getPackedMaskElts(Src);
+        Elts && Dst->isIntegerTy())
+      return InstructionCost(MaskPackCostPerElt) * *Elts *
+             getFullRateInstrCost();
+    if (std::optional<unsigned> Elts = getPackedMaskElts(Dst);
+        Elts && Src->isIntegerTy())
+      return InstructionCost(MaskUnpackCostPerElt) * *Elts *
+             getFullRateInstrCost();
+  }
+
+  return BaseT::getCastInstrCost(Opcode, Dst, Src, CCH, CostKind, I);
+}
+
 InstructionCost
 GCNTTIImpl::getArithmeticReductionCost(unsigned Opcode, VectorType *Ty,
                                        std::optional<FastMathFlags> FMF,
                                        TTI::TargetCostKind CostKind) const {
   if (TTI::requiresOrderedReduction(FMF))
     return BaseT::getArithmeticReductionCost(Opcode, Ty, FMF, CostKind);
 
+  // These three reductions over a vector of i1 go through the packed form of
+  // the mask, and the packing dominates their cost.
+  if (Opcode == Instruction::Add || Opcode == Instruction::And ||
+      Opcode == Instruction::Or) {
----------------
michaelselehov wrote:

The premise is right, and the conclusion does not follow. Let me split the two.

A direct add reduction over a vector of i1 really is a parity chain, and it
never packs. I measured it from 2 to 16 elements on gfx1030: 1.0 to 1.88
instructions per element, so charging 4 overprices it. The patch says so, in the
code and in the description.

Removing `Add` brings the bug back. SLP does not build the intrinsic. In
`emitReduction()` it builds `ctpop(bitcast <n x i1> to iN)` itself, and it prices
that tree through `getArithmeticReductionCost`, not through
`getExtendedReductionCost`. The dispatch is in `SLPVectorizer.cpp`: the extended
hook is used only when `RType != RedTy`, and for this tree the two are equal.

I tested your suggestion, with the new bitcast cost in place:

| with `Add` in the hook | without it |
|---|---|
| cost of `reduce.add` over `<8 x i1>`: 32 | 15 |
| packing sites in the hot rocSPARSE `csrgemm` kernel: 0 | 3 |
| the SLP test keeps the scalar chain | it builds the `ctpop` form again |

I also tested a cast-only version earlier, with no reduction override at all: the
whole `csrgemm` module came out byte for byte the same as with no patch, 516
packing sites both ways. A cast cost never reaches this decision.

So SLP prices one form and emits another. That is the real defect, it needs a fix
in SLP, and it cannot be fixed on the cost side alone, because every target
answers that query with a different number today. Until then the target has to
price the form that SLP actually emits.


https://github.com/llvm/llvm-project/pull/217327


More information about the llvm-commits mailing list