From 73f438e18537ec301b9257179f4c50eac38d34a6 Mon Sep 17 00:00:00 2001 From: dmcdiar Date: Fri, 14 May 2021 19:10:25 -0700 Subject: [PATCH] Changed the descriptor a shared_ptr. --- .../DiffuseProbeGridRayTracingPass.cpp | 6 ++--- .../DiffuseProbeGridRayTracingPass.h | 1 - .../Code/Include/Atom/RHI/FrameScheduler.h | 4 ++-- .../Include/Atom/RHI/RayTracingShaderTable.h | 5 ++-- .../Code/Source/RHI/RayTracingShaderTable.cpp | 1 - .../Code/Source/RHI/RayTracingShaderTable.cpp | 24 +++++++++---------- .../Code/Source/RHI/RayTracingShaderTable.cpp | 18 +++++++------- 7 files changed, 27 insertions(+), 32 deletions(-) diff --git a/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.cpp b/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.cpp index 540b50f3ba..b2b5f60660 100644 --- a/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.cpp +++ b/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.cpp @@ -275,12 +275,12 @@ namespace AZ // scene changed, need to rebuild the shader table m_rayTracingRevision = rayTracingRevision; - m_rayTracingShaderTableDescriptor = AZStd::make_shared(); + AZStd::shared_ptr descriptor = AZStd::make_shared(); if (rayTracingFeatureProcessor->GetSubMeshCount()) { // build the ray tracing shader table descriptor - RHI::RayTracingShaderTableDescriptor* descriptorBuild = m_rayTracingShaderTableDescriptor->Build(AZ::Name("RayTracingShaderTable"), m_rayTracingPipelineState) + RHI::RayTracingShaderTableDescriptor* descriptorBuild = descriptor->Build(AZ::Name("RayTracingShaderTable"), m_rayTracingPipelineState) ->RayGenerationRecord(AZ::Name("RayGen")) ->MissRecord(AZ::Name("Miss")); @@ -291,7 +291,7 @@ namespace AZ } } - m_rayTracingShaderTable->Build(m_rayTracingShaderTableDescriptor); + m_rayTracingShaderTable->Build(descriptor); } } diff --git a/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.h b/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.h index 67e5503825..bb35803e51 100644 --- a/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.h +++ b/Gems/Atom/Feature/Common/Code/Source/DiffuseProbeGrid/DiffuseProbeGridRayTracingPass.h @@ -62,7 +62,6 @@ namespace AZ // ray tracing shader table RHI::Ptr m_rayTracingShaderTable; - AZStd::shared_ptr m_rayTracingShaderTableDescriptor; // ray tracing global shader resource group asset and pipeline state Data::Asset m_globalSrgAsset; diff --git a/Gems/Atom/RHI/Code/Include/Atom/RHI/FrameScheduler.h b/Gems/Atom/RHI/Code/Include/Atom/RHI/FrameScheduler.h index cc603b89ae..4cb0280837 100644 --- a/Gems/Atom/RHI/Code/Include/Atom/RHI/FrameScheduler.h +++ b/Gems/Atom/RHI/Code/Include/Atom/RHI/FrameScheduler.h @@ -18,6 +18,7 @@ #include #include #include +#include #include #include #include @@ -31,7 +32,6 @@ namespace AZ { class ShaderResourceGroupPool; class FrameGraphExecuteGroup; - class RayTracingShaderTable; //! @brief Fill this descriptor when initializing a FrameScheduler instance. struct FrameSchedulerDescriptor @@ -231,7 +231,7 @@ namespace AZ AZStd::unordered_map m_scopeProducerLookup; // list of RayTracingShaderTables that should be built this frame - AZStd::vector m_rayTracingShaderTablesToBuild; + AZStd::vector> m_rayTracingShaderTablesToBuild; }; } } diff --git a/Gems/Atom/RHI/Code/Include/Atom/RHI/RayTracingShaderTable.h b/Gems/Atom/RHI/Code/Include/Atom/RHI/RayTracingShaderTable.h index bb16005ef5..57ab9b87fc 100644 --- a/Gems/Atom/RHI/Code/Include/Atom/RHI/RayTracingShaderTable.h +++ b/Gems/Atom/RHI/Code/Include/Atom/RHI/RayTracingShaderTable.h @@ -99,12 +99,13 @@ namespace AZ static RHI::Ptr CreateRHIRayTracingShaderTable(); void Init(Device& device, const RayTracingBufferPools& rayTracingBufferPools); - //! Queues this RayTracingShaderTable to be built by the FrameScheduler + //! Queues this RayTracingShaderTable to be built by the FrameScheduler. + //! Note that the descriptor must be heap allocated, preferably using make_shared. void Build(const AZStd::shared_ptr descriptor); protected: - AZStd::weak_ptr m_descriptor; + AZStd::shared_ptr m_descriptor; const RayTracingBufferPools* m_bufferPools = nullptr; private: diff --git a/Gems/Atom/RHI/Code/Source/RHI/RayTracingShaderTable.cpp b/Gems/Atom/RHI/Code/Source/RHI/RayTracingShaderTable.cpp index 07204a26f0..5934cb2e22 100644 --- a/Gems/Atom/RHI/Code/Source/RHI/RayTracingShaderTable.cpp +++ b/Gems/Atom/RHI/Code/Source/RHI/RayTracingShaderTable.cpp @@ -99,7 +99,6 @@ namespace AZ void RayTracingShaderTable::Validate() { AZ_Assert(m_isQueuedForBuild, "Attempting to build a RayTracingShaderTable that is not queued."); - AZ_Assert(!m_descriptor.expired(), "RayTracingShaderTable descriptor is no longer valid, make sure it is not freed after calling Build."); AZ_Assert(m_bufferPools, "RayTracingBufferPools pointer is null."); } } diff --git a/Gems/Atom/RHI/DX12/Code/Source/RHI/RayTracingShaderTable.cpp b/Gems/Atom/RHI/DX12/Code/Source/RHI/RayTracingShaderTable.cpp index 7899e7dc21..f7d32f9392 100644 --- a/Gems/Atom/RHI/DX12/Code/Source/RHI/RayTracingShaderTable.cpp +++ b/Gems/Atom/RHI/DX12/Code/Source/RHI/RayTracingShaderTable.cpp @@ -128,10 +128,8 @@ namespace AZ m_currentBufferIndex = (m_currentBufferIndex + 1) % BufferCount; ShaderTableBuffers& buffers = m_buffers[m_currentBufferIndex]; - AZStd::shared_ptr descriptor = m_descriptor.lock(); - // clear the shader table if the descriptor has no ray generation shader - if (descriptor->GetRayGenerationRecord().empty()) + if (m_descriptor->GetRayGenerationRecord().empty()) { buffers.m_rayGenerationTable = nullptr; buffers.m_rayGenerationTableSize = 0; @@ -146,7 +144,7 @@ namespace AZ // retrieve the ID3D12StateObjectProperties interface from the raytracing pipeline state object // this is needed to get the shader identifiers to put in the table - const RayTracingPipelineState* rayTracingPipelineState = static_cast(descriptor->GetPipelineState().get()); + const RayTracingPipelineState* rayTracingPipelineState = static_cast(m_descriptor->GetPipelineState().get()); Microsoft::WRL::ComPtr stateObjectProperties; [[maybe_unused]] HRESULT hr = rayTracingPipelineState->Get()->QueryInterface(IID_GRAPHICS_PPV_ARGS(stateObjectProperties.GetAddressOf())); @@ -155,28 +153,28 @@ namespace AZ // ray generation shader table { // RayGeneration table must have one and only one record - AZ_Assert(descriptor->GetRayGenerationRecord().size() == 1, "Descriptor must contain one and only one RayGeneration record"); - uint32_t shaderRecordSize = RHI::AlignUp(FindLargestRecordSize(descriptor->GetRayGenerationRecord()), D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT); + AZ_Assert(m_descriptor->GetRayGenerationRecord().size() == 1, "Descriptor must contain one and only one RayGeneration record"); + uint32_t shaderRecordSize = RHI::AlignUp(FindLargestRecordSize(m_descriptor->GetRayGenerationRecord()), D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT); - buffers.m_rayGenerationTable = BuildTable(GetDevice(), *m_bufferPools, descriptor->GetRayGenerationRecord(), shaderRecordSize, L"Ray Generation Shader Table", stateObjectProperties); + buffers.m_rayGenerationTable = BuildTable(GetDevice(), *m_bufferPools, m_descriptor->GetRayGenerationRecord(), shaderRecordSize, L"Ray Generation Shader Table", stateObjectProperties); buffers.m_rayGenerationTableSize = shaderRecordSize; } // miss shader table { - uint32_t shaderRecordSize = RHI::AlignUp(FindLargestRecordSize(descriptor->GetMissRecords()), D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT); + uint32_t shaderRecordSize = RHI::AlignUp(FindLargestRecordSize(m_descriptor->GetMissRecords()), D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT); - buffers.m_missTable = BuildTable(GetDevice(), *m_bufferPools, descriptor->GetMissRecords(), shaderRecordSize, L"Miss Shader Table", stateObjectProperties); - buffers.m_missTableSize = shaderRecordSize * static_cast(descriptor->GetMissRecords().size()); + buffers.m_missTable = BuildTable(GetDevice(), *m_bufferPools, m_descriptor->GetMissRecords(), shaderRecordSize, L"Miss Shader Table", stateObjectProperties); + buffers.m_missTableSize = shaderRecordSize * static_cast(m_descriptor->GetMissRecords().size()); buffers.m_missTableStride = shaderRecordSize; } // hit group shader table { - uint32_t shaderRecordSize = RHI::AlignUp(FindLargestRecordSize(descriptor->GetHitGroupRecords()), D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT); + uint32_t shaderRecordSize = RHI::AlignUp(FindLargestRecordSize(m_descriptor->GetHitGroupRecords()), D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT); - buffers.m_hitGroupTable = BuildTable(GetDevice(), *m_bufferPools, descriptor->GetHitGroupRecords(), shaderRecordSize, L"HitGroup Shader Table", stateObjectProperties); - buffers.m_hitGroupTableSize = shaderRecordSize * static_cast(descriptor->GetHitGroupRecords().size()); + buffers.m_hitGroupTable = BuildTable(GetDevice(), *m_bufferPools, m_descriptor->GetHitGroupRecords(), shaderRecordSize, L"HitGroup Shader Table", stateObjectProperties); + buffers.m_hitGroupTableSize = shaderRecordSize * static_cast(m_descriptor->GetHitGroupRecords().size()); buffers.m_hitGroupTableStride = shaderRecordSize; } #endif diff --git a/Gems/Atom/RHI/Vulkan/Code/Source/RHI/RayTracingShaderTable.cpp b/Gems/Atom/RHI/Vulkan/Code/Source/RHI/RayTracingShaderTable.cpp index c76cd9573a..d6d6f3b3c0 100644 --- a/Gems/Atom/RHI/Vulkan/Code/Source/RHI/RayTracingShaderTable.cpp +++ b/Gems/Atom/RHI/Vulkan/Code/Source/RHI/RayTracingShaderTable.cpp @@ -83,14 +83,12 @@ namespace AZ uint32_t shaderHandleSize = rayTracingPipelineProperties.shaderGroupHandleSize; uint32_t alignedShaderHandleSize = RHI::AlignUp(shaderHandleSize, rayTracingPipelineProperties.shaderGroupBaseAlignment); - AZStd::shared_ptr descriptor = m_descriptor.lock(); - // advance to the next buffer m_currentBufferIndex = (m_currentBufferIndex + 1) % BufferCount; ShaderTableBuffers& buffers = m_buffers[m_currentBufferIndex]; // clear the shader table if the descriptor has no ray generation shader - if (descriptor->GetRayGenerationRecord().empty()) + if (m_descriptor->GetRayGenerationRecord().empty()) { buffers.m_rayGenerationTable = nullptr; buffers.m_rayGenerationTableStride = 0; @@ -110,18 +108,18 @@ namespace AZ buffers.m_hitGroupTableStride = RHI::AlignUp(alignedShaderHandleSize, rayTracingPipelineProperties.shaderGroupBaseAlignment); // calculate sub-table sizes - buffers.m_rayGenerationTableSize = buffers.m_rayGenerationTableStride * aznumeric_cast(descriptor->GetRayGenerationRecord().size()); - buffers.m_missTableSize = buffers.m_missTableStride * aznumeric_cast(descriptor->GetMissRecords().size()); - buffers.m_hitGroupTableSize = buffers.m_hitGroupTableStride * aznumeric_cast(descriptor->GetHitGroupRecords().size()); + buffers.m_rayGenerationTableSize = buffers.m_rayGenerationTableStride * aznumeric_cast(m_descriptor->GetRayGenerationRecord().size()); + buffers.m_missTableSize = buffers.m_missTableStride * aznumeric_cast(m_descriptor->GetMissRecords().size()); + buffers.m_hitGroupTableSize = buffers.m_hitGroupTableStride * aznumeric_cast(m_descriptor->GetHitGroupRecords().size()); - const RayTracingPipelineState* rayTracingPipelineState = static_cast(descriptor->GetPipelineState().get()); + const RayTracingPipelineState* rayTracingPipelineState = static_cast(m_descriptor->GetPipelineState().get()); // build sub-tables buffers.m_rayGenerationTable = BuildTable( rayTracingPipelineProperties, rayTracingPipelineState, *m_bufferPools, - descriptor->GetRayGenerationRecord(), + m_descriptor->GetRayGenerationRecord(), buffers.m_rayGenerationTableStride, "RayGenerationTable"); @@ -129,7 +127,7 @@ namespace AZ rayTracingPipelineProperties, rayTracingPipelineState, *m_bufferPools, - descriptor->GetMissRecords(), + m_descriptor->GetMissRecords(), buffers.m_missTableStride, "MissTable"); @@ -137,7 +135,7 @@ namespace AZ rayTracingPipelineProperties, rayTracingPipelineState, *m_bufferPools, - descriptor->GetHitGroupRecords(), + m_descriptor->GetHitGroupRecords(), buffers.m_hitGroupTableStride, "HitGroupTable");