mirror of
https://github.com/KytyPS5/KytyPS5.git
synced 2026-08-18 22:42:23 +00:00
shader: resolve subtractive dword SETPC tables
This commit is contained in:
@@ -153,6 +153,9 @@ bool InstructionWritesScalarCode(const Instruction& inst, uint32_t code) {
|
||||
return code >= dst_code && code < dst_code + count;
|
||||
}
|
||||
|
||||
bool ScalarCodeWrittenInRange(const Decoder::Program& program, uint32_t begin_index,
|
||||
uint32_t end_index, uint32_t code);
|
||||
|
||||
bool FindPreviousGetpc(const Decoder::Program& program, uint32_t before_index, uint32_t dst_code,
|
||||
uint32_t* index) {
|
||||
for (uint32_t i = before_index; i > 0; i--) {
|
||||
@@ -193,11 +196,46 @@ bool ResolvePcRelativeBase(const Decoder::Program& program, uint32_t before_inde
|
||||
const Decoder::Operand* offset = nullptr;
|
||||
int32_t imm = 0;
|
||||
uint32_t getpc_index = 0;
|
||||
if (!OtherScalarSource(candidate, base_code, &offset) || !IsImmediateSigned(*offset, imm) ||
|
||||
if (adds) {
|
||||
if (!OtherScalarSource(candidate, base_code, &offset)) {
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
if (!IsScalarCode(candidate.src0, base_code)) {
|
||||
return false;
|
||||
}
|
||||
offset = &candidate.src1;
|
||||
}
|
||||
if (!IsImmediateSigned(*offset, imm) ||
|
||||
!FindPreviousGetpc(program, candidate_index, base_code, &getpc_index)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (candidate_index + 1u >= before_index) {
|
||||
return false;
|
||||
}
|
||||
const auto& high = program.instructions[candidate_index + 1u];
|
||||
const Decoder::Operand* high_other = nullptr;
|
||||
uint32_t zero = 0;
|
||||
if (high.opcode != (adds ? Opcode::SAddcU32 : Opcode::SSubbU32) ||
|
||||
!IsScalarCode(high.dst, base_code + 1u)) {
|
||||
return false;
|
||||
}
|
||||
if (adds) {
|
||||
if (!OtherScalarSource(high, base_code + 1u, &high_other)) {
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
if (!IsScalarCode(high.src0, base_code + 1u)) {
|
||||
return false;
|
||||
}
|
||||
high_other = &high.src1;
|
||||
}
|
||||
if (!IsImmediate(*high_other, zero) || zero != 0u ||
|
||||
ScalarCodeWrittenInRange(program, candidate_index + 2u, before_index, base_code + 1u)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const auto base = InstructionEndPc(program.instructions[getpc_index]);
|
||||
*pc = adds ? base + static_cast<uint32_t>(imm) : base - static_cast<uint32_t>(imm);
|
||||
*pc &= ~3u;
|
||||
@@ -206,13 +244,22 @@ bool ResolvePcRelativeBase(const Decoder::Program& program, uint32_t before_inde
|
||||
return false;
|
||||
}
|
||||
|
||||
bool FindPreviousScalarLoadPair(const Decoder::Program& program, uint32_t before_index,
|
||||
uint32_t dst_code, uint32_t* index) {
|
||||
bool AddSignedByteOffset(uint32_t base, uint32_t encoded_offset, uint32_t* address) {
|
||||
const auto value = static_cast<int64_t>(base) + static_cast<int32_t>(encoded_offset);
|
||||
if (address == nullptr || value < 0 || value > UINT32_MAX) {
|
||||
return false;
|
||||
}
|
||||
*address = static_cast<uint32_t>(value);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool FindPreviousScalarLoad(const Decoder::Program& program, uint32_t before_index,
|
||||
uint32_t dst_code, Opcode opcode, uint32_t* index) {
|
||||
for (uint32_t i = before_index; i > 0; i--) {
|
||||
const uint32_t candidate_index = i - 1u;
|
||||
const auto& candidate = program.instructions[candidate_index];
|
||||
if (InstructionWritesScalarCode(candidate, dst_code)) {
|
||||
if (candidate.opcode == Opcode::SLoadDwordx2 && IsScalarCode(candidate.dst, dst_code)) {
|
||||
if (candidate.opcode == opcode && IsScalarCode(candidate.dst, dst_code)) {
|
||||
if (index != nullptr) {
|
||||
*index = candidate_index;
|
||||
}
|
||||
@@ -225,8 +272,8 @@ bool FindPreviousScalarLoadPair(const Decoder::Program& program, uint32_t before
|
||||
}
|
||||
|
||||
bool ResolveJumpTableEntryCount(const Decoder::Program& program, uint32_t before_index,
|
||||
const Decoder::Operand& byte_offset_operand,
|
||||
uint32_t* entry_count) {
|
||||
const Decoder::Operand& byte_offset_operand, uint32_t stride_shift,
|
||||
uint32_t* entry_count) {
|
||||
if (entry_count == nullptr) {
|
||||
return false;
|
||||
}
|
||||
@@ -248,7 +295,7 @@ bool ResolveJumpTableEntryCount(const Decoder::Program& program, uint32_t before
|
||||
const Decoder::Operand* shift_amount_operand = nullptr;
|
||||
uint32_t shift_amount = 0;
|
||||
if (!OtherScalarSource(shift, byte_offset_code, &shift_amount_operand) ||
|
||||
!IsImmediate(*shift_amount_operand, shift_amount) || shift_amount != 3u) {
|
||||
!IsImmediate(*shift_amount_operand, shift_amount) || shift_amount != stride_shift) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -264,10 +311,8 @@ bool ResolveJumpTableEntryCount(const Decoder::Program& program, uint32_t before
|
||||
return false;
|
||||
}
|
||||
|
||||
const Decoder::Operand* clamp_other = nullptr;
|
||||
uint32_t max_index = 0;
|
||||
if (!OtherScalarSource(clamp, index_code, &clamp_other) ||
|
||||
!IsImmediate(*clamp_other, max_index)) {
|
||||
uint32_t max_index = 0;
|
||||
if (!IsImmediate(clamp.src0, max_index) && !IsImmediate(clamp.src1, max_index)) {
|
||||
return false;
|
||||
}
|
||||
*entry_count = max_index + 1u;
|
||||
@@ -332,7 +377,8 @@ bool ResolveSetpcJumpTable(const Decoder::Program& program, uint32_t setpc_index
|
||||
}
|
||||
|
||||
uint32_t load_index = 0;
|
||||
if (!FindPreviousScalarLoadPair(program, setpc_index - 3u, offset_low_code, &load_index)) {
|
||||
if (!FindPreviousScalarLoad(program, setpc_index - 3u, offset_low_code, Opcode::SLoadDwordx2,
|
||||
&load_index)) {
|
||||
return false;
|
||||
}
|
||||
const auto& load = program.instructions[load_index];
|
||||
@@ -342,12 +388,13 @@ bool ResolveSetpcJumpTable(const Decoder::Program& program, uint32_t setpc_index
|
||||
}
|
||||
|
||||
uint32_t table_pc = 0;
|
||||
if (!ResolvePcRelativeBase(program, load_index, table_base_code, &table_pc)) {
|
||||
if (!ResolvePcRelativeBase(program, load_index, table_base_code, &table_pc) ||
|
||||
!AddSignedByteOffset(table_pc, load.offset, &table_pc)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
uint32_t entry_count = 0;
|
||||
if (!ResolveJumpTableEntryCount(program, load_index, load.src1, &entry_count)) {
|
||||
if (!ResolveJumpTableEntryCount(program, load_index, load.src1, 3u, &entry_count)) {
|
||||
return false;
|
||||
}
|
||||
uint32_t selector_code = UINT32_MAX;
|
||||
@@ -357,8 +404,10 @@ bool ResolveSetpcJumpTable(const Decoder::Program& program, uint32_t setpc_index
|
||||
}
|
||||
|
||||
const uint32_t target_base = InstructionEndPc(getpc);
|
||||
const uint32_t table_word = table_pc / 4u;
|
||||
if ((table_pc & 3u) != 0 || table_word + entry_count * 2u > program.code.size()) {
|
||||
const size_t table_word = table_pc / 4u;
|
||||
const size_t table_words = static_cast<size_t>(entry_count) * 2u;
|
||||
if ((table_pc & 3u) != 0 || table_word > program.code.size() ||
|
||||
table_words > program.code.size() - table_word) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -390,6 +439,92 @@ bool ResolveSetpcJumpTable(const Decoder::Program& program, uint32_t setpc_index
|
||||
return true;
|
||||
}
|
||||
|
||||
bool ResolveSetpcDwordJumpTable(const Decoder::Program& program, uint32_t setpc_index,
|
||||
SetpcTargetInfo& info) {
|
||||
if (setpc_index < 3u || setpc_index >= program.instructions.size()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const auto& setpc = program.instructions[setpc_index];
|
||||
uint32_t pc_reg = 0;
|
||||
if (setpc.opcode != Opcode::SSetpcB64 || setpc.src0.kind != Decoder::OperandKind::Sgpr ||
|
||||
!ScalarOperandCode(setpc.src0, pc_reg)) {
|
||||
return false;
|
||||
}
|
||||
const auto& low_sub = program.instructions[setpc_index - 2u];
|
||||
const auto& high_sub = program.instructions[setpc_index - 1u];
|
||||
uint32_t offset_code = 0;
|
||||
uint32_t zero = 0;
|
||||
if (low_sub.opcode != Opcode::SSubU32 || !IsScalarCode(low_sub.dst, pc_reg) ||
|
||||
!IsScalarCode(low_sub.src0, pc_reg) || !ScalarOperandCode(low_sub.src1, offset_code) ||
|
||||
high_sub.opcode != Opcode::SSubbU32 || !IsScalarCode(high_sub.dst, pc_reg + 1u) ||
|
||||
!IsScalarCode(high_sub.src0, pc_reg + 1u) || !IsImmediate(high_sub.src1, zero) ||
|
||||
zero != 0u || offset_code == pc_reg || offset_code == pc_reg + 1u) {
|
||||
return false;
|
||||
}
|
||||
|
||||
uint32_t load_index = 0;
|
||||
if (!FindPreviousScalarLoad(program, setpc_index - 2u, offset_code, Opcode::SLoadDword,
|
||||
&load_index)) {
|
||||
return false;
|
||||
}
|
||||
const auto& load = program.instructions[load_index];
|
||||
if (!IsScalarCode(load.src0, pc_reg)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
uint32_t target_base = 0;
|
||||
if (!ResolvePcRelativeBase(program, load_index, pc_reg, &target_base)) {
|
||||
return false;
|
||||
}
|
||||
uint32_t table_pc = 0;
|
||||
if (!AddSignedByteOffset(target_base, load.offset, &table_pc)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
uint32_t entry_count = 0;
|
||||
if (!ResolveJumpTableEntryCount(program, load_index, load.src1, 2u, &entry_count)) {
|
||||
return false;
|
||||
}
|
||||
uint32_t selector_code = UINT32_MAX;
|
||||
if (!ScalarOperandCode(load.src1, selector_code) || selector_code == offset_code ||
|
||||
ScalarCodeWrittenInRange(program, load_index + 1u, setpc_index, selector_code)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const size_t table_word = table_pc / 4u;
|
||||
if ((table_pc & 3u) != 0 || table_word > program.code.size() ||
|
||||
entry_count > program.code.size() - table_word) {
|
||||
return false;
|
||||
}
|
||||
|
||||
std::vector<uint32_t> targets;
|
||||
std::vector<uint32_t> selector_values;
|
||||
std::vector<uint32_t> selector_target_pcs;
|
||||
for (uint32_t i = 0; i < entry_count; i++) {
|
||||
const uint32_t offset = program.code[table_word + i];
|
||||
if (offset > target_base) {
|
||||
return false;
|
||||
}
|
||||
const auto target_pc = (target_base - offset) & ~3u;
|
||||
AddUniqueTargetPc(targets, target_pc);
|
||||
selector_values.push_back(i * 4u);
|
||||
selector_target_pcs.push_back(target_pc);
|
||||
}
|
||||
if (targets.empty()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
info.indirect = true;
|
||||
info.pc_sgpr = pc_reg;
|
||||
info.selector_code = selector_code;
|
||||
info.table_load_pc = load.pc;
|
||||
info.target_pcs = std::move(targets);
|
||||
info.selector_values = std::move(selector_values);
|
||||
info.selector_target_pcs = std::move(selector_target_pcs);
|
||||
return true;
|
||||
}
|
||||
|
||||
bool ResolveSetpcTarget(const Decoder::Program& program, uint32_t setpc_index, uint32_t& target) {
|
||||
if (setpc_index >= program.instructions.size()) {
|
||||
return false;
|
||||
@@ -436,6 +571,9 @@ bool ResolveSetpcTarget(const Decoder::Program& program, uint32_t setpc_index, u
|
||||
bool ResolveSetpcTargets(const Decoder::Program& program, uint32_t setpc_index,
|
||||
SetpcTargetInfo& info) {
|
||||
info = {};
|
||||
if (ResolveSetpcDwordJumpTable(program, setpc_index, info)) {
|
||||
return true;
|
||||
}
|
||||
if (ResolveSetpcJumpTable(program, setpc_index, info)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -7507,6 +7507,70 @@ void TestNewShaderRecompilerSetpcJumpTable() {
|
||||
CheckSpirvBinaryValidates(result.spirv);
|
||||
}
|
||||
|
||||
void TestNewShaderRecompilerSetpcDwordJumpTable() {
|
||||
const uint32_t shader[] = {
|
||||
EncodeSop2(0x07, 106, 0, 130), // s_min_u32 vcc_lo, s0, 2
|
||||
EncodeSop2(0x1e, 106, 106, 130), // s_lshl_b32 vcc_lo, vcc_lo, 2
|
||||
EncodeSop1(0x1f, 4, 0), // s_getpc_b64 s[4:5]
|
||||
EncodeSop2(0x00, 4, 4, 255), // s_add_u32 s4, s4, literal
|
||||
0x00000034u,
|
||||
EncodeSop2(0x04, 5, 5, 128), // s_addc_u32 s5, s5, 0
|
||||
EncodeSmem0(0x00, 6, 2),
|
||||
(106u << 25u) | 4u, // s_load_dword s6, s[4:5], vcc_lo offset:4
|
||||
EncodeSopp(0x0c, 0), // s_waitcnt 0
|
||||
EncodeSop2(0x01, 4, 4, 6), // s_sub_u32 s4, s4, s6
|
||||
EncodeSop2(0x05, 5, 5, 128), // s_subb_u32 s5, s5, 0
|
||||
EncodeSop1(0x20, 0, 4), // s_setpc_b64 s[4:5]
|
||||
EncodeSMovB32(1, 129), // case 0
|
||||
EncodeSopp(0x02, 1),
|
||||
EncodeSMovB32(2, 129), // case 1
|
||||
EncodeSopp(0x01, 0),
|
||||
0u,
|
||||
0x00000010u,
|
||||
0x00000008u,
|
||||
0x00000010u, // table at pc 0x44, targets relative backward from pc 0x40
|
||||
};
|
||||
|
||||
auto options = MakeCompileOptions(ShaderType::Compute);
|
||||
options.dump_ir = true;
|
||||
|
||||
ShaderRecompiler::CompileResult result;
|
||||
std::string error;
|
||||
Check(ShaderRecompiler::TryRecompile(shader, options, result, &error),
|
||||
error.c_str());
|
||||
Check(Common::ContainsStr(result.ir_dump, "mode=dispatcher"),
|
||||
"subtractive S_SETPC_B64 table did not select dispatcher fallback");
|
||||
const auto jump = std::find_if(
|
||||
result.program.values->block_info.begin(),
|
||||
result.program.values->block_info.end(), [](const auto &block) {
|
||||
return !block.terminator.indirect_targets.empty();
|
||||
});
|
||||
const auto block_pc = [&](uint32_t id) {
|
||||
const auto block =
|
||||
std::find_if(result.program.values->block_info.begin(),
|
||||
result.program.values->block_info.end(),
|
||||
[=](const auto &info) { return info.id == id; });
|
||||
return block != result.program.values->block_info.end() ? block->start_pc
|
||||
: UINT32_MAX;
|
||||
};
|
||||
Check(jump != result.program.values->block_info.end() &&
|
||||
jump->terminator.indirect_targets.size() == 2 &&
|
||||
jump->terminator.indirect_target_pcs ==
|
||||
std::vector<uint32_t>({0x30u, 0x38u}) &&
|
||||
jump->terminator.indirect_selector_values ==
|
||||
std::vector<uint32_t>({0u, 4u, 8u}) &&
|
||||
jump->terminator.indirect_selector_targets.size() == 3 &&
|
||||
block_pc(jump->terminator.indirect_selector_targets[0]) == 0x30u &&
|
||||
block_pc(jump->terminator.indirect_selector_targets[1]) == 0x38u &&
|
||||
block_pc(jump->terminator.indirect_selector_targets[2]) == 0x30u,
|
||||
"subtractive S_SETPC_B64 table targets or selector mapping changed");
|
||||
Check(!Common::ContainsStr(result.ir_dump, "SLoadDword"),
|
||||
"subtractive S_SETPC_B64 table load reached normal IR");
|
||||
Check(SpirvContainsOpcode(result.spirv, 251),
|
||||
"subtractive S_SETPC_B64 dispatcher lacks OpSwitch");
|
||||
CheckSpirvBinaryValidates(result.spirv);
|
||||
}
|
||||
|
||||
void TestNewShaderRecompilerExpVertexOutputs() {
|
||||
const uint32_t shader[] = {
|
||||
EncodeExp0(0x0c, 0xf), EncodeExp1(0, 1, 2, 3), // POS0
|
||||
@@ -9280,6 +9344,7 @@ int main() {
|
||||
TestNewShaderRecompilerBranchConditionForms();
|
||||
TestNewShaderRecompilerSetpcBranch();
|
||||
TestNewShaderRecompilerSetpcJumpTable();
|
||||
TestNewShaderRecompilerSetpcDwordJumpTable();
|
||||
TestNewShaderRecompilerZeroInitialRegisterState();
|
||||
TestNewShaderRecompilerVertexExportUsesLaneExecMask();
|
||||
TestNewShaderRecompilerPerInvocationU64Complement();
|
||||
|
||||
Reference in New Issue
Block a user