CustomDX12API.cpp 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370
  1. #include "CustomDX12API.h"
  2. #include <DX12Buffer.h>
  3. #include <DX12CommandQueue.h>
  4. #include <DX12Shader.h>
  5. #include <DX12TLAS.h>
  6. #include "CustomChunkAnyHitShader.h"
  7. #include "CustomChunkClosestHitShader.h"
  8. #include "CustomChunkIntersectionShader.h"
  9. #include "CustomCustomMissShader.h"
  10. #include "CustomCustomRayGenShader.h"
  11. #include "CustomDefaultAnyHitShader.h"
  12. #include "CustomDefaultClosestHitShader.h"
  13. #include "CustomSimpleBlocksAnyHitShader.h"
  14. #include "CustomSimpleBlocksClosestHitShader.h"
  15. #include "Dimension.h"
  16. #include "DX12ChunkData.h"
  17. using namespace Framework;
  18. CustomDX12API::CustomDX12API()
  19. : DirectX12()
  20. {}
  21. CustomDX12API::~CustomDX12API() {}
  22. void CustomDX12API::initializePipeline()
  23. {
  24. if (pipeline->getShaders().getEntryCount() == 0)
  25. { // add default shaders
  26. DX12Shader* rayGenShader = new DX12Shader(
  27. CustomCustomRayGenShader, sizeof(CustomCustomRayGenShader));
  28. // RayGen from RayGen.hlsl
  29. DX12ShaderSignature* rayGenSignature = new DX12ShaderSignature();
  30. rayGenSignature->addRegisterUsageLinkedToDescriptorHeap(
  31. 0, DX12_SHADER_REGISTER_U_UNORDERED_ACCESS, 0);
  32. rayGenSignature->addRegisterUsageLinkedToDescriptorHeap(
  33. 1, DX12_SHADER_REGISTER_U_UNORDERED_ACCESS, 1);
  34. rayGenSignature->addRegisterUsageLinkedToDescriptorHeap(
  35. 2, DX12_SHADER_REGISTER_B_CONST_BUFFER, 0);
  36. rayGenSignature->addRegisterUsageLinkedToDescriptorHeap(
  37. 3, DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 0);
  38. defaultRayGenerationShaderFunction = new DX12ShaderFunction(
  39. "RayGen", rayGenSignature, DX12_SHADER_FUNCTION_TYPE_RAY_GEN);
  40. rayGenShader->addFunction(defaultRayGenerationShaderFunction);
  41. pipeline->addShader(rayGenShader);
  42. DX12Shader* missShader = new DX12Shader(
  43. CustomCustomMissShader, sizeof(CustomCustomMissShader));
  44. // Miss from Miss.hlsl
  45. missShader->addFunction(new DX12ShaderFunction(
  46. "Miss", new DX12ShaderSignature(), DX12_SHADER_FUNCTION_TYPE_MISS));
  47. pipeline->addShader(missShader);
  48. DX12ShaderSignature* hitSignature = new DX12ShaderSignature();
  49. hitSignature->addRegisterUsageLinkedToDescriptorHeap(
  50. 0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP);
  51. hitSignature->addRegisterUsageLinkedToDescriptorHeap(0,
  52. DX12_SHADER_REGISTER_T_SHADER_RESOURCE,
  53. 0,
  54. 1,
  55. TEXTURE_DESCRIPTOR_HEAP,
  56. 1);
  57. sbtTextureIdBufferOffset
  58. = hitSignature->addRegisterUsageLinkedToShaderBindingTable(
  59. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2);
  60. sbtIndexBufferOffset
  61. = hitSignature->addRegisterUsageLinkedToShaderBindingTable(
  62. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 3);
  63. sbtVertexDataBufferOffset
  64. = hitSignature->addRegisterUsageLinkedToShaderBindingTable(
  65. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 4);
  66. sbtPolygonSizeBufferOffset
  67. = hitSignature->addRegisterUsageLinkedToShaderBindingTable(
  68. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 5);
  69. DX12Shader* anyHitShader = new DX12Shader(
  70. CustomDefaultAnyHitShader, sizeof(CustomDefaultAnyHitShader));
  71. DX12ShaderFunction* anyHitFunction = new DX12ShaderFunction(
  72. "AnyHit", hitSignature, DX12_SHADER_FUNCTION_TYPE_ANY_HIT);
  73. // AnyHit from AnyHit.hlsl
  74. anyHitShader->addFunction(anyHitFunction);
  75. pipeline->addShader(anyHitShader);
  76. DX12Shader* hitShader = new DX12Shader(CustomDefaultClosestHitShader,
  77. sizeof(CustomDefaultClosestHitShader));
  78. // ClosestHit from Hit.hlsl
  79. DX12ShaderFunction* closestHitFunction
  80. = new DX12ShaderFunction("ClosestHit",
  81. dynamic_cast<DX12ShaderSignature*>(hitSignature->getThis()),
  82. DX12_SHADER_FUNCTION_TYPE_CLOSEST_HIT);
  83. hitShader->addFunction(closestHitFunction);
  84. pipeline->addShader(hitShader);
  85. defaultHitGroup = new DX12ShaderHitGroup("HitGroup");
  86. defaultHitGroup->setAttributeSize(8);
  87. defaultHitGroup->setPayloadSize(28 * DEFAULT_MAX_TRANSPARENT_HITS + 4);
  88. defaultHitGroup->setAnyHitShaderFunction(anyHitFunction);
  89. defaultHitGroup->setClosestHitShaderFunction(closestHitFunction);
  90. pipeline->addHitGroup(defaultHitGroup);
  91. // Chunk ray traversal shaders
  92. DX12ShaderSignature* chunkHitSignature = new DX12ShaderSignature();
  93. chunkHitSignature->addRegisterUsageLinkedToDescriptorHeap(
  94. 0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP);
  95. chunkHitSignature->addRegisterUsageLinkedToDescriptorHeap(0,
  96. DX12_SHADER_REGISTER_T_SHADER_RESOURCE,
  97. 0,
  98. 1,
  99. TEXTURE_DESCRIPTOR_HEAP,
  100. 1);
  101. sbtChunkTextureIdBufferOffset
  102. = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable(
  103. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2);
  104. sbtChunkIndexBufferOffset
  105. = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable(
  106. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 3);
  107. sbtChunkLightBufferOffset
  108. = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable(
  109. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 4);
  110. sbtChunkDataBufferOffset
  111. = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable(
  112. DX12_SHADER_REGISTER_B_CONST_BUFFER, 0, 2);
  113. DX12Shader* chunkIntersectionShader
  114. = new DX12Shader(CustomChunkIntersectionShader,
  115. sizeof(CustomChunkIntersectionShader));
  116. DX12ShaderFunction* chunkIntersectionFunction
  117. = new DX12ShaderFunction("ChunkIntersection",
  118. chunkHitSignature,
  119. DX12_SHADER_FUNCTION_TYPE_INTERSECTION);
  120. chunkIntersectionShader->addFunction(chunkIntersectionFunction);
  121. pipeline->addShader(chunkIntersectionShader);
  122. DX12Shader* chunkAnyHitShader = new DX12Shader(
  123. CustomChunkAnyHitShader, sizeof(CustomChunkAnyHitShader));
  124. DX12ShaderFunction* chunkAnyHitFunction = new DX12ShaderFunction(
  125. "ChunkAnyHit",
  126. dynamic_cast<DX12ShaderSignature*>(chunkHitSignature->getThis()),
  127. DX12_SHADER_FUNCTION_TYPE_ANY_HIT);
  128. chunkAnyHitShader->addFunction(chunkAnyHitFunction);
  129. pipeline->addShader(chunkAnyHitShader);
  130. DX12Shader* chunkClosestHitShader = new DX12Shader(
  131. CustomChunkClosestHitShader, sizeof(CustomChunkClosestHitShader));
  132. DX12ShaderFunction* chunkClosestHitFunction = new DX12ShaderFunction(
  133. "ChunkClosestHit",
  134. dynamic_cast<DX12ShaderSignature*>(chunkHitSignature->getThis()),
  135. DX12_SHADER_FUNCTION_TYPE_CLOSEST_HIT);
  136. chunkClosestHitShader->addFunction(chunkClosestHitFunction);
  137. pipeline->addShader(chunkClosestHitShader);
  138. chunkHitGroup = new DX12ShaderHitGroup("ChunkHitGroup");
  139. chunkHitGroup->setIntersectionShaderFunction(chunkIntersectionFunction);
  140. chunkHitGroup->setAnyHitShaderFunction(chunkAnyHitFunction);
  141. chunkHitGroup->setClosestHitShaderFunction(chunkClosestHitFunction);
  142. chunkHitGroup->setAttributeSize(20);
  143. chunkHitGroup->setPayloadSize(28 * DEFAULT_MAX_TRANSPARENT_HITS + 4);
  144. pipeline->addHitGroup(chunkHitGroup);
  145. // simple blocks hit shaders
  146. DX12ShaderSignature* simpleBlocksHitSignature
  147. = new DX12ShaderSignature();
  148. simpleBlocksHitSignature->addRegisterUsageLinkedToDescriptorHeap(
  149. 0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP);
  150. simpleBlocksHitSignature->addRegisterUsageLinkedToDescriptorHeap(0,
  151. DX12_SHADER_REGISTER_T_SHADER_RESOURCE,
  152. 0,
  153. 1,
  154. TEXTURE_DESCRIPTOR_HEAP,
  155. 1);
  156. sbtSimpleBlocksTextureIdBufferOffset
  157. = simpleBlocksHitSignature
  158. ->addRegisterUsageLinkedToShaderBindingTable(
  159. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2);
  160. sbtSimpleBlocksIndexBufferOffset
  161. = simpleBlocksHitSignature
  162. ->addRegisterUsageLinkedToShaderBindingTable(
  163. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 3);
  164. sbtSimpleBlocksVertexDataBufferOffset
  165. = simpleBlocksHitSignature
  166. ->addRegisterUsageLinkedToShaderBindingTable(
  167. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 4);
  168. sbtSimpleBlocksIndexOffsetBufferOffset
  169. = simpleBlocksHitSignature
  170. ->addRegisterUsageLinkedToShaderBindingTable(
  171. DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 5);
  172. DX12Shader* simpleBlocksAnyHitShader
  173. = new DX12Shader(CustomSimpleBlocksAnyHitShader,
  174. sizeof(CustomSimpleBlocksAnyHitShader));
  175. DX12ShaderFunction* simpleBlocksAnyHitFunction
  176. = new DX12ShaderFunction("SimpleBlocksAnyHit",
  177. simpleBlocksHitSignature,
  178. DX12_SHADER_FUNCTION_TYPE_ANY_HIT);
  179. simpleBlocksAnyHitShader->addFunction(simpleBlocksAnyHitFunction);
  180. pipeline->addShader(simpleBlocksAnyHitShader);
  181. DX12Shader* simpleBlocksClosestHitShader
  182. = new DX12Shader(CustomSimpleBlocksClosestHitShader,
  183. sizeof(CustomSimpleBlocksClosestHitShader));
  184. DX12ShaderFunction* simpleBlocksClosestHitFunction
  185. = new DX12ShaderFunction("SimpleBlocksClosestHit",
  186. dynamic_cast<DX12ShaderSignature*>(
  187. simpleBlocksHitSignature->getThis()),
  188. DX12_SHADER_FUNCTION_TYPE_CLOSEST_HIT);
  189. simpleBlocksClosestHitShader->addFunction(
  190. simpleBlocksClosestHitFunction);
  191. pipeline->addShader(simpleBlocksClosestHitShader);
  192. simpleBlocksHitGroup = new DX12ShaderHitGroup("SimpleBlocksHitGroup");
  193. simpleBlocksHitGroup->setAnyHitShaderFunction(
  194. simpleBlocksAnyHitFunction);
  195. simpleBlocksHitGroup->setClosestHitShaderFunction(
  196. simpleBlocksClosestHitFunction);
  197. simpleBlocksHitGroup->setAttributeSize(8);
  198. simpleBlocksHitGroup->setPayloadSize(
  199. 28 * DEFAULT_MAX_TRANSPARENT_HITS + 4);
  200. pipeline->addHitGroup(simpleBlocksHitGroup);
  201. pipeline->setMaxRecursionDepth(10);
  202. }
  203. pipeline->createPipelineState(device, pfnD3D12SerializeRootSignature);
  204. }
  205. void CustomDX12API::initializeGlobalDescriptorHeap()
  206. {
  207. DirectX12::initializeGlobalDescriptorHeap();
  208. }
  209. void CustomDX12API::fillShaderBindingTable(
  210. Framework::DX12ShaderBindingTable* zShaderBindingTable,
  211. Framework::Model3D* zModel,
  212. int objectIndex,
  213. const Framework::DX12BLAS* zBLAS,
  214. int& lastHitGroupIndex)
  215. {
  216. DirectX12::fillShaderBindingTable(
  217. zShaderBindingTable, zModel, objectIndex, zBLAS, lastHitGroupIndex);
  218. }
  219. void CustomDX12API::renderWorld(Framework::World3D* zWorld,
  220. DX12TLAS* zTLAS,
  221. DX12ShaderBindingTable* zSBT,
  222. int& objectIndex)
  223. {
  224. Mat4<float> identity = Mat4<float>::identity();
  225. for (const Model3DCollection* collection : zWorld->getCollections())
  226. {
  227. const Dimension* dim = dynamic_cast<const Dimension*>(collection);
  228. if (dim)
  229. {
  230. for (const Chunk* chunk : dim->getChunks())
  231. {
  232. DX12ChunkData* chunkData = chunk->zDX12Data();
  233. chunkData->updateBuffers(device, directCommandQueue);
  234. if (chunkData->hasCustomBlocks())
  235. {
  236. bool changed
  237. = chunkData->wasCustomBufferChanged()
  238. || chunkData->getLastCustomObjectIndex() != objectIndex;
  239. if (changed)
  240. {
  241. chunkData->setLastCustomObjectIndex(objectIndex);
  242. D3D12_RAYTRACING_INSTANCE_DESC* desc
  243. = zTLAS->nextInstanceDesc();
  244. desc->InstanceID = objectIndex;
  245. desc->InstanceContributionToHitGroupIndex = objectIndex;
  246. desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
  247. desc->InstanceMask = 0xFF;
  248. desc->AccelerationStructure
  249. = chunkData->zCustomBlocksBlasBuffer()
  250. ->zBuffer()
  251. ->GetGPUVirtualAddress();
  252. Framework::Mat4<float> transform
  253. = Framework::Mat4<float>::translation(
  254. Vec3<float>((float)chunk->getCenter().x,
  255. (float)chunk->getCenter().y,
  256. 0.f));
  257. memcpy(desc->Transform, &transform, sizeof(float) * 12);
  258. }
  259. else
  260. {
  261. zTLAS->nextInstanceDesc();
  262. }
  263. int hitGroupOffset = zSBT->addHitGroup(simpleBlocksHitGroup,
  264. chunkData->getLastCustomHitGroupIndex());
  265. if (chunkData->getLastCustomHitGroupIndex()
  266. != hitGroupOffset
  267. || chunkData->wasCustomBufferChanged())
  268. {
  269. chunkData->setLastCustomHitGroupIndex(hitGroupOffset);
  270. zSBT->setHitGroupShaderInput(hitGroupOffset,
  271. sbtSimpleBlocksTextureIdBufferOffset,
  272. chunkData->zCustomBlocksTextureBuffer()
  273. ->zBuffer()
  274. ->GetGPUVirtualAddress());
  275. zSBT->setHitGroupShaderInput(hitGroupOffset,
  276. sbtSimpleBlocksIndexBufferOffset,
  277. chunkData->zCustomBlocksCombinedIndexBuffer()
  278. ->zBuffer()
  279. ->GetGPUVirtualAddress());
  280. zSBT->setHitGroupShaderInput(hitGroupOffset,
  281. sbtSimpleBlocksVertexDataBufferOffset,
  282. chunkData->zCustomBlocksVertexDataBuffer()
  283. ->zBuffer()
  284. ->GetGPUVirtualAddress());
  285. zSBT->setHitGroupShaderInput(hitGroupOffset,
  286. sbtSimpleBlocksIndexOffsetBufferOffset,
  287. chunkData->zCustomBlocksIndexOffsetBuffer()
  288. ->zBuffer()
  289. ->GetGPUVirtualAddress());
  290. }
  291. chunkData->setCustomBufferChanged(0);
  292. ++objectIndex;
  293. }
  294. bool changed = chunkData->wasBufferChanged()
  295. || chunkData->getLastObjectIndex() != objectIndex;
  296. if (changed)
  297. {
  298. chunkData->setLastObjectIndex(objectIndex);
  299. D3D12_RAYTRACING_INSTANCE_DESC* desc
  300. = zTLAS->nextInstanceDesc();
  301. desc->InstanceID = objectIndex;
  302. desc->InstanceContributionToHitGroupIndex = objectIndex;
  303. desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
  304. desc->InstanceMask = 0xFF;
  305. desc->AccelerationStructure = chunkData->zBlasBuffer()
  306. ->zBuffer()
  307. ->GetGPUVirtualAddress();
  308. memcpy(desc->Transform, &identity, sizeof(float) * 12);
  309. }
  310. else
  311. {
  312. zTLAS->nextInstanceDesc();
  313. }
  314. int hitGroupOffset = zSBT->addHitGroup(
  315. chunkHitGroup, chunkData->getLastHitGroupIndex());
  316. if (chunkData->getLastHitGroupIndex() != hitGroupOffset
  317. || chunkData->wasBufferChanged())
  318. {
  319. chunkData->setLastHitGroupIndex(hitGroupOffset);
  320. zSBT->setHitGroupShaderInput(hitGroupOffset,
  321. sbtChunkIndexBufferOffset,
  322. chunkData->zBlockIndexBuffer()
  323. ->zBuffer()
  324. ->GetGPUVirtualAddress());
  325. zSBT->setHitGroupShaderInput(hitGroupOffset,
  326. sbtChunkTextureIdBufferOffset,
  327. chunkData->zTextureIdBuffer()
  328. ->zBuffer()
  329. ->GetGPUVirtualAddress());
  330. zSBT->setHitGroupShaderInput(hitGroupOffset,
  331. sbtChunkDataBufferOffset,
  332. chunkData->zChunkInfoBuffer()
  333. ->zBuffer()
  334. ->GetGPUVirtualAddress());
  335. zSBT->setHitGroupShaderInput(hitGroupOffset,
  336. sbtChunkLightBufferOffset,
  337. chunkData->zLightBuffer()
  338. ->zBuffer()
  339. ->GetGPUVirtualAddress());
  340. }
  341. chunkData->setBufferChanged(0);
  342. ++objectIndex;
  343. }
  344. }
  345. }
  346. DirectX12::renderWorld(zWorld, zTLAS, zSBT, objectIndex);
  347. }