DX12TLAS.cpp 8.3 KB

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