summaryrefslogtreecommitdiff
path: root/source/slang/slang-ir-specialize-address-space.cpp
diff options
context:
space:
mode:
authorYong He <yonghe@outlook.com>2024-07-10 16:17:10 -0700
committerGitHub <noreply@github.com>2024-07-10 16:17:10 -0700
commit746d47bb491e0b97e35ab373b4b78d33b9a61164 (patch)
tree74e0936472d911d8c6c561ca4b21e800306c5f51 /source/slang/slang-ir-specialize-address-space.cpp
parent82f308ca692878bfe9844b86629c6536b4cd0f0a (diff)
Specialize address space during spirv legalization. (#4600)
* Specialize address space during spirv legalization. * Fix. * Fix building doc. * Fix cmake. * Update assert.
Diffstat (limited to 'source/slang/slang-ir-specialize-address-space.cpp')
-rw-r--r--source/slang/slang-ir-specialize-address-space.cpp84
1 files changed, 30 insertions, 54 deletions
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<IRInst*, AddressSpace> mapInstToAddrSpace;
+ InitialAddressSpaceAssigner* addrSpaceAssigner;
- AddressSpaceContext(IRModule* inModule)
+ AddressSpaceContext(IRModule* inModule, InitialAddressSpaceAssigner* inAddrSpaceAssigner)
: module(inModule)
+ , addrSpaceAssigner(inAddrSpaceAssigner)
{
}
AddressSpace getAddressSpaceFromVarType(IRInst* type)
{
- if (as<IRUniformParameterGroupType>(type))
- {
- return AddressSpace::Uniform;
- }
- if (as<IRByteAddressBufferTypeBase>(type))
- {
- return AddressSpace::Global;
- }
- if (as<IRHLSLStructuredBufferTypeBase>(type))
- {
- return AddressSpace::Global;
- }
- if (as<IRGLSLShaderStorageBufferType>(type))
- {
- return AddressSpace::Global;
- }
- if (auto ptrType = as<IRPtrTypeBase>(type))
- {
- if (ptrType->hasAddressSpace())
- return (AddressSpace)ptrType->getAddressSpace();
- return AddressSpace::Global;
- }
- return AddressSpace::Generic;
+ return addrSpaceAssigner->getAddressSpaceFromVarType(type);
}
AddressSpace getLeafInstAddressSpace(IRInst* inst)
{
- if (as<IRGroupSharedRate>(inst->getRate()))
- return AddressSpace::GroupShared;
- switch (inst->getOp())
- {
- case kIROp_RWStructuredBufferGetElementPtr:
- return AddressSpace::Global;
- case kIROp_Var:
- if (as<IRBlock>(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<IRPtrTypeBase>(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();
}
}