Some work around CAS

This commit is contained in:
CamilleLaVey
2026-09-29 16:13:50 -04:00
parent ef7cde0c2b
commit 2d6f28c50e
10 changed files with 347 additions and 265 deletions
@@ -486,8 +486,7 @@ EmitContext::EmitContext(const Profile& profile_, const RuntimeInfo& runtime_inf
DefineTextures(program.info, texture_binding, bindings.texture_scaling_index);
DefineImages(program.info, image_binding, bindings.image_scaling_index);
DefineAttributeMemAccess(program.info);
DefineWriteStorageCasLoopFunction(program.info);
DefineGlobalMemoryFunctions(program.info);
DefineGlobalMemoryFunctions(program);
DefineRescalingInput(program.info);
DefineRenderArea(program.info);
}
@@ -532,6 +531,24 @@ Id EmitContext::BitOffset16(const IR::Value& offset) {
return OpBitwiseAnd(U32[1], OpShiftLeftLogical(U32[1], Def(offset), Const(3u)), Const(16u));
}
Id EmitContext::CallGlobalMemory(IR::Opcode opcode, Id result_type, Id address, Id value) {
if (profile.support_int64) {
return OpFunctionCall(result_type, global_memory_funcs.at(opcode), address, value);
}
if (result_type.value == void_id.value) {
return Id{};
}
return ConstantNull(result_type);
}
void EmitContext::AtomicBitFieldInsert(Id pointer, Id value, Id offset, Id count) {
const Id scope{Const(static_cast<u32>(spv::Scope::Device))};
const Id mask{OpBitFieldInsert(U32[1], u32_zero_value, Const(0xFFFFFFFFU), offset, count)};
const Id bits{OpBitFieldInsert(U32[1], u32_zero_value, value, offset, count)};
OpAtomicAnd(U32[1], pointer, scope, u32_zero_value, OpNot(U32[1], mask));
OpAtomicOr(U32[1], pointer, scope, u32_zero_value, bits);
}
void EmitContext::DefineCommonTypes(const Info& info) {
void_id = TypeVoid();
@@ -889,142 +906,241 @@ void EmitContext::DefineAttributeMemAccess(const Info& info) {
}
}
void EmitContext::DefineWriteStorageCasLoopFunction(const Info& info) {
if (profile.support_int8 && profile.support_int16) {
return;
}
if (!info.uses_int8 && !info.uses_int16) {
return;
}
AddCapability(spv::Capability::VariablePointersStorageBuffer);
const Id ptr_type{TypePointer(spv::StorageClass::StorageBuffer, U32[1])};
const Id func_type{TypeFunction(void_id, ptr_type, U32[1], U32[1], U32[1])};
const Id func{OpFunction(void_id, spv::FunctionControlMask::MaskNone, func_type)};
const Id pointer{OpFunctionParameter(ptr_type)};
const Id value{OpFunctionParameter(U32[1])};
const Id bit_offset{OpFunctionParameter(U32[1])};
const Id bit_count{OpFunctionParameter(U32[1])};
AddLabel();
const Id scope_device{Const(1u)};
const Id ordering_relaxed{u32_zero_value};
const Id body_label{OpLabel()};
const Id continue_label{OpLabel()};
const Id endloop_label{OpLabel()};
const Id beginloop_label{OpLabel()};
OpBranch(beginloop_label);
AddLabel(beginloop_label);
OpLoopMerge(endloop_label, continue_label, spv::LoopControlMask::MaskNone);
OpBranch(body_label);
AddLabel(body_label);
const Id expected_value{OpLoad(U32[1], pointer)};
const Id desired_value{OpBitFieldInsert(U32[1], expected_value, value, bit_offset, bit_count)};
const Id actual_value{OpAtomicCompareExchange(U32[1], pointer, scope_device, ordering_relaxed,
ordering_relaxed, desired_value, expected_value)};
const Id store_successful{OpIEqual(U1, expected_value, actual_value)};
OpBranchConditional(store_successful, endloop_label, continue_label);
AddLabel(endloop_label);
OpReturn();
AddLabel(continue_label);
OpBranch(beginloop_label);
OpFunctionEnd();
write_storage_cas_loop_func = func;
}
void EmitContext::DefineGlobalMemoryFunctions(const Info& info) {
void EmitContext::DefineGlobalMemoryFunctions(const IR::Program& program) {
const Info& info{program.info};
if (!info.uses_global_memory || !profile.support_int64) {
return;
}
using DefPtr = Id StorageDefinitions::*;
const Id zero{u32_zero_value};
const auto define_body{[&](DefPtr ssbo_member, Id addr, Id element_pointer, u32 shift,
auto&& callback) {
const Id scope{Const(static_cast<u32>(spv::Scope::Device))};
const Id align_mask{Const(~(static_cast<u32>(profile.min_ssbo_alignment) - 1U))};
const auto word_pointer{[&](Id ssbo, Id word, u32 element) {
return OpAccessChain(storage_types.U32.element, ssbo, zero,
OpIAdd(U32[1], word, Const(element)));
}};
const auto cbuf_word{[&](u32 index, u32 offset) {
if (profile.support_descriptor_aliasing) {
return OpLoad(U32[1], OpAccessChain(uniform_types.U32, cbufs[index].U32, zero,
Const(offset / 4)));
}
const Id vector{OpLoad(U32[4], OpAccessChain(uniform_types.U32x4, cbufs[index].U32x4,
zero, Const(offset / 16)))};
return OpCompositeExtract(U32[1], vector, (offset / 4) % 4);
}};
const auto bits{[&](Id offset, u32 count) {
return OpBitwiseAnd(U32[1], OpShiftLeftLogical(U32[1], offset, Const(3U)),
Const(32U - count));
}};
const auto define{[&](IR::Opcode opcode, Id result_type, Id value_type, auto&& callback) {
const std::array<Id, 2> params{U64, value_type};
const Id func{OpFunction(result_type, spv::FunctionControlMask::MaskNone,
TypeFunction(result_type, params))};
const Id addr{OpFunctionParameter(U64)};
const Id value{OpFunctionParameter(value_type)};
const bool returns_value{result_type.value != void_id.value};
AddLabel();
const Id addr_words{OpBitcast(U32[2], addr)};
const Id addr_low{OpCompositeExtract(U32[1], addr_words, 0U)};
const Id addr_high{OpCompositeExtract(U32[1], addr_words, 1U)};
const Id align_mask{Const(~(static_cast<u32>(profile.min_ssbo_alignment) - 1U))};
const size_t num_buffers{info.storage_buffers_descriptors.size()};
for (size_t index = 0; index < num_buffers; ++index) {
const auto& ssbo{info.storage_buffers_descriptors[index]};
if (!ssbo.is_global_fallback) {
for (size_t index = 0; index < info.storage_buffers_descriptors.size(); ++index) {
const auto& desc{info.storage_buffers_descriptors[index]};
if (!desc.is_global_fallback) {
continue;
}
const Id ssbo_addr_cbuf_offset{Const(ssbo.cbuf_offset / 8)};
const Id ssbo_size_cbuf_offset{Const(ssbo.cbuf_offset / 4 + 2)};
const Id ssbo_addr_pointer{OpAccessChain(
uniform_types.U32x2, cbufs[ssbo.cbuf_index].U32x2, zero, ssbo_addr_cbuf_offset)};
const Id ssbo_size_pointer{OpAccessChain(uniform_types.U32, cbufs[ssbo.cbuf_index].U32,
zero, ssbo_size_cbuf_offset)};
const Id ssbo_addr{OpLoad(U32[2], ssbo_addr_pointer)};
const Id ssbo_low{
OpBitwiseAnd(U32[1], OpCompositeExtract(U32[1], ssbo_addr, 0U), align_mask)};
const Id ssbo_high{OpCompositeExtract(U32[1], ssbo_addr, 1U)};
const Id ssbo_size{OpLoad(U32[1], ssbo_size_pointer)};
const Id ssbo_offset{OpISub(U32[1], addr_low, ssbo_low)};
OpBitwiseAnd(U32[1], cbuf_word(desc.cbuf_index, desc.cbuf_offset), align_mask)};
const Id ssbo_high{cbuf_word(desc.cbuf_index, desc.cbuf_offset + 4)};
const Id ssbo_size{cbuf_word(desc.cbuf_index, desc.cbuf_offset + 8)};
const Id offset{OpISub(U32[1], addr_low, ssbo_low)};
const Id borrow{
OpSelect(U32[1], OpULessThan(U1, addr_low, ssbo_low), Const(1U), zero)};
const Id cond{
OpLogicalAnd(U1, OpULessThan(U1, ssbo_offset, ssbo_size),
OpIEqual(U1, OpISub(U32[1], addr_high, borrow), ssbo_high))};
const Id cond{OpLogicalAnd(U1, OpULessThan(U1, offset, ssbo_size),
OpIEqual(U1, OpISub(U32[1], addr_high, borrow), ssbo_high))};
const Id then_label{OpLabel()};
const Id else_label{OpLabel()};
OpSelectionMerge(else_label, spv::SelectionControlMask::MaskNone);
OpBranchConditional(cond, then_label, else_label);
AddLabel(then_label);
const Id ssbo_id{ssbos[index].*ssbo_member};
const Id ssbo_index{OpShiftRightLogical(U32[1], ssbo_offset, Const(shift))};
const Id ssbo_pointer{OpAccessChain(element_pointer, ssbo_id, zero, ssbo_index)};
callback(ssbo_pointer);
const Id word{OpShiftRightLogical(U32[1], offset, Const(2U))};
const Id result{callback(ssbos[index].U32, word, offset, value)};
if (returns_value) {
OpReturnValue(result);
} else {
OpReturn();
}
AddLabel(else_label);
}
}};
const auto define_load{[&](DefPtr ssbo_member, Id element_pointer, Id type, u32 shift) {
const Id function_type{TypeFunction(type, U64)};
const Id func_id{OpFunction(type, spv::FunctionControlMask::MaskNone, function_type)};
const Id addr{OpFunctionParameter(U64)};
define_body(ssbo_member, addr, element_pointer, shift,
[&](Id ssbo_pointer) { OpReturnValue(OpLoad(type, ssbo_pointer)); });
OpReturnValue(ConstantNull(type));
OpFunctionEnd();
return func_id;
}};
const auto define_write{[&](DefPtr ssbo_member, Id element_pointer, Id type, u32 shift) {
const Id function_type{TypeFunction(void_id, U64, type)};
const Id func_id{OpFunction(void_id, spv::FunctionControlMask::MaskNone, function_type)};
const Id addr{OpFunctionParameter(U64)};
const Id data{OpFunctionParameter(type)};
define_body(ssbo_member, addr, element_pointer, shift, [&](Id ssbo_pointer) {
OpStore(ssbo_pointer, data);
if (returns_value) {
OpReturnValue(ConstantNull(result_type));
} else {
OpReturn();
});
OpReturn();
}
OpFunctionEnd();
return func_id;
global_memory_funcs.emplace(opcode, func);
}};
const auto define{
[&](DefPtr ssbo_member, const StorageTypeDefinition& type_def, Id type, size_t size) {
const Id element_type{type_def.element};
const u32 shift{static_cast<u32>(std::countr_zero(size))};
const Id load_func{define_load(ssbo_member, element_type, type, shift)};
const Id write_func{define_write(ssbo_member, element_type, type, shift)};
return std::make_pair(load_func, write_func);
}};
std::tie(load_global_func_u32, write_global_func_u32) =
define(&StorageDefinitions::U32, storage_types.U32, U32[1], sizeof(u32));
std::tie(load_global_func_u32x2, write_global_func_u32x2) =
define(&StorageDefinitions::U32x2, storage_types.U32x2, U32[2], sizeof(u32[2]));
std::tie(load_global_func_u32x4, write_global_func_u32x4) =
define(&StorageDefinitions::U32x4, storage_types.U32x4, U32[4], sizeof(u32[4]));
const auto load{[&](Id type, u32 count) {
return [&, type, count](Id ssbo, Id word, Id, Id) {
std::array<Id, 4> words{};
for (u32 element = 0; element < count; ++element) {
words[element] = OpLoad(U32[1], word_pointer(ssbo, word, element));
}
if (count == 1) {
return words[0];
}
return OpCompositeConstruct(type, std::span<const Id>(words.data(), count));
};
}};
const auto store{[&](u32 count) {
return [&, count](Id ssbo, Id word, Id, Id value) {
if (count == 1) {
OpStore(word_pointer(ssbo, word, 0), value);
return Id{};
}
for (u32 element = 0; element < count; ++element) {
OpStore(word_pointer(ssbo, word, element),
OpCompositeExtract(U32[1], value, element));
}
return Id{};
};
}};
const auto extract{[&](bool is_signed, u32 count) {
return [&, is_signed, count](Id ssbo, Id word, Id offset, Id) {
const Id loaded{OpLoad(U32[1], word_pointer(ssbo, word, 0))};
if (is_signed) {
return OpBitFieldSExtract(U32[1], loaded, bits(offset, count), Const(count));
}
return OpBitFieldUExtract(U32[1], loaded, bits(offset, count), Const(count));
};
}};
const auto insert{[&](u32 count) {
return [&, count](Id ssbo, Id word, Id offset, Id value) {
AtomicBitFieldInsert(word_pointer(ssbo, word, 0), value, bits(offset, count),
Const(count));
return Id{};
};
}};
const auto atomic{[&](Id (Sirit::Module::*func)(Id, Id, Id, Id, Id)) {
return [&, func](Id ssbo, Id word, Id, Id value) {
return (this->*func)(U32[1], word_pointer(ssbo, word, 0), scope, zero, value);
};
}};
const auto cas{[&](Id type, Id helper) {
return [&, type, helper](Id ssbo, Id word, Id, Id value) {
return OpFunctionCall(type, helper, word, value, ssbo);
};
}};
const auto packed{[&](bool is_half, Id helper) {
return [&, is_half, helper](Id ssbo, Id word, Id, Id value) {
if (is_half) {
return OpBitcast(U32[1], OpFunctionCall(F16[2], helper, word, value, ssbo));
}
return OpPackHalf2x16(U32[1], OpFunctionCall(F32[2], helper, word, value, ssbo));
};
}};
for (const IR::Block* const block : program.post_order_blocks) {
for (const IR::Inst& inst : block->Instructions()) {
const IR::Opcode opcode{inst.GetOpcode()};
if (global_memory_funcs.contains(opcode)) {
continue;
}
switch (opcode) {
case IR::Opcode::LoadGlobalU8:
define(opcode, U32[1], U32[1], extract(false, 8));
break;
case IR::Opcode::LoadGlobalS8:
define(opcode, U32[1], U32[1], extract(true, 8));
break;
case IR::Opcode::LoadGlobalU16:
define(opcode, U32[1], U32[1], extract(false, 16));
break;
case IR::Opcode::LoadGlobalS16:
define(opcode, U32[1], U32[1], extract(true, 16));
break;
case IR::Opcode::LoadGlobal32:
define(opcode, U32[1], U32[1], load(U32[1], 1));
break;
case IR::Opcode::LoadGlobal64:
define(opcode, U32[2], U32[1], load(U32[2], 2));
break;
case IR::Opcode::LoadGlobal128:
define(opcode, U32[4], U32[1], load(U32[4], 4));
break;
case IR::Opcode::WriteGlobalU8:
case IR::Opcode::WriteGlobalS8:
define(opcode, void_id, U32[1], insert(8));
break;
case IR::Opcode::WriteGlobalU16:
case IR::Opcode::WriteGlobalS16:
define(opcode, void_id, U32[1], insert(16));
break;
case IR::Opcode::WriteGlobal32:
define(opcode, void_id, U32[1], store(1));
break;
case IR::Opcode::WriteGlobal64:
define(opcode, void_id, U32[2], store(2));
break;
case IR::Opcode::WriteGlobal128:
define(opcode, void_id, U32[4], store(4));
break;
case IR::Opcode::GlobalAtomicIAdd32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicIAdd));
break;
case IR::Opcode::GlobalAtomicSMin32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicSMin));
break;
case IR::Opcode::GlobalAtomicUMin32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicUMin));
break;
case IR::Opcode::GlobalAtomicSMax32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicSMax));
break;
case IR::Opcode::GlobalAtomicUMax32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicUMax));
break;
case IR::Opcode::GlobalAtomicAnd32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicAnd));
break;
case IR::Opcode::GlobalAtomicOr32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicOr));
break;
case IR::Opcode::GlobalAtomicXor32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicXor));
break;
case IR::Opcode::GlobalAtomicExchange32:
define(opcode, U32[1], U32[1], atomic(&Sirit::Module::OpAtomicExchange));
break;
case IR::Opcode::GlobalAtomicInc32:
define(opcode, U32[1], U32[1], cas(U32[1], increment_cas_ssbo));
break;
case IR::Opcode::GlobalAtomicDec32:
define(opcode, U32[1], U32[1], cas(U32[1], decrement_cas_ssbo));
break;
case IR::Opcode::GlobalAtomicAddF32:
define(opcode, F32[1], F32[1], cas(F32[1], f32_add_cas));
break;
case IR::Opcode::GlobalAtomicAddF16x2:
define(opcode, U32[1], F16[2], packed(true, f16x2_add_cas));
break;
case IR::Opcode::GlobalAtomicMinF16x2:
define(opcode, U32[1], F16[2], packed(true, f16x2_min_cas));
break;
case IR::Opcode::GlobalAtomicMaxF16x2:
define(opcode, U32[1], F16[2], packed(true, f16x2_max_cas));
break;
case IR::Opcode::GlobalAtomicAddF32x2:
define(opcode, U32[1], F32[2], packed(false, f32x2_add_cas));
break;
case IR::Opcode::GlobalAtomicMinF32x2:
define(opcode, U32[1], F32[2], packed(false, f32x2_min_cas));
break;
case IR::Opcode::GlobalAtomicMaxF32x2:
define(opcode, U32[1], F32[2], packed(false, f32x2_max_cas));
break;
default:
break;
}
}
}
}
void EmitContext::DefineRescalingInput(const Info& info) {