DX12TLAS.cpp 8.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233
  1. #include "DX12TLAS.h"
  2. #include "DX12CommandQueue.h"
  3. Framework::DX12TLAS::DX12TLAS(
  4. ID3D12Device5* zDevice, Framework::DX12DirectCommandQueue* zDirectQueue)
  5. : ReferenceCounter(),
  6. scratchBuffer(0),
  7. resultBuffer(0),
  8. descriptorBuffer(0),
  9. previousResultBuffer(0),
  10. zDevice(zDevice),
  11. zDirectQueue(zDirectQueue),
  12. overflowInstanceIterator(0, 0, 0, 0),
  13. currentInstanceIndex(0),
  14. lastInstanceCount(0),
  15. mappedDescriptorBuffer(0)
  16. {}
  17. Framework::DX12TLAS::~DX12TLAS()
  18. {
  19. if (scratchBuffer)
  20. {
  21. scratchBuffer->release();
  22. }
  23. if (resultBuffer)
  24. {
  25. resultBuffer->release();
  26. }
  27. if (descriptorBuffer)
  28. {
  29. descriptorBuffer->release();
  30. }
  31. if (previousResultBuffer)
  32. {
  33. previousResultBuffer->release();
  34. }
  35. for (const D3D12_RAYTRACING_INSTANCE_DESC* desc : overflowInstanceDescs)
  36. {
  37. delete desc;
  38. }
  39. }
  40. void Framework::DX12TLAS::startUpdate()
  41. {
  42. currentInstanceIndex = -1;
  43. if (descriptorBuffer)
  44. {
  45. descriptorBuffer->zBuffer()->Map(0, 0, (void**)&mappedDescriptorBuffer);
  46. }
  47. else
  48. {
  49. mappedDescriptorBuffer = 0;
  50. }
  51. overflowInstanceIterator = overflowInstanceDescs.begin();
  52. }
  53. D3D12_RAYTRACING_INSTANCE_DESC* Framework::DX12TLAS::nextInstanceDesc()
  54. {
  55. currentInstanceIndex++;
  56. if (mappedDescriptorBuffer
  57. && currentInstanceIndex < descriptorBuffer->getElementCount())
  58. {
  59. return mappedDescriptorBuffer + currentInstanceIndex;
  60. }
  61. else if (overflowInstanceIterator)
  62. {
  63. return *overflowInstanceIterator++;
  64. }
  65. else
  66. {
  67. D3D12_RAYTRACING_INSTANCE_DESC* desc
  68. = new D3D12_RAYTRACING_INSTANCE_DESC();
  69. overflowInstanceDescs.add(desc);
  70. return desc;
  71. }
  72. }
  73. void Framework::DX12TLAS::endUpdate()
  74. {
  75. if (currentInstanceIndex < 0) // at least one instance ust be added
  76. {
  77. D3D12_RAYTRACING_INSTANCE_DESC* desc = nextInstanceDesc();
  78. desc->InstanceContributionToHitGroupIndex = 0;
  79. desc->InstanceID = 0;
  80. desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
  81. desc->AccelerationStructure = 0;
  82. desc->InstanceMask = 0; // Mark unused instances with a mask of 0
  83. }
  84. if (descriptorBuffer)
  85. {
  86. while (currentInstanceIndex < lastInstanceCount - 1)
  87. {
  88. D3D12_RAYTRACING_INSTANCE_DESC* desc = nextInstanceDesc();
  89. desc->InstanceMask = 0; // Mark unused instances with a mask of 0
  90. }
  91. descriptorBuffer->zBuffer()->Unmap(0, nullptr);
  92. }
  93. mappedDescriptorBuffer = 0;
  94. if (!descriptorBuffer || currentInstanceIndex >= lastInstanceCount)
  95. {
  96. // recalculate the size of the descriptor buffer and reallocate it
  97. D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_INPUTS
  98. prebuildDesc = {};
  99. prebuildDesc.Type
  100. = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL;
  101. prebuildDesc.DescsLayout = D3D12_ELEMENTS_LAYOUT_ARRAY;
  102. prebuildDesc.NumDescs = (unsigned)currentInstanceIndex + 1;
  103. prebuildDesc.Flags
  104. = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_ALLOW_UPDATE;
  105. D3D12_RAYTRACING_ACCELERATION_STRUCTURE_PREBUILD_INFO info = {};
  106. zDevice->GetRaytracingAccelerationStructurePrebuildInfo(
  107. &prebuildDesc, &info);
  108. // create new buffers witch fit the TLAS
  109. if (resultBuffer)
  110. {
  111. if (previousResultBuffer)
  112. {
  113. previousResultBuffer->release();
  114. }
  115. previousResultBuffer = resultBuffer;
  116. }
  117. resultBuffer = new DX12Buffer(1,
  118. zDevice,
  119. dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
  120. D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS);
  121. resultBuffer->setLength(
  122. ROUND_UP_POWER_OF_2(info.ResultDataMaxSizeInBytes, 256));
  123. resultBuffer->createBufferWithoutData(
  124. D3D12_RESOURCE_STATE_RAYTRACING_ACCELERATION_STRUCTURE);
  125. resultBuffer->zBuffer()->SetName(L"TLAS Result Buffer");
  126. if (!scratchBuffer)
  127. {
  128. scratchBuffer = new DX12Buffer(1,
  129. zDevice,
  130. dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
  131. D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS);
  132. }
  133. scratchBuffer->setLength(
  134. ROUND_UP_POWER_OF_2(info.ScratchDataSizeInBytes, 256));
  135. scratchBuffer->createBufferWithoutData(
  136. D3D12_RESOURCE_STATE_UNORDERED_ACCESS);
  137. scratchBuffer->zBuffer()->SetName(L"TLAS Scratch Buffer");
  138. ID3D12Resource* oldDescriptorBuffer
  139. = descriptorBuffer ? descriptorBuffer->zBuffer() : 0;
  140. __int64 oldDescriptorBufferElementCount
  141. = descriptorBuffer ? descriptorBuffer->getElementCount() : 0;
  142. if (oldDescriptorBuffer)
  143. {
  144. oldDescriptorBuffer->AddRef();
  145. oldDescriptorBuffer->Map(0, 0, (void**)&mappedDescriptorBuffer);
  146. }
  147. if (!descriptorBuffer)
  148. {
  149. descriptorBuffer
  150. = new DX12Buffer(sizeof(D3D12_RAYTRACING_INSTANCE_DESC),
  151. zDevice,
  152. dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
  153. D3D12_RESOURCE_FLAG_NONE);
  154. }
  155. descriptorBuffer->setLength(ROUND_UP_POWER_OF_2(
  156. (currentInstanceIndex + 1) * sizeof(D3D12_RAYTRACING_INSTANCE_DESC),
  157. 256));
  158. descriptorBuffer->createBufferWithoutData(
  159. D3D12_RESOURCE_STATE_GENERIC_READ, D3D12_HEAP_TYPE_UPLOAD);
  160. D3D12_RAYTRACING_INSTANCE_DESC* newMappedDescriptorBuffer = 0;
  161. // copy the old and new instance descriptions to the new buffer
  162. descriptorBuffer->zBuffer()->Map(
  163. 0, 0, (void**)&newMappedDescriptorBuffer);
  164. if (mappedDescriptorBuffer)
  165. {
  166. memcpy(newMappedDescriptorBuffer,
  167. mappedDescriptorBuffer,
  168. oldDescriptorBufferElementCount);
  169. D3D12_RANGE range = {0, 0}; // do not write to the old buffer
  170. oldDescriptorBuffer->Unmap(0, &range);
  171. oldDescriptorBuffer->Release();
  172. }
  173. overflowInstanceIterator = overflowInstanceDescs.begin();
  174. for (__int64 i = oldDescriptorBufferElementCount;
  175. i <= currentInstanceIndex;
  176. i++)
  177. {
  178. memcpy(newMappedDescriptorBuffer + i,
  179. overflowInstanceIterator.val(),
  180. sizeof(D3D12_RAYTRACING_INSTANCE_DESC));
  181. overflowInstanceIterator++;
  182. }
  183. descriptorBuffer->zBuffer()->Unmap(0, nullptr);
  184. }
  185. D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_DESC buildDesc = {};
  186. buildDesc.Inputs.Type
  187. = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL;
  188. buildDesc.Inputs.DescsLayout = D3D12_ELEMENTS_LAYOUT_ARRAY;
  189. buildDesc.Inputs.InstanceDescs
  190. = descriptorBuffer->zBuffer()->GetGPUVirtualAddress();
  191. buildDesc.Inputs.NumDescs = (unsigned)currentInstanceIndex + 1;
  192. buildDesc.DestAccelerationStructureData
  193. = {resultBuffer->zBuffer()->GetGPUVirtualAddress()};
  194. buildDesc.ScratchAccelerationStructureData
  195. = {scratchBuffer->zBuffer()->GetGPUVirtualAddress()};
  196. buildDesc.SourceAccelerationStructureData
  197. = previousResultBuffer
  198. ? previousResultBuffer->zBuffer()->GetGPUVirtualAddress()
  199. : 0;
  200. buildDesc.Inputs.Flags
  201. = previousResultBuffer
  202. ? D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_PERFORM_UPDATE
  203. : D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_ALLOW_UPDATE;
  204. // Build the top-level AS
  205. zDirectQueue->zCommandList()->BuildRaytracingAccelerationStructure(
  206. &buildDesc, 0, nullptr);
  207. // Wait for the builder to complete by setting a barrier on the resulting
  208. // buffer. This can be important in case the rendering is triggered
  209. // immediately afterwards, without executing the command list
  210. D3D12_RESOURCE_BARRIER uavBarrier;
  211. uavBarrier.Type = D3D12_RESOURCE_BARRIER_TYPE_UAV;
  212. uavBarrier.UAV.pResource = resultBuffer->zBuffer();
  213. uavBarrier.Flags = D3D12_RESOURCE_BARRIER_FLAG_NONE;
  214. zDirectQueue->zCommandList()->ResourceBarrier(1, &uavBarrier);
  215. lastInstanceCount = currentInstanceIndex + 1;
  216. }
  217. Framework::DX12Buffer* Framework::DX12TLAS::zResultBuffer() const
  218. {
  219. return resultBuffer;
  220. }