CustomDX12API.cpp 16 KB

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