Ver código fonte

fix problems with descriptor heap creation

Kolja Strohm 4 semanas atrás
pai
commit
2ad1e21403
5 arquivos alterados com 443 adições e 267 exclusões
  1. 102 59
      DX12GraphicsApi.cpp
  2. 11 3
      DX12GraphicsApi.h
  3. 255 179
      DX12Shader.cpp
  4. 67 21
      DX12Shader.h
  5. 8 5
      Framework.vcxproj

+ 102 - 59
DX12GraphicsApi.cpp

@@ -10,6 +10,7 @@
 #include "DLLRegister.h"
 #include "DX12BLASModel.h"
 #include "DX12CommandQueue.h"
+#include "DX12Shader.h"
 #include "DX12Texture.h"
 #include "DX12TLAS.h"
 #include "Globals.h"
@@ -17,7 +18,6 @@
 #include "Model3D.h"
 #include "Model3DList.h"
 #include "Screen.h"
-#include "Shader.h"
 #include "TextureList.h"
 #include "TextureModel.h"
 #include "Window.h"
@@ -48,6 +48,7 @@ DirectX12::DirectX12()
       texturRegister(new TextureList()),
       blasModels(0),
       worldTLAS(0),
+      worldShaderBindingTables(0),
       lastTLASId(-1),
       lastModelId(-1),
       defaultRenderTarget(0)
@@ -74,6 +75,15 @@ DirectX12::~DirectX12()
         }
         delete[] worldTLAS;
     }
+    if (worldShaderBindingTables)
+    {
+        for (int i = 0; i <= lastTLASId; i++)
+        {
+            if (worldShaderBindingTables[i])
+                worldShaderBindingTables[i]->release();
+        }
+        delete[] worldShaderBindingTables;
+    }
     if (directCommandQueue)
     {
         directCommandQueue->flush();
@@ -142,6 +152,81 @@ void DirectX12::updateBottomLevelAccelerationStructure()
     }
 }
 
+void Framework::DirectX12::renderKamera(
+    Cam3D* zKamera, DX12Texture* zTarget, bool guiVisible)
+{
+    // TODO
+    // directCommandQueue->getCommandList()->RSSetViewports(
+    //     1, (D3D12_VIEWPORT*)zKamera->zViewPort());
+    Mat4<float> identity = Mat4<float>::identity();
+
+    World3D* w = zKamera->zWorld();
+
+    if (w->getId() < 0)
+    {
+        w->setId(++lastTLASId);
+        DX12TLAS** tmp = new DX12TLAS*[lastTLASId + 1];
+        if (lastTLASId)
+        {
+            memcpy(tmp, worldTLAS, sizeof(DX12TLAS*) * lastTLASId);
+        }
+        tmp[lastTLASId] = new DX12TLAS(device, directCommandQueue);
+        delete[] worldTLAS;
+        worldTLAS = tmp;
+        DX12ShaderBindingTable** tmpsbts
+            = new DX12ShaderBindingTable*[lastTLASId + 1];
+        if (lastTLASId)
+        {
+            memcpy(tmpsbts,
+                worldShaderBindingTables,
+                sizeof(DX12ShaderBindingTable*) * lastTLASId);
+        }
+        tmpsbts[lastTLASId] = pipeline->createShaderBindingTable();
+        delete[] worldShaderBindingTables;
+        worldShaderBindingTables = tmpsbts;
+    }
+    else if (worldShaderBindingTables[w->getId()]->zPipeline() != pipeline)
+    {
+        worldShaderBindingTables[w->getId()]->release();
+        worldShaderBindingTables[w->getId()]
+            = pipeline->createShaderBindingTable();
+    }
+    DX12TLAS* tlas = worldTLAS[w->getId()];
+    DX12ShaderBindingTable* sbt = worldShaderBindingTables[w->getId()];
+    tlas->startUpdate();
+    sbt->startUpdate();
+    int objectIndex = 0;
+    int instanceIndex = 0;
+    w->render(
+        [this, &tlas, &objectIndex, &instanceIndex, &identity](Model3D* obj) {
+            obj->calculateMatrices(identity, matrixBuffer);
+            int modelId = obj->zModelData()->getId();
+            DX12BLASModel* blasModel = blasModels[modelId];
+            ArrayIterator<int> boneIds = blasModel->zBoneIds()->begin();
+            for (const DX12BLAS* blas : *blasModel->zBLAS())
+            {
+                D3D12_RAYTRACING_INSTANCE_DESC* desc = tlas->nextInstanceDesc();
+                desc->InstanceID = objectIndex;
+                desc->InstanceContributionToHitGroupIndex = objectIndex;
+                desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
+                desc->InstanceMask = 0xFF;
+                desc->AccelerationStructure
+                    = blas->zResultBuffer()->zBuffer()->GetGPUVirtualAddress();
+                memcpy(desc->Transform,
+                    &matrixBuffer[boneIds.val()],
+                    sizeof(float) * 12);
+                // TODO: sbt->setHitGroupShaderInputs(instanceIndex, ...)
+                boneIds++;
+                instanceIndex++;
+            }
+            objectIndex++;
+        });
+    tlas->endUpdate();
+    // TODO: setup ray gen and miss params
+    sbt->endUpdate(device);
+    // TODO: call ray tracing
+}
+
 typedef HRESULT(__stdcall* CreateDXGIFactory2Function)(UINT, REFIID, void**);
 
 typedef HRESULT(__stdcall* D3D12CreateDeviceFunction)(
@@ -735,60 +820,12 @@ void DirectX12::beginFrame(bool fill2D, bool fill3D, int fillColor)
 
 void DirectX12::renderKamera(Cam3D* zKamera)
 {
-    // TODO
-    // directCommandQueue->getCommandList()->RSSetViewports(
-    //     1, (D3D12_VIEWPORT*)zKamera->zViewPort());
-    Mat4<float> identity = Mat4<float>::identity();
-
-    World3D* w = zKamera->zWorld();
-
-    if (w->getId() < 0)
-    {
-        w->setId(++lastTLASId);
-        DX12TLAS** tmp = new DX12TLAS*[lastTLASId + 1];
-        if (lastTLASId)
-        {
-            memcpy(tmp, worldTLAS, sizeof(DX12TLAS*) * lastTLASId);
-        }
-        tmp[lastTLASId] = new DX12TLAS(device, directCommandQueue);
-        delete[] worldTLAS;
-        worldTLAS = tmp;
-    }
-    DX12TLAS* tlas = worldTLAS[w->getId()];
-    tlas->startUpdate();
-    int objectIndex = 0;
-    int instanceIndex = 0;
-    w->render(
-        [this, &tlas, &objectIndex, &instanceIndex, &identity](Model3D* obj) {
-            obj->calculateMatrices(identity, matrixBuffer);
-            int modelId = obj->zModelData()->getId();
-            DX12BLASModel* blasModel = blasModels[modelId];
-            ArrayIterator<int> boneIds = blasModel->zBoneIds()->begin();
-            for (const DX12BLAS* blas : *blasModel->zBLAS())
-            {
-                D3D12_RAYTRACING_INSTANCE_DESC* desc = tlas->nextInstanceDesc();
-                desc->InstanceID = objectIndex;
-                desc->InstanceContributionToHitGroupIndex = objectIndex;
-                desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
-                desc->InstanceMask = 0xFF;
-                desc->AccelerationStructure
-                    = blas->zResultBuffer()->zBuffer()->GetGPUVirtualAddress();
-                memcpy(desc->Transform,
-                    &matrixBuffer[boneIds.val()],
-                    sizeof(float) * 12);
-                boneIds++;
-                instanceIndex++;
-            }
-            objectIndex++;
-        });
-    tlas->endUpdate();
-
-    // TODO: call ray tracing
+    renderKamera(zKamera, defaultRenderTarget, true);
 }
 
 void Framework::DirectX12::renderKamera(Cam3D* zKamera, Texture* zTarget)
 {
-    // TODO: implement rendering to a texture target
+    renderKamera(zKamera, dynamic_cast<DX12Texture*>(zTarget), false);
 }
 
 void DirectX12::presentFrame()
@@ -847,6 +884,20 @@ Image* DirectX12::zUIRenderImage() const
     return uiTexture ? uiTexture->zImage() : 0;
 }
 
+DXBuffer* DirectX12::createStructuredBuffer(int eSize)
+{
+    return new DX12Buffer(eSize, device, D3D12_RESOURCE_FLAG_NONE);
+}
+
+void Framework::DirectX12::setPipeline(DX12Pipeline* pipeline)
+{
+    if (this->pipeline)
+    {
+        this->pipeline->release();
+    }
+    this->pipeline = pipeline;
+}
+
 bool DirectX12::isAvailable()
 {
     HINSTANCE dxgiDLL = getDLLRegister()->loadDLL("dxgi.dll", "dxgi.dll");
@@ -929,12 +980,4 @@ bool DirectX12::isAvailable()
     getDLLRegister()->releaseDLL("dxgi.dll");
     getDLLRegister()->releaseDLL("d3d12.dll");
     return 0;
-}
-
-DXBuffer* DirectX12::createStructuredBuffer(int eSize)
-{
-    throw "Sorry, support for structured buffers will come eventualy"; // TODO:
-                                                                       // support
-                                                                       // structured
-                                                                       // buffers
 }

+ 11 - 3
DX12GraphicsApi.h

@@ -37,6 +37,7 @@ namespace Framework
     class TextureModel;
     class DX12TLAS;
     class DX12Texture;
+    class DX12Pipeline;
 
     class DirectX12 : public GraphicsApi
     {
@@ -66,15 +67,22 @@ namespace Framework
         TextureList* texturRegister;
         DX12BLASModel** blasModels;
         DX12TLAS** worldTLAS;
+        DX12ShaderBindingTable** worldShaderBindingTables;
         int lastTLASId;
         int lastModelId;
         DX12Texture* defaultRenderTarget;
-
-        DLLEXPORT void updateBottomLevelAccelerationStructure();
+        DX12Pipeline* pipeline;
 
     public:
         DLLEXPORT DirectX12();
         DLLEXPORT ~DirectX12();
+
+    private:
+        DLLEXPORT void updateBottomLevelAccelerationStructure();
+        DLLEXPORT void renderKamera(
+            Cam3D* zKamera, DX12Texture* zTarget, bool guiVisible);
+
+    public:
         DLLEXPORT void initialize(NativeWindow* fenster,
             Vec2<int> backBufferSize,
             bool fullScreen) override;
@@ -90,7 +98,7 @@ namespace Framework
             const char* name, Image* b, TextureDirection dir) override;
         DLLEXPORT Image* zUIRenderImage() const override;
         DLLEXPORT virtual DXBuffer* createStructuredBuffer(int eSize) override;
-
+        DLLEXPORT void setPipeline(DX12Pipeline* pipeline);
         DLLEXPORT static bool isAvailable();
     };
 } // namespace Framework

+ 255 - 179
DX12Shader.cpp

@@ -20,8 +20,18 @@ Framework::DX12ShaderSignature::~DX12ShaderSignature()
     }
 }
 
-void Framework::DX12ShaderSignature::addRegisterUsage(
+void Framework::DX12ShaderSignature::addRegisterUsageLinkedToShaderBindingTable(
     DX12ShaderRegister registerType, int registerIndex, int spaceIndex)
+{
+    addRegisterUsageLinkedToDescriptorHeap(
+        -1, registerType, registerIndex, spaceIndex);
+}
+
+void Framework::DX12ShaderSignature::addRegisterUsageLinkedToDescriptorHeap(
+    int descriptorHeapIndex,
+    DX12ShaderRegister registerType,
+    int registerIndex,
+    int spaceIndex)
 {
     // register usages sould be sorted by registerType -> spaceIndex ->
     // registerIndex
@@ -55,11 +65,20 @@ void Framework::DX12ShaderSignature::addRegisterUsage(
     }
     if (found)
     {
-        it.addBefore({registerType, registerIndex, spaceIndex});
+        if (it.val().registerIndex == registerIndex)
+        {
+            Logging::error()
+                << "Duplicate register usage in root signature: "
+                << registerType << " " << registerIndex << " " << spaceIndex;
+            return;
+        }
+        it.addBefore(
+            {registerType, registerIndex, spaceIndex, descriptorHeapIndex});
     }
     else
     {
-        registerUsages.add({registerType, registerIndex, spaceIndex});
+        registerUsages.add(
+            {registerType, registerIndex, spaceIndex, descriptorHeapIndex});
     }
     changed = 1;
 }
@@ -77,56 +96,88 @@ void Framework::DX12ShaderSignature::createSignature(ID3D12Device5* zDevice,
         signature->Release();
         signature = 0;
     }
-    D3D12_ROOT_PARAMETER descriptorTable;
-    descriptorTable.ParameterType = D3D12_ROOT_PARAMETER_TYPE_DESCRIPTOR_TABLE;
-    descriptorTable.ShaderVisibility = D3D12_SHADER_VISIBILITY_ALL;
-    descriptorTable.DescriptorTable.NumDescriptorRanges
-        = registerUsages.getEntryCount();
+    D3D12_ROOT_PARAMETER* descriptorTable
+        = new D3D12_ROOT_PARAMETER[registerUsages.getEntryCount()];
     D3D12_DESCRIPTOR_RANGE* descriptorRanges
         = new D3D12_DESCRIPTOR_RANGE[registerUsages.getEntryCount()];
-    int index = 0;
     ArrayIterator<DX12ShaderRegisterUsage> it = registerUsages.begin();
+    int index = 0;
     while (it)
     {
         const auto& usage = it.val();
-        D3D12_DESCRIPTOR_RANGE* range = &descriptorRanges[index];
-        switch (usage.registerType)
+        if (usage.descriptorHeapIndex >= 0)
         {
-        case DX12_SHADER_REGISTER_B_CONST_BUFFER:
-            range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_CBV;
-            break;
-        case DX12_SHADER_REGISTER_T_SHADER_RESOURCE:
-            range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_SRV;
-            break;
-        case DX12_SHADER_REGISTER_U_UNORDERED_ACCESS:
-            range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_UAV;
-            break;
-        default:
-            Logging::error() << "Unknown register type for root signature: "
-                             << usage.registerType;
-            range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_SRV;
+            descriptorTable[index].ParameterType
+                = D3D12_ROOT_PARAMETER_TYPE_DESCRIPTOR_TABLE;
+            descriptorTable[index].ShaderVisibility
+                = D3D12_SHADER_VISIBILITY_ALL;
+            descriptorTable[index].DescriptorTable.NumDescriptorRanges
+                = registerUsages.getEntryCount();
+            D3D12_DESCRIPTOR_RANGE* range = &descriptorRanges[index];
+            switch (usage.registerType)
+            {
+            case DX12_SHADER_REGISTER_B_CONST_BUFFER:
+                range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_CBV;
+                break;
+            case DX12_SHADER_REGISTER_T_SHADER_RESOURCE:
+                range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_SRV;
+                break;
+            case DX12_SHADER_REGISTER_U_UNORDERED_ACCESS:
+                range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_UAV;
+                break;
+            default:
+                Logging::error() << "Unknown register type for root signature: "
+                                 << usage.registerType;
+                range->RangeType = D3D12_DESCRIPTOR_RANGE_TYPE_SRV;
+            }
+            ArrayIterator<DX12ShaderRegisterUsage> next = it.next();
+            int size = 1;
+            while (next && next.val().registerType == usage.registerType
+                   && next.val().spaceIndex == usage.spaceIndex
+                   && next.val().registerIndex == usage.registerIndex + size
+                   && next.val().descriptorHeapIndex
+                          == usage.descriptorHeapIndex + size)
+            {
+                ++size;
+                it = next;
+                ++next;
+            }
+            range->NumDescriptors = size;
+            range->BaseShaderRegister = usage.registerIndex;
+            range->RegisterSpace = usage.spaceIndex;
+            range->OffsetInDescriptorsFromTableStart
+                = usage.descriptorHeapIndex;
+            descriptorTable[index].DescriptorTable.pDescriptorRanges = range;
         }
-        ArrayIterator<DX12ShaderRegisterUsage> next = it.next();
-        int size = 1;
-        while (next && next.val().registerType == usage.registerType
-               && next.val().spaceIndex == usage.spaceIndex
-               && next.val().registerIndex == usage.registerIndex + size)
+        else
         {
-            ++size;
-            it = next;
-            ++next;
+            switch (usage.registerType)
+            {
+            case DX12_SHADER_REGISTER_B_CONST_BUFFER:
+                descriptorTable[index].ParameterType
+                    = D3D12_ROOT_PARAMETER_TYPE_CBV;
+                break;
+            case DX12_SHADER_REGISTER_T_SHADER_RESOURCE:
+                descriptorTable[index].ParameterType
+                    = D3D12_ROOT_PARAMETER_TYPE_SRV;
+                break;
+            case DX12_SHADER_REGISTER_U_UNORDERED_ACCESS:
+                descriptorTable[index].ParameterType
+                    = D3D12_ROOT_PARAMETER_TYPE_UAV;
+                break;
+            }
+            descriptorTable[index].ShaderVisibility
+                = D3D12_SHADER_VISIBILITY_ALL;
+            descriptorTable[index].Descriptor.ShaderRegister
+                = usage.registerIndex;
+            descriptorTable[index].Descriptor.RegisterSpace = usage.spaceIndex;
         }
-        range->NumDescriptors = size;
-        range->BaseShaderRegister = usage.registerIndex;
-        range->RegisterSpace = usage.spaceIndex;
-        range->OffsetInDescriptorsFromTableStart = index;
         ++it;
-        index++;
+        ++index;
     }
-    descriptorTable.DescriptorTable.pDescriptorRanges = descriptorRanges;
     D3D12_ROOT_SIGNATURE_DESC rootDesc = {};
-    rootDesc.NumParameters = 1;
-    rootDesc.pParameters = &descriptorTable;
+    rootDesc.NumParameters = index;
+    rootDesc.pParameters = descriptorTable;
     rootDesc.Flags = D3D12_ROOT_SIGNATURE_FLAG_LOCAL_ROOT_SIGNATURE;
     ID3DBlob* pSigBlob = 0;
     ID3DBlob* pErrorBlob = 0;
@@ -159,16 +210,11 @@ ID3D12RootSignature* Framework::DX12ShaderSignature::zSignature() const
 }
 
 const Array<DX12ShaderRegisterUsage>&
-Framework::DX12ShaderSignature::getRegisterUsagesOrder() const
+Framework::DX12ShaderSignature::getRegisterUsages() const
 {
     return registerUsages;
 }
 
-ShaderHeap* Framework::DX12ShaderSignature::createShaderHeap()
-{
-    return new ShaderHeap(dynamic_cast<DX12ShaderSignature*>(getThis()));
-}
-
 Framework::DX12ShaderFunction::DX12ShaderFunction(
     const Text& functionName, DX12ShaderSignature* signature)
     : ReferenceCounter(),
@@ -832,16 +878,26 @@ ID3D12StateObject* Framework::DX12Pipeline::zPipelineState() const
     return pipelineState;
 }
 
-Framework::ShaderHeap::ShaderHeap(DX12ShaderSignature* signature)
+DX12ShaderBindingTable* Framework::DX12Pipeline::createShaderBindingTable()
+{
+    return new DX12ShaderBindingTable(dynamic_cast<DX12Pipeline*>(getThis()));
+}
+
+DX12GlobalDescriptorHeap* Framework::DX12Pipeline::createGlobalDescriptorHeap()
+{
+    return new DX12GlobalDescriptorHeap(dynamic_cast<DX12Pipeline*>(getThis()));
+}
+
+Framework::DX12GlobalDescriptorHeap::DX12GlobalDescriptorHeap(
+    DX12Pipeline* pipeline)
     : ReferenceCounter(),
-      signature(signature),
       descriptorHeap(0),
       lastDescriptorHeapSize(0)
 {}
 
-Framework::ShaderHeap::~ShaderHeap()
+Framework::DX12GlobalDescriptorHeap::~DX12GlobalDescriptorHeap()
 {
-    signature->release();
+    pipeline->release();
     if (descriptorHeap)
     {
         descriptorHeap->Release();
@@ -853,182 +909,202 @@ Framework::ShaderHeap::~ShaderHeap()
     }
 }
 
-void Framework::ShaderHeap::setRegisterInput(DX12ShaderRegister registerType,
-    int registerIndex,
-    int spaceIndex,
-    ReferenceCounter* inputResource)
+void Framework::DX12GlobalDescriptorHeap::addInput(
+    DX12ShaderRegister type, ReferenceCounter* inputResource)
 {
-    for (DX12ShaderRegisterInput* input : registerInputs)
+    bool found = 0;
+    for (DX12Shader* shader : pipeline->getShaders())
     {
-        if (input->registerIndex == registerIndex
-            && input->spaceIndex == spaceIndex
-            && input->registerType == registerType)
+        for (DX12ShaderFunction* function : shader->getFunctions())
         {
-            input->inputResource->release();
-            input->inputResource = inputResource->getThis();
-            return;
+            for (const DX12ShaderRegisterUsage& usage :
+                function->zSignature()->getRegisterUsages())
+            {
+                if (usage.descriptorHeapIndex == registerInputs.getEntryCount())
+                {
+                    if (usage.registerType != type)
+                    {
+                        Logging::error()
+                            << "Register type mismatch for register index "
+                            << usage.registerIndex << ", space index "
+                            << usage.spaceIndex << ". Expected register type: "
+                            << usage.registerType
+                            << ", given register type: " << type
+                            << ". The register type is specified in the "
+                               "signature of shader function '"
+                            << function->getFunctionName() << "'";
+                        throw std::logic_error(
+                            "Register type mismatch for shader input");
+                    }
+                    else
+                    {
+                        found = 1;
+                        break;
+                    }
+                }
+            }
+            if (found)
+            {
+                break;
+            }
         }
-    }
-    bool found = 0;
-    for (const DX12ShaderRegisterUsage& usage :
-        signature->getRegisterUsagesOrder())
-    {
-        if (usage.registerIndex == registerIndex
-            && usage.spaceIndex == spaceIndex
-            && usage.registerType == registerType)
+        if (found)
         {
-            found = 1;
             break;
         }
     }
-    if (found)
-    {
-        registerInputs.add(new DX12ShaderRegisterInput{
-            registerType, registerIndex, spaceIndex, inputResource->getThis()});
-    }
-    else
-    {
-        Logging::error() << "Register type " << registerType
-                         << ", register index " << registerIndex
-                         << ", space index " << spaceIndex
-                         << " is not used in the shader signature. The given "
-                            "input will be ignored.";
-    }
+    registerInputs.add(
+        new DX12ShaderRegisterInput{type, inputResource->getThis()});
 }
 
-void Framework::ShaderHeap::setRegisterInput(
-    Texture* zTexture, int registerIndex, int spaceIndex)
+void Framework::DX12GlobalDescriptorHeap::addTextureInput(
+    DX12ShaderRegister type, Texture* zTexture)
 {
-    setRegisterInput(DX12_SHADER_REGISTER_U_UNORDERED_ACCESS,
-        registerIndex,
-        spaceIndex,
-        zTexture);
+    addInput(type, zTexture);
 }
 
-void Framework::ShaderHeap::setRegisterInput(
-    DXBuffer* zBuffer, int registerIndex, int spaceIndex)
+void Framework::DX12GlobalDescriptorHeap::addBufferInput(
+    DX12ShaderRegister type, DXBuffer* zBuffer)
 {
-    setRegisterInput(DX12_SHADER_REGISTER_B_CONST_BUFFER,
-        registerIndex,
-        spaceIndex,
-        zBuffer);
+    addInput(type, zBuffer);
 }
 
-void Framework::ShaderHeap::setRegisterInput(
-    DX12TLAS* zTLAS, int registerIndex, int spaceIndex)
+void Framework::DX12GlobalDescriptorHeap::addTLASInput(
+    DX12ShaderRegister type, DX12TLAS* zTLAS)
 {
-    setRegisterInput(DX12_SHADER_REGISTER_T_SHADER_RESOURCE,
-        registerIndex,
-        spaceIndex,
-        zTLAS);
+    addInput(type, zTLAS);
 }
 
-void Framework::ShaderHeap::updateDescriptorHeap(ID3D12Device5* zDevice)
+void Framework::DX12GlobalDescriptorHeap::updateDescriptorHeap(
+    ID3D12Device5* zDevice)
 {
-    const Array<DX12ShaderRegisterUsage>& registerUsages
-        = signature->getRegisterUsagesOrder();
     if (!descriptorHeap
-        || lastDescriptorHeapSize
-               != signature->getRegisterUsagesOrder().getEntryCount())
+        || lastDescriptorHeapSize != registerInputs.getEntryCount())
     {
         if (descriptorHeap)
         {
             descriptorHeap->Release();
         }
         D3D12_DESCRIPTOR_HEAP_DESC desc = {};
-        desc.NumDescriptors = registerUsages.getEntryCount();
+        desc.NumDescriptors = registerInputs.getEntryCount();
         desc.Type = D3D12_DESCRIPTOR_HEAP_TYPE_CBV_SRV_UAV;
         desc.Flags = D3D12_DESCRIPTOR_HEAP_FLAG_SHADER_VISIBLE;
         desc.NodeMask = 0;
 
         HRESULT r = zDevice->CreateDescriptorHeap(
             &desc, __uuidof(ID3D12DescriptorHeap), (void**)&descriptorHeap);
-        lastDescriptorHeapSize = registerUsages.getEntryCount();
+        lastDescriptorHeapSize = registerInputs.getEntryCount();
     }
 
     D3D12_CPU_DESCRIPTOR_HANDLE descriptorHeapHandle
         = descriptorHeap->GetCPUDescriptorHandleForHeapStart();
-    ArrayIterator<DX12ShaderRegisterUsage> it = registerUsages.begin();
-    while (it)
+    for (const DX12ShaderRegisterInput* input : registerInputs)
     {
-        const DX12ShaderRegisterUsage& usage = it.val();
-        ArrayIterator<DX12ShaderRegisterInput*> inputIt
-            = registerInputs.begin();
-        bool found = 0;
-        while (inputIt)
+        DX12TLAS* zTLAS = dynamic_cast<DX12TLAS*>(input->inputResource);
+        DX12Texture* zTexture
+            = dynamic_cast<DX12Texture*>(input->inputResource);
+        DX12Buffer* zBuffer = dynamic_cast<DX12Buffer*>(input->inputResource);
+        switch (input->registerType)
         {
-            if (inputIt->registerType == usage.registerType
-                && inputIt->registerIndex == usage.registerIndex
-                && inputIt->spaceIndex == usage.spaceIndex)
+        case DX12_SHADER_REGISTER_B_CONST_BUFFER:
+            D3D12_CONSTANT_BUFFER_VIEW_DESC cbvDesc = {};
+            if (!zBuffer)
             {
-                switch (inputIt->registerType)
-                {
-                case DX12_SHADER_REGISTER_B_CONST_BUFFER:
-                    D3D12_CONSTANT_BUFFER_VIEW_DESC cbvDesc = {};
-                    DX12Buffer* buffer
-                        = dynamic_cast<DX12Buffer*>(inputIt->inputResource);
-                    cbvDesc.BufferLocation
-                        = buffer->zBuffer()->GetGPUVirtualAddress();
-                    cbvDesc.SizeInBytes = buffer->getElementCount()
-                                        * buffer->getElementLength();
-                    zDevice->CreateConstantBufferView(
-                        &cbvDesc, descriptorHeapHandle);
-                    break;
-                case DX12_SHADER_REGISTER_T_SHADER_RESOURCE:
-                    D3D12_SHADER_RESOURCE_VIEW_DESC srvDesc;
-                    srvDesc.Format = DXGI_FORMAT_UNKNOWN;
-                    srvDesc.ViewDimension
-                        = D3D12_SRV_DIMENSION_RAYTRACING_ACCELERATION_STRUCTURE;
-                    srvDesc.Shader4ComponentMapping
-                        = D3D12_DEFAULT_SHADER_4_COMPONENT_MAPPING;
-                    srvDesc.RaytracingAccelerationStructure.Location
-                        = dynamic_cast<DX12TLAS*>(inputIt->inputResource)
-                              ->zResultBuffer()
-                              ->zBuffer()
-                              ->GetGPUVirtualAddress();
-                    zDevice->CreateShaderResourceView(
-                        0, &srvDesc, descriptorHeapHandle);
-                    break;
-                case DX12_SHADER_REGISTER_U_UNORDERED_ACCESS:
-                    D3D12_UNORDERED_ACCESS_VIEW_DESC uavDesc = {};
-                    uavDesc.ViewDimension = D3D12_UAV_DIMENSION_TEXTURE2D;
-                    uavDesc.Format = DXGI_FORMAT_UNKNOWN;
-                    uavDesc.Texture2D.MipSlice = 0;
-                    uavDesc.Texture2D.PlaneSlice = 0;
-                    zDevice->CreateUnorderedAccessView(
-                        dynamic_cast<DX12Texture*>(inputIt->inputResource)
-                            ->zResource(),
-                        0,
-                        &uavDesc,
-                        descriptorHeapHandle);
-                    break;
-                default:
-                    Logging::error()
-                        << "Unknown register type for descriptor heap: "
-                        << inputIt->registerType;
-                }
-                found = 1;
-                break;
+                Logging::error()
+                    << "Expected a buffer resource for register type "
+                    << input->registerType;
+                throw std::logic_error(
+                    "Expected a buffer resource for register type "
+                    + std::to_string(input->registerType));
             }
-            ++inputIt;
-        }
-        if (!found)
-        {
-            Logging::error()
-                << "No input resource found for register type "
-                << usage.registerType << ", register index "
-                << usage.registerIndex << ", space index " << usage.spaceIndex
-                << ". This will result in an uninitialized descriptor in the "
-                   "descriptor heap. Access to the register in the shader "
-                   "might lead to undefined behaviour.";
+            cbvDesc.BufferLocation = zBuffer->zBuffer()->GetGPUVirtualAddress();
+            cbvDesc.SizeInBytes
+                = zBuffer->getElementCount() * zBuffer->getElementLength();
+            zDevice->CreateConstantBufferView(&cbvDesc, descriptorHeapHandle);
+            break;
+        case DX12_SHADER_REGISTER_T_SHADER_RESOURCE:
+            D3D12_SHADER_RESOURCE_VIEW_DESC srvDesc;
+            srvDesc.Format = DXGI_FORMAT_UNKNOWN;
+            if (zTLAS)
+            {
+                srvDesc.ViewDimension
+                    = D3D12_SRV_DIMENSION_RAYTRACING_ACCELERATION_STRUCTURE;
+                srvDesc.RaytracingAccelerationStructure.Location
+                    = zTLAS->zResultBuffer()->zBuffer()->GetGPUVirtualAddress();
+            }
+            else if (zTexture)
+            {
+                srvDesc.ViewDimension = D3D12_SRV_DIMENSION_TEXTURE2D;
+                srvDesc.Texture2D.MipLevels = 0;
+                srvDesc.Texture2D.MostDetailedMip = 0;
+                srvDesc.Texture2D.PlaneSlice = 0;
+                srvDesc.Texture2D.ResourceMinLODClamp = 0.0f;
+            }
+            else if (zBuffer)
+            {
+                srvDesc.ViewDimension = D3D12_SRV_DIMENSION_BUFFER;
+                srvDesc.Buffer.FirstElement = 0;
+                srvDesc.Buffer.NumElements = zBuffer->getElementCount();
+                srvDesc.Buffer.StructureByteStride
+                    = zBuffer->getElementLength();
+                srvDesc.Buffer.Flags = D3D12_BUFFER_SRV_FLAG_NONE;
+            }
+            srvDesc.Shader4ComponentMapping
+                = D3D12_DEFAULT_SHADER_4_COMPONENT_MAPPING;
+            zDevice->CreateShaderResourceView(
+                zTexture ? zTexture->zResource()
+                         : (zBuffer ? zBuffer->zBuffer() : 0),
+                &srvDesc,
+                descriptorHeapHandle);
+            break;
+        case DX12_SHADER_REGISTER_U_UNORDERED_ACCESS:
+            D3D12_UNORDERED_ACCESS_VIEW_DESC uavDesc = {};
+            if (zTexture)
+            {
+                uavDesc.ViewDimension = D3D12_UAV_DIMENSION_TEXTURE2D;
+                uavDesc.Format = DXGI_FORMAT_UNKNOWN;
+                uavDesc.Texture2D.MipSlice = 0;
+                uavDesc.Texture2D.PlaneSlice = 0;
+            }
+            else if (zBuffer)
+            {
+                uavDesc.ViewDimension = D3D12_UAV_DIMENSION_BUFFER;
+                uavDesc.Buffer.FirstElement = 0;
+                uavDesc.Buffer.NumElements = zBuffer->getElementCount();
+                uavDesc.Buffer.StructureByteStride
+                    = zBuffer->getElementLength();
+                uavDesc.Buffer.CounterOffsetInBytes = 0;
+                uavDesc.Buffer.Flags = D3D12_BUFFER_UAV_FLAG_NONE;
+            }
+            else
+            {
+                Logging::error() << "Expected a texture or buffer resource for "
+                                    "register type "
+                                 << input->registerType;
+                throw std::logic_error(
+                    "Expected a texture or buffer resource for register type "
+                    + std::to_string(input->registerType));
+            }
+            zDevice->CreateUnorderedAccessView(
+                zTexture ? zTexture->zResource() : zBuffer->zBuffer(),
+                0,
+                &uavDesc,
+                descriptorHeapHandle);
+            break;
         }
         descriptorHeapHandle.ptr += zDevice->GetDescriptorHandleIncrementSize(
             D3D12_DESCRIPTOR_HEAP_TYPE_CBV_SRV_UAV);
-        ++it;
     }
 }
 
-ID3D12DescriptorHeap* Framework::ShaderHeap::zDescriptorHeap() const
+ID3D12DescriptorHeap*
+Framework::DX12GlobalDescriptorHeap::zDescriptorHeap() const
 {
     return descriptorHeap;
 }
+
+Framework::DX12ShaderBindingTable::DX12ShaderBindingTable(
+    DX12Pipeline* pipeline)
+    : ReferenceCounter(),
+      pipeline(pipeline)
+{}

+ 67 - 21
DX12Shader.h

@@ -27,9 +27,10 @@ namespace Framework
         DX12ShaderRegister registerType;
         int registerIndex;
         int spaceIndex;
+        int descriptorHeapIndex;
     };
 
-    class ShaderHeap;
+    class DX12ShaderHeap;
 
     class DX12ShaderSignature : public ReferenceCounter
     {
@@ -42,7 +43,10 @@ namespace Framework
         DX12ShaderSignature();
         ~DX12ShaderSignature();
         /**
-         * needs to be called for each datastructure with : register(...)
+         * for each datastructure with : register(...) in the shader code,
+         * either this function or addRegisterUsageLinkedToDescriptorHeap must
+         * be called to link the register usage to the shader binding table or
+         * descriptor heap.
          *
          * \param registerType the register type e.g.
          * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0)
@@ -50,7 +54,26 @@ namespace Framework
          * \param spaceIndex the optional space index e.g. 3 for : register(b1,
          * space3)
          */
-        void addRegisterUsage(DX12ShaderRegister registerType,
+        void addRegisterUsageLinkedToShaderBindingTable(
+            DX12ShaderRegister registerType,
+            int registerIndex,
+            int spaceIndex = 0);
+        /**
+         * for each datastructure with : register(...) in the shader code,
+         * either this function or addRegisterUsageLinkedToDescriptorHeap must
+         * be called to link the register usage to the shader binding table or
+         * descriptor heap.
+         *
+         * \param descriptorHeapIndex the index in the descriptor heap witch
+         * contains the resource for this register usage
+         * \param registerType the register type e.g.
+         * DX12_SHADER_REGISTER_B_CONST_BUFFER for : register(b0)
+         * \param registerIndex the register index e.g. 1 for : register(b1)
+         * \param spaceIndex the optional space index e.g. 3 for : register(b1,
+         * space3)
+         */
+        void addRegisterUsageLinkedToDescriptorHeap(int descriptorHeapIndex,
+            DX12ShaderRegister registerType,
             int registerIndex,
             int spaceIndex = 0);
         /**
@@ -59,8 +82,7 @@ namespace Framework
         void createSignature(ID3D12Device5* zDevice,
             PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature);
         ID3D12RootSignature* zSignature() const;
-        const Array<DX12ShaderRegisterUsage>& getRegisterUsagesOrder() const;
-        ShaderHeap* createShaderHeap();
+        const Array<DX12ShaderRegisterUsage>& getRegisterUsages() const;
     };
 
     class DX12ShaderFunction : public ReferenceCounter
@@ -127,6 +149,9 @@ namespace Framework
         D3D12_HIT_GROUP_DESC* zHitGroupDesc() const;
     };
 
+    class DX12ShaderBindingTable;
+    class DX12GlobalDescriptorHeap;
+
     class DX12Pipeline : public ReferenceCounter
     {
     private:
@@ -146,43 +171,64 @@ namespace Framework
         void createPipelineState(ID3D12Device5* zDevice,
             PFN_D3D12_SERIALIZE_ROOT_SIGNATURE pfnD3D12SerializeRootSignature);
         ID3D12StateObject* zPipelineState() const;
+        DX12ShaderBindingTable* createShaderBindingTable();
+        DX12GlobalDescriptorHeap* createGlobalDescriptorHeap();
+        const RCArray<DX12Shader>& getShaders() const;
     }; // namespace Framework
 
     struct DX12ShaderRegisterInput
     {
         DX12ShaderRegister registerType;
-        int registerIndex;
-        int spaceIndex;
         ReferenceCounter*
             inputResource; // Can be Texture*, DXBuffer*, or DX12TLAS*
     };
 
-    class ShaderHeap : public ReferenceCounter
+    class DX12GlobalDescriptorHeap : public ReferenceCounter
     {
     private:
-        DX12ShaderSignature* signature;
+        DX12Pipeline* pipeline;
         ID3D12DescriptorHeap* descriptorHeap;
         Array<DX12ShaderRegisterInput*> registerInputs;
         int lastDescriptorHeapSize;
 
     public:
-        ShaderHeap(DX12ShaderSignature* signature);
-        ~ShaderHeap();
+        DX12GlobalDescriptorHeap(DX12Pipeline* pipeline);
+        ~DX12GlobalDescriptorHeap();
 
     private:
-        void setRegisterInput(DX12ShaderRegister registerType,
-            int registerIndex,
-            int spaceIndex,
-            ReferenceCounter* inputResource);
+        void addInput(DX12ShaderRegister type, ReferenceCounter* inputResource);
 
     public:
-        void setRegisterInput(
-            Texture* zTexture, int registerIndex, int spaceIndex = 0);
-        void setRegisterInput(
-            DXBuffer* zBuffer, int registerIndex, int spaceIndex = 0);
-        void setRegisterInput(
-            DX12TLAS* zTLAS, int registerIndex, int spaceIndex = 0);
+        void addTextureInput(DX12ShaderRegister type, Texture* zTexture);
+        void addBufferInput(DX12ShaderRegister type, DXBuffer* zBuffer);
+        void addTLASInput(DX12ShaderRegister type, DX12TLAS* zTLAS);
         void updateDescriptorHeap(ID3D12Device5* zDevice);
         ID3D12DescriptorHeap* zDescriptorHeap() const;
     };
+
+    class DX12ShaderBindingTable : public ReferenceCounter
+    {
+    private:
+        DX12Pipeline* pipeline;
+        DX12Buffer* shaderBindingTableBuffer;
+        int rayGenRecordSize;
+        int rayGenCount;
+        int missRecordSize;
+        int missCount;
+        int hitGroupRecordSize;
+        int hitGroupCount;
+        bool changed;
+
+    public:
+        DX12ShaderBindingTable(DX12Pipeline* pipeline);
+        ~DX12ShaderBindingTable();
+
+        void startUpdate();
+        void setHitGroupShaderInputs(int instanceIndex,
+            DX12ShaderHitGroup* zHitGroup,
+            std::initializer_list<unsigned __int64> gpuAddresses);
+        void endUpdate(ID3D12Device5* zDevice);
+        void fillDispatchRaysDesc(D3D12_DISPATCH_RAYS_DESC* dispatchRaysDesc);
+        DX12Pipeline* zPipeline() const;
+    };
 } // namespace Framework

+ 8 - 5
Framework.vcxproj

@@ -476,7 +476,7 @@ copy "x64\Release\Framework.dll" "..\..\Spiele Platform\SMP\Fertig\x64\framework
     <FxCompile Include="Hit.hlsl">
       <EntryPointName Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">
       </EntryPointName>
-      <ShaderModel Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">6.3</ShaderModel>
+      <ShaderModel Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">6.8</ShaderModel>
       <HeaderFileOutput Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">DX12HitShader.h</HeaderFileOutput>
       <ObjectFileOutput Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">
       </ObjectFileOutput>
@@ -484,13 +484,14 @@ copy "x64\Release\Framework.dll" "..\..\Spiele Platform\SMP\Fertig\x64\framework
       <Command Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">dxc.exe /T lib_5_0 /Zi /E"closesthit" /Od /Fh"DX12HitShader.h" /nologo "%(FullPath)"</Command>
       <Outputs Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">DX12HitShader.h</Outputs>
       <AdditionalInputs Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">Common.hlsl</AdditionalInputs>
-      <AdditionalOptions Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">/T lib_6_3  -Fd "x64\Debug\DX12HitShader.pdb" %(AdditionalOptions)</AdditionalOptions>
+      <AdditionalOptions Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">-Fd "x64\Debug\DX12HitShader.pdb" %(AdditionalOptions)</AdditionalOptions>
+      <ShaderType Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">Library</ShaderType>
     </FxCompile>
     <FxCompile Include="Miss.hlsl">
       <EntryPointName Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">
       </EntryPointName>
-      <ShaderModel Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">6.3</ShaderModel>
-      <AdditionalOptions Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">/T lib_6_3 -Fd "x64\Debug\DX12MissShader.pdb" %(AdditionalOptions)</AdditionalOptions>
+      <ShaderModel Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">6.8</ShaderModel>
+      <AdditionalOptions Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">-Fd "x64\Debug\DX12MissShader.pdb" %(AdditionalOptions)</AdditionalOptions>
       <HeaderFileOutput Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">DX12MissShader.h</HeaderFileOutput>
       <ObjectFileOutput Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">
       </ObjectFileOutput>
@@ -498,12 +499,13 @@ copy "x64\Release\Framework.dll" "..\..\Spiele Platform\SMP\Fertig\x64\framework
       <Command Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">dxc.exe /T lib_5_0 /Zi /E"miss" /Od /Fh"DX12MissShader.h" /nologo "%(FullPath)"</Command>
       <Outputs Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">DX12MissShader.h</Outputs>
       <AdditionalInputs Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">Common.hlsl</AdditionalInputs>
+      <ShaderType Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">Library</ShaderType>
     </FxCompile>
     <FxCompile Include="RayGen.hlsl">
       <EntryPointName Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">
       </EntryPointName>
       <ShaderModel Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">6.3</ShaderModel>
-      <AdditionalOptions Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">/T lib_6_3 -Fd "x64\Debug\DX12RayGenShader.pdb" %(AdditionalOptions)</AdditionalOptions>
+      <AdditionalOptions Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">-Fd "x64\Debug\DX12RayGenShader.pdb" %(AdditionalOptions)</AdditionalOptions>
       <HeaderFileOutput Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">DX12RayGenShader.h</HeaderFileOutput>
       <ObjectFileOutput Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">
       </ObjectFileOutput>
@@ -511,6 +513,7 @@ copy "x64\Release\Framework.dll" "..\..\Spiele Platform\SMP\Fertig\x64\framework
       <Command Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">dxc.exe /T lib_5_0 /Zi /E"raygen" /Od /Fh"DX12RayGenShader.h" /nologo "%(FullPath)"</Command>
       <Outputs Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">DX12RayGenShader.h</Outputs>
       <AdditionalInputs Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">Common.hlsl</AdditionalInputs>
+      <ShaderType Condition="'$(Configuration)|$(Platform)'=='Debug|x64'">Library</ShaderType>
     </FxCompile>
   </ItemGroup>
   <ItemGroup>