diff options
Diffstat (limited to 'source/slang/slang-emit-spirv.cpp')
| -rw-r--r-- | source/slang/slang-emit-spirv.cpp | 40 |
1 files changed, 28 insertions, 12 deletions
diff --git a/source/slang/slang-emit-spirv.cpp b/source/slang/slang-emit-spirv.cpp index d8c479cd1..1407404ad 100644 --- a/source/slang/slang-emit-spirv.cpp +++ b/source/slang/slang-emit-spirv.cpp @@ -1311,6 +1311,8 @@ struct SPIRVEmitContext : public SourceEmitterBase, public SPIRVEmitSharedContex return SpvStorageClassImage; case AddressSpace::UserPointer: return SpvStorageClassPhysicalStorageBuffer; + case AddressSpace::NodePayloadAMDX: + return SpvStorageClassNodePayloadAMDX; case AddressSpace::Global: case AddressSpace::MetalObjectData: case AddressSpace::SpecializationConstant: @@ -1504,13 +1506,22 @@ struct SPIRVEmitContext : public SourceEmitterBase, public SPIRVEmitSharedContex SLANG_ASSERT(ptrType); if (ptrType->hasAddressSpace()) storageClass = addressSpaceToStorageClass(ptrType->getAddressSpace()); - if (storageClass == SpvStorageClassStorageBuffer) + + switch (storageClass) + { + case SpvStorageClassStorageBuffer: ensureExtensionDeclaration( UnownedStringSlice("SPV_KHR_storage_buffer_storage_class")); - if (storageClass == SpvStorageClassPhysicalStorageBuffer) - { + break; + case SpvStorageClassPhysicalStorageBuffer: requirePhysicalStorageAddressing(); + break; + case SpvStorageClassNodePayloadAMDX: + requireSPIRVCapability(SpvCapabilityShaderEnqueueAMDX); + ensureExtensionDeclaration(UnownedStringSlice("SPV_AMDX_shader_enqueue")); + break; } + auto valueType = ptrType->getValueType(); // If we haven't emitted the inner type yet, we need to emit a forward declaration. bool useForwardDeclaration = @@ -1524,17 +1535,20 @@ struct SPIRVEmitContext : public SourceEmitterBase, public SPIRVEmitSharedContex builder.setInsertBefore(valueType); valueTypeId = getID(ensureInst(builder.getUIntType())); } + else if (useForwardDeclaration) + { + valueTypeId = getIRInstSpvID(valueType); + } + else if (storageClass == SpvStorageClassNodePayloadAMDX) + { + auto spvValueType = ensureInst(valueType); + auto spvNodePayloadType = emitOpTypeNodePayloadArray(inst, spvValueType); + valueTypeId = getID(spvNodePayloadType); + } else { - if (useForwardDeclaration) - { - valueTypeId = getIRInstSpvID(valueType); - } - else - { - auto spvValueType = ensureInst(valueType); - valueTypeId = getID(spvValueType); - } + auto spvValueType = ensureInst(valueType); + valueTypeId = getID(spvValueType); } auto resultSpvType = emitOpTypePointer(inst, storageClass, valueTypeId); @@ -7564,6 +7578,8 @@ struct SPIRVEmitContext : public SourceEmitterBase, public SPIRVEmitSharedContex case SpvOpMemberDecorate: case SpvOpMemberDecorateString: return getSection(SpvLogicalSectionID::Annotations); + case SpvOpTypeNodePayloadArrayAMDX: + return getSection(SpvLogicalSectionID::ConstantsAndTypes); default: return defaultParent; } |
