DX12Shader.h 5.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151
  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. enum DX12ShaderRegister
  12. {
  13. DX12_SHADER_REGISTER_B_CONST_BUFFER = 0,
  14. DX12_SHADER_REGISTER_T_SHADER_RESOURCE = 1,
  15. DX12_SHADER_REGISTER_U_UNORDERED_ACCESS = 2,
  16. // TODO: Do we need Sampler?
  17. };
  18. struct DX12ShaderRegisterUsage
  19. {
  20. DX12ShaderRegister registerType;
  21. int registerIndex;
  22. int spaceIndex;
  23. };
  24. class DX12ShaderSignature : public ReferenceCounter
  25. {
  26. private:
  27. ID3D12RootSignature* signature;
  28. Array<DX12ShaderRegisterUsage> registerUsages;
  29. bool changed;
  30. public:
  31. DX12ShaderSignature();
  32. ~DX12ShaderSignature();
  33. /**
  34. * needs to be called for each datastructure with : register(...)
  35. *
  36. * \param registerType the register type e.g.
  37. * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0)
  38. * \param registerIndex the register index e.g. 1 for : register(b1)
  39. */
  40. void addRegisterUsage(
  41. DX12ShaderRegister registerType, int registerIndex);
  42. /**
  43. * needs to be called for each datastructure with : register(...)
  44. *
  45. * \param registerType the register type e.g.
  46. * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0, space1)
  47. * \param registerIndex the register index e.g. 1 for : register(b1,
  48. * space2)
  49. * \param spaceIndex the space index e.g. 3 for : register(b1, space3)
  50. */
  51. void addRegisterUsage(
  52. DX12ShaderRegister registerType, int registerIndex, int spaceIndex);
  53. /**
  54. * Creates the root signature.
  55. */
  56. void createSignature(ID3D12Device5* zDevice,
  57. PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature);
  58. ID3D12RootSignature* zSignature() const;
  59. };
  60. class DX12ShaderFunction : public ReferenceCounter
  61. {
  62. private:
  63. Text functionName;
  64. DX12ShaderSignature* signature;
  65. D3D12_EXPORT_DESC* exportDesc;
  66. public:
  67. DX12ShaderFunction(
  68. const Text& functionName, DX12ShaderSignature* signature);
  69. ~DX12ShaderFunction();
  70. const Text& getFunctionName() const;
  71. DX12ShaderSignature* zSignature() const;
  72. D3D12_EXPORT_DESC* zExportDesc() const;
  73. };
  74. class DX12Shader : public ReferenceCounter
  75. {
  76. private:
  77. RCArray<DX12ShaderFunction> functions;
  78. const char* shaderBytes;
  79. int shaderBytesSize;
  80. D3D12_DXIL_LIBRARY_DESC* libraryDesc;
  81. public:
  82. DX12Shader(const char* shaderBytes, int shaderBytesSize);
  83. ~DX12Shader();
  84. void addFunction(DX12ShaderFunction* function);
  85. int getShaderBytesSize() const;
  86. const char* getShaderBytes() const;
  87. const RCArray<DX12ShaderFunction>& getFunctions() const;
  88. D3D12_DXIL_LIBRARY_DESC* zLibraryDesc() const;
  89. };
  90. class DX12ShaderHitGroup : public ReferenceCounter
  91. {
  92. private:
  93. Text name;
  94. DX12ShaderFunction* closestHitShaderFunction;
  95. DX12ShaderFunction* anyHitShaderFunction;
  96. DX12ShaderFunction* intersectionShaderFunction;
  97. int payloadSize;
  98. int attributeSize;
  99. D3D12_HIT_GROUP_DESC* hitGroupDesc;
  100. public:
  101. DX12ShaderHitGroup(const Text name);
  102. ~DX12ShaderHitGroup();
  103. void setClosestHitShaderFunction(
  104. DX12ShaderFunction* closestHitShaderFunction);
  105. void setAnyHitShaderFunction(DX12ShaderFunction* anyHitShaderFunction);
  106. void setIntersectionShaderFunction(
  107. DX12ShaderFunction* intersectionShaderFunction);
  108. void setPayloadSize(int payloadSize);
  109. void setAttributeSize(int attributeSize);
  110. const Text& getName() const;
  111. DX12ShaderFunction* zClosestHitShaderFunction() const;
  112. DX12ShaderFunction* zAnyHitShaderFunction() const;
  113. DX12ShaderFunction* zIntersectionShaderFunction() const;
  114. int getPayloadSize() const;
  115. int getAttributeSize() const;
  116. D3D12_HIT_GROUP_DESC* zHitGroupDesc() const;
  117. };
  118. class DX12Pipeline : public ReferenceCounter
  119. {
  120. private:
  121. RCArray<DX12Shader> shaders;
  122. RCArray<DX12ShaderHitGroup> hitGroups;
  123. ID3D12RootSignature* emptyGlobalRootSignature;
  124. ID3D12RootSignature* emptyLocalRootSignature;
  125. ID3D12StateObject* pipelineState;
  126. int maxRecursionDepth;
  127. public:
  128. DX12Pipeline();
  129. ~DX12Pipeline();
  130. void addShader(DX12Shader* shader);
  131. void addHitGroup(DX12ShaderHitGroup* hitGroup);
  132. void setMaxRecursionDepth(int maxRecursionDepth);
  133. void createPipelineState(ID3D12Device5* zDevice,
  134. PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature);
  135. ID3D12StateObject* zPipelineState() const;
  136. }; // namespace Framework
  137. } // namespace Framework