DX12Shader.h 5.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188
  1. #pragma once
  2. #include "Array.h"
  3. #include "DX12Buffer.h"
  4. struct ID3D12Device5;
  5. struct ID3D12GraphicsCommandList;
  6. struct D3D12_INPUT_ELEMENT_DESC;
  7. struct D3D12_ROOT_PARAMETER1;
  8. struct D3D12_CONSTANT_BUFFER_VIEW_DESC;
  9. namespace Framework
  10. {
  11. class Texture;
  12. class DX12TLAS;
  13. enum DX12ShaderRegister
  14. {
  15. DX12_SHADER_REGISTER_B_CONST_BUFFER = 0,
  16. DX12_SHADER_REGISTER_T_SHADER_RESOURCE = 1,
  17. DX12_SHADER_REGISTER_U_UNORDERED_ACCESS = 2,
  18. // TODO: Do we need Sampler?
  19. };
  20. struct DX12ShaderRegisterUsage
  21. {
  22. DX12ShaderRegister registerType;
  23. int registerIndex;
  24. int spaceIndex;
  25. };
  26. class ShaderHeap;
  27. class DX12ShaderSignature : public ReferenceCounter
  28. {
  29. private:
  30. ID3D12RootSignature* signature;
  31. Array<DX12ShaderRegisterUsage> registerUsages;
  32. bool changed;
  33. public:
  34. DX12ShaderSignature();
  35. ~DX12ShaderSignature();
  36. /**
  37. * needs to be called for each datastructure with : register(...)
  38. *
  39. * \param registerType the register type e.g.
  40. * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0)
  41. * \param registerIndex the register index e.g. 1 for : register(b1)
  42. * \param spaceIndex the optional space index e.g. 3 for : register(b1,
  43. * space3)
  44. */
  45. void addRegisterUsage(DX12ShaderRegister registerType,
  46. int registerIndex,
  47. int spaceIndex = 0);
  48. /**
  49. * Creates the root signature.
  50. */
  51. void createSignature(ID3D12Device5* zDevice,
  52. PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature);
  53. ID3D12RootSignature* zSignature() const;
  54. const Array<DX12ShaderRegisterUsage>& getRegisterUsagesOrder() const;
  55. ShaderHeap* createShaderHeap();
  56. };
  57. class DX12ShaderFunction : public ReferenceCounter
  58. {
  59. private:
  60. Text functionName;
  61. DX12ShaderSignature* signature;
  62. D3D12_EXPORT_DESC* exportDesc;
  63. public:
  64. DX12ShaderFunction(
  65. const Text& functionName, DX12ShaderSignature* signature);
  66. ~DX12ShaderFunction();
  67. const Text& getFunctionName() const;
  68. DX12ShaderSignature* zSignature() const;
  69. D3D12_EXPORT_DESC* zExportDesc() const;
  70. };
  71. class DX12Shader : public ReferenceCounter
  72. {
  73. private:
  74. RCArray<DX12ShaderFunction> functions;
  75. const char* shaderBytes;
  76. int shaderBytesSize;
  77. D3D12_DXIL_LIBRARY_DESC* libraryDesc;
  78. public:
  79. DX12Shader(const char* shaderBytes, int shaderBytesSize);
  80. ~DX12Shader();
  81. void addFunction(DX12ShaderFunction* function);
  82. int getShaderBytesSize() const;
  83. const char* getShaderBytes() const;
  84. const RCArray<DX12ShaderFunction>& getFunctions() const;
  85. D3D12_DXIL_LIBRARY_DESC* zLibraryDesc() const;
  86. };
  87. class DX12ShaderHitGroup : public ReferenceCounter
  88. {
  89. private:
  90. Text name;
  91. DX12ShaderFunction* closestHitShaderFunction;
  92. DX12ShaderFunction* anyHitShaderFunction;
  93. DX12ShaderFunction* intersectionShaderFunction;
  94. int payloadSize;
  95. int attributeSize;
  96. D3D12_HIT_GROUP_DESC* hitGroupDesc;
  97. public:
  98. DX12ShaderHitGroup(const Text name);
  99. ~DX12ShaderHitGroup();
  100. void setClosestHitShaderFunction(
  101. DX12ShaderFunction* closestHitShaderFunction);
  102. void setAnyHitShaderFunction(DX12ShaderFunction* anyHitShaderFunction);
  103. void setIntersectionShaderFunction(
  104. DX12ShaderFunction* intersectionShaderFunction);
  105. void setPayloadSize(int payloadSize);
  106. void setAttributeSize(int attributeSize);
  107. const Text& getName() const;
  108. DX12ShaderFunction* zClosestHitShaderFunction() const;
  109. DX12ShaderFunction* zAnyHitShaderFunction() const;
  110. DX12ShaderFunction* zIntersectionShaderFunction() const;
  111. int getPayloadSize() const;
  112. int getAttributeSize() const;
  113. D3D12_HIT_GROUP_DESC* zHitGroupDesc() const;
  114. };
  115. class DX12Pipeline : public ReferenceCounter
  116. {
  117. private:
  118. RCArray<DX12Shader> shaders;
  119. RCArray<DX12ShaderHitGroup> hitGroups;
  120. ID3D12RootSignature* emptyGlobalRootSignature;
  121. ID3D12RootSignature* emptyLocalRootSignature;
  122. ID3D12StateObject* pipelineState;
  123. int maxRecursionDepth;
  124. public:
  125. DX12Pipeline();
  126. ~DX12Pipeline();
  127. void addShader(DX12Shader* shader);
  128. void addHitGroup(DX12ShaderHitGroup* hitGroup);
  129. void setMaxRecursionDepth(int maxRecursionDepth);
  130. void createPipelineState(ID3D12Device5* zDevice,
  131. PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature);
  132. ID3D12StateObject* zPipelineState() const;
  133. }; // namespace Framework
  134. struct DX12ShaderRegisterInput
  135. {
  136. DX12ShaderRegister registerType;
  137. int registerIndex;
  138. int spaceIndex;
  139. ReferenceCounter*
  140. inputResource; // Can be Texture*, DXBuffer*, or DX12TLAS*
  141. };
  142. class ShaderHeap : public ReferenceCounter
  143. {
  144. private:
  145. DX12ShaderSignature* signature;
  146. ID3D12DescriptorHeap* descriptorHeap;
  147. Array<DX12ShaderRegisterInput*> registerInputs;
  148. int lastDescriptorHeapSize;
  149. public:
  150. ShaderHeap(DX12ShaderSignature* signature);
  151. ~ShaderHeap();
  152. private:
  153. void setRegisterInput(DX12ShaderRegister registerType,
  154. int registerIndex,
  155. int spaceIndex,
  156. ReferenceCounter* inputResource);
  157. public:
  158. void setRegisterInput(
  159. Texture* zTexture, int registerIndex, int spaceIndex = 0);
  160. void setRegisterInput(
  161. DXBuffer* zBuffer, int registerIndex, int spaceIndex = 0);
  162. void setRegisterInput(
  163. DX12TLAS* zTLAS, int registerIndex, int spaceIndex = 0);
  164. void updateDescriptorHeap(ID3D12Device5* zDevice);
  165. ID3D12DescriptorHeap* zDescriptorHeap() const;
  166. };
  167. } // namespace Framework