From 70c6dd32a226f8d97d5bc4c5cd312a35232e8c01 Mon Sep 17 00:00:00 2001 From: baldurk Date: Thu, 5 Dec 2024 13:26:03 +0000 Subject: [PATCH] Implement auditing of indirect RT dispatches --- renderdoc/data/hlsl/hlsl_cbuffers.h | 3 +- renderdoc/data/hlsl/raytracing.hlsl | 26 +- renderdoc/driver/d3d12/d3d12_manager.cpp | 544 +++++++++++++++++++---- renderdoc/driver/d3d12/d3d12_manager.h | 15 +- 4 files changed, 498 insertions(+), 90 deletions(-) diff --git a/renderdoc/data/hlsl/hlsl_cbuffers.h b/renderdoc/data/hlsl/hlsl_cbuffers.h index 6491f4bdd..473c80b6d 100644 --- a/renderdoc/data/hlsl/hlsl_cbuffers.h +++ b/renderdoc/data/hlsl/hlsl_cbuffers.h @@ -338,7 +338,8 @@ struct LocalRootSigData cbuffer RayIndirectDispatchCB REG(b0) { - GPUAddress scratchBuffer; + GPUAddress destBuffer; + GPUAddress destBufferEnd; uint commandSigDispatchOffset; uint commandSigStride; diff --git a/renderdoc/data/hlsl/raytracing.hlsl b/renderdoc/data/hlsl/raytracing.hlsl index 104352215..1b0aed53c 100644 --- a/renderdoc/data/hlsl/raytracing.hlsl +++ b/renderdoc/data/hlsl/raytracing.hlsl @@ -265,11 +265,24 @@ void PatchTable(uint byteOffset) } // Each SV_GroupId corresponds to one shader record to patch -[numthreads(RECORD_PATCH_THREADS, 1, 1)] void RENDERDOC_PatchRayDispatchCS(uint3 dispatchThread +[numthreads(RECORD_PATCH_THREADS, 1, 1)] void RENDERDOC_PatchShaderTableCS(uint3 dispatchThread : SV_DispatchThreadID) { if(dispatchThread.x < shaderrecord_count) PatchTable(shaderrecord_stride * dispatchThread.x); -} +}; + +// Each SV_GroupId corresponds to one shader record to patch +[numthreads(RECORD_PATCH_THREADS, 1, 1)] void RENDERDOC_CopyShaderTableCS(uint3 dispatchThread + : SV_DispatchThreadID) { + if(dispatchThread.x < shaderrecord_count) + { + for(uint b = 0; b < shaderrecord_stride;) + { + b = CopyData(patchSource, patchDest, shaderrecord_stride * dispatchThread.x, b, + shaderrecord_stride); + } + } +}; // define these structs in hlsl for simplicity @@ -327,7 +340,7 @@ GPUAddress AlignRecordAddress(GPUAddress x) numCommands = min(numCommands, applicationCountBuffer.Load(0)); } - GPUAddress outputBufferLocation = scratchBuffer; + GPUAddress outputBufferLocation = destBuffer; uint dispatchIndex = 0; PatchingExecute execute = (PatchingExecute)0xccddeeff; @@ -443,6 +456,13 @@ GPUAddress AlignRecordAddress(GPUAddress x) patchedExecuteArguments.Store(commandSigDispatchOffset + commandOffset + 6 * 16, raw.b); } + // check for buffer overrun + if(!lessEqual(outputBufferLocation, destBufferEnd)) + { + // error! don't patch, this will fail. Needs to be caught by auditing + dispatchIndex = 0; + } + // store the number of patching indirect dispatches we'll do, up to 4 per the application's number internalExecuteCount.Store(0, dispatchIndex); } diff --git a/renderdoc/driver/d3d12/d3d12_manager.cpp b/renderdoc/driver/d3d12/d3d12_manager.cpp index be6d5e5a4..eabe7385e 100644 --- a/renderdoc/driver/d3d12/d3d12_manager.cpp +++ b/renderdoc/driver/d3d12/d3d12_manager.cpp @@ -701,6 +701,21 @@ D3D12Descriptor *DescriptorFromPortableHandle(D3D12ResourceManager *manager, Por return NULL; } +static uint32_t GuessTablePatchRequirements(UINT MaxCommandCount) +{ + uint32_t patchDataSize = RDCMAX(64U, 20 * 1024 * 1024 * MaxCommandCount); + + if(D3D12_Debug_RT_IndirectEstimateOverride() > 0) + patchDataSize = D3D12_Debug_RT_IndirectEstimateOverride() * MaxCommandCount; + + return patchDataSize; +} + +static UINT64 GetTableCount(D3D12_GPU_VIRTUAL_ADDRESS_RANGE_AND_STRIDE table) +{ + return table.StrideInBytes == 0 ? 1ULL : table.SizeInBytes / table.StrideInBytes; +} + // debugging logging for barriers #if 0 #define BARRIER_DBG RDCLOG @@ -726,8 +741,9 @@ D3D12RTManager::~D3D12RTManager() SAFE_RELEASE(m_TLASCopyingData.CopyPipe); SAFE_RELEASE(m_TLASCopyingData.RootSig); SAFE_RELEASE(m_TLASCopyingData.IndirectSig); - SAFE_RELEASE(m_RayPatchingData.descPatchRootSig); - SAFE_RELEASE(m_RayPatchingData.descPatchPipe); + SAFE_RELEASE(m_RayPatchingData.shaderTablePatchRootSig); + SAFE_RELEASE(m_RayPatchingData.shaderTablePatchPipe); + SAFE_RELEASE(m_RayPatchingData.shaderTableCopyPipe); SAFE_RELEASE(m_RayPatchingData.indirectComSig); SAFE_RELEASE(m_RayPatchingData.indirectPrepPipe); SAFE_RELEASE(m_RayPatchingData.indirectPrepRootSig); @@ -780,24 +796,6 @@ 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) { @@ -811,53 +809,246 @@ void D3D12RTManager::Verify(PatchedRayDispatch &r) resHeap = heap; } - if(r.desc.RayGenerationShaderRecord.StartAddress) + byte *data = (byte *)r.resources.readbackBuffer->Map(); + + // simple case, single non-indirect dispatch + if(r.comSig == NULL) { - VerifyRecord(r.desc.RayGenerationShaderRecord.SizeInBytes, data + raygenOffs, - data + r.resources.patchScratchBuffer->Size() + raygenOffs, resHeap, sampHeap); + VerifyDispatch(r.desc, data, data + r.resources.patchScratchBuffer->Size(), resHeap, sampHeap); + return; } - if(r.desc.MissShaderTable.StartAddress) + // indirect dispatches need extra work (and extra verification!) + + const uint32_t indirectPatchReservationSize = GuessTablePatchRequirements(r.MaxCommands); + const uint64_t applicationArgsSize = AlignUp16(r.MaxCommands * r.comSig->sig.ByteStride); + const uint64_t patchingArgsSize = AlignUp16(r.MaxCommands * 4 * AlignUp16(sizeof(PatchingExecute))); + + byte *wrappedExecArgs = data + indirectPatchReservationSize * 2; + UINT *sourceCount = (UINT *)(wrappedExecArgs + applicationArgsSize); + byte *unwrappedExecArgs = ((byte *)sourceCount) + 4; + PatchingExecute *internalExecs = (PatchingExecute *)(unwrappedExecArgs + applicationArgsSize); + UINT *patchExecCount = (UINT *)(((byte *)internalExecs) + patchingArgsSize); + + UINT cmdCount = *sourceCount; + RDCASSERT(cmdCount < r.MaxCommands, cmdCount, r.MaxCommands); + // if we didn't have a dynamic count this should be 0, revert to the fixed CPU-side count + if(!r.HasDynamicCount) { - 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); + RDCASSERTEQUAL(cmdCount, 0); + cmdCount = r.MaxCommands; } - if(r.desc.HitGroupTable.StartAddress) + byte *wrappedRecords = data; + byte *unwrappedRecords = wrappedRecords + indirectPatchReservationSize; + + UINT internalExecsRequired = 0; + + // this loop emulates/verifies RENDERDOC_PrepareRayIndirectExecuteCS + for(UINT cmd = 0; cmd < cmdCount; cmd++) { - 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); + D3D12_DISPATCH_RAYS_DESC *wrappedDisp = + (D3D12_DISPATCH_RAYS_DESC *)(wrappedExecArgs + cmd * r.comSig->sig.ByteStride + + r.comSig->sig.PackedByteSize - sizeof(D3D12_DISPATCH_RAYS_DESC)); + D3D12_DISPATCH_RAYS_DESC *unwrappedDisp = + (D3D12_DISPATCH_RAYS_DESC *)(unwrappedExecArgs + cmd * r.comSig->sig.ByteStride + + r.comSig->sig.PackedByteSize - sizeof(D3D12_DISPATCH_RAYS_DESC)); + + // sizes should have been propagated trivially + RDCASSERTEQUAL(wrappedDisp->Width, unwrappedDisp->Width); + RDCASSERTEQUAL(wrappedDisp->Height, unwrappedDisp->Height); + RDCASSERTEQUAL(wrappedDisp->Depth, unwrappedDisp->Depth); + + uint64_t offset = 0; + + // raygen is required. Check it w as placed correctly + RDCASSERT(wrappedDisp->RayGenerationShaderRecord.SizeInBytes); + RDCASSERTEQUAL(unwrappedDisp->RayGenerationShaderRecord.StartAddress, + r.resources.patchScratchBuffer->Address() + offset); + RDCASSERTEQUAL(wrappedDisp->RayGenerationShaderRecord.SizeInBytes, + unwrappedDisp->RayGenerationShaderRecord.SizeInBytes); + offset += AlignUp(wrappedDisp->RayGenerationShaderRecord.SizeInBytes, 256ULL); + + // also check the internal indirect patch was properly configured + RDCASSERTEQUAL(internalExecs[internalExecsRequired].sourceData, + wrappedDisp->RayGenerationShaderRecord.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].destData, + unwrappedDisp->RayGenerationShaderRecord.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_count, 1); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_stride, + unwrappedDisp->RayGenerationShaderRecord.SizeInBytes); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.x, 1); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.y, 1); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.z, 1); + + internalExecsRequired++; + + // for each optional table check its placement as well and the corresponding indirect patch + + if(wrappedDisp->MissShaderTable.SizeInBytes) + { + if(unwrappedDisp->HitGroupTable.StrideInBytes == 0) + unwrappedDisp->HitGroupTable.StrideInBytes = unwrappedDisp->HitGroupTable.SizeInBytes; + + RDCASSERTEQUAL(unwrappedDisp->MissShaderTable.StartAddress, + r.resources.patchScratchBuffer->Address() + offset); + RDCASSERTEQUAL(wrappedDisp->MissShaderTable.SizeInBytes, + unwrappedDisp->MissShaderTable.SizeInBytes); + RDCASSERTEQUAL(wrappedDisp->MissShaderTable.StrideInBytes, + unwrappedDisp->MissShaderTable.StrideInBytes); + offset += AlignUp(wrappedDisp->MissShaderTable.SizeInBytes, 256ULL); + + RDCASSERTEQUAL(internalExecs[internalExecsRequired].sourceData, + wrappedDisp->MissShaderTable.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].destData, + unwrappedDisp->MissShaderTable.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_count, + GetTableCount(wrappedDisp->MissShaderTable)); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_stride, + unwrappedDisp->MissShaderTable.StrideInBytes); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.x, + AlignUp(internalExecs[internalExecsRequired].shaderrecord_count, + (uint32_t)RECORD_PATCH_THREADS) / + RECORD_PATCH_THREADS); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.y, 1); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.z, 1); + + internalExecsRequired++; + } + + if(wrappedDisp->HitGroupTable.SizeInBytes) + { + if(wrappedDisp->HitGroupTable.StrideInBytes == 0) + wrappedDisp->HitGroupTable.StrideInBytes = wrappedDisp->HitGroupTable.SizeInBytes; + + RDCASSERTEQUAL(unwrappedDisp->HitGroupTable.StartAddress, + r.resources.patchScratchBuffer->Address() + offset); + RDCASSERTEQUAL(wrappedDisp->HitGroupTable.SizeInBytes, + unwrappedDisp->HitGroupTable.SizeInBytes); + RDCASSERTEQUAL(wrappedDisp->HitGroupTable.StrideInBytes, + unwrappedDisp->HitGroupTable.StrideInBytes); + offset += AlignUp(wrappedDisp->HitGroupTable.SizeInBytes, 256ULL); + + RDCASSERTEQUAL(internalExecs[internalExecsRequired].sourceData, + wrappedDisp->HitGroupTable.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].destData, + unwrappedDisp->HitGroupTable.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_count, + GetTableCount(wrappedDisp->HitGroupTable)); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_stride, + unwrappedDisp->HitGroupTable.StrideInBytes); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.x, + AlignUp(internalExecs[internalExecsRequired].shaderrecord_count, + (uint32_t)RECORD_PATCH_THREADS) / + RECORD_PATCH_THREADS); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.y, 1); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.z, 1); + + internalExecsRequired++; + } + + if(wrappedDisp->CallableShaderTable.SizeInBytes) + { + if(wrappedDisp->CallableShaderTable.StrideInBytes == 0) + wrappedDisp->CallableShaderTable.StrideInBytes = wrappedDisp->CallableShaderTable.SizeInBytes; + + RDCASSERTEQUAL(unwrappedDisp->CallableShaderTable.StartAddress, + r.resources.patchScratchBuffer->Address() + offset); + RDCASSERTEQUAL(wrappedDisp->CallableShaderTable.SizeInBytes, + unwrappedDisp->CallableShaderTable.SizeInBytes); + RDCASSERTEQUAL(wrappedDisp->CallableShaderTable.StrideInBytes, + unwrappedDisp->CallableShaderTable.StrideInBytes); + offset += AlignUp(wrappedDisp->CallableShaderTable.SizeInBytes, 256ULL); + + RDCASSERTEQUAL(internalExecs[internalExecsRequired].sourceData, + wrappedDisp->CallableShaderTable.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].destData, + unwrappedDisp->CallableShaderTable.StartAddress); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_count, + GetTableCount(wrappedDisp->CallableShaderTable)); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].shaderrecord_stride, + unwrappedDisp->CallableShaderTable.StrideInBytes); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.x, + AlignUp(internalExecs[internalExecsRequired].shaderrecord_count, + (uint32_t)RECORD_PATCH_THREADS) / + RECORD_PATCH_THREADS); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.y, 1); + RDCASSERTEQUAL(internalExecs[internalExecsRequired].dispatchDim.z, 1); + + internalExecsRequired++; + } + + RDCASSERTMSG("Indirect reservation size too small", offset < indirectPatchReservationSize, + offset, indirectPatchReservationSize); + + VerifyDispatch(*unwrappedDisp, wrappedRecords, unwrappedRecords, resHeap, sampHeap); + + wrappedRecords += offset; + unwrappedRecords += offset; } - if(r.desc.CallableShaderTable.StartAddress) - { - if(r.desc.CallableShaderTable.StrideInBytes == 0) - r.desc.CallableShaderTable.StrideInBytes = r.desc.CallableShaderTable.SizeInBytes; - for(UINT64 i = 0; - i < r.desc.CallableShaderTable.SizeInBytes / r.desc.CallableShaderTable.StrideInBytes; i++) - VerifyRecord(r.desc.CallableShaderTable.StrideInBytes, - data + callOffs + r.desc.CallableShaderTable.StrideInBytes * i, - data + r.resources.patchScratchBuffer->Size() + callOffs + - r.desc.CallableShaderTable.StrideInBytes * i, - resHeap, sampHeap); - } + RDCASSERTEQUAL(internalExecsRequired, *patchExecCount); r.resources.readbackBuffer->Unmap(); } +void D3D12RTManager::VerifyDispatch(D3D12_DISPATCH_RAYS_DESC desc, byte *wrappedRecords, + byte *unwrappedRecords, WrappedID3D12DescriptorHeap *resHeap, + WrappedID3D12DescriptorHeap *sampHeap) +{ + // raygen is not optional + RDCASSERT(desc.RayGenerationShaderRecord.SizeInBytes >= D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES); + RDCASSERT(desc.RayGenerationShaderRecord.StartAddress); + if(desc.RayGenerationShaderRecord.StartAddress) + { + VerifyRecord(desc.RayGenerationShaderRecord.SizeInBytes, wrappedRecords, unwrappedRecords, + resHeap, sampHeap); + } + + wrappedRecords += AlignUp(desc.RayGenerationShaderRecord.SizeInBytes, 256ULL); + unwrappedRecords += AlignUp(desc.RayGenerationShaderRecord.SizeInBytes, 256ULL); + + if(desc.MissShaderTable.StartAddress) + { + if(desc.MissShaderTable.StrideInBytes == 0) + desc.MissShaderTable.StrideInBytes = desc.MissShaderTable.SizeInBytes; + for(UINT64 i = 0; i < desc.MissShaderTable.SizeInBytes / desc.MissShaderTable.StrideInBytes; i++) + VerifyRecord(desc.MissShaderTable.StrideInBytes, + wrappedRecords + desc.MissShaderTable.StrideInBytes * i, + unwrappedRecords + desc.MissShaderTable.StrideInBytes * i, resHeap, sampHeap); + } + + wrappedRecords += AlignUp(desc.MissShaderTable.SizeInBytes, 256ULL); + unwrappedRecords += AlignUp(desc.MissShaderTable.SizeInBytes, 256ULL); + + if(desc.HitGroupTable.StartAddress) + { + if(desc.HitGroupTable.StrideInBytes == 0) + desc.HitGroupTable.StrideInBytes = desc.HitGroupTable.SizeInBytes; + for(UINT64 i = 0; i < desc.HitGroupTable.SizeInBytes / desc.HitGroupTable.StrideInBytes; i++) + VerifyRecord(desc.HitGroupTable.StrideInBytes, + wrappedRecords + desc.HitGroupTable.StrideInBytes * i, + unwrappedRecords + desc.HitGroupTable.StrideInBytes * i, resHeap, sampHeap); + } + + wrappedRecords += AlignUp(desc.HitGroupTable.SizeInBytes, 256ULL); + unwrappedRecords += AlignUp(desc.HitGroupTable.SizeInBytes, 256ULL); + + if(desc.CallableShaderTable.StartAddress) + { + if(desc.CallableShaderTable.StrideInBytes == 0) + desc.CallableShaderTable.StrideInBytes = desc.CallableShaderTable.SizeInBytes; + for(UINT64 i = 0; + i < desc.CallableShaderTable.SizeInBytes / desc.CallableShaderTable.StrideInBytes; i++) + VerifyRecord(desc.CallableShaderTable.StrideInBytes, + wrappedRecords + desc.CallableShaderTable.StrideInBytes * i, + unwrappedRecords + desc.CallableShaderTable.StrideInBytes * i, resHeap, sampHeap); + } + + wrappedRecords += AlignUp(desc.CallableShaderTable.SizeInBytes, 256ULL); + unwrappedRecords += AlignUp(desc.CallableShaderTable.SizeInBytes, 256ULL); +} + void D3D12RTManager::VerifyRecord(const uint64_t recordSize, byte *wrappedRecord, byte *unwrappedRef, WrappedID3D12DescriptorHeap *resHeap, WrappedID3D12DescriptorHeap *sampHeap) @@ -889,9 +1080,17 @@ void D3D12RTManager::VerifyRecord(const uint64_t recordSize, byte *wrappedRecord } else { + // we should not get NULL objects back for records with data + RDCASSERT(ident->id == ResourceId()); memset(record.data(), 0, D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES); } + // copy remaining data - it will be overwritten below by an VA/handle unwraps, otherwise it could + // be constant data or just padding + memcpy(record.data() + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES, + wrappedRecord + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES, + recordSize - D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES); + if(localIdx != 0xffff) { rdcarray &rootConfig = m_UniqueLocalRootSigs[localIdx]; @@ -1450,15 +1649,19 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * const uint32_t raygenOffs = patchDataSize; patchDataSize = (uint32_t)desc.RayGenerationShaderRecord.SizeInBytes; - patchDataSize = AlignUp(patchDataSize, (uint32_t)D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT); + patchDataSize = AlignUp(patchDataSize, 256U); const uint32_t missOffs = patchDataSize; patchDataSize += (uint32_t)desc.MissShaderTable.SizeInBytes; - patchDataSize = AlignUp(patchDataSize, (uint32_t)D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT); + patchDataSize = AlignUp(patchDataSize, 256U); const uint32_t hitOffs = patchDataSize; patchDataSize += (uint32_t)desc.HitGroupTable.SizeInBytes; - patchDataSize = AlignUp(patchDataSize, (uint32_t)D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT); + patchDataSize = AlignUp(patchDataSize, 256U); + + // use 256 to match the indirect patching in raytracing.hlsl + RDCCOMPILE_ASSERT(256U >= D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, + "expect 256 to satisfy shader table alignment"); const uint32_t callOffs = patchDataSize; patchDataSize += (uint32_t)desc.CallableShaderTable.SizeInBytes; @@ -1508,8 +1711,8 @@ PatchedRayDispatch D3D12RTManager::PatchRayDispatch(ID3D12GraphicsCommandList4 * // set up general patching data - lookup buffers and so on - unwrappedCmd->SetPipelineState(m_RayPatchingData.descPatchPipe); - unwrappedCmd->SetComputeRootSignature(m_RayPatchingData.descPatchRootSig); + unwrappedCmd->SetPipelineState(m_RayPatchingData.shaderTablePatchPipe); + unwrappedCmd->SetComputeRootSignature(m_RayPatchingData.shaderTablePatchRootSig); unwrappedCmd->SetComputeRoot32BitConstants((UINT)D3D12PatchRayDispatchParam::GeneralCB, sizeof(cbufferData) / sizeof(uint32_t), &cbufferData, 0); unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::StateObjectData, @@ -1685,14 +1888,7 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( // in the largest TLAS, multiplied by the largest local root signature size. // // minimum of 64 bytes to account for MaxCommandCount so that we don't allocate 0-sized buffers - uint32_t patchDataSize = RDCMAX(64U, 20 * 1024 * 1024 * MaxCommandCount); - - if(D3D12_Debug_RT_IndirectEstimateOverride() > 0) - patchDataSize = D3D12_Debug_RT_IndirectEstimateOverride() * MaxCommandCount; - - m_GPUBufferAllocator.Alloc(D3D12GpuBufferHeapType::DefaultHeapWithUav, - D3D12GpuBufferHeapMemoryFlag::Default, patchDataSize, - D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, &scratchBuffer); + uint32_t patchDataSize = GuessTablePatchRequirements(MaxCommandCount); WrappedID3D12CommandSignature *comSig = (WrappedID3D12CommandSignature *)pCommandSignature; @@ -1704,14 +1900,56 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( uint64_t applicationArgsSize = AlignUp16(MaxCommandCount * comSig->sig.ByteStride); uint64_t patchingArgsSize = AlignUp16(MaxCommandCount * 4 * AlignUp16(sizeof(PatchingExecute))); - m_GPUBufferAllocator.Alloc(D3D12GpuBufferHeapType::DefaultHeapWithUav, - D3D12GpuBufferHeapMemoryFlag::Default, - applicationArgsSize + patchingArgsSize + 4, - D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, &argsBuffer); + ret.MaxCommands = MaxCommandCount; + ret.comSig = comSig; + ret.HasDynamicCount = pCountBuffer != NULL; + + uint64_t totalArgsSize = + // unwrapped exec args + applicationArgsSize + + // unwrapping internal exec + patchingArgsSize + + // count + 4; + uint64_t totalScratchSize = patchDataSize; + + uint64_t auditCopyArgsOffset = totalArgsSize; RDCCOMPILE_ASSERT(WRAPPED_DESCRIPTOR_STRIDE == sizeof(D3D12Descriptor), "Shader descriptor stride is wrong"); + ret.resources.readbackBuffer = NULL; + + if(IsReplayMode(m_wrappedDevice->GetState()) && D3D12_Debug_RT_Auditing()) + { + m_GPUBufferAllocator.Alloc( + D3D12GpuBufferHeapType::ReadBackHeap, D3D12GpuBufferHeapMemoryFlag::Default, + // unpatched records + patchDataSize + + // patched records + patchDataSize + + // app exec args + optional count + (applicationArgsSize + 4) + + // patched exec args + internal args + count + applicationArgsSize + patchingArgsSize + 4, + D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, &ret.resources.readbackBuffer); + + // double the args we need internally to make room for a second internal indirect dispatch to + // copy the unwrapped arguments + totalArgsSize *= 2; + + // we can't write directly to our readback buffer with our indirect copy, need to copy into scratch first + totalScratchSize *= 2; + } + + m_GPUBufferAllocator.Alloc(D3D12GpuBufferHeapType::DefaultHeapWithUav, + D3D12GpuBufferHeapMemoryFlag::Default, totalScratchSize, + D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, &scratchBuffer); + + m_GPUBufferAllocator.Alloc(D3D12GpuBufferHeapType::DefaultHeapWithUav, + D3D12GpuBufferHeapMemoryFlag::Default, totalArgsSize, + D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT, &argsBuffer); + RayDispatchPatchCB cbufferData = {}; for(ResourceId heapId : heaps) @@ -1744,7 +1982,8 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( prepInfo.commandSigSize = comSig->sig.PackedByteSize; prepInfo.commandSigStride = comSig->sig.ByteStride; prepInfo.maxCommandCount = MaxCommandCount; - prepInfo.scratchBuffer = scratchBuffer->Address(); + prepInfo.destBuffer = scratchBuffer->Address(); + prepInfo.destBufferEnd = scratchBuffer->Address() + patchDataSize; if(pCountBuffer == NULL) prepInfo.maxCommandCount |= 0x80000000U; @@ -1781,6 +2020,28 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( // patching the actual arguments buffer we'll return unwrappedCmd->Dispatch(1, 1, 1); + // if we're reading back, do another dispatch to write args which just purely copy + if(ret.resources.readbackBuffer) + { + prepInfo.destBuffer = scratchBuffer->Address() + patchDataSize; + prepInfo.destBufferEnd = scratchBuffer->Address() + patchDataSize * 2; + unwrappedCmd->SetComputeRoot32BitConstants((UINT)D3D12IndirectPrepParam::GeneralCB, + sizeof(prepInfo) / sizeof(uint32_t), &prepInfo, 0); + + // unused, we don't need this but it keeps the shader the same + unwrappedCmd->SetComputeRootUnorderedAccessView((UINT)D3D12IndirectPrepParam::PatchedExecuteArgs, + argsBuffer->Address() + auditCopyArgsOffset); + // this will be the arguments for our internal copy-dispatch + unwrappedCmd->SetComputeRootUnorderedAccessView( + (UINT)D3D12IndirectPrepParam::InternalExecuteArgs, + argsBuffer->Address() + auditCopyArgsOffset + applicationArgsSize); + unwrappedCmd->SetComputeRootUnorderedAccessView( + (UINT)D3D12IndirectPrepParam::InternalExecuteCount, + argsBuffer->Address() + auditCopyArgsOffset + applicationArgsSize + patchingArgsSize); + + unwrappedCmd->Dispatch(1, 1, 1); + } + barrier.Transition.pResource = Unwrap(pArgumentBuffer); barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_NON_PIXEL_SHADER_RESOURCE; barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_INDIRECT_ARGUMENT; @@ -1792,8 +2053,8 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_INDIRECT_ARGUMENT; unwrappedCmd->ResourceBarrier(1, &barrier); - unwrappedCmd->SetPipelineState(m_RayPatchingData.descPatchPipe); - unwrappedCmd->SetComputeRootSignature(m_RayPatchingData.descPatchRootSig); + unwrappedCmd->SetPipelineState(m_RayPatchingData.shaderTablePatchPipe); + unwrappedCmd->SetComputeRootSignature(m_RayPatchingData.shaderTablePatchRootSig); unwrappedCmd->SetComputeRoot32BitConstants((UINT)D3D12PatchRayDispatchParam::GeneralCB, sizeof(cbufferData) / sizeof(uint32_t), &cbufferData, 0); unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::StateObjectData, @@ -1823,9 +2084,98 @@ PatchedRayDispatch D3D12RTManager::PatchIndirectRayDispatch( // scratch buffer has now been patched and is ready to use as well 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) + { + barrier.Transition.pResource = Unwrap(pArgumentBuffer); + barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_INDIRECT_ARGUMENT; + barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_COPY_SOURCE; + unwrappedCmd->ResourceBarrier(1, &barrier); + + uint64_t offset = 0; + + unwrappedCmd->SetPipelineState(m_RayPatchingData.shaderTableCopyPipe); + unwrappedCmd->SetComputeRootSignature(m_RayPatchingData.shaderTablePatchRootSig); + unwrappedCmd->SetComputeRoot32BitConstants((UINT)D3D12PatchRayDispatchParam::GeneralCB, + sizeof(cbufferData) / sizeof(uint32_t), &cbufferData, + 0); + unwrappedCmd->SetComputeRootShaderResourceView( + (UINT)D3D12PatchRayDispatchParam::StateObjectData, m_LookupAddrs[0]); + unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::RecordData, + m_LookupAddrs[1]); + unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::RootSigData, + m_LookupAddrs[2]); + unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::AddrPatchData, + m_LookupAddrs[3]); + + // these will be overwritten by the execute indirect, but set them to something to be safe + unwrappedCmd->SetComputeRoot32BitConstants((UINT)D3D12PatchRayDispatchParam::RecordCB, + sizeof(recordInfo) / sizeof(uint32_t), &recordInfo, 0); + unwrappedCmd->SetComputeRootShaderResourceView((UINT)D3D12PatchRayDispatchParam::SourceBuffer, + m_LookupAddrs[0]); + unwrappedCmd->SetComputeRootUnorderedAccessView((UINT)D3D12PatchRayDispatchParam::DestBuffer, + scratchBuffer->Address()); + + unwrappedCmd->ExecuteIndirect( + m_RayPatchingData.indirectComSig, MaxCommandCount * 4, argsBuffer->Resource(), + argsBuffer->Offset() + auditCopyArgsOffset + applicationArgsSize, argsBuffer->Resource(), + argsBuffer->Offset() + auditCopyArgsOffset + applicationArgsSize + patchingArgsSize); + + // unpatched data + unwrappedCmd->CopyBufferRegion( + ret.resources.readbackBuffer->Resource(), ret.resources.readbackBuffer->Offset() + offset, + scratchBuffer->Resource(), scratchBuffer->Offset() + patchDataSize, patchDataSize); + offset += patchDataSize; + + // patched data + unwrappedCmd->CopyBufferRegion( + ret.resources.readbackBuffer->Resource(), ret.resources.readbackBuffer->Offset() + offset, + scratchBuffer->Resource(), scratchBuffer->Offset(), patchDataSize); + offset += patchDataSize; + + // unpatched execute args + unwrappedCmd->CopyBufferRegion(ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + patchDataSize * 2, + Unwrap(pArgumentBuffer), ArgumentBufferOffset, + applicationArgsSize); + offset += applicationArgsSize; + + // source count + if(pCountBuffer) + { + unwrappedCmd->CopyBufferRegion(ret.resources.readbackBuffer->Resource(), + ret.resources.readbackBuffer->Offset() + offset, + Unwrap(pCountBuffer), CountBufferOffset, applicationArgsSize); + } + offset += 4; + + barrier.Transition.pResource = argsBuffer->Resource(); + barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_INDIRECT_ARGUMENT; + barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_COPY_SOURCE; + unwrappedCmd->ResourceBarrier(1, &barrier); + + // patched execute args, our internal executes, count + unwrappedCmd->CopyBufferRegion( + ret.resources.readbackBuffer->Resource(), ret.resources.readbackBuffer->Offset() + offset, + argsBuffer->Resource(), argsBuffer->Offset(), applicationArgsSize + patchingArgsSize + 4); + offset += applicationArgsSize + patchingArgsSize + 4; + + RDCASSERTEQUAL(offset, ret.resources.readbackBuffer->Size()); + + barrier.Transition.pResource = argsBuffer->Resource(); + barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_COPY_SOURCE; + barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_INDIRECT_ARGUMENT; + unwrappedCmd->ResourceBarrier(1, &barrier); + + barrier.Transition.pResource = Unwrap(pArgumentBuffer); + barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_COPY_SOURCE; + barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_INDIRECT_ARGUMENT; + unwrappedCmd->ResourceBarrier(1, &barrier); + } + // 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(); @@ -2450,7 +2800,7 @@ void D3D12RTManager::InitRayDispatchPatchingResources() { HRESULT result = m_wrappedDevice->GetReal()->CreateRootSignature( 0, rootSig.data(), rootSig.size(), __uuidof(ID3D12RootSignature), - (void **)&m_RayPatchingData.descPatchRootSig); + (void **)&m_RayPatchingData.shaderTablePatchRootSig); if(!SUCCEEDED(result)) RDCERR("Unable to create root signature for dispatch patching"); @@ -2458,7 +2808,7 @@ void D3D12RTManager::InitRayDispatchPatchingResources() // PipelineState ID3DBlob *shader = NULL; rdcstr hlsl = GetEmbeddedResource(raytracing_hlsl); - shaderCache->GetShaderBlob(hlsl.c_str(), "RENDERDOC_PatchRayDispatchCS", + shaderCache->GetShaderBlob(hlsl.c_str(), "RENDERDOC_PatchShaderTableCS", D3DCOMPILE_WARNINGS_ARE_ERRORS, {}, "cs_5_0", &shader); if(shader) @@ -2468,20 +2818,48 @@ void D3D12RTManager::InitRayDispatchPatchingResources() pipeline.NodeMask = 0; pipeline.CS = {(void *)shader->GetBufferPointer(), shader->GetBufferSize()}; pipeline.CachedPSO = {NULL, 0}; - pipeline.pRootSignature = m_RayPatchingData.descPatchRootSig; + pipeline.pRootSignature = m_RayPatchingData.shaderTablePatchRootSig; result = m_wrappedDevice->GetReal()->CreateComputePipelineState( - &pipeline, __uuidof(ID3D12PipelineState), (void **)&m_RayPatchingData.descPatchPipe); + &pipeline, __uuidof(ID3D12PipelineState), (void **)&m_RayPatchingData.shaderTablePatchPipe); if(!SUCCEEDED(result)) RDCERR("Unable to create pipeline for dispatch patching"); - if(m_RayPatchingData.descPatchPipe) - m_RayPatchingData.descPatchPipe->SetName(L"RENDERDOC_PatchRayDispatchCS"); + if(m_RayPatchingData.shaderTablePatchPipe) + m_RayPatchingData.shaderTablePatchPipe->SetName(L"RENDERDOC_PatchShaderTableCS"); } else { - RDCERR("Failed to get shader for dispatch patching"); + RDCERR("Failed to get shader for record patching in dispatch"); + } + + SAFE_RELEASE(shader); + + shaderCache->GetShaderBlob(hlsl.c_str(), "RENDERDOC_CopyShaderTableCS", + D3DCOMPILE_WARNINGS_ARE_ERRORS, {}, "cs_5_0", &shader); + + if(shader) + { + D3D12_COMPUTE_PIPELINE_STATE_DESC pipeline; + pipeline.Flags = D3D12_PIPELINE_STATE_FLAG_NONE; + pipeline.NodeMask = 0; + pipeline.CS = {(void *)shader->GetBufferPointer(), shader->GetBufferSize()}; + pipeline.CachedPSO = {NULL, 0}; + pipeline.pRootSignature = m_RayPatchingData.shaderTablePatchRootSig; + + result = m_wrappedDevice->GetReal()->CreateComputePipelineState( + &pipeline, __uuidof(ID3D12PipelineState), (void **)&m_RayPatchingData.shaderTableCopyPipe); + + if(!SUCCEEDED(result)) + RDCERR("Unable to create pipeline for dispatch patching"); + + if(m_RayPatchingData.shaderTableCopyPipe) + m_RayPatchingData.shaderTableCopyPipe->SetName(L"RENDERDOC_CopyShaderTableCS"); + } + else + { + RDCERR("Failed to get shader for record copying in indirect auditing"); } SAFE_RELEASE(shader); @@ -2635,7 +3013,7 @@ void D3D12RTManager::InitRayDispatchPatchingResources() desc.pArgumentDescs = args; HRESULT hr = m_wrappedDevice->GetReal()->CreateCommandSignature( - &desc, m_RayPatchingData.descPatchRootSig, __uuidof(ID3D12CommandSignature), + &desc, m_RayPatchingData.shaderTablePatchRootSig, __uuidof(ID3D12CommandSignature), (void **)&m_RayPatchingData.indirectComSig); if(!SUCCEEDED(hr)) diff --git a/renderdoc/driver/d3d12/d3d12_manager.h b/renderdoc/driver/d3d12/d3d12_manager.h index eba571485..2705b8260 100644 --- a/renderdoc/driver/d3d12/d3d12_manager.h +++ b/renderdoc/driver/d3d12/d3d12_manager.h @@ -1036,6 +1036,8 @@ struct D3D12AccStructPatchInfo ID3D12PipelineState *m_pipeline = NULL; }; +class WrappedID3D12CommandSignature; + struct PatchedRayDispatch { struct Resources @@ -1078,6 +1080,10 @@ struct PatchedRayDispatch // the patched dispatch descriptor D3D12_DISPATCH_RAYS_DESC desc = {}; rdcarray heaps; + // for auditing, from an indirect RT dispatch + UINT MaxCommands = 0; + WrappedID3D12CommandSignature *comSig = NULL; + bool HasDynamicCount = false; }; struct D3D12ShaderExportDatabase; @@ -1334,7 +1340,9 @@ public: void Verify(PatchedRayDispatch &r); - void VerifyRecord(const uint64_t recordSize, byte *table, byte *ref, + void VerifyDispatch(D3D12_DISPATCH_RAYS_DESC desc, byte *wrappedRecords, byte *unwrappedRecords, + WrappedID3D12DescriptorHeap *resHeap, WrappedID3D12DescriptorHeap *sampHeap); + void VerifyRecord(const uint64_t recordSize, byte *wrappedRecord, byte *unwrappedRef, WrappedID3D12DescriptorHeap *resHeap, WrappedID3D12DescriptorHeap *sampHeap); void AddDispatchTimer(uint32_t q); @@ -1393,8 +1401,9 @@ private: // pipeline data for patching ray dispatches struct { - ID3D12RootSignature *descPatchRootSig = NULL; - ID3D12PipelineState *descPatchPipe = NULL; + ID3D12RootSignature *shaderTablePatchRootSig = NULL; + ID3D12PipelineState *shaderTablePatchPipe = NULL; + ID3D12PipelineState *shaderTableCopyPipe = NULL; ID3D12RootSignature *indirectPrepRootSig = NULL; ID3D12PipelineState *indirectPrepPipe = NULL; ID3D12CommandSignature *indirectComSig = NULL;