DX12TLAS.cpp 9.1 KB

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