[llvm] [X86][CostModel] Add per-shape gather/scatter cost tables for AMD znver4+ (PR #199488)
Sumukh J Bharadwaj via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 17 03:47:14 PDT 2026
================
@@ -6588,11 +6590,191 @@ InstructionCost X86TTIImpl::getCFInstrCost(unsigned Opcode,
return TTI::TCC_Free;
}
+// Pick a representative masked gather/scatter opcode for a data and index
+// shape, used to read the body cost from the schedule model. The integer and FP
+// variants of a shape share a scheduling class, so the integer form is
+// returned. Returns 0 for unsupported shapes.
+static unsigned getAVX512GSRepresentativeOpcode(bool IsLoad, MVT VT,
+ unsigned IndexBits) {
+ if (!VT.isVector())
+ return 0;
+ unsigned NumElts = VT.getVectorNumElements();
+ unsigned EltBits = VT.getScalarSizeInBits();
+ if (IndexBits != 32 && IndexBits != 64)
+ return 0;
+
+ if (IsLoad) {
+ if (EltBits == 32) {
+ if (IndexBits == 32)
+ return NumElts == 4 ? X86::VPGATHERDDZ128rm
+ : NumElts == 8 ? X86::VPGATHERDDZ256rm
+ : NumElts == 16 ? X86::VPGATHERDDZrm
+ : 0;
+ return NumElts == 4 ? X86::VPGATHERQDZ256rm
+ : NumElts == 8 ? X86::VPGATHERQDZrm
+ : 0;
+ }
+ if (EltBits == 64) {
+ if (IndexBits == 32)
+ return NumElts == 4 ? X86::VPGATHERDQZ256rm
+ : NumElts == 8 ? X86::VPGATHERDQZrm
+ : 0;
+ return NumElts == 4 ? X86::VPGATHERQQZ256rm
+ : NumElts == 8 ? X86::VPGATHERQQZrm
+ : 0;
+ }
+ return 0;
+ }
+ if (EltBits == 32) {
+ if (IndexBits == 32)
+ return NumElts == 4 ? X86::VPSCATTERDDZ128mr
+ : NumElts == 8 ? X86::VPSCATTERDDZ256mr
+ : NumElts == 16 ? X86::VPSCATTERDDZmr
+ : 0;
+ return NumElts == 4 ? X86::VPSCATTERQDZ256mr
+ : NumElts == 8 ? X86::VPSCATTERQDZmr
+ : 0;
+ }
+ if (EltBits == 64) {
+ if (IndexBits == 32)
+ return NumElts == 4 ? X86::VPSCATTERDQZ256mr
+ : NumElts == 8 ? X86::VPSCATTERDQZmr
+ : 0;
+ return NumElts == 4 ? X86::VPSCATTERQQZ256mr
+ : NumElts == 8 ? X86::VPSCATTERQQZmr
+ : 0;
+ }
+ return 0;
+}
+
+// Read an opcode's reciprocal-throughput body cost from this subtarget's
+// schedule model, or nullopt (caller falls back) when there is no
+// per-instruction model, the class is invalid or a variant (variants need a
+// real MachineInstr), or there is no real per-shape override.
+static std::optional<unsigned> getSchedModelGSBody(unsigned Opc,
+ const X86Subtarget *ST) {
+ const MCSchedModel &SM = ST->getSchedModel();
+ if (!SM.hasInstrSchedModel())
+ return std::nullopt;
+
+ const auto *TII = ST->getInstrInfo();
+ auto ReciprocalThroughputOf = [&](unsigned Opcode) -> std::optional<double> {
+ unsigned SClassID = TII->get(Opcode).getSchedClass();
+ const MCSchedClassDesc *SCDesc = SM.getSchedClassDesc(SClassID);
+ if (!SCDesc || !SCDesc->isValid() || SCDesc->isVariant())
+ return std::nullopt;
+ return MCSchedModel::getReciprocalThroughput(*ST, *SCDesc);
+ };
+
+ std::optional<double> RThru = ReciprocalThroughputOf(Opc);
+ if (!RThru)
+ return std::nullopt;
+
+ // A valid class alone is not enough: an unmodelled op also has one, and
+ // would round to 0. The cheapest a real gather/scatter can be is a plain
+ // vector load, so anything at or below the load's throughput is the generic
+ // default in disguise. Reading the baseline from the model avoids a magic
+ // floor.
+ std::optional<double> LoadRThru = ReciprocalThroughputOf(X86::VMOVUPSZrm);
+ if (!LoadRThru || *RThru <= *LoadRThru)
+ return std::nullopt;
+
+ return static_cast<unsigned>(std::lround(*RThru));
+}
+
+// Hardware body cost of a single native-width masked gather/scatter, read live
+// from the schedule model; getZenGSOverhead adds the profitability premium.
+std::optional<unsigned>
+X86TTIImpl::getModeledGSInstrCost(bool IsLoad, Type *SrcVTy,
+ unsigned IndexSize,
+ TTI::TargetCostKind CostKind) const {
+ if (CostKind != TTI::TCK_RecipThroughput || !ST->hasAVX512() ||
+ !ST->hasPreferGSCostTable() || !SrcVTy)
+ return std::nullopt;
+ EVT VT = TLI->getValueType(DL, SrcVTy);
+ if (!VT.isSimple())
+ return std::nullopt;
+ // An encoding whose data and index operands are both sub-512-bit requires
+ // AVX512VL. A v8i32 operation with i64 indices still uses a full zmm index
+ // and therefore does not.
+ unsigned IndexVectorBits = IndexSize * VT.getVectorNumElements();
+ if (VT.getSizeInBits() < 512 && IndexVectorBits < 512 && !ST->hasVLX())
+ return std::nullopt;
----------------
amd-subharad wrote:
this has been fixed
https://github.com/llvm/llvm-project/pull/199488
More information about the llvm-commits
mailing list