#ifndef GEMMSTONE_GENERATOR_PIECES_STATE_HPP
#define GEMMSTONE_GENERATOR_PIECES_STATE_HPP
#include "internal/ngen_includes.hpp"
#include "gemmstone/driver_info.hpp"
#include "gemmstone/problem.hpp"
#include "gemmstone/strategy.hpp"
#include "gemmstone/type.hpp"
#include "register_layout.hpp"
#include "allocators.hpp"
#include "ngen_emulation.hpp"
#include "grf_multirange.hpp"
GEMMSTONE_NAMESPACE_START
class SubregisterPair {
protected:
ngen::Subregister regs[2];
bool negative;
public:
SubregisterPair() : SubregisterPair(ngen::Subregister()) {}
SubregisterPair(ngen::Subregister reg0, ngen::Subregister reg1) : regs{reg0, reg1}, negative(false) {}
explicit SubregisterPair(ngen::Subregister reg) : SubregisterPair(reg, reg) {}
operator ngen::Subregister() const { return regs[0]; }
SubregisterPair &operator=(ngen::Subregister reg) { regs[0] = regs[1] = reg; negative = false; return *this; }
ngen::Subregister getReg(int idx) const;
ngen::Subregister getRegAvoiding(ngen::HW hw, const ngen::RegData &rd) const;
bool isValid() const { return regs[0].isValid() && regs[1].isValid(); }
bool isInvalid() const { return !isValid(); }
void invalidate() { regs[0].invalidate(); regs[1].invalidate(); }
bool isDuplicated() const { return regs[0] != regs[1]; }
SubregisterPair operator-() const {
auto copy = *this;
copy.negative = !copy.negative;
return copy;
}
};
class MultishiftSubregister {
protected:
static constexpr int maxShift = 5;
ngen::Subregister regs[maxShift + 1] = {ngen::Subregister()};
bool neg = false;
public:
MultishiftSubregister operator-() const {
auto copy = *this;
copy.neg = !copy.neg;
return copy;
}
ngen::Subregister operator>>(int shift) const {
ngen::RegData sub = ngen::Subregister{};
if (shift >= 0 && shift <= maxShift)
sub = regs[shift];
if (neg)
sub = -sub;
return *reinterpret_cast<ngen::Subregister *>(&sub);
}
void set(int shift, ngen::Subregister reg) {
regs[shift] = reg;
}
};
struct LDMultiples {
ngen::GRFRange range;
bool a64 = false;
int count = 0;
};
struct Address2DParams {
ngen::Subregister rows, cols;
ngen::Subregister offR, offC;
ngen::Subregister remR, remC;
int fixedRows = 0, fixedCols = 0;
};
struct LDIncrements {
using value_type = std::pair<int, SubregisterPair>;
LDIncrements(const MatrixAddressingStrategy &s): type(s.base.isA64() ? ngen::DataType::q : ngen::DataType::ud) {}
std::vector<value_type>::iterator begin() { return increments.begin(); }
std::vector<value_type>::const_iterator begin()const { return increments.begin(); }
std::vector<value_type>::iterator end() { return increments.end(); }
std::vector<value_type>::const_iterator end()const { return increments.end(); }
void push_back(const value_type &v) { increments.push_back(v); }
void clear() { increments.clear(); }
ngen::DataType type;
std::vector<value_type> increments;
};
struct MaskAssignment {
MaskInfo mask; VirtualFlag flag; LoopType var; uint16_t offset;
bool compatible(const MaskAssignment &other) const {
return mask == other.mask && var == other.var && offset == other.offset;
}
void reverse(int width) {
offset = width - offset - mask.variable.rsize;
mask.variable.reverse = !mask.variable.reverse;
}
};
struct CommonState {
ngen::RegisterAllocator ra;
ngen::GRF signChange, selectImag;
GRFMultirange vflagStorage;
std::array<VirtualFlag, 8> activeVFlags;
VirtualFlagAllocator raVFlag;
TokenAllocator tokenAllocator;
std::vector<std::pair<uint8_t, int8_t>> tokenMap;
ngen::Subregister readFailures;
ngen::Subregister fusedID;
ngen::Subregister lsDescConstant[4];
ngen::FlagRegister flagSwizzle;
ngen::EmulationState emulate;
ngen::GRFRange eatomicAddRegs[2];
ngen::GRFRange remaskRegs[3];
VirtualFlag vflagEAtomicAdd;
VirtualFlag blockEMask;
ngen::Label blockDone;
ngen::Subregister all1s;
ngen::RegData r0_info;
bool movedR0 = false;
ngen::Subregister lid0;
GRFMultirange indexVec; int ivEntries = 0;
struct {
ngen::GRF zero, one;
ngen::GRFRange src1Storage;
ngen::GRF src1, srcR1, srcI1, r, d;
GRFMultirange mathTemp;
ngen::GRF temp;
std::array<ngen::FlagRegister, 2> tempFlags;
ngen::Subregister flagStore; ngen::Label label;
int simd;
ngen::Subregister callStorageSub, callStorage;
bool use = false;
} invertSub;
CommonState(ngen::HW hw) : ra(hw), raVFlag(hw), tokenAllocator(hw) {}
VirtualFlag allocVFlag(ngen::HW hw, int n = 1);
void wipeActiveVFlags();
bool vflagsEnabled() const { return !vflagStorage.empty(); }
void usePhysicalFlag(ngen::FlagRegister flag) { activeVFlags[flag.index()] = flag; }
void allocEmulate64Temp(const ngen::EmulationStrategy &estrategy);
};
struct GEMMState : public CommonState {
struct Inputs {
ngen::Subregister A, B, C[2], base, tempC; ngen::Subregister ao, bo, abo, co; ngen::Subregister aoPtr, boPtr, coPtr; ngen::Subregister aScalePtr, bScalePtr, cScalePtr; ngen::Subregister agPtr, bgPtr; ngen::Subregister offsetA, offsetB, offsetC[2]; ngen::Subregister offsetAO, offsetBO, offsetCO; ngen::Subregister offsetAScale, offsetBScale; ngen::Subregister offsetAg, offsetBg; ngen::Subregister offsetAq, offsetBq; ngen::Subregister lda, ldb, ldc[2]; ngen::Subregister ldao, ldbo, ldco; ngen::Subregister ldaScale, ldbScale, ldcScale; ngen::Subregister ldag, ldbg; ngen::Subregister ldaq, ldbq, ldcq; ngen::Subregister m, n, k, k0; SubregisterPair alpha_real, alpha_imag; SubregisterPair beta_real, beta_imag; ngen::Subregister alphaPtr, betaPtr; ngen::Subregister groupIDM, groupIDN, groupIDK; ngen::Subregister groupIDMN; ngen::GRF localIDM, localIDN, localIDK; ngen::Subregister localSizeM, localSizeN, localSizeK; ngen::Subregister groupCountM, groupCountN, groupCountK; ngen::Subregister groupCountMN; ngen::Subregister gcMNRecip, gcMRecip, gcNRecip; ngen::Subregister wgStride[2]; ngen::Subregister groupStride; ngen::Subregister kvConfig, kRecip; ngen::Subregister kSlicedTiles, kSyncSlabs; ngen::Subregister hilbertVD, hilbertUVDRecip; ngen::Subregister hilbertBail; ngen::Subregister bslice, bthresh; ngen::Subregister flags; ngen::Subregister diagA, diagB, diagC; ngen::Subregister statusBuffer; uint8_t surfaceA, surfaceAO, surfaceAScale; uint8_t surfaceB, surfaceBO, surfaceBScale; uint8_t surfaceAg, surfaceBg; uint8_t surfaceC[2], surfaceCO, surfaceTempC; uint8_t surfaceCScale; std::vector<ngen::Subregister> strideA; std::vector<ngen::Subregister> strideB; std::vector<ngen::Subregister> strideC; std::vector<ngen::Subregister> strideScaleA; std::vector<ngen::Subregister> strideScaleB; std::vector<ngen::Subregister> strideScaleC; std::vector<ngen::Subregister> strideOffsetA; std::vector<ngen::Subregister> strideOffsetB; std::vector<ngen::Subregister> strideGroupSumsA; std::vector<ngen::Subregister> strideGroupSumsB; std::vector<ngen::Subregister> batchSize; std::vector<ngen::Subregister> recipBatchSize; ngen::Subregister offsetBatch; ngen::Subregister incr_a_array, incr_b_array; ngen::Subregister incr_alpha, incr_beta; ngen::Subregister alpha_array, beta_array; ngen::Subregister slmBase; std::vector<ngen::Subregister> binarySrcs; std::vector<ngen::Subregister> binaryOffsets; std::vector<ngen::Subregister> binaryLDs; std::vector<std::vector<ngen::Subregister>> binaryStrides; std::vector<uint8_t> binarySurfaces;
ngen::Subregister sroundSeedPtr; ngen::Subregister sroundSeed; } inputs;
Type Ta_load, Tb_load; Type Tacc; ngen::Subregister persistentGroupID; ngen::Subregister batchID[4]; ngen::Subregister offsetA, offsetB, offsetC[2];
ngen::Subregister offsetAs, offsetBs, offsetCs;
ngen::Subregister offsetAo, offsetBo;
ngen::Subregister offsetAp, offsetBp, offsetCp;
ngen::Subregister offsetCO;
ngen::Subregister saveOffsetA, saveOffsetB, saveOffsetC[2];
ngen::Subregister saveOffsetCO;
ngen::Subregister fullK;
ngen::Subregister effA, effB, effC[2], effCO, effTempC; ngen::Subregister effAi, effBi;
ngen::Subregister effAo, effBo;
ngen::Subregister effAp, effBp, effCp;
ngen::Subregister effAs, effBs;
std::vector<ngen::GRFRange> A_addrs, B_addrs, C_addrs[2];
std::vector<ngen::GRFRange> A_addrsRem, B_addrsRem;
std::vector<ngen::GRFRange> A_addrsAlt, B_addrsAlt;
std::vector<ngen::GRFRange> A_addrsAltRem, B_addrsAltRem;
std::vector<ngen::GRFRange> Ai_addrs, Bi_addrs;
std::vector<std::vector<ngen::GRFRange>> Ai_addrsK, Bi_addrsK;
std::vector<ngen::GRFRange> Ai_addrsRem, Bi_addrsRem;
std::vector<ngen::GRFRange> Ao_addrs, Bo_addrs;
std::vector<ngen::GRFRange> Ap_addrs, Bp_addrs, Cp_addrs;
std::vector<ngen::GRFRange> Ap_addrsAlt, Bp_addrsAlt;
std::vector<ngen::GRFRange> A_offsetAddrs, B_offsetAddrs;
std::vector<ngen::GRFRange> A_scaleAddrs, B_scaleAddrs, C_scaleAddrs;
std::vector<ngen::GRFRange> Ag_addrs, Bg_addrs;
std::vector<GRFMultirange> A_regs, B_regs, C_regs;
GRFMultirange Ar_regs, Br_regs; GRFMultirange Cr_regs; std::vector<GRFMultirange> Ai_regs, Bi_regs; std::vector<GRFMultirange> Ai_regsRem, Bi_regsRem;
GRFMultirange Ao_regs, Bo_regs; GRFMultirange Ao_regsRem, Bo_regsRem;
GRFMultirange As_regs, Bs_regs; GRFMultirange Asr_regs, Bsr_regs; GRFMultirange Ap_regs, Bp_regs, Cp_regs; GRFMultirange A_offsetRegs, B_offsetRegs; GRFMultirange A_scaleRegs, B_scaleRegs; GRFMultirange Ag_regs, Bg_regs; GRFMultirange Ar_offsetRegs, Br_offsetRegs; GRFMultirange Ar_scaleRegs, Br_scaleRegs; GRFMultirange Agr_regs, Bgr_regs; std::vector<MaskAssignment> AB_masks, AB_masksCoop;
std::vector<ngen::GRFRange> tempMul_regs;
ngen::Subregister groupCountMN, groupIDMN; ngen::Subregister i0, j0, h0; ngen::Subregister wgI0, wgJ0; ngen::Subregister unclampedI0; ngen::Subregister ctiShiftJ0; ngen::Subregister threadK0, k0Rem, wgK; ngen::Subregister remainders[3]; ngen::Subregister remaindersFused[2]; ngen::Subregister remaindersWG[2]; ngen::Subregister remaindersCoop[3]; ngen::Subregister remFusedStorage; ngen::Subregister diagA, diagB, diagC; SubregisterPair lda, ldb;
SubregisterPair ldao, ldbo, ldag, ldbg, ldaScale, ldbScale, ldcScale;
LDIncrements ldaIncrements, ldbIncrements; LDIncrements ldaoIncrements, ldboIncrements;
LDIncrements ldasIncrements, ldbsIncrements;
LDIncrements ldagIncrements, ldbgIncrements;
LDMultiples ldaMultiples, ldbMultiples, ldcMultiples[2];
ngen::Subregister k, K; ngen::Subregister kNoBarrierStart, kNoBarrierEnd; ngen::FlagRegister flagAP;
ngen::Subregister beta1; ngen::Subregister add64; ngen::Subregister lidM, lidN, lidStorage; ngen::Subregister lidK, lszK, lidszKStorage; ngen::Subregister ia0_slm, jb0_slm; ngen::Subregister postRemA, postRemB; ngen::Subregister postRemAi, postRemBi; ngen::Subregister postRemAo, postRemBo; ngen::Subregister isCompute; ngen::GRF sysSumAll1s; ngen::GRF betaCheckReturn;
ngen::GRF tmpCScales; ngen::Subregister statusFlagAddr; bool systolicSumA = false, systolicSumB = false;
bool useBDPAS = false;
bool upConvertATo8Bit = false;
bool upConvertBTo8Bit = false;
bool lateKLoopCheck = false;
bool splitBarrierAlways = false;
int ka_loadRem, kb_loadRem;
bool A_lateKRem, B_lateKRem;
bool A_descRem, B_descRem;
bool Ai_hasKRem, Bi_hasKRem;
bool Ai_lateKRem, Bi_lateKRem;
bool Ai_incrementalRem, Bi_incrementalRem;
bool Ai_remIncrCopy, Bi_remIncrCopy;
int ma_slm, ka_slm, kb_slm, nb_slm;
int ma_prefetch, ka_prefetch, kb_prefetch, nb_prefetch;
CoopSplit effCoopA = CoopSplit::K;
CoopSplit effCoopB = CoopSplit::K;
ngen::Subregister kSLMA, kSLMB, kSLMStorage; bool kSLMCountUp = false;
int kaq = 0, kbq = 0, kaqStride, kbqStride, kaqLate = 0, kbqLate = 0;
bool lateScale2DA = false, lateScale2DB = false;
RegisterLayout A_layout, B_layout, C_layout;
RegisterLayout A_layoutRem, B_layoutRem;
RegisterLayout A_layoutAlt, B_layoutAlt;
RegisterLayout A_layoutAltRem, B_layoutAltRem;
RegisterLayout Ar_layout, Br_layout;
RegisterLayout Ai_layout, Bi_layout;
std::vector<RegisterLayout> Ai_layoutK, Bi_layoutK;
RegisterLayout Ai_layoutRem, Bi_layoutRem;
RegisterLayout Ao_layout, Bo_layout;
RegisterLayout As_layout, Bs_layout;
RegisterLayout Asr_layout, Bsr_layout;
RegisterLayout Ap_layout, Bp_layout, Cp_layout;
RegisterLayout Ap_layoutAlt, Bp_layoutAlt;
RegisterLayout A_offsetLayout, B_offsetLayout;
RegisterLayout A_scaleLayout, B_scaleLayout, C_scaleLayout;
RegisterLayout Ag_layout, Bg_layout;
RegisterLayout Ar_offsetLayout, Br_offsetLayout;
RegisterLayout Ar_scaleLayout, Br_scaleLayout;
RegisterLayout Agr_layout, Bgr_layout;
RegisterLayout Cr_layout;
RegisterLayout C_layoutReduced;
RegisterLayout C_layoutExt, C_layoutExtUnmasked, C_layoutExtNonatomicUnmasked;
Address2DParams A_params, B_params;
Address2DParams Ai_params, Bi_params;
Address2DParams Ap_params, Bp_params;
int Ai_regCount = 0, Bi_regCount = 0;
bool aioShare = false, bioShare = false;
bool aioShareRem = false, bioShareRem = false;
bool aoReuseA = false, boReuseB = false;
Type Tao_int, Ta_scaleInt, Tag_int;
Type Tbo_int, Tb_scaleInt, Tbg_int;
MatrixAddressing Ai, Bi, Ao, Bo, tempC;
MatrixAddressingStrategy Ai_strategy, Bi_strategy;
MatrixAddressingStrategy Ao_strategy, Bo_strategy;
MatrixAddressingStrategy Cext_strategy, tempCStrategy;
ngen::FlagRegister panelMaskA, panelMaskB;
int8_t tokenBarrierFence[2];
ngen::InstructionModifier modBarrierFence[2];
bool barrierReady = false;
ngen::GRF barrierHeader;
ngen::GRF barrierHeaderM, barrierHeaderN;
ngen::FlagRegister barrierM, barrierN;
bool firstKLoopSegment;
bool isNested = false;
int C_accCount;
bool cSwapActive = false;
bool haveCSwap = false;
int C_count = 1;
int C_buffers = 1;
bool allocedAo = false, allocedBo = false;
bool allowEmptyC = false;
bool copyC = false;
bool useTempC = false;
bool repackA = false, repackB = false;
bool repackARem = false, repackBRem = false;
int ka_repack, ka_repackRem, kb_repackRem;
int cRepackPeriod = 0;
bool remActiveA, remActiveB, remActiveSLM;
std::vector<MaskAssignment> kMasksA, kMasksB, kMasksAi, kMasksBi;
int initSLMKOffset = 0;
bool slmRemaskA = false, slmRemaskB = false;
bool slmASums = false, slmBSums = false;
bool doLateExit = false;
bool needBRFallback = true;
ngen::GRF emulate64TempSave[2];
bool simd32KMasks = false;
int lastThresh = 0;
ngen::Subregister nextGroupIDM, nextGroupIDN;
ngen::Subregister nextFlagL3PFA, nextFlagL3PFB;
ngen::FlagRegister flagL3PFA, flagL3PFB;
RegisterLayout Apl3_layout, Bpl3_layout;
std::vector<ngen::GRFRange> Apl3_addrs, Bpl3_addrs;
std::vector<ngen::Subregister> effBinary;
struct {
ngen::InstructionModifier depAddr[4];
} sysgemm;
GEMMState(ngen::HW hw, const GEMMStrategy &strategy) : CommonState(hw),
ldaIncrements(strategy.A), ldbIncrements(strategy.B),
ldaoIncrements(strategy.AO), ldboIncrements(strategy.BO),
ldasIncrements(strategy.A_scale), ldbsIncrements(strategy.B_scale),
ldagIncrements(strategy.Ag), ldbgIncrements(strategy.Bg) {}
int internalSIMD() const { return simd32KMasks ? 32 : 16; }
ngen::GRF r0InfoGRF() const;
void setTacc(Type T);
};
GEMMSTONE_NAMESPACE_END
#endif