DX12TLAS.cpp 9.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268
  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. }
  98. descriptorBuffer->zBuffer()->Unmap(0, nullptr);
  99. }
  100. mappedDescriptorBuffer = 0;
  101. if (!descriptorBuffer || currentInstanceIndex >= lastInstanceCount)
  102. {
  103. // recalculate the size of the descriptor buffer and reallocate it
  104. D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_INPUTS
  105. prebuildDesc = {};
  106. prebuildDesc.Type
  107. = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL;
  108. prebuildDesc.DescsLayout = D3D12_ELEMENTS_LAYOUT_ARRAY;
  109. prebuildDesc.NumDescs = (unsigned)currentInstanceIndex + 1;
  110. prebuildDesc.Flags
  111. = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_ALLOW_UPDATE;
  112. D3D12_RAYTRACING_ACCELERATION_STRUCTURE_PREBUILD_INFO info = {};
  113. zDevice->GetRaytracingAccelerationStructurePrebuildInfo(
  114. &prebuildDesc, &info);
  115. // create new buffers witch fit the TLAS
  116. if (scratchBuffer)
  117. {
  118. if (lastScratchBuffer)
  119. {
  120. lastScratchBuffer->Release();
  121. }
  122. lastScratchBuffer = scratchBuffer->zBuffer();
  123. lastScratchBuffer->AddRef();
  124. }
  125. if (!resultBuffer)
  126. {
  127. resultBuffer = new DX12Buffer(1,
  128. zDevice,
  129. dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
  130. D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS);
  131. }
  132. resultBuffer->setLength(
  133. ROUND_UP_POWER_OF_2(info.ResultDataMaxSizeInBytes, 256));
  134. resultBuffer->createBufferWithoutData(
  135. D3D12_RESOURCE_STATE_RAYTRACING_ACCELERATION_STRUCTURE);
  136. resultBuffer->zBuffer()->SetName(L"TLAS Result Buffer");
  137. if (!scratchBuffer)
  138. {
  139. scratchBuffer = new DX12Buffer(1,
  140. zDevice,
  141. dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
  142. D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS);
  143. }
  144. scratchBuffer->setLength(
  145. ROUND_UP_POWER_OF_2(info.ScratchDataSizeInBytes, 256));
  146. scratchBuffer->createBufferWithoutData(
  147. D3D12_RESOURCE_STATE_UNORDERED_ACCESS);
  148. scratchBuffer->zBuffer()->SetName(L"TLAS Scratch Buffer");
  149. ID3D12Resource* oldDescriptorBuffer
  150. = descriptorBuffer ? descriptorBuffer->zBuffer() : 0;
  151. __int64 oldDescriptorBufferElementCount
  152. = descriptorBuffer ? descriptorBuffer->getElementCount() : 0;
  153. if (oldDescriptorBuffer)
  154. {
  155. oldDescriptorBuffer->AddRef();
  156. oldDescriptorBuffer->Map(0, 0, (void**)&mappedDescriptorBuffer);
  157. }
  158. if (!descriptorBuffer)
  159. {
  160. descriptorBuffer
  161. = new DX12Buffer(sizeof(D3D12_RAYTRACING_INSTANCE_DESC),
  162. zDevice,
  163. dynamic_cast<DX12CommandQueue*>(zDirectQueue->getThis()),
  164. D3D12_RESOURCE_FLAG_NONE);
  165. }
  166. descriptorBuffer->setLength(ROUND_UP_POWER_OF_2(
  167. (currentInstanceIndex + 1) * sizeof(D3D12_RAYTRACING_INSTANCE_DESC),
  168. 256));
  169. descriptorBuffer->createBufferWithoutData(
  170. D3D12_RESOURCE_STATE_GENERIC_READ, D3D12_HEAP_TYPE_UPLOAD);
  171. D3D12_RAYTRACING_INSTANCE_DESC* newMappedDescriptorBuffer = 0;
  172. // copy the old and new instance descriptions to the new buffer
  173. descriptorBuffer->zBuffer()->Map(
  174. 0, 0, (void**)&newMappedDescriptorBuffer);
  175. if (mappedDescriptorBuffer)
  176. {
  177. memcpy(newMappedDescriptorBuffer,
  178. mappedDescriptorBuffer,
  179. oldDescriptorBufferElementCount
  180. * sizeof(D3D12_RAYTRACING_INSTANCE_DESC));
  181. D3D12_RANGE range = {0, 0}; // do not write to the old buffer
  182. oldDescriptorBuffer->Unmap(0, &range);
  183. oldDescriptorBuffer->Release();
  184. }
  185. overflowInstanceIterator = overflowInstanceDescs.begin();
  186. for (__int64 i = oldDescriptorBufferElementCount;
  187. i <= currentInstanceIndex;
  188. i++)
  189. {
  190. memcpy(newMappedDescriptorBuffer + i,
  191. overflowInstanceIterator.val(),
  192. sizeof(D3D12_RAYTRACING_INSTANCE_DESC));
  193. overflowInstanceIterator++;
  194. }
  195. descriptorBuffer->zBuffer()->Unmap(0, nullptr);
  196. bufferChanged = 1;
  197. }
  198. else
  199. {
  200. if (resultBuffer)
  201. {
  202. if (lastResultBuffer)
  203. {
  204. lastResultBuffer->Release();
  205. }
  206. lastResultBuffer = resultBuffer->zBuffer();
  207. lastResultBuffer->AddRef();
  208. }
  209. }
  210. D3D12_BUILD_RAYTRACING_ACCELERATION_STRUCTURE_DESC buildDesc = {};
  211. buildDesc.Inputs.Type
  212. = D3D12_RAYTRACING_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL;
  213. buildDesc.Inputs.DescsLayout = D3D12_ELEMENTS_LAYOUT_ARRAY;
  214. buildDesc.Inputs.InstanceDescs
  215. = descriptorBuffer->zBuffer()->GetGPUVirtualAddress();
  216. buildDesc.Inputs.NumDescs = (unsigned)currentInstanceIndex + 1;
  217. buildDesc.DestAccelerationStructureData
  218. = {resultBuffer->zBuffer()->GetGPUVirtualAddress()};
  219. buildDesc.ScratchAccelerationStructureData
  220. = {scratchBuffer->zBuffer()->GetGPUVirtualAddress()};
  221. buildDesc.SourceAccelerationStructureData
  222. = lastResultBuffer && currentInstanceIndex < lastInstanceCount
  223. ? lastResultBuffer->GetGPUVirtualAddress()
  224. : 0;
  225. buildDesc.Inputs.Flags
  226. = lastResultBuffer && currentInstanceIndex < lastInstanceCount
  227. ? D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_PERFORM_UPDATE
  228. : D3D12_RAYTRACING_ACCELERATION_STRUCTURE_BUILD_FLAG_ALLOW_UPDATE;
  229. // Build the top-level AS
  230. zDirectQueue->zCommandList()->BuildRaytracingAccelerationStructure(
  231. &buildDesc, 0, nullptr);
  232. // Wait for the builder to complete by setting a barrier on the resulting
  233. // buffer. This can be important in case the rendering is triggered
  234. // immediately afterwards, without executing the command list
  235. D3D12_RESOURCE_BARRIER uavBarrier;
  236. uavBarrier.Type = D3D12_RESOURCE_BARRIER_TYPE_UAV;
  237. uavBarrier.UAV.pResource = resultBuffer->zBuffer();
  238. uavBarrier.Flags = D3D12_RESOURCE_BARRIER_FLAG_NONE;
  239. zDirectQueue->zCommandList()->ResourceBarrier(1, &uavBarrier);
  240. lastInstanceCount = currentInstanceIndex + 1;
  241. }
  242. Framework::DX12Buffer* Framework::DX12TLAS::zResultBuffer() const
  243. {
  244. return resultBuffer;
  245. }
  246. bool Framework::DX12TLAS::hasBufferChanged() const
  247. {
  248. return bufferChanged;
  249. }
  250. void Framework::DX12TLAS::setBufferChanged(bool changed)
  251. {
  252. bufferChanged = changed;
  253. }