Changed the descriptor a shared_ptr.
This commit is contained in:
+3
-3
@@ -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");
|
||||
|
||||
|
||||
Reference in New Issue
Block a user