From 746d47bb491e0b97e35ab373b4b78d33b9a61164 Mon Sep 17 00:00:00 2001 From: Yong He Date: Wed, 10 Jul 2024 16:17:10 -0700 Subject: Specialize address space during spirv legalization. (#4600) * Specialize address space during spirv legalization. * Fix. * Fix building doc. * Fix cmake. * Update assert. --- source/slang/slang-ir-specialize-address-space.cpp | 84 ++++++++-------------- 1 file changed, 30 insertions(+), 54 deletions(-) (limited to 'source/slang/slang-ir-specialize-address-space.cpp') diff --git a/source/slang/slang-ir-specialize-address-space.cpp b/source/slang/slang-ir-specialize-address-space.cpp index 55d61d527..1d899e240 100644 --- a/source/slang/slang-ir-specialize-address-space.cpp +++ b/source/slang/slang-ir-specialize-address-space.cpp @@ -7,66 +7,30 @@ namespace Slang { - struct AddressSpaceContext + struct AddressSpaceContext : public AddressSpaceSpecializationContext { IRModule* module; Dictionary mapInstToAddrSpace; + InitialAddressSpaceAssigner* addrSpaceAssigner; - AddressSpaceContext(IRModule* inModule) + AddressSpaceContext(IRModule* inModule, InitialAddressSpaceAssigner* inAddrSpaceAssigner) : module(inModule) + , addrSpaceAssigner(inAddrSpaceAssigner) { } AddressSpace getAddressSpaceFromVarType(IRInst* type) { - if (as(type)) - { - return AddressSpace::Uniform; - } - if (as(type)) - { - return AddressSpace::Global; - } - if (as(type)) - { - return AddressSpace::Global; - } - if (as(type)) - { - return AddressSpace::Global; - } - if (auto ptrType = as(type)) - { - if (ptrType->hasAddressSpace()) - return (AddressSpace)ptrType->getAddressSpace(); - return AddressSpace::Global; - } - return AddressSpace::Generic; + return addrSpaceAssigner->getAddressSpaceFromVarType(type); } AddressSpace getLeafInstAddressSpace(IRInst* inst) { - if (as(inst->getRate())) - return AddressSpace::GroupShared; - switch (inst->getOp()) - { - case kIROp_RWStructuredBufferGetElementPtr: - return AddressSpace::Global; - case kIROp_Var: - if (as(inst->getParent())) - return AddressSpace::ThreadLocal; - break; - default: - break; - } - auto type = unwrapAttributedType(inst->getDataType()); - if (!type) - return AddressSpace::Generic; - return getAddressSpaceFromVarType(type); + return addrSpaceAssigner->getLeafInstAddressSpace(inst); } - AddressSpace getAddrSpace(IRInst* inst) + AddressSpace getAddrSpace(IRInst* inst) override { auto addrSpace = mapInstToAddrSpace.tryGetValue(inst); if (addrSpace) @@ -186,20 +150,29 @@ namespace Slang continue; } + // If the inst already has a pointer type with explicit address space, then use it. + if (auto ptrType = as(inst->getDataType())) + { + if (ptrType->hasAddressSpace()) + { + mapInstToAddrSpace[inst] = (AddressSpace)ptrType->getAddressSpace(); + continue; + } + } + + // Otherwise, try to assign an address space based on the instruction type. switch (inst->getOp()) { case kIROp_Var: - { - // All local variables should be in the thread-local address space. - mapInstToAddrSpace[inst] = AddressSpace::ThreadLocal; - changed = true; - break; - } case kIROp_RWStructuredBufferGetElementPtr: { - // The address space of the result of RWStructuredBufferGetElementPtr is always global. - mapInstToAddrSpace[inst] = AddressSpace::Global; - changed = true; + // The address space of these insts should be assigned by the initial address space assigner. + AddressSpace addrSpace = AddressSpace::Generic; + if (addrSpaceAssigner->tryAssignAddressSpace(inst, addrSpace)) + { + mapInstToAddrSpace[inst] = addrSpace; + changed = true; + } break; } case kIROp_GetElementPtr: @@ -340,7 +313,10 @@ namespace Slang { auto rate = inst->getRate(); if (!rate) + { inst->setFullType(dataType); + return; + } IRBuilder builder(inst); builder.setInsertBefore(inst); @@ -405,9 +381,9 @@ namespace Slang } }; - void specializeAddressSpace(IRModule* module) + void specializeAddressSpace(IRModule* module, InitialAddressSpaceAssigner* addrSpaceAssigner) { - AddressSpaceContext context(module); + AddressSpaceContext context(module, addrSpaceAssigner); context.processModule(); } } -- cgit v1.2.3