From 9c6d21cc6b790433a1d397b29d212ad4b99904da Mon Sep 17 00:00:00 2001 From: baldurk Date: Thu, 21 Nov 2024 17:19:23 +0000 Subject: [PATCH] Add CPU-side verification of ray dispatch patching --- renderdoc/driver/d3d12/d3d12_device.cpp | 1 + renderdoc/driver/d3d12/d3d12_manager.cpp | 225 ++++++++++++++++++++++- renderdoc/driver/d3d12/d3d12_manager.h | 5 + 3 files changed, 230 insertions(+), 1 deletion(-) diff --git a/renderdoc/driver/d3d12/d3d12_device.cpp b/renderdoc/driver/d3d12/d3d12_device.cpp index b7014cbac..b10fd97ff 100644 --- a/renderdoc/driver/d3d12/d3d12_device.cpp +++ b/renderdoc/driver/d3d12/d3d12_device.cpp @@ -5321,6 +5321,7 @@ void WrappedID3D12Device::ReplayLog(uint32_t startEventID, uint32_t endEventID, for(PatchedRayDispatch &r : cmd.m_RayDispatches) { + GetResourceManager()->GetRTManager()->Verify(r); r.resources.Release(); } cmd.m_RayDispatches.clear(); diff --git a/renderdoc/driver/d3d12/d3d12_manager.cpp b/renderdoc/driver/d3d12/d3d12_manager.cpp index 6419d6119..d5c1b2bcf 100644 --- a/renderdoc/driver/d3d12/d3d12_manager.cpp +++ b/renderdoc/driver/d3d12/d3d12_manager.cpp @@ -49,6 +49,7 @@ RDOC_CONFIG( RDOC_CONFIG(uint32_t, D3D12_Debug_RTASCacheThreshold, 5000, "How many milliseconds to wait before caching an AS to disk if it has been unmodified " "for that long"); +RDOC_CONFIG(bool, D3D12_Debug_RTAuditing, false, "Audit RT work during capture and replay."); // batch 50 at a time, if we have one check per frame this would cache 5000 BLASs in 100 frames // which is a reasonable background pace @@ -809,6 +810,179 @@ void D3D12RTManager::ResizeSerialisationBuffer(UINT64 size) } } +void D3D12RTManager::Verify(PatchedRayDispatch &r) +{ + if(!r.resources.readbackBuffer) + return; + + byte *data = (byte *)r.resources.readbackBuffer->Map(); + + uint32_t patchDataSize = 0; + const uint32_t raygenOffs = patchDataSize; + patchDataSize = (uint32_t)r.desc.RayGenerationShaderRecord.SizeInBytes; + patchDataSize = AlignUp(patchDataSize, (uint32_t)D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT); + + const uint32_t missOffs = patchDataSize; + patchDataSize += (uint32_t)r.desc.MissShaderTable.SizeInBytes; + patchDataSize = AlignUp(patchDataSize, (uint32_t)D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT); + + const uint32_t hitOffs = patchDataSize; + patchDataSize += (uint32_t)r.desc.HitGroupTable.SizeInBytes; + patchDataSize = AlignUp(patchDataSize, (uint32_t)D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT); + + const uint32_t callOffs = patchDataSize; + patchDataSize += (uint32_t)r.desc.CallableShaderTable.SizeInBytes; + + WrappedID3D12DescriptorHeap *sampHeap = NULL, *resHeap = NULL; + for(ResourceId heapId : r.heaps) + { + WrappedID3D12DescriptorHeap *heap = + (WrappedID3D12DescriptorHeap *)m_wrappedDevice->GetResourceManager() + ->GetCurrentAs(heapId); + + if(heap->GetDescriptors()->GetType() == D3D12DescriptorType::Sampler) + sampHeap = heap; + else + resHeap = heap; + } + + if(r.desc.RayGenerationShaderRecord.StartAddress) + { + VerifyRecord(r.desc.RayGenerationShaderRecord.SizeInBytes, data + raygenOffs, + data + r.resources.patchScratchBuffer->Size() + raygenOffs, resHeap, sampHeap); + } + + if(r.desc.MissShaderTable.StartAddress) + { + if(r.desc.MissShaderTable.StrideInBytes == 0) + r.desc.MissShaderTable.StrideInBytes = r.desc.MissShaderTable.SizeInBytes; + for(UINT64 i = 0; i < r.desc.MissShaderTable.SizeInBytes / r.desc.MissShaderTable.StrideInBytes; + i++) + VerifyRecord(r.desc.MissShaderTable.StrideInBytes, + data + missOffs + r.desc.MissShaderTable.StrideInBytes * i, + data + r.resources.patchScratchBuffer->Size() + missOffs + + r.desc.MissShaderTable.StrideInBytes * i, + resHeap, sampHeap); + } + + if(r.desc.HitGroupTable.StartAddress) + { + if(r.desc.HitGroupTable.StrideInBytes == 0) + r.desc.HitGroupTable.StrideInBytes = r.desc.HitGroupTable.SizeInBytes; + for(UINT64 i = 0; i < r.desc.HitGroupTable.SizeInBytes / r.desc.HitGroupTable.StrideInBytes; i++) + VerifyRecord(r.desc.HitGroupTable.StrideInBytes, + data + hitOffs + r.desc.HitGroupTable.StrideInBytes * i, + data + r.resources.patchScratchBuffer->Size() + hitOffs + + r.desc.HitGroupTable.StrideInBytes * i, + resHeap, sampHeap); + } + + r.resources.readbackBuffer->Unmap(); +} + +void D3D12RTManager::VerifyRecord(const uint64_t recordSize, byte *wrappedRecord, + byte *unwrappedRef, WrappedID3D12DescriptorHeap *resHeap, + WrappedID3D12DescriptorHeap *sampHeap) +{ + bytebuf record; + + record.resize(recordSize); + memcpy(record.data(), wrappedRecord, recordSize); + + struct ShaderIdentifier + { + ResourceId id; // the object which has the actual identifier in its ownExports array + uint32_t index; // the index in the object's ownExports array + uint32_t pad[5]; // padding up to D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES + }; + + ShaderIdentifier *ident = (ShaderIdentifier *)wrappedRecord; + + WrappedID3D12StateObject *obj = + m_wrappedDevice->GetResourceManager()->GetLiveAs(ident->id); + + uint16_t localIdx = 0xffff; + + if(obj) + { + memcpy(record.data(), obj->exports->ownExports[ident->index].real, + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES); + localIdx = obj->exports->ownExports[ident->index].localRootSigIndex; + } + else + { + memset(record.data(), 0, D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES); + } + + if(localIdx != 0xffff) + { + rdcarray &rootConfig = m_UniqueLocalRootSigs[localIdx]; + + for(uint32_t offs : rootConfig) + { + bool isVA = (offs & 0x80000000U) != 0; + offs &= ~0x80000000U; + + if(isVA) + { + D3D12_GPU_VIRTUAL_ADDRESS *va = (uint64_t *)(wrappedRecord + offs); + + ResourceId id; + uint64_t resoffs; + m_wrappedDevice->GetResIDFromOrigAddr(*va, id, resoffs); + + D3D12_GPU_VIRTUAL_ADDRESS unwrappedVA = 0; + if(*va == 0) + { + } + else if(id == ResourceId()) + { + RDCWARN("Invalid VA %llx, setting to 0", *va); + } + else + { + ID3D12Resource *res = m_wrappedDevice->GetResourceManager()->GetLiveAs(id); + + unwrappedVA = res->GetGPUVirtualAddress() + resoffs; + } + + memcpy(record.data() + offs, &unwrappedVA, sizeof(unwrappedVA)); + } + else + { + uint64_t wrappedHandle = *(uint64_t *)(wrappedRecord + offs); + + D3D12_GPU_DESCRIPTOR_HANDLE unwrappedHandle = {}; + + if(resHeap && wrappedHandle >= resHeap->GetOriginalGPUBase() && + wrappedHandle < resHeap->GetOriginalGPUBase() + + resHeap->GetNumDescriptors() * sizeof(D3D12Descriptor)) + { + uint32_t idx = + uint32_t((wrappedHandle - resHeap->GetOriginalGPUBase()) / sizeof(D3D12Descriptor)); + unwrappedHandle = resHeap->GetGPU(idx); + } + else if(sampHeap && wrappedHandle >= sampHeap->GetOriginalGPUBase() && + wrappedHandle < +sampHeap->GetOriginalGPUBase() + + sampHeap->GetNumDescriptors() * sizeof(D3D12Descriptor)) + { + uint32_t idx = + uint32_t((wrappedHandle - sampHeap->GetOriginalGPUBase()) / sizeof(D3D12Descriptor)); + unwrappedHandle = sampHeap->GetGPU(idx); + } + else + { + RDCWARN("Invalid descriptor, setting to 0"); + } + + memcpy(record.data() + offs, &unwrappedHandle, sizeof(unwrappedHandle)); + } + } + } + + RDCASSERT(memcmp(record.data(), unwrappedRef, record.size()) == 0); +} + void D3D12RTManager::AddPendingASBuilds(ID3D12Fence *fence, UINT64 waitValue, const rdcarray> &callbacks) { @@ -1039,6 +1213,14 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * ret.resources.readbackBuffer = NULL; + if(IsReplayMode(m_wrappedDevice->GetState()) && D3D12_Debug_RTAuditing()) + { + m_GPUBufferAllocator.Alloc(D3D12GpuBufferHeapType::ReadBackHeap, + D3D12GpuBufferHeapMemoryFlag::Default, patchDataSize * 2, + D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, + &ret.resources.readbackBuffer); + } + RDCCOMPILE_ASSERT(WRAPPED_DESCRIPTOR_STRIDE == sizeof(D3D12Descriptor), "Shader descriptor stride is wrong"); @@ -1094,6 +1276,15 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * sizeof(recordInfo) / sizeof(uint32_t), &recordInfo, 0); unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::SourceBuffer, ret.desc.RayGenerationShaderRecord.StartAddress); + + if(ret.resources.readbackBuffer) + { + CopyFromVA(unwrappedCmd, ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + raygenOffs, + ret.desc.RayGenerationShaderRecord.StartAddress, + desc.RayGenerationShaderRecord.SizeInBytes); + } + ret.desc.RayGenerationShaderRecord.StartAddress = scratchBuffer->Address() + raygenOffs; unwrappedCmd->SetComputeRootUnorderedAccessView((UINT)D3D12PatchRayDispatchParam::DestBuffer, ret.desc.RayGenerationShaderRecord.StartAddress); @@ -1111,6 +1302,14 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * sizeof(recordInfo) / sizeof(uint32_t), &recordInfo, 0); unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::SourceBuffer, ret.desc.MissShaderTable.StartAddress); + + if(ret.resources.readbackBuffer) + { + CopyFromVA(unwrappedCmd, ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + missOffs, + ret.desc.MissShaderTable.StartAddress, desc.MissShaderTable.SizeInBytes); + } + ret.desc.MissShaderTable.StartAddress = scratchBuffer->Address() + missOffs; unwrappedCmd->SetComputeRootUnorderedAccessView((UINT)D3D12PatchRayDispatchParam::DestBuffer, ret.desc.MissShaderTable.StartAddress); @@ -1130,6 +1329,14 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * sizeof(recordInfo) / sizeof(uint32_t), &recordInfo, 0); unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::SourceBuffer, ret.desc.HitGroupTable.StartAddress); + + if(ret.resources.readbackBuffer) + { + CopyFromVA(unwrappedCmd, ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + hitOffs, + ret.desc.HitGroupTable.StartAddress, desc.HitGroupTable.SizeInBytes); + } + ret.desc.HitGroupTable.StartAddress = scratchBuffer->Address() + hitOffs; unwrappedCmd->SetComputeRootUnorderedAccessView((UINT)D3D12PatchRayDispatchParam::DestBuffer, ret.desc.HitGroupTable.StartAddress); @@ -1150,6 +1357,14 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * sizeof(recordInfo) / sizeof(uint32_t), &recordInfo, 0); unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::SourceBuffer, ret.desc.CallableShaderTable.StartAddress); + + if(ret.resources.readbackBuffer) + { + CopyFromVA(unwrappedCmd, ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + callOffs, + ret.desc.CallableShaderTable.StartAddress, desc.CallableShaderTable.SizeInBytes); + } + ret.desc.CallableShaderTable.StartAddress = scratchBuffer->Address() + callOffs; unwrappedCmd->SetComputeRootUnorderedAccessView((UINT)D3D12PatchRayDispatchParam::DestBuffer, ret.desc.CallableShaderTable.StartAddress); @@ -1162,9 +1377,17 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * barrier.Type = D3D12_RESOURCE_BARRIER_TYPE_TRANSITION; barrier.Transition.pResource = scratchBuffer->Resource(); barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_UNORDERED_ACCESS; - barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_NON_PIXEL_SHADER_RESOURCE; + barrier.Transition.StateAfter = + D3D12_RESOURCE_STATE_NON_PIXEL_SHADER_RESOURCE | D3D12_RESOURCE_STATE_COPY_SOURCE; unwrappedCmd->ResourceBarrier(1, &barrier); + if(ret.resources.readbackBuffer) + { + unwrappedCmd->CopyBufferRegion(ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + patchDataSize, + scratchBuffer->Resource(), scratchBuffer->Offset(), patchDataSize); + } + // we have our own ref, the patch data has its ref too that will be held while the list is // submittable. Each submission will also get a ref to keep this referenced lookup buffer alive until then m_LookupBuffer->AddRef(); diff --git a/renderdoc/driver/d3d12/d3d12_manager.h b/renderdoc/driver/d3d12/d3d12_manager.h index 07c766812..9d84f2db0 100644 --- a/renderdoc/driver/d3d12/d3d12_manager.h +++ b/renderdoc/driver/d3d12/d3d12_manager.h @@ -1268,6 +1268,11 @@ public: double GetCurrentASTimestamp() { return m_Timestamp.GetMilliseconds(); } + void Verify(PatchedRayDispatch &r); + + void VerifyRecord(const uint64_t recordSize, byte *table, byte *ref, + WrappedID3D12DescriptorHeap *resHeap, WrappedID3D12DescriptorHeap *sampHeap); + private: void InitRayDispatchPatchingResources(); void InitTLASInstanceCopyingResources();