CustomDX12API.cpp 19 KB

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