Переглянути джерело

improve rendering performance of raytracing by only updating tlas and sbt when an object was changed

Kolja Strohm 2 тижнів тому
батько
коміт
6a0eca2a6f
8 змінених файлів з 191 додано та 51 видалено
  1. 104 30
      DX12GraphicsApi.cpp
  2. 2 1
      DX12GraphicsApi.h
  3. 44 18
      DX12Shader.cpp
  4. 7 1
      DX12Shader.h
  5. 9 1
      Drawing3D.cpp
  6. 3 0
      Drawing3D.h
  7. 17 0
      Model3D.cpp
  8. 5 0
      Model3D.h

+ 104 - 30
DX12GraphicsApi.cpp

@@ -32,6 +32,26 @@
 
 using namespace Framework;
 
+class DX12RenderingData : public Framework::ReferenceCounter
+{
+public:
+    int lastObjectIndex;
+    int lastBlasCount;
+    int* lastHitGroupIndices;
+
+    DX12RenderingData()
+        : Framework::ReferenceCounter(),
+          lastObjectIndex(-1),
+          lastBlasCount(0),
+          lastHitGroupIndices(0)
+    {}
+
+    ~DX12RenderingData()
+    {
+        if (lastHitGroupIndices) delete[] lastHitGroupIndices;
+    }
+};
+
 DirectX12::DirectX12()
     : GraphicsApi(DIRECTX12),
       debug(0),
@@ -240,26 +260,68 @@ void Framework::DirectX12::renderKamera(
     sbt->setSamplerDescriptorHeap(samplerDescriptorHeap);
     w->render([this, &tlas, &objectIndex, &instanceIndex, &identity, &sbt](
                   Model3D* obj) {
-        obj->calculateMatrices(identity, matrixBuffer);
+        DX12RenderingData* renderingData
+            = dynamic_cast<DX12RenderingData*>(obj->zRenderingData());
+        if (!renderingData)
+        {
+            renderingData = new DX12RenderingData();
+            obj->setRenderingData(renderingData);
+        }
         int modelId = obj->zModelData()->getId();
         DX12BLASModel* blasModel = blasModels[modelId];
-        blasModel->calculateBuffers();
+        bool changed
+            = obj->getLastTickReturn()
+           || renderingData->lastObjectIndex != objectIndex
+           || renderingData->lastBlasCount != blasModel->getBufferCount()
+           || obj->zModelData()->wasIndexBufferChanged()
+           || obj->zModelData()->wasVertexBufferChanged();
+        renderingData->lastObjectIndex = objectIndex;
+        if (changed)
+        {
+            obj->calculateMatrices(identity, matrixBuffer);
+            blasModel->calculateBuffers();
+        }
+        if (renderingData->lastBlasCount < blasModel->getBufferCount())
+        {
+            if (renderingData->lastHitGroupIndices)
+            {
+                delete[] renderingData->lastHitGroupIndices;
+            }
+            renderingData->lastHitGroupIndices
+                = new int[blasModel->getBufferCount()];
+            memset(renderingData->lastHitGroupIndices,
+                -1,
+                sizeof(int) * blasModel->getBufferCount());
+            renderingData->lastBlasCount = blasModel->getBufferCount();
+        }
         ArrayIterator<int> boneIds = blasModel->zBoneIds()->begin();
+        instanceIndex = 0;
         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 * (obj->getSize() > 0);
-            desc->AccelerationStructure
-                = blas->zResultBuffer()->zBuffer()->GetGPUVirtualAddress();
-            memcpy(desc->Transform,
-                &matrixBuffer[boneIds.val()],
-                sizeof(float) * 12);
+            if (changed)
+            {
+                D3D12_RAYTRACING_INSTANCE_DESC* desc = tlas->nextInstanceDesc();
+                desc->InstanceID = objectIndex;
+                desc->InstanceContributionToHitGroupIndex = objectIndex;
+                desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE;
+                desc->InstanceMask = 0xFF * (obj->getSize() > 0);
+                desc->AccelerationStructure
+                    = blas->zResultBuffer()->zBuffer()->GetGPUVirtualAddress();
+                memcpy(desc->Transform,
+                    &matrixBuffer[boneIds.val()],
+                    sizeof(float) * 12);
+            }
+            else
+            {
+                tlas->nextInstanceDesc();
+            }
             boneIds++;
+            fillShaderBindingTable(sbt,
+                obj,
+                objectIndex,
+                blas,
+                renderingData->lastHitGroupIndices[instanceIndex]);
             instanceIndex++;
-            fillShaderBindingTable(sbt, obj, objectIndex, blas);
             objectIndex++;
         }
     });
@@ -427,26 +489,34 @@ void Framework::DirectX12::fillShaderBindingTable(
     DX12ShaderBindingTable* zShaderBindingTable,
     Model3D* zModel,
     int instanceIndex,
-    const DX12BLAS* zBLAS)
+    const DX12BLAS* zBLAS,
+    int& lastHitGroupIndex)
 {
     if (defaultHitGroup)
     {
-        int hitGroupOffset = zShaderBindingTable->addHitGroup(defaultHitGroup);
-        zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
-            sbtIndexBufferOffset,
-            zBLAS->zIndexBuffer()->zBuffer()->GetGPUVirtualAddress());
-        zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
-            sbtVertexDataBufferOffset,
-            zBLAS->zVertexDataBuffer()->zBuffer()->GetGPUVirtualAddress());
-        zModel->zTexture()->updateTextureIndexBuffer(this);
-        zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
-            sbtTextureIdBufferOffset,
-            dynamic_cast<DX12Buffer*>(zModel->zTexture()->zTextureIndexBuffer())
-                ->zBuffer()
-                ->GetGPUVirtualAddress());
-        zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
-            sbtPolygonSizeBufferOffset,
-            zBLAS->zPolygonSizeBuffer()->zBuffer()->GetGPUVirtualAddress());
+        int hitGroupOffset = zShaderBindingTable->addHitGroup(
+            defaultHitGroup, lastHitGroupIndex);
+        bool changed = hitGroupOffset != lastHitGroupIndex;
+        lastHitGroupIndex = hitGroupOffset;
+        if (changed)
+        { // TODO: check if the buffers are valid and have been recreated
+            zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
+                sbtIndexBufferOffset,
+                zBLAS->zIndexBuffer()->zBuffer()->GetGPUVirtualAddress());
+            zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
+                sbtVertexDataBufferOffset,
+                zBLAS->zVertexDataBuffer()->zBuffer()->GetGPUVirtualAddress());
+            zModel->zTexture()->updateTextureIndexBuffer(this);
+            zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
+                sbtTextureIdBufferOffset,
+                dynamic_cast<DX12Buffer*>(
+                    zModel->zTexture()->zTextureIndexBuffer())
+                    ->zBuffer()
+                    ->GetGPUVirtualAddress());
+            zShaderBindingTable->setHitGroupShaderInput(hitGroupOffset,
+                sbtPolygonSizeBufferOffset,
+                zBLAS->zPolygonSizeBuffer()->zBuffer()->GetGPUVirtualAddress());
+        }
     }
 }
 
@@ -1008,6 +1078,10 @@ void DirectX12::presentFrame()
 
     swapChain->Present(0, DXGI_PRESENT_ALLOW_TEARING);
 
+    globalDescriptorHeap->setHeapChanged(0);
+    textureDescriptorHeap->setHeapChanged(0);
+    samplerDescriptorHeap->setHeapChanged(0);
+
     backBufferIndex = swapChain->GetCurrentBackBufferIndex();
 
     cs.unlock();

+ 2 - 1
DX12GraphicsApi.h

@@ -116,7 +116,8 @@ namespace Framework
             DX12ShaderBindingTable* zShaderBindingTable,
             Model3D* zModel,
             int objectIndex,
-            const DX12BLAS* zBLAS);
+            const DX12BLAS* zBLAS,
+            int& lastHitGroupIndex);
 
     public:
         DLLEXPORT void initialize(NativeWindow* fenster,

+ 44 - 18
DX12Shader.cpp

@@ -1099,7 +1099,8 @@ Framework::DX12GlobalDescriptorHeap::DX12GlobalDescriptorHeap(
       descriptorHeap(0),
       lastDescriptorHeapSize(0),
       zDevice(0),
-      type(type)
+      type(type),
+      heapChanged(0)
 {}
 
 Framework::DX12GlobalDescriptorHeap::~DX12GlobalDescriptorHeap()
@@ -1331,6 +1332,7 @@ void Framework::DX12GlobalDescriptorHeap::updateDescriptorHeap(
         HRESULT r = zDevice->CreateDescriptorHeap(
             &desc, __uuidof(ID3D12DescriptorHeap), (void**)&descriptorHeap);
         lastDescriptorHeapSize = registerInputs.getEntryCount();
+        heapChanged = 1;
     }
 
     D3D12_CPU_DESCRIPTOR_HANDLE descriptorHeapHandle
@@ -1510,6 +1512,16 @@ Framework::DX12GlobalDescriptorHeap::zDescriptorHeap() const
     return descriptorHeap;
 }
 
+void Framework::DX12GlobalDescriptorHeap::setHeapChanged(bool changed)
+{
+    heapChanged = changed;
+}
+
+bool Framework::DX12GlobalDescriptorHeap::wasHeapChanged() const
+{
+    return heapChanged;
+}
+
 Framework::DX12ShaderBindingTable::DX12ShaderBindingTable(
     DX12Pipeline* pipeline)
     : ReferenceCounter(),
@@ -1776,36 +1788,50 @@ void Framework::DX12ShaderBindingTable::setShaderInput(
 }
 
 int Framework::DX12ShaderBindingTable::addHitGroup(
-    DX12ShaderHitGroup* zHitGroup)
+    DX12ShaderHitGroup* zHitGroup, int lastIndex)
 {
     int index = nextHitGroupOffset;
-    set(index,
-        stateObjectProperties->GetShaderIdentifier(
-            zHitGroup->zHitGroupDesc()->HitGroupExport),
-        D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT);
+    bool changed
+        = index + D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT != lastIndex;
+    if (changed)
+    {
+        set(index,
+            stateObjectProperties->GetShaderIdentifier(
+                zHitGroup->zHitGroupDesc()->HitGroupExport),
+            D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT);
+    }
     index += D3D12_RAYTRACING_SHADER_RECORD_BYTE_ALIGNMENT;
     if (zHitGroup->zSignature()->doesUseGlobalDescriptorHeap())
     {
-        D3D12_GPU_DESCRIPTOR_HANDLE gpuAddress
-            = globalDescriptorHeap->zDescriptorHeap()
-                  ->GetGPUDescriptorHandleForHeapStart();
-        set(index, &gpuAddress.ptr, sizeof(__int64));
+        if ((changed || globalDescriptorHeap->wasHeapChanged()))
+        {
+            D3D12_GPU_DESCRIPTOR_HANDLE gpuAddress
+                = globalDescriptorHeap->zDescriptorHeap()
+                      ->GetGPUDescriptorHandleForHeapStart();
+            set(index, &gpuAddress.ptr, sizeof(__int64));
+        }
         index += sizeof(__int64);
     }
     if (zHitGroup->zSignature()->doesUseTextureDescriptorHeap())
     {
-        D3D12_GPU_DESCRIPTOR_HANDLE gpuAddress
-            = textureDescriptorHeap->zDescriptorHeap()
-                  ->GetGPUDescriptorHandleForHeapStart();
-        set(index, &gpuAddress.ptr, sizeof(__int64));
+        if ((changed || textureDescriptorHeap->wasHeapChanged()))
+        {
+            D3D12_GPU_DESCRIPTOR_HANDLE gpuAddress
+                = textureDescriptorHeap->zDescriptorHeap()
+                      ->GetGPUDescriptorHandleForHeapStart();
+            set(index, &gpuAddress.ptr, sizeof(__int64));
+        }
         index += sizeof(__int64);
     }
     if (zHitGroup->zSignature()->doesUseSamplerDescriptorHeap())
     {
-        D3D12_GPU_DESCRIPTOR_HANDLE gpuAddress
-            = samplerDescriptorHeap->zDescriptorHeap()
-                  ->GetGPUDescriptorHandleForHeapStart();
-        set(index, &gpuAddress.ptr, sizeof(__int64));
+        if ((changed || samplerDescriptorHeap->wasHeapChanged()))
+        {
+            D3D12_GPU_DESCRIPTOR_HANDLE gpuAddress
+                = samplerDescriptorHeap->zDescriptorHeap()
+                      ->GetGPUDescriptorHandleForHeapStart();
+            set(index, &gpuAddress.ptr, sizeof(__int64));
+        }
         index += sizeof(__int64);
     }
     hitGroupCount++;

+ 7 - 1
DX12Shader.h

@@ -253,6 +253,7 @@ namespace Framework
         int lastDescriptorHeapSize;
         ID3D12Device5* zDevice;
         DX12DescriptorHeapType type;
+        bool heapChanged;
 
     public:
         DX12GlobalDescriptorHeap(
@@ -277,6 +278,8 @@ namespace Framework
         DLLEXPORT void updateDescriptorHeap(ID3D12Device5* zDevice);
         DLLEXPORT DX12Pipeline* zPipeline() const;
         DLLEXPORT ID3D12DescriptorHeap* zDescriptorHeap() const;
+        DLLEXPORT void setHeapChanged(bool changed);
+        DLLEXPORT bool wasHeapChanged() const;
     };
 
     class DX12ShaderBindingTable : public ReferenceCounter
@@ -340,11 +343,14 @@ namespace Framework
          *
          * \param zHitGroup the hit group to be added. The hit group must be
          * part of the pipeline.
+         * \param lastIndex the index of the hit group when it was added at the
+         * last frame
          * \return the index of the hit group in the shader binding table. This
          * index can be used to set the inputs for this hit group by calling
          * setHitGroupShaderInput.
          */
-        DLLEXPORT int addHitGroup(DX12ShaderHitGroup* zHitGroup);
+        DLLEXPORT int addHitGroup(
+            DX12ShaderHitGroup* zHitGroup, int lastIndex = -1);
         /**
          * sets a specific input for a previously added hitgroup.
          *

+ 9 - 1
Drawing3D.cpp

@@ -11,6 +11,7 @@ Drawable3D::Drawable3D()
     welt = welt.identity();
     pos = Vec3<float>(0, 0, 0);
     angle = Vec3<float>(0, 0, 0);
+    lastTickReturn = 0;
     rend = 0;
     alpha = 0;
     radius = 0;
@@ -158,8 +159,10 @@ bool Drawable3D::tick(double tickval)
              * welt.rotationX(angle.x) * welt.rotationY(angle.y)
              * welt.scaling(size);
         rend = 0;
+        lastTickReturn = 1;
         return 1;
     }
+    lastTickReturn = 0;
     return 0;
 }
 
@@ -236,4 +239,9 @@ Vec3<float> Drawable3D::applyWorldTransformation(
     const Vec3<float>& modelPos) const
 {
     return welt * modelPos;
-}
+}
+
+bool Framework::Drawable3D::getLastTickReturn() const
+{
+    return lastTickReturn;
+}

+ 3 - 0
Drawing3D.h

@@ -20,6 +20,7 @@ namespace Framework
         bool alpha;   //! Stores whether the object contains partially or fully
                       //! transparent areas
         bool rend;
+        bool lastTickReturn;
         float size;
 
     public:
@@ -112,5 +113,7 @@ namespace Framework
         //! drawing coordinates by applying rotation, scaling and translation
         DLLEXPORT Vec3<float> applyWorldTransformation(
             const Vec3<float>& modelPos) const;
+        //! returns whether the object should be rendered because it was changed
+        DLLEXPORT bool getLastTickReturn() const;
     };
 } // namespace Framework

+ 17 - 0
Model3D.cpp

@@ -765,6 +765,7 @@ void Framework::Model3DTexture::updateTextureIndexBuffer(GraphicsApi* zApi)
     textureIndexBuffer->setData(textureIndexList, changed);
     textureIndexBuffer->setLength(textureCount * sizeof(int));
     textureIndexBuffer->copyToGPU();
+    changed = 0;
 }
 
 DXBuffer* Framework::Model3DTexture::zTextureIndexBuffer() const
@@ -785,6 +786,7 @@ Model3D::Model3D()
     model = 0;
     texture = 0;
     skelett = 0;
+    renderingData = 0;
     ambientFactor = 1.f;
     diffusFactor = 0.f;
     specularFactor = 0.f;
@@ -796,6 +798,7 @@ Model3D::~Model3D()
     if (model) model->release();
     if (texture) texture->release();
     if (skelett) skelett->release();
+    if (renderingData) renderingData->release();
 }
 
 // Sets the model data
@@ -1092,4 +1095,18 @@ Texture* Model3D::zEffectTexture()
 float Model3D::getEffectPercentage()
 {
     return 0;
+}
+
+void Framework::Model3D::setRenderingData(ReferenceCounter* data)
+{
+    if (renderingData)
+    {
+        renderingData->release();
+    }
+    renderingData = data;
+}
+
+ReferenceCounter* Framework::Model3D::zRenderingData() const
+{
+    return renderingData;
 }

+ 5 - 0
Model3D.h

@@ -315,6 +315,7 @@ namespace Framework
         float ambientFactor;
         float diffusFactor;
         float specularFactor;
+        ReferenceCounter* renderingData;
 
     public:
         //! Constructor
@@ -402,5 +403,9 @@ namespace Framework
         DLLEXPORT virtual bool needRenderPolygon(int index);
         DLLEXPORT virtual Texture* zEffectTexture();
         DLLEXPORT virtual float getEffectPercentage();
+        // Sets additional data associated with the model
+        DLLEXPORT void setRenderingData(ReferenceCounter* data);
+        // Returns additional data associated with the model without increasing
+        DLLEXPORT ReferenceCounter* zRenderingData() const;
     };
 } // namespace Framework