DX12TLAS.cpp 7.3 KB

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