| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269 |
- #include "DX12TLAS.h"
- #include "DX12CommandQueue.h"
- #include "Logging.h"
- Framework::DX12TLAS::DX12TLAS(
- ID3D12Device5* zDevice, Framework::DX12DirectCommandQueue* zDirectQueue)
- : ReferenceCounter(),
- scratchBuffer(0),
- resultBuffer(0),
- descriptorBuffer(0),
- zDevice(zDevice),
- zDirectQueue(zDirectQueue),
- overflowInstanceIterator(0, 0, 0, 0),
- currentInstanceIndex(0),
- lastInstanceCount(0),
- mappedDescriptorBuffer(0),
- lastResultBuffer(0),
- lastScratchBuffer(0),
- bufferChanged(0)
- {}
- Framework::DX12TLAS::~DX12TLAS()
- {
- if (scratchBuffer)
- {
- scratchBuffer->release();
- }
- if (resultBuffer)
- {
- resultBuffer->release();
- }
- if (descriptorBuffer)
- {
- descriptorBuffer->release();
- }
- if (lastResultBuffer)
- {
- lastResultBuffer->Release();
- }
- if (lastScratchBuffer)
- {
- lastScratchBuffer->Release();
- }
- for (const D3D12_RAYTRACING_INSTANCE_DESC* desc : overflowInstanceDescs)
- {
- delete desc;
- }
- }
- void Framework::DX12TLAS::startUpdate()
- {
- currentInstanceIndex = -1;
- if (descriptorBuffer)
- {
- descriptorBuffer->zBuffer()->Map(0, 0, (void**)&mappedDescriptorBuffer);
- }
- else
- {
- mappedDescriptorBuffer = 0;
- }
- overflowInstanceIterator = overflowInstanceDescs.begin();
- }
- D3D12_RAYTRACING_INSTANCE_DESC* Framework::DX12TLAS::nextInstanceDesc()
- {
- currentInstanceIndex++;
- if (mappedDescriptorBuffer
- && currentInstanceIndex < descriptorBuffer->getElementCount())
- {
- return mappedDescriptorBuffer + currentInstanceIndex;
- }
- else if (overflowInstanceIterator)
- {
- return *overflowInstanceIterator++;
- }
- else
- {
- D3D12_RAYTRACING_INSTANCE_DESC* desc
- = new D3D12_RAYTRACING_INSTANCE_DESC();
- overflowInstanceDescs.add(desc);
- return desc;
- }
- }
- void Framework::DX12TLAS::endUpdate()
- {
- if (currentInstanceIndex < 0) // at least one instance ust be added
- {
- D3D12_RAYTRACING_INSTANCE_DESC* desc = nextInstanceDesc();
- desc->InstanceContributionToHitGroupIndex = 0;
- desc->InstanceID = 0;
- desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
- desc->AccelerationStructure = 0;
- desc->InstanceMask = 0; // Mark unused instances with a mask of 0
- }
- if (descriptorBuffer)
- {
- while (currentInstanceIndex < lastInstanceCount - 1)
- {
- D3D12_RAYTRACING_INSTANCE_DESC* desc = nextInstanceDesc();
- desc->InstanceMask = 0; // Mark unused instances with a mask of 0
- desc->InstanceContributionToHitGroupIndex = 0;
- }
- descriptorBuffer->zBuffer()->Unmap(0, nullptr);
- }
- mappedDescriptorBuffer = 0;
- if (!descriptorBuffer || currentInstanceIndex >= lastInstanceCount)
- {
- // recalculate the size of the descriptor buffer and reallocate it
- D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_INPUTS
- prebuildDesc = {};
- prebuildDesc.Type
- = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL;
- prebuildDesc.DescsLayout = D3D12_ELEMENTS_LAYOUT_ARRAY;
- prebuildDesc.NumDescs = (unsigned)currentInstanceIndex + 1;
- prebuildDesc.Flags
- = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_ALLOW_UPDATE;
- D3D12_RAYTRACING_ACCELERATION_STRUCTURE_PREBUILD_INFO info = {};
- zDevice->GetRaytracingAccelerationStructurePrebuildInfo(
- &prebuildDesc, &info);
- // create new buffers witch fit the TLAS
- if (scratchBuffer)
- {
- if (lastScratchBuffer)
- {
- lastScratchBuffer->Release();
- }
- lastScratchBuffer = scratchBuffer->zBuffer();
- lastScratchBuffer->AddRef();
- }
- if (!resultBuffer)
- {
- resultBuffer = new DX12Buffer(1,
- zDevice,
- dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
- D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS);
- }
- resultBuffer->setLength(
- ROUND_UP_POWER_OF_2(info.ResultDataMaxSizeInBytes, 256));
- resultBuffer->createBufferWithoutData(
- D3D12_RESOURCE_STATE_RAYTRACING_ACCELERATION_STRUCTURE);
- resultBuffer->zBuffer()->SetName(L"TLAS Result Buffer");
- if (!scratchBuffer)
- {
- scratchBuffer = new DX12Buffer(1,
- zDevice,
- dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
- D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS);
- }
- scratchBuffer->setLength(
- ROUND_UP_POWER_OF_2(info.ScratchDataSizeInBytes, 256));
- scratchBuffer->createBufferWithoutData(
- D3D12_RESOURCE_STATE_UNORDERED_ACCESS);
- scratchBuffer->zBuffer()->SetName(L"TLAS Scratch Buffer");
- ID3D12Resource* oldDescriptorBuffer
- = descriptorBuffer ? descriptorBuffer->zBuffer() : 0;
- __int64 oldDescriptorBufferElementCount
- = descriptorBuffer ? descriptorBuffer->getElementCount() : 0;
- if (oldDescriptorBuffer)
- {
- oldDescriptorBuffer->AddRef();
- oldDescriptorBuffer->Map(0, 0, (void**)&mappedDescriptorBuffer);
- }
- if (!descriptorBuffer)
- {
- descriptorBuffer
- = new DX12Buffer(sizeof(D3D12_RAYTRACING_INSTANCE_DESC),
- zDevice,
- dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
- D3D12_RESOURCE_FLAG_NONE);
- }
- descriptorBuffer->setLength(ROUND_UP_POWER_OF_2(
- (currentInstanceIndex + 1) * sizeof(D3D12_RAYTRACING_INSTANCE_DESC),
- 256));
- descriptorBuffer->createBufferWithoutData(
- D3D12_RESOURCE_STATE_GENERIC_READ, D3D12_HEAP_TYPE_UPLOAD);
- D3D12_RAYTRACING_INSTANCE_DESC* newMappedDescriptorBuffer = 0;
- // copy the old and new instance descriptions to the new buffer
- descriptorBuffer->zBuffer()->Map(
- 0, 0, (void**)&newMappedDescriptorBuffer);
- if (mappedDescriptorBuffer)
- {
- memcpy(newMappedDescriptorBuffer,
- mappedDescriptorBuffer,
- oldDescriptorBufferElementCount
- * sizeof(D3D12_RAYTRACING_INSTANCE_DESC));
- D3D12_RANGE range = {0, 0}; // do not write to the old buffer
- oldDescriptorBuffer->Unmap(0, &range);
- oldDescriptorBuffer->Release();
- }
- overflowInstanceIterator = overflowInstanceDescs.begin();
- for (__int64 i = oldDescriptorBufferElementCount;
- i <= currentInstanceIndex;
- i++)
- {
- memcpy(newMappedDescriptorBuffer + i,
- overflowInstanceIterator.val(),
- sizeof(D3D12_RAYTRACING_INSTANCE_DESC));
- overflowInstanceIterator++;
- }
- descriptorBuffer->zBuffer()->Unmap(0, nullptr);
- bufferChanged = 1;
- }
- else
- {
- if (resultBuffer)
- {
- if (lastResultBuffer)
- {
- lastResultBuffer->Release();
- }
- lastResultBuffer = resultBuffer->zBuffer();
- lastResultBuffer->AddRef();
- }
- }
- D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_DESC buildDesc = {};
- buildDesc.Inputs.Type
- = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL;
- buildDesc.Inputs.DescsLayout = D3D12_ELEMENTS_LAYOUT_ARRAY;
- buildDesc.Inputs.InstanceDescs
- = descriptorBuffer->zBuffer()->GetGPUVirtualAddress();
- buildDesc.Inputs.NumDescs = (unsigned)currentInstanceIndex + 1;
- buildDesc.DestAccelerationStructureData
- = {resultBuffer->zBuffer()->GetGPUVirtualAddress()};
- buildDesc.ScratchAccelerationStructureData
- = {scratchBuffer->zBuffer()->GetGPUVirtualAddress()};
- buildDesc.SourceAccelerationStructureData
- = lastResultBuffer && currentInstanceIndex < lastInstanceCount
- ? lastResultBuffer->GetGPUVirtualAddress()
- : 0;
- buildDesc.Inputs.Flags
- = lastResultBuffer && currentInstanceIndex < lastInstanceCount
- ? D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_PERFORM_UPDATE
- : D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_ALLOW_UPDATE;
- // Build the top-level AS
- zDirectQueue->zCommandList()->BuildRaytracingAccelerationStructure(
- &buildDesc, 0, nullptr);
- // Wait for the builder to complete by setting a barrier on the resulting
- // buffer. This can be important in case the rendering is triggered
- // immediately afterwards, without executing the command list
- D3D12_RESOURCE_BARRIER uavBarrier;
- uavBarrier.Type = D3D12_RESOURCE_BARRIER_TYPE_UAV;
- uavBarrier.UAV.pResource = resultBuffer->zBuffer();
- uavBarrier.Flags = D3D12_RESOURCE_BARRIER_FLAG_NONE;
- zDirectQueue->zCommandList()->ResourceBarrier(1, &uavBarrier);
- lastInstanceCount = currentInstanceIndex + 1;
- }
- Framework::DX12Buffer* Framework::DX12TLAS::zResultBuffer() const
- {
- return resultBuffer;
- }
- bool Framework::DX12TLAS::hasBufferChanged() const
- {
- return bufferChanged;
- }
- void Framework::DX12TLAS::setBufferChanged(bool changed)
- {
- bufferChanged = changed;
- }
|