CustomDX12API.cpp 18 KB

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