DX12TLAS.cpp 8.4 KB

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