Prechádzať zdrojové kódy

add support for global root signature to pass shader arguments globally once for all shaders

Kolja Strohm 1 týždeň pred
rodič
commit
965721c63b
8 zmenil súbory, kde vykonal 85 pridanie a 98 odobranie
  1. 0 3
      AnyHit.hlsl
  2. 4 0
      Common.hlsl
  3. 55 11
      DX12GraphicsApi.cpp
  4. 1 0
      DX12GraphicsApi.h
  5. 20 77
      DX12Shader.cpp
  6. 4 3
      DX12Shader.h
  7. 0 3
      Hit.hlsl
  8. 1 1
      RayGen.hlsl

+ 0 - 3
AnyHit.hlsl

@@ -1,8 +1,5 @@
 #include "Common.hlsl"
 
-Texture2D<float4> textures[] : register(t0, space1);
-SamplerState gSampler : register(s0, space0);
-
 struct VertexData
 {
     float2 texcoord;

+ 4 - 0
Common.hlsl

@@ -20,3 +20,7 @@ struct Attributes
 {
     float2 bary;
 };
+
+// global root signature resources
+Texture2D<float4> textures[] : register(t0, space1);
+SamplerState gSampler : register(s0, space0);

+ 55 - 11
DX12GraphicsApi.cpp

@@ -271,15 +271,14 @@ void Framework::DirectX12::renderKamera(
             3, DX12_SHADER_REGISTER_T_SHADER_RESOURCE, tlas);
     }
     sbt->endUpdate(device, directCommandQueue);
-    ID3D12DescriptorHeap* heaps[] = {globalDescriptorHeap->zDescriptorHeap(),
-        samplerDescriptorHeap->zDescriptorHeap()};
-    directCommandQueue->zCommandList()->SetDescriptorHeaps(2, heaps);
 
     D3D12_DISPATCH_RAYS_DESC desc;
     sbt->fillDispatchRaysDesc(&desc);
     desc.Width = zTarget->zImage()->getWidth();
     desc.Height = zTarget->zImage()->getHeight();
     desc.Depth = 1;
+    // order of definition in DX12DescriptorHeapType must be matched
+    fillGlobalShaderParams(0);
     directCommandQueue->zCommandList()->SetPipelineState1(
         pipeline->zPipelineState());
     directCommandQueue->zCommandList()->DispatchRays(&desc);
@@ -365,6 +364,15 @@ void Framework::DirectX12::initializePipeline()
 {
     if (pipeline->getShaders().getEntryCount() == 0)
     { // add default shaders
+        pipeline->zGlobalSignature()->addRegisterUsageLinkedToDescriptorHeap(
+            0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP);
+        pipeline->zGlobalSignature()->addRegisterUsageLinkedToDescriptorHeap(0,
+            DX12_SHADER_REGISTER_T_SHADER_RESOURCE,
+            0,
+            1,
+            TEXTURE_DESCRIPTOR_HEAP,
+            1);
+
         DX12Shader* rayGenShader = new DX12Shader(
             DX12DefaultRayGenShaderBytes, sizeof(DX12DefaultRayGenShaderBytes));
         // RayGen from RayGen.hlsl
@@ -390,14 +398,6 @@ void Framework::DirectX12::initializePipeline()
         pipeline->addShader(missShader);
 
         DX12ShaderSignature* hitSignature = new DX12ShaderSignature();
-        hitSignature->addRegisterUsageLinkedToDescriptorHeap(
-            0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP);
-        hitSignature->addRegisterUsageLinkedToDescriptorHeap(0,
-            DX12_SHADER_REGISTER_T_SHADER_RESOURCE,
-            0,
-            1,
-            TEXTURE_DESCRIPTOR_HEAP,
-            1);
         sbtTextureIdBufferOffset
             = hitSignature->addRegisterUsageLinkedToShaderBindingTable(
                 DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2);
@@ -525,6 +525,50 @@ void Framework::DirectX12::fillShaderBindingTable(
     }
 }
 
+int Framework::DirectX12::fillGlobalShaderParams(int startIndex)
+{
+    ID3D12DescriptorHeap* heaps[]
+        = {pipeline->zGlobalSignature()->doesUseTextureDescriptorHeap()
+                ? textureDescriptorHeap->zDescriptorHeap()
+                : globalDescriptorHeap->zDescriptorHeap(),
+            samplerDescriptorHeap->zDescriptorHeap()};
+    directCommandQueue->zCommandList()->SetDescriptorHeaps(2, heaps);
+    if (pipeline->zGlobalSignature()->gerShaderBindingTableParamCount()
+            + pipeline->zGlobalSignature()
+                ->getDescriptorHeapBindings()
+                .getEntryCount()
+        > startIndex)
+    {
+        if (startIndex == 0)
+        {
+            directCommandQueue->zCommandList()->SetComputeRootSignature(
+                pipeline->zGlobalSignature()->zSignature());
+        }
+        if (pipeline->zGlobalSignature()->doesUseGlobalDescriptorHeap())
+        {
+            directCommandQueue->zCommandList()->SetComputeRootDescriptorTable(
+                startIndex++,
+                globalDescriptorHeap->zDescriptorHeap()
+                    ->GetGPUDescriptorHandleForHeapStart());
+        }
+        if (pipeline->zGlobalSignature()->doesUseTextureDescriptorHeap())
+        {
+            directCommandQueue->zCommandList()->SetComputeRootDescriptorTable(
+                startIndex++,
+                textureDescriptorHeap->zDescriptorHeap()
+                    ->GetGPUDescriptorHandleForHeapStart());
+        }
+        if (pipeline->zGlobalSignature()->doesUseSamplerDescriptorHeap())
+        {
+            directCommandQueue->zCommandList()->SetComputeRootDescriptorTable(
+                startIndex++,
+                samplerDescriptorHeap->zDescriptorHeap()
+                    ->GetGPUDescriptorHandleForHeapStart());
+        }
+    }
+    return startIndex;
+}
+
 typedef HRESULT(__stdcall* CreateDXGIFactory2Function)(UINT, REFIID, void**);
 
 typedef HRESULT(__stdcall* D3D12CreateDeviceFunction)(

+ 1 - 0
DX12GraphicsApi.h

@@ -124,6 +124,7 @@ namespace Framework
             int objectIndex,
             const DX12BLAS* zBLAS,
             int& lastHitGroupIndex);
+        DLLEXPORT virtual int fillGlobalShaderParams(int startIndex = 0);
 
     public:
         DLLEXPORT void initialize(NativeWindow* fenster,

+ 20 - 77
DX12Shader.cpp

@@ -8,11 +8,12 @@
 
 using namespace Framework;
 
-Framework::DX12ShaderSignature::DX12ShaderSignature()
+Framework::DX12ShaderSignature::DX12ShaderSignature(bool global)
     : ReferenceCounter(),
       signature(0),
       changed(1),
-      shaderBindingTableParamCount(0)
+      shaderBindingTableParamCount(0),
+      global(global)
 {}
 
 Framework::DX12ShaderSignature::~DX12ShaderSignature()
@@ -312,7 +313,8 @@ void Framework::DX12ShaderSignature::createSignature(ID3D12Device5* zDevice,
     D3D12_ROOT_SIGNATURE_DESC rootDesc = {};
     rootDesc.NumParameters = shaderBindingTableParamCount;
     rootDesc.pParameters = descriptorTable;
-    rootDesc.Flags = D3D12_ROOT_SIGNATURE_FLAG_LOCAL_ROOT_SIGNATURE;
+    rootDesc.Flags = global ? D3D12_ROOT_SIGNATURE_FLAG_NONE
+                            : D3D12_ROOT_SIGNATURE_FLAG_LOCAL_ROOT_SIGNATURE;
     ID3DBlob* pSigBlob = 0;
     ID3DBlob* pErrorBlob = 0;
     HRESULT hr = pfnD3D12SerializeRootSignature(
@@ -740,21 +742,16 @@ DX12ShaderSignature* Framework::DX12ShaderHitGroup::zSignature() const
 
 Framework::DX12Pipeline::DX12Pipeline()
     : ReferenceCounter(),
-      emptyGlobalRootSignature(0),
-      emptyLocalRootSignature(0),
+      globalSignature(new DX12ShaderSignature(true)),
       pipelineState(0),
       maxRecursionDepth(0)
 {}
 
 Framework::DX12Pipeline::~DX12Pipeline()
 {
-    if (emptyGlobalRootSignature)
+    if (globalSignature)
     {
-        emptyGlobalRootSignature->Release();
-    }
-    if (emptyLocalRootSignature)
-    {
-        emptyLocalRootSignature->Release();
+        globalSignature->release();
     }
     if (pipelineState)
     {
@@ -780,66 +777,9 @@ void Framework::DX12Pipeline::setMaxRecursionDepth(int maxRecursionDepth)
 void Framework::DX12Pipeline::createPipelineState(ID3D12Device5* zDevice,
     PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature)
 {
-    if (!emptyGlobalRootSignature)
-    {
-        D3D12_ROOT_SIGNATURE_DESC rootDesc = {};
-        rootDesc.NumParameters = 0;
-        rootDesc.pParameters = 0;
-        rootDesc.Flags = D3D12_ROOT_SIGNATURE_FLAG_NONE;
-        ID3DBlob* pSigBlob = 0;
-        ID3DBlob* pErrorBlob = 0;
-        HRESULT hr = pfnD3D12SerializeRootSignature(
-            &rootDesc, D3D_ROOT_SIGNATURE_VERSION_1_0, &pSigBlob, &pErrorBlob);
-        if (pSigBlob)
-        {
-            zDevice->CreateRootSignature(0,
-                pSigBlob->GetBufferPointer(),
-                pSigBlob->GetBufferSize(),
-                __uuidof(ID3D12RootSignature),
-                (void**)&emptyGlobalRootSignature);
-            pSigBlob->Release();
-        }
-        if (pErrorBlob)
-        {
-            std::string errorMessage(
-                static_cast<const char*>(pErrorBlob->GetBufferPointer()),
-                pErrorBlob->GetBufferSize());
-            Logging::error()
-                << "Failed to serialize empty root signature: " << errorMessage;
-            pErrorBlob->Release();
-        }
-    }
-    if (!emptyLocalRootSignature)
-    {
-        D3D12_ROOT_SIGNATURE_DESC rootDesc = {};
-        rootDesc.NumParameters = 0;
-        rootDesc.pParameters = 0;
-        rootDesc.Flags = D3D12_ROOT_SIGNATURE_FLAG_LOCAL_ROOT_SIGNATURE;
-        ID3DBlob* pSigBlob = 0;
-        ID3DBlob* pErrorBlob = 0;
-        HRESULT hr = pfnD3D12SerializeRootSignature(
-            &rootDesc, D3D_ROOT_SIGNATURE_VERSION_1_0, &pSigBlob, &pErrorBlob);
-        if (pSigBlob)
-        {
-            zDevice->CreateRootSignature(0,
-                pSigBlob->GetBufferPointer(),
-                pSigBlob->GetBufferSize(),
-                __uuidof(ID3D12RootSignature),
-                (void**)&emptyLocalRootSignature);
-            pSigBlob->Release();
-        }
-        if (pErrorBlob)
-        {
-            std::string errorMessage(
-                static_cast<const char*>(pErrorBlob->GetBufferPointer()),
-                pErrorBlob->GetBufferSize());
-            Logging::error()
-                << "Failed to serialize empty root signature: " << errorMessage;
-            pErrorBlob->Release();
-        }
-    }
+    globalSignature->createSignature(zDevice, pfnD3D12SerializeRootSignature);
     unsigned int subobjectCount
-        = shaders.getEntryCount() + hitGroups.getEntryCount() + 5;
+        = shaders.getEntryCount() + hitGroups.getEntryCount() + 4;
     Array<DX12ShaderSignature*> distinctSignatures;
     for (const DX12Shader* shader : shaders)
     {
@@ -1034,12 +974,10 @@ void Framework::DX12Pipeline::createPipelineState(ID3D12Device5* zDevice,
         rootSignatureIndex++;
     }
 
+    D3D12_GLOBAL_ROOT_SIGNATURE globalRootSignature = {};
+    globalRootSignature.pGlobalRootSignature = globalSignature->zSignature();
     subobjects[index].Type = D3D12_STATE_SUBOBJECT_TYPE_GLOBAL_ROOT_SIGNATURE;
-    subobjects[index].pDesc = &emptyGlobalRootSignature;
-    index++;
-
-    subobjects[index].Type = D3D12_STATE_SUBOBJECT_TYPE_LOCAL_ROOT_SIGNATURE;
-    subobjects[index].pDesc = &emptyLocalRootSignature;
+    subobjects[index].pDesc = &globalRootSignature;
     index++;
 
     D3D12_RAYTRACING_PIPELINE_CONFIG pipelineConfig = {};
@@ -1063,14 +1001,14 @@ void Framework::DX12Pipeline::createPipelineState(ID3D12Device5* zDevice,
         throw std::logic_error("Could not create the raytracing state object");
     }
 
-    /* delete[] functionAndHitGroupNames;
+    delete[] functionAndHitGroupNames;
     delete[] localRootSignatures;
     for (int i = 0; i < distinctSignatures.getEntryCount(); i++)
     {
         delete[] rootSignatureExports[i];
     }
     delete[] rootSignatureExports;
-    delete[] localRootAssociations;*/
+    delete[] localRootAssociations;
 }
 
 ID3D12StateObject* Framework::DX12Pipeline::zPipelineState() const
@@ -1106,6 +1044,11 @@ const RCArray<DX12ShaderHitGroup>& Framework::DX12Pipeline::getHitGroups() const
     return hitGroups;
 }
 
+DX12ShaderSignature* Framework::DX12Pipeline::zGlobalSignature() const
+{
+    return globalSignature;
+}
+
 Framework::DX12GlobalDescriptorHeap::DX12GlobalDescriptorHeap(
     DX12Pipeline* pipeline, DX12DescriptorHeapType type)
     : ReferenceCounter(),

+ 4 - 3
DX12Shader.h

@@ -57,9 +57,10 @@ namespace Framework
         bool useGlobalDescriptorHeap;
         bool useTextureDescriptorHeap;
         bool useSamplerDescriptorHeap;
+        bool global;
 
     public:
-        DLLEXPORT DX12ShaderSignature();
+        DLLEXPORT DX12ShaderSignature(bool global = false);
         DLLEXPORT ~DX12ShaderSignature();
         /**
          * for each datastructure with : register(...) in the shader code,
@@ -214,8 +215,7 @@ namespace Framework
         RCArray<DX12Shader> shaders;
         RCArray<DX12ShaderHitGroup> hitGroups;
         Array<const DX12ShaderFunction*> functionsWithoutHitGroups;
-        ID3D12RootSignature* emptyGlobalRootSignature;
-        ID3D12RootSignature* emptyLocalRootSignature;
+        DX12ShaderSignature* globalSignature;
         ID3D12StateObject* pipelineState;
         int maxRecursionDepth;
 
@@ -235,6 +235,7 @@ namespace Framework
         DLLEXPORT const Array<const DX12ShaderFunction*>&
         getFunctionsWithoutHitGroups() const;
         DLLEXPORT const RCArray<DX12ShaderHitGroup>& getHitGroups() const;
+        DLLEXPORT DX12ShaderSignature* zGlobalSignature() const;
     }; // namespace Framework
 
     struct DX12ShaderRegisterInput

+ 0 - 3
Hit.hlsl

@@ -1,8 +1,5 @@
 #include "Common.hlsl"
 
-Texture2D<float4> textures[] : register(t0, space1);
-SamplerState gSampler : register(s0, space0);
-
 struct VertexData
 {
     float2 texcoord;

+ 1 - 1
RayGen.hlsl

@@ -6,7 +6,7 @@ RWTexture2D<float4> gOutput : register(u0);
 RWTexture2D<float4> guiTexture : register(u1);
 
 // Raytracing acceleration structure, accessed as a SRV
-RaytracingAccelerationStructure TLAS : register(t0);
+RaytracingAccelerationStructure TLAS : register(t0, space0);
 
 cbuffer RayGenerationSettings : register(b0)
 {