Changed the descriptor a shared_ptr.

This commit is contained in:
dmcdiar
2021-05-14 19:10:25 -07:00
parent 3371483315
commit 73f438e185
7 changed files with 27 additions and 32 deletions
@@ -275,12 +275,12 @@ namespace AZ
// scene changed, need to rebuild the shader table
m_rayTracingRevision = rayTracingRevision;
m_rayTracingShaderTableDescriptor = AZStd::make_shared<RHI::RayTracingShaderTableDescriptor>();
AZStd::shared_ptr<RHI::RayTracingShaderTableDescriptor> descriptor = AZStd::make_shared<RHI::RayTracingShaderTableDescriptor>();
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);
}
}
@@ -62,7 +62,6 @@ namespace AZ
// ray tracing shader table
RHI::Ptr<RHI::RayTracingShaderTable> m_rayTracingShaderTable;
AZStd::shared_ptr<RHI::RayTracingShaderTableDescriptor> m_rayTracingShaderTableDescriptor;
// ray tracing global shader resource group asset and pipeline state
Data::Asset<RPI::ShaderResourceGroupAsset> m_globalSrgAsset;
@@ -18,6 +18,7 @@
#include <Atom/RHI/FrameGraphExecuter.h>
#include <Atom/RHI/FrameGraphCompiler.h>
#include <Atom/RHI/FrameGraph.h>
#include <Atom/RHI/RayTracingShaderTable.h>
#include <Atom/RHI/ScopeProducer.h>
#include <Atom/RHI/ScopeProducerEmpty.h>
#include <Atom/RHI/TransientAttachmentPool.h>
@@ -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<ScopeId, ScopeProducer*> m_scopeProducerLookup;
// list of RayTracingShaderTables that should be built this frame
AZStd::vector<RayTracingShaderTable*> m_rayTracingShaderTablesToBuild;
AZStd::vector<RHI::Ptr<RayTracingShaderTable>> m_rayTracingShaderTablesToBuild;
};
}
}
@@ -99,12 +99,13 @@ namespace AZ
static RHI::Ptr<RHI::RayTracingShaderTable> 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<RayTracingShaderTableDescriptor> descriptor);
protected:
AZStd::weak_ptr<RayTracingShaderTableDescriptor> m_descriptor;
AZStd::shared_ptr<RayTracingShaderTableDescriptor> m_descriptor;
const RayTracingBufferPools* m_bufferPools = nullptr;
private:
@@ -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.");
}
}
@@ -128,10 +128,8 @@ namespace AZ
m_currentBufferIndex = (m_currentBufferIndex + 1) % BufferCount;
ShaderTableBuffers& buffers = m_buffers[m_currentBufferIndex];
AZStd::shared_ptr<RHI::RayTracingShaderTableDescriptor> 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<const RayTracingPipelineState*>(descriptor->GetPipelineState().get());
const RayTracingPipelineState* rayTracingPipelineState = static_cast<const RayTracingPipelineState*>(m_descriptor->GetPipelineState().get());
Microsoft::WRL::ComPtr<ID3D12StateObjectProperties> 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<uint32_t>(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<uint32_t>(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<uint32_t>(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<uint32_t>(m_descriptor->GetHitGroupRecords().size());
buffers.m_hitGroupTableStride = shaderRecordSize;
}
#endif
@@ -83,14 +83,12 @@ namespace AZ
uint32_t shaderHandleSize = rayTracingPipelineProperties.shaderGroupHandleSize;
uint32_t alignedShaderHandleSize = RHI::AlignUp(shaderHandleSize, rayTracingPipelineProperties.shaderGroupBaseAlignment);
AZStd::shared_ptr<RHI::RayTracingShaderTableDescriptor> 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<uint32_t>(descriptor->GetRayGenerationRecord().size());
buffers.m_missTableSize = buffers.m_missTableStride * aznumeric_cast<uint32_t>(descriptor->GetMissRecords().size());
buffers.m_hitGroupTableSize = buffers.m_hitGroupTableStride * aznumeric_cast<uint32_t>(descriptor->GetHitGroupRecords().size());
buffers.m_rayGenerationTableSize = buffers.m_rayGenerationTableStride * aznumeric_cast<uint32_t>(m_descriptor->GetRayGenerationRecord().size());
buffers.m_missTableSize = buffers.m_missTableStride * aznumeric_cast<uint32_t>(m_descriptor->GetMissRecords().size());
buffers.m_hitGroupTableSize = buffers.m_hitGroupTableStride * aznumeric_cast<uint32_t>(m_descriptor->GetHitGroupRecords().size());
const RayTracingPipelineState* rayTracingPipelineState = static_cast<const RayTracingPipelineState*>(descriptor->GetPipelineState().get());
const RayTracingPipelineState* rayTracingPipelineState = static_cast<const RayTracingPipelineState*>(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");