#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(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(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(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; }