CustomDX12API.cpp 17 KB

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