Enable RA of SVE Predicate Registers

This commit is contained in:
Paulo Matos committed 2024-12-02 18:35:31 +01:00
1 parent 731e4d6271
commit fcbf0de05a
11 files changed
+109 -8

No files matched your search

+6 -4
View File
@@ -101,7 +101,8 @@ def is_ssa_type(type):
if (type == "SSA" or
type == "GPR" or
type == "GPRPair" or
type == "FPR"):
type == "FPR" or
type == "PRED"):
return True
return False
@@ -150,8 +151,8 @@ def parse_ops(ops):
RHS += f", {DType}:$Out{Name}"
else:
# Single anonymous destination
if LHS not in ["SSA", "GPR", "GPRPair", "FPR"]:
ExitError(f"Unknown destination class type {LHS}. Needs to be one of SSA, GPR, GPRPair, FPR")
if LHS not in ["SSA", "GPR", "GPRPair", "FPR", "PRED"]:
ExitError(f"Unknown destination class type {LHS}. Needs to be one of SSA, GPR, GPRPair, FPR, PRED")
OpDef.HasDest = True
OpDef.DestType = LHS
@@ -221,7 +222,8 @@ def parse_ops(ops):
if (OpArg.IsSSA and
(OpArg.Type == "GPR" or
OpArg.Type == "GPRPair" or
OpArg.Type == "FPR")):
OpArg.Type == "FPR" or
OpArg.Type == "PR")):
OpDef.EmitValidation.append(f"GetOpRegClass({ArgName}) == InvalidClass || WalkFindRegClass({ArgName}) == {OpArg.Type}Class")
OpArg.Name = ArgName
@@ -57,6 +57,12 @@ namespace x64 {
ARMEmitter::Reg::r24, ARMEmitter::Reg::r25, ARMEmitter::Reg::r30, ARMEmitter::Reg::r18,
};
// p6 and p7 registers are used as temporaries no not added here for RA
// See PREF_TMP_16B and PREF_TMP_32B
// Also p8-p15 cannot be used can only encode p0-p7, so we're left with p0-p5.
constexpr std::array<ARMEmitter::PRegister, 6> PR = {ARMEmitter::PReg::p0, ARMEmitter::PReg::p1, ARMEmitter::PReg::p2,
ARMEmitter::PReg::p3, ARMEmitter::PReg::p4, ARMEmitter::PReg::p5};
constexpr unsigned RAPairs = 6;
// All are caller saved
@@ -103,6 +109,12 @@ namespace x64 {
ARMEmitter::Reg::r16, ARMEmitter::Reg::r17, ARMEmitter::Reg::r30,
};
// p6 and p7 registers are used as temporaries no not added here for RA
// See PREF_TMP_16B and PREF_TMP_32B
// Also p8-p15 cannot be used can only encode p0-p7, so we're left with p0-p5.
constexpr std::array<ARMEmitter::PRegister, 6> PR = {ARMEmitter::PReg::p0, ARMEmitter::PReg::p1, ARMEmitter::PReg::p2,
ARMEmitter::PReg::p3, ARMEmitter::PReg::p4, ARMEmitter::PReg::p5};
constexpr unsigned RAPairs = 6;
constexpr std::array<ARMEmitter::VRegister, 16> SRAFPR = {
@@ -234,6 +246,12 @@ namespace x32 {
constexpr unsigned RAPairs = 12;
// p6 and p7 registers are used as temporaries no not added here for RA
// See PREF_TMP_16B and PREF_TMP_32B
// Also p8-p15 cannot be used can only encode p0-p7, so we're left with p0-p5.
constexpr std::array<ARMEmitter::PRegister, 6> PR = {ARMEmitter::PReg::p0, ARMEmitter::PReg::p1, ARMEmitter::PReg::p2,
ARMEmitter::PReg::p3, ARMEmitter::PReg::p4, ARMEmitter::PReg::p5};
// All are caller saved
constexpr std::array<ARMEmitter::VRegister, 8> SRAFPR = {
ARMEmitter::VReg::v16, ARMEmitter::VReg::v17, ARMEmitter::VReg::v18, ARMEmitter::VReg::v19,
@@ -357,6 +375,7 @@ Arm64Emitter::Arm64Emitter(FEXCore::Context::ContextImpl* ctx, void* EmissionPtr
GeneralRegisters = x64::RA;
StaticFPRegisters = x64::SRAFPR;
GeneralFPRegisters = x64::RAFPR;
PredicateRegisters = x64::PR;
PairRegisters = x64::RAPairs;
#ifdef _M_ARM_64EC
ConfiguredDynamicRegisterBase = std::span(x64::RA.begin(), 7);
@@ -370,6 +389,8 @@ Arm64Emitter::Arm64Emitter(FEXCore::Context::ContextImpl* ctx, void* EmissionPtr
StaticFPRegisters = x32::SRAFPR;
GeneralFPRegisters = x32::RAFPR;
PredicateRegisters = x32::PR;
}
}
@@ -94,6 +94,7 @@ protected:
std::span<const ARMEmitter::Register> ConfiguredDynamicRegisterBase {};
std::span<const ARMEmitter::Register> StaticRegisters {};
std::span<const ARMEmitter::Register> GeneralRegisters {};
std::span<const ARMEmitter::PRegister> PredicateRegisters {};
std::span<const ARMEmitter::VRegister> StaticFPRegisters {};
std::span<const ARMEmitter::VRegister> GeneralFPRegisters {};
uint32_t PairRegisters = 0;
@@ -534,6 +534,7 @@ Arm64JITCore::Arm64JITCore(FEXCore::Context::ContextImpl* ctx, FEXCore::Core::In
RAPass->AddRegisters(FEXCore::IR::GPRFixedClass, StaticRegisters.size());
RAPass->AddRegisters(FEXCore::IR::FPRClass, GeneralFPRegisters.size());
RAPass->AddRegisters(FEXCore::IR::FPRFixedClass, StaticFPRegisters.size());
RAPass->AddRegisters(FEXCore::IR::PREDClass, PredicateRegisters.size());
RAPass->PairRegs = PairRegisters;
{
@@ -94,6 +94,19 @@ private:
FEX_UNREACHABLE;
}
[[nodiscard]]
ARMEmitter::PRegister GetPReg(IR::NodeID Node) const {
const auto Reg = GetPhys(Node);
LOGMAN_THROW_AA_FMT(Reg.Class == IR::PREDClass.Val, "Unexpected Class: {}", Reg.Class);
if (Reg.Class == IR::PREDClass.Val) {
return PredicateRegisters[Reg.Reg];
}
FEX_UNREACHABLE;
}
[[nodiscard]]
FEXCore::IR::RegisterClassType GetRegClass(IR::NodeID Node) const;
@@ -10,6 +10,7 @@ $end_info$
#include "Interface/Context/Context.h"
#include "Interface/Core/CPUID.h"
#include "Interface/Core/JIT/JITClass.h"
#include "Interface/IR/IR.h"
#include <FEXCore/Utils/CompilerDefs.h>
#include <FEXCore/Utils/MathUtils.h>
@@ -1551,6 +1552,45 @@ DEF_OP(StoreMem) {
}
}
DEF_OP(InitPredicate) {
const auto Op = IROp->C<IR::IROp_InitPredicate>();
const auto OpSize = IROp->Size;
ptrue(ConvertSubRegSize16(OpSize), GetPReg(Node), static_cast<ARMEmitter::PredicatePattern>(Op->Pattern));
}
DEF_OP(StoreMemPredicate) {
const auto Op = IROp->C<IR::IROp_StoreMemPredicate>();
// FIXME: Dont we need OpSize then?
const auto Predicate = GetPReg(Op->Mask.ID());
const auto RegData = GetVReg(Op->Value.ID());
const auto MemReg = GetReg(Op->Addr.ID());
LOGMAN_THROW_A_FMT(HostSupportsSVE128 || HostSupportsSVE256, "StoreMemPredicate needs SVE support");
const auto MemDst = ARMEmitter::SVEMemOperand(MemReg.X(), 0);
switch (IROp->ElementSize) {
case IR::OpSize::i8Bit: {
st1b<ARMEmitter::SubRegSize::i8Bit>(RegData.Z(), Predicate, MemDst);
break;
}
case IR::OpSize::i16Bit: {
st1h<ARMEmitter::SubRegSize::i16Bit>(RegData.Z(), Predicate, MemDst);
break;
}
case IR::OpSize::i32Bit: {
st1w<ARMEmitter::SubRegSize::i32Bit>(RegData.Z(), Predicate, MemDst);
break;
}
case IR::OpSize::i64Bit: {
st1d(RegData.Z(), Predicate, MemDst);
break;
}
default: break;
}
}
DEF_OP(StoreMemPair) {
const auto Op = IROp->C<IR::IROp_StoreMemPair>();
const auto OpSize = IROp->Size;
+18 -2
View File
@@ -7,6 +7,7 @@
" SSA = untyped",
" GPR = GPR class type",
" FPR = FPR class type",
" PRED = Predicate register class type",
"Declaring the SSA types correctly will allow validation passes to ensure the op is getting passed correct arguments",
"",
"Arguments must always follow a particular order. <Type>:<Prefix><Name>",
@@ -83,6 +84,7 @@
"constexpr FEXCore::IR::RegisterClassType GPRFixedClass {1}",
"constexpr FEXCore::IR::RegisterClassType FPRClass {2}",
"constexpr FEXCore::IR::RegisterClassType FPRFixedClass {3}",
"constexpr FEXCore::IR::RegisterClassType PREDClass {4}",
"constexpr FEXCore::IR::RegisterClassType ComplexClass {5}",
"constexpr FEXCore::IR::RegisterClassType InvalidClass {7}",
"",
@@ -148,6 +150,7 @@
"SSA": "OrderedNode*",
"GPR": "OrderedNode*",
"FPR": "OrderedNode*",
"PRED": "OrderedNode*",
"FenceType": "FenceType",
"RegisterClass": "RegisterClassType",
"CondClass": "CondClassType",
@@ -560,8 +563,21 @@
"HasSideEffects": true,
"DestSize": "Size",
"EmitValidation": [
"WalkFindRegClass($Value1) == $Class",
"WalkFindRegClass($Value2) == $Class"
"WalkFindRegClass($Value1) == $Class"
]
},
"PRED = InitPredicate OpSize:#Size, u8:$Pattern": {
"Desc": ["Initialize predicate register from Pattern"],
"DestSize": "Size"
},
"StoreMemPredicate RegisterClass:$Class, OpSize:#Size, SSA:$Value, PRED:$Mask, GPR:$Addr": {
"Desc": [ "Stores a value to memory using SVE predicate mask." ],
"HasSideEffects": true,
"DestSize": "Size",
"EmitValidation": [
"WalkFindRegClass($Value) == $Class"
]
},
+3
View File
@@ -77,6 +77,8 @@ static void PrintArg(fextl::stringstream* out, [[maybe_unused]] const IRListView
*out << "FPR";
} else if (Arg == FPRFixedClass.Val) {
*out << "FPRFixed";
} else if (Arg == PREDClass.Val) {
*out << "PRED";
} else {
*out << "Unknown Registerclass " << Arg;
}
@@ -98,6 +100,7 @@ static void PrintArg(fextl::stringstream* out, const IRListView* IR, OrderedNode
case FEXCore::IR::GPRFixedClass.Val: *out << "(GPRFixed"; break;
case FEXCore::IR::FPRClass.Val: *out << "(FPR"; break;
case FEXCore::IR::FPRFixedClass.Val: *out << "(FPRFixed"; break;
case FEXCore::IR::PREDClass.Val: *out << "(PRED"; break;
case FEXCore::IR::ComplexClass.Val: *out << "(Complex"; break;
case FEXCore::IR::InvalidClass.Val: *out << "(Invalid"; break;
default: *out << "(Unknown"; break;
+1 -1
View File
@@ -70,7 +70,7 @@ void PassManager::AddDefaultPasses(FEXCore::Context::ContextImpl* ctx) {
FEX_CONFIG_OPT(DisablePasses, O0);
if (!DisablePasses()) {
InsertPass(CreateX87StackOptimizationPass());
InsertPass(CreateX87StackOptimizationPass(ctx->HostFeatures));
InsertPass(CreateConstProp(ctx->HostFeatures.SupportsTSOImm9, &ctx->CPUID));
InsertPass(CreateDeadFlagCalculationEliminination());
}
+2 -1
View File
@@ -5,6 +5,7 @@
namespace FEXCore {
class CPUIDEmu;
struct HostFeatures;
}
namespace FEXCore::Utils {
@@ -19,7 +20,7 @@ class RegisterAllocationData;
fextl::unique_ptr<FEXCore::IR::Pass> CreateConstProp(bool SupportsTSOImm9, const FEXCore::CPUIDEmu* CPUID);
fextl::unique_ptr<FEXCore::IR::Pass> CreateDeadFlagCalculationEliminination();
fextl::unique_ptr<FEXCore::IR::RegisterAllocationPass> CreateRegisterAllocationPass();
fextl::unique_ptr<FEXCore::IR::Pass> CreateX87StackOptimizationPass();
fextl::unique_ptr<FEXCore::IR::Pass> CreateX87StackOptimizationPass(const FEXCore::HostFeatures&);
namespace Validation {
fextl::unique_ptr<FEXCore::IR::Pass> CreateIRValidation();
@@ -47,6 +47,7 @@ struct RegState {
// On arm64, there are 16 Fixed and 12 normal
FPRsFixed[Reg.Reg] = ssa;
return true;
case PREDClass: PREGs[Reg.Reg] = ssa; return true;
}
return false;
}
@@ -59,6 +60,7 @@ struct RegState {
case GPRFixedClass: return GPRsFixed[Reg.Reg];
case FPRClass: return FPRs[Reg.Reg];
case FPRFixedClass: return FPRsFixed[Reg.Reg];
case PREDClass: return PREGs[Reg.Reg];
}
return InvalidReg;
}
@@ -82,6 +84,7 @@ private:
std::array<IR::NodeID, 32> FPRsFixed = {};
std::array<IR::NodeID, 32> GPRs = {};
std::array<IR::NodeID, 32> FPRs = {};
std::array<IR::NodeID, 32> PREGs = {};
fextl::unordered_map<uint32_t, IR::NodeID> Spills;
};