diff options
| author | azhirnov <zh1dron@gmail.com> | 2020-11-03 10:52:24 +0000 |
|---|---|---|
| committer | azhirnov <zh1dron@gmail.com> | 2020-11-03 11:16:11 +0000 |
| commit | efa43e2bd2475a4dec6771bf9759f6a99f7d77ed (patch) | |
| tree | bbb257c3825ff07078626e3e137468f99d235a07 /Graphics | |
| parent | Few improvements to ray tracing tests (diff) | |
| download | DiligentCore-efa43e2bd2475a4dec6771bf9759f6a99f7d77ed.tar.gz DiligentCore-efa43e2bd2475a4dec6771bf9759f6a99f7d77ed.zip | |
fixed resource state transitions, some improvements for ray tracing
Diffstat (limited to 'Graphics')
27 files changed, 536 insertions, 297 deletions
diff --git a/Graphics/GraphicsEngine/include/BottomLevelASBase.hpp b/Graphics/GraphicsEngine/include/BottomLevelASBase.hpp index 41f2fdf3..2bdd51dc 100644 --- a/Graphics/GraphicsEngine/include/BottomLevelASBase.hpp +++ b/Graphics/GraphicsEngine/include/BottomLevelASBase.hpp @@ -184,6 +184,18 @@ public: return (this->m_State & State) == State; } +#ifdef DILIGENT_DEVELOPMENT + void UpdateVersion() + { + m_Version.fetch_add(1); + } + + Uint32 GetVersion() const + { + return m_Version.load(); + } +#endif + protected: static void ValidateBottomLevelASDesc(const BottomLevelASDesc& Desc) { @@ -215,6 +227,10 @@ protected: std::unordered_map<HashMapStringKey, Uint32, HashMapStringKey::Hasher> m_NameToIndex; StringPool m_StringPool; + +#ifdef DILIGENT_DEVELOPMENT + std::atomic<Uint32> m_Version{0}; +#endif }; } // namespace Diligent diff --git a/Graphics/GraphicsEngine/include/DeviceContextBase.hpp b/Graphics/GraphicsEngine/include/DeviceContextBase.hpp index 829416d2..ba586f4b 100644 --- a/Graphics/GraphicsEngine/include/DeviceContextBase.hpp +++ b/Graphics/GraphicsEngine/include/DeviceContextBase.hpp @@ -1855,13 +1855,23 @@ void DeviceContextBase<BaseInterface, ImplementationTraits>:: DEV_CHECK_ERR(OldState != RESOURCE_STATE_UNKNOWN, "The state of buffer '", BuffDesc.Name, "' is unknown to the engine and is not explicitly specified in the barrier"); DEV_CHECK_ERR(VerifyResourceStates(OldState, false), "Invlaid old state specified for buffer '", BuffDesc.Name, "'"); } - else if (RefCntAutoPtr<IBottomLevelAS> pBLAS{Barrier.pResource, IID_BottomLevelAS}) + else if (RefCntAutoPtr<IBottomLevelAS> pBottomLevelAS{Barrier.pResource, IID_BottomLevelAS}) { - // AZ TODO + const auto& BLASDesc = pBottomLevelAS->GetDesc(); + OldState = Barrier.OldState != RESOURCE_STATE_UNKNOWN ? Barrier.OldState : pBottomLevelAS->GetState(); + DEV_CHECK_ERR(OldState != RESOURCE_STATE_UNKNOWN, "The state of BLAS '", BLASDesc.Name, "' is unknown to the engine and is not explicitly specified in the barrier"); + DEV_CHECK_ERR(Barrier.NewState == RESOURCE_STATE_BUILD_AS_READ || Barrier.NewState == RESOURCE_STATE_BUILD_AS_WRITE || Barrier.NewState == RESOURCE_STATE_RAY_TRACING, + "Invlaid new state specified for BLAS '", BLASDesc.Name, "'"); + DEV_CHECK_ERR(Barrier.TransitionType != STATE_TRANSITION_TYPE_IMMEDIATE, "Split barriers are not supported for BLAS"); } - else if (RefCntAutoPtr<ITopLevelAS> pTLAS{Barrier.pResource, IID_TopLevelAS}) + else if (RefCntAutoPtr<ITopLevelAS> pTopLevelAS{Barrier.pResource, IID_TopLevelAS}) { - // AZ TODO + const auto& TLASDesc = pTopLevelAS->GetDesc(); + OldState = Barrier.OldState != RESOURCE_STATE_UNKNOWN ? Barrier.OldState : pTopLevelAS->GetState(); + DEV_CHECK_ERR(OldState != RESOURCE_STATE_UNKNOWN, "The state of TLAS '", TLASDesc.Name, "' is unknown to the engine and is not explicitly specified in the barrier"); + DEV_CHECK_ERR(Barrier.NewState == RESOURCE_STATE_BUILD_AS_READ || Barrier.NewState == RESOURCE_STATE_BUILD_AS_WRITE || Barrier.NewState == RESOURCE_STATE_RAY_TRACING, + "Invlaid new state specified for TLAS '", TLASDesc.Name, "'"); + DEV_CHECK_ERR(Barrier.TransitionType != STATE_TRANSITION_TYPE_IMMEDIATE, "Split barriers are not supported for TLAS"); } else { @@ -1943,6 +1953,12 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>:: template <typename BaseInterface, typename ImplementationTraits> bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildBLAS(const BLASBuildAttribs& Attribs, int) { + if (m_pActiveRenderPass != nullptr) + { + LOG_ERROR_MESSAGE("BuildBLAS command must be performed outside of render pass"); + return false; + } + if (Attribs.pBLAS == nullptr) { LOG_ERROR_MESSAGE("IDeviceContext::BuildBLAS: pBLAS must not be null"); @@ -2090,6 +2106,7 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildBLAS(const BLA return false; } } +#endif // DILIGENT_DEVELOPMENT const auto& BLASDesc = Attribs.pBLAS->GetDesc(); @@ -2113,7 +2130,7 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildBLAS(const BLA return false; } - if (ScratchDesc.uiSizeInBytes - Attribs.ScratchBufferOffset > Attribs.pBLAS->GetScratchBufferSizes().Build) + if (ScratchDesc.uiSizeInBytes - Attribs.ScratchBufferOffset < Attribs.pBLAS->GetScratchBufferSizes().Build) { LOG_ERROR_MESSAGE("IDeviceContext::BuildBLAS: pScratchBuffer size is too small, use pBLAS->GetScratchBufferSizes().Build to get required size for scratch buffer"); return false; @@ -2124,7 +2141,6 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildBLAS(const BLA LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pScratchBuffer must be created with BIND_RAY_TRACING flag"); return false; } -#endif // DILIGENT_DEVELOPMENT return true; } @@ -2132,6 +2148,12 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildBLAS(const BLA template <typename BaseInterface, typename ImplementationTraits> bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLASBuildAttribs& Attribs, int) { + if (m_pActiveRenderPass != nullptr) + { + LOG_ERROR_MESSAGE("BuildTLAS command must be performed outside of render pass"); + return false; + } + if (Attribs.pTLAS == nullptr) { LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pTLAS must not be null"); @@ -2162,7 +2184,6 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLA return false; } -#ifdef DILIGENT_DEVELOPMENT const auto& TLASDesc = Attribs.pTLAS->GetDesc(); if (Attribs.InstanceCount > TLASDesc.MaxInstanceCount) @@ -2171,9 +2192,11 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLA return false; } - const auto& InstDesc = Attribs.pInstanceBuffer->GetDesc(); - const size_t InstDataSize = Attribs.InstanceCount * TLAS_INSTANCE_DATA_SIZE; - Uint32 AutoOffsetCounter = 0; + const auto& InstDesc = Attribs.pInstanceBuffer->GetDesc(); + const size_t InstDataSize = Attribs.InstanceCount * TLAS_INSTANCE_DATA_SIZE; + +#ifdef DILIGENT_DEVELOPMENT + Uint32 AutoOffsetCounter = 0; // calculate instance data size for (Uint32 i = 0; i < Attribs.InstanceCount; ++i) @@ -2203,6 +2226,7 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLA LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: exactly all pInstances[i].ContributionToHitGroupIndex must be TLAS_INSTANCE_OFFSET_AUTO or not"); return false; } +#endif // DILIGENT_DEVELOPMENT if (Attribs.InstanceBufferOffset > InstDesc.uiSizeInBytes) { @@ -2210,15 +2234,15 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLA return false; } - if (InstDesc.uiSizeInBytes - Attribs.InstanceBufferOffset > InstDataSize) + if (InstDesc.uiSizeInBytes - Attribs.InstanceBufferOffset < InstDataSize) { - LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pInstanceaBuffer size is too small, ..."); + LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pInstanceBuffer size is too small, ..."); return false; } if ((InstDesc.BindFlags & BIND_RAY_TRACING) != BIND_RAY_TRACING) { - LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pInstanceaBuffer must be created with BIND_RAY_TRACING flag"); + LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pInstanceBuffer must be created with BIND_RAY_TRACING flag"); return false; } @@ -2230,7 +2254,7 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLA return false; } - if (ScratchDesc.uiSizeInBytes - Attribs.ScratchBufferOffset > Attribs.pTLAS->GetScratchBufferSizes().Build) + if (ScratchDesc.uiSizeInBytes - Attribs.ScratchBufferOffset < Attribs.pTLAS->GetScratchBufferSizes().Build) { LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pScratchBuffer size is too small, use pTLAS->GetScratchBufferSizes().Build to get required size for scratch buffer"); return false; @@ -2241,7 +2265,6 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::BuildTLAS(const TLA LOG_ERROR_MESSAGE("IDeviceContext::BuildTLAS: pScratchBuffer must be created with BIND_RAY_TRACING flag"); return false; } -#endif // DILIGENT_DEVELOPMENT return true; } @@ -2261,6 +2284,12 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::CopyBLAS(const Copy return false; } + if (m_pActiveRenderPass != nullptr) + { + LOG_ERROR_MESSAGE("CopyBLAS command must be performed outside of render pass"); + return false; + } + #ifdef DILIGENT_DEVELOPMENT if (Attribs.Mode == COPY_AS_MODE_CLONE) { @@ -2338,7 +2367,19 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::CopyTLAS(const Copy return false; } + if (m_pActiveRenderPass != nullptr) + { + LOG_ERROR_MESSAGE("CopyTLAS command must be performed outside of render pass"); + return false; + } + #ifdef DILIGENT_DEVELOPMENT + if (!ValidatedCast<TopLevelASType>(Attribs.pSrc)->CheckBLASVersion()) + { + LOG_ERROR_MESSAGE("IDeviceContext::CopyTLAS: pSrc must be rebuilded to apply BLAS changes before being copied to another TLAS"); + return false; + } + if (Attribs.Mode == COPY_AS_MODE_CLONE) { auto& SrcDesc = Attribs.pSrc->GetDesc(); @@ -2370,6 +2411,33 @@ bool DeviceContextBase<BaseInterface, ImplementationTraits>::TraceRays(const Tra return false; } +#ifdef DILIGENT_DEVELOPMENT + if (!Attribs.pSBT->Verify()) + { + LOG_ERROR_MESSAGE("IDeviceContext::TraceRays: pSBT content is not valid"); + return false; + } +#endif // DILIGENT_DEVELOPMENT + + if (!m_pPipelineState) + { + LOG_ERROR_MESSAGE("IDeviceContext::TraceRays command arguments are invalid: no pipeline state is bound."); + return false; + } + + if (!m_pPipelineState->GetDesc().IsRayTracingPipeline()) + { + LOG_ERROR_MESSAGE("IDeviceContext::TraceRays command arguments are invalid: pipeline state '", m_pPipelineState->GetDesc().Name, "' is not a ray tracing pipeline."); + return false; + } + + if (Attribs.pSBT->GetDesc().pPSO != m_pPipelineState) + { + LOG_ERROR_MESSAGE("IDeviceContext::TraceRays command arguments are invalid: currently bound pipeline ", m_pPipelineState->GetDesc().Name, + "doesn't match the pipeline ", Attribs.pSBT->GetDesc().pPSO->GetDesc().Name, " that was used in ShaderBindingTable"); + return false; + } + if (Attribs.DimensionX == 0) LOG_WARNING_MESSAGE("IDeviceContext::TraceRays command arguments are invalid: DimensionX is zero."); diff --git a/Graphics/GraphicsEngine/include/ShaderBase.hpp b/Graphics/GraphicsEngine/include/ShaderBase.hpp index 24ad92ee..8b6e9efa 100644 --- a/Graphics/GraphicsEngine/include/ShaderBase.hpp +++ b/Graphics/GraphicsEngine/include/ShaderBase.hpp @@ -79,6 +79,9 @@ public: if ((ShdrDesc.ShaderType == SHADER_TYPE_AMPLIFICATION || ShdrDesc.ShaderType == SHADER_TYPE_MESH) && !deviceFeatures.MeshShaders) LOG_ERROR_AND_THROW("Mesh shaders are not supported by this device"); + + if ((ShdrDesc.ShaderType >= SHADER_TYPE_RAY_GEN && ShdrDesc.ShaderType <= SHADER_TYPE_CALLABLE) && !deviceFeatures.RayTracing) + LOG_ERROR_AND_THROW("Ray tracing shaders are not supported by this device"); } IMPLEMENT_QUERY_INTERFACE_IN_PLACE(IID_Shader, TDeviceObjectBase) diff --git a/Graphics/GraphicsEngine/include/ShaderBindingTableBase.hpp b/Graphics/GraphicsEngine/include/ShaderBindingTableBase.hpp index 35371958..b5642a75 100644 --- a/Graphics/GraphicsEngine/include/ShaderBindingTableBase.hpp +++ b/Graphics/GraphicsEngine/include/ShaderBindingTableBase.hpp @@ -65,41 +65,78 @@ public: TDeviceObjectBase{pRefCounters, pDevice, Desc, bIsDeviceInternal} { ValidateShaderBindingTableDesc(Desc); + + this->m_pPSO = ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO); + this->m_ShaderRecordSize = this->m_pPSO->GetRayTracingPipelineDesc().ShaderRecordSize; + this->m_ShaderRecordStride = this->m_ShaderRecordSize + this->m_pDevice->GetShaderGroupHandleSize(); } ~ShaderBindingTableBase() { } - void BindRayGenShader(const char* ShaderGroupName, const void* Data, Uint32 DataSize) override final + void DILIGENT_CALL_TYPE Reset(const ShaderBindingTableDesc& Desc) override final { - VERIFY(Data == nullptr && DataSize == 0, "not supported yet"); + this->m_RayGenShaderRecord.clear(); + this->m_MissShadersRecord.clear(); + this->m_CallableShadersRecord.clear(); + this->m_HitGroupsRecord.clear(); + this->m_Changed = true; + this->m_pPSO = nullptr; + this->m_Desc = {}; + + try + { + ValidateShaderBindingTableDesc(Desc); + } + catch (const std::runtime_error&) + { + return; + } - this->m_RayGenShaderRecord.resize(this->m_ShaderRecordStride); - ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO)->CopyShaderHandle(ShaderGroupName, this->m_RayGenShaderRecord.data(), this->m_ShaderRecordStride); + this->m_Desc = Desc; + this->m_pPSO = ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO); + this->m_ShaderRecordSize = this->m_pPSO->GetRayTracingPipelineDesc().ShaderRecordSize; + this->m_ShaderRecordStride = this->m_ShaderRecordSize + this->m_pDevice->GetShaderGroupHandleSize(); + } + + void DILIGENT_CALL_TYPE BindRayGenShader(const char* ShaderGroupName, const void* Data, Uint32 DataSize) override final + { + VERIFY_EXPR((Data == nullptr) == (DataSize == 0)); + VERIFY_EXPR(Data == nullptr || (DataSize == this->m_ShaderRecordSize)); + + this->m_RayGenShaderRecord.resize(this->m_ShaderRecordStride, EmptyElem); + this->m_pPSO->CopyShaderHandle(ShaderGroupName, this->m_RayGenShaderRecord.data(), this->m_ShaderRecordStride); + + const Uint32 GroupSize = this->m_pDevice->GetShaderGroupHandleSize(); + std::memcpy(this->m_RayGenShaderRecord.data() + GroupSize, Data, DataSize); this->m_Changed = true; } - void BindMissShader(const char* ShaderGroupName, Uint32 MissIndex, const void* Data, Uint32 DataSize) override final + void DILIGENT_CALL_TYPE BindMissShader(const char* ShaderGroupName, Uint32 MissIndex, const void* Data, Uint32 DataSize) override final { - VERIFY(Data == nullptr && DataSize == 0, "not supported yet"); + VERIFY_EXPR((Data == nullptr) == (DataSize == 0)); + VERIFY_EXPR(Data == nullptr || (DataSize == this->m_ShaderRecordSize)); - const Uint32 Offset = MissIndex * this->m_ShaderRecordStride; - this->m_MissShadersRecord.resize(std::max<size_t>(this->m_MissShadersRecord.size(), Offset + this->m_ShaderRecordStride)); + const Uint32 GroupSize = this->m_pDevice->GetShaderGroupHandleSize(); + const Uint32 Offset = MissIndex * this->m_ShaderRecordStride; + this->m_MissShadersRecord.resize(std::max<size_t>(this->m_MissShadersRecord.size(), Offset + this->m_ShaderRecordStride), EmptyElem); - ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO)->CopyShaderHandle(ShaderGroupName, this->m_MissShadersRecord.data() + Offset, this->m_ShaderRecordStride); + this->m_pPSO->CopyShaderHandle(ShaderGroupName, this->m_MissShadersRecord.data() + Offset, this->m_ShaderRecordStride); + std::memcpy(this->m_MissShadersRecord.data() + Offset + GroupSize, Data, DataSize); this->m_Changed = true; } - void BindHitGroup(ITopLevelAS* pTLAS, - const char* InstanceName, - const char* GeometryName, - Uint32 RayOffsetInHitGroupIndex, - const char* ShaderGroupName, - const void* Data, - Uint32 DataSize) override final + void DILIGENT_CALL_TYPE BindHitGroup(ITopLevelAS* pTLAS, + const char* InstanceName, + const char* GeometryName, + Uint32 RayOffsetInHitGroupIndex, + const char* ShaderGroupName, + const void* Data, + Uint32 DataSize) override final { - VERIFY(Data == nullptr && DataSize == 0, "not supported yet"); + VERIFY_EXPR((Data == nullptr) == (DataSize == 0)); + VERIFY_EXPR(Data == nullptr || (DataSize == this->m_ShaderRecordSize)); VERIFY_EXPR(pTLAS != nullptr); VERIFY_EXPR(RayOffsetInHitGroupIndex < this->m_Desc.HitShadersPerInstance); VERIFY_EXPR(pTLAS->GetDesc().BindingMode == SHADER_BINDING_MODE_PER_GEOMETRY); @@ -111,21 +148,23 @@ public: const Uint32 GeometryIndex = Desc.pBLAS->GetGeometryIndex(GeometryName); const Uint32 Index = InstanceIndex + GeometryIndex * this->m_Desc.HitShadersPerInstance + RayOffsetInHitGroupIndex; const Uint32 Offset = Index * this->m_ShaderRecordStride; + const Uint32 GroupSize = this->m_pDevice->GetShaderGroupHandleSize(); - this->m_HitGroupsRecord.resize(std::max<size_t>(this->m_HitGroupsRecord.size(), Offset + this->m_ShaderRecordStride)); + this->m_HitGroupsRecord.resize(std::max<size_t>(this->m_HitGroupsRecord.size(), Offset + this->m_ShaderRecordStride), EmptyElem); - ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO)->CopyShaderHandle(ShaderGroupName, this->m_HitGroupsRecord.data() + Offset, this->m_ShaderRecordStride); + this->m_pPSO->CopyShaderHandle(ShaderGroupName, this->m_HitGroupsRecord.data() + Offset, this->m_ShaderRecordStride); + std::memcpy(this->m_HitGroupsRecord.data() + Offset + GroupSize, Data, DataSize); this->m_Changed = true; } - void BindHitGroups(ITopLevelAS* pTLAS, - const char* InstanceName, - Uint32 RayOffsetInHitGroupIndex, - const char* ShaderGroupName, - const void* Data, - Uint32 DataSize) override final + void DILIGENT_CALL_TYPE BindHitGroups(ITopLevelAS* pTLAS, + const char* InstanceName, + Uint32 RayOffsetInHitGroupIndex, + const char* ShaderGroupName, + const void* Data, + Uint32 DataSize) override final { - VERIFY(Data == nullptr && DataSize == 0, "not supported yet"); + VERIFY_EXPR((Data == nullptr) == (DataSize == 0)); VERIFY_EXPR(pTLAS != nullptr); VERIFY_EXPR(RayOffsetInHitGroupIndex < this->m_Desc.HitShadersPerInstance); VERIFY_EXPR(pTLAS->GetDesc().BindingMode == SHADER_BINDING_MODE_PER_GEOMETRY || @@ -134,39 +173,64 @@ public: const auto Desc = pTLAS->GetInstanceDesc(InstanceName); VERIFY_EXPR(Desc.pBLAS != nullptr); - const Uint32 InstanceIndex = Desc.ContributionToHitGroupIndex; - const auto& GeometryDesc = Desc.pBLAS->GetDesc(); - const Uint32 GeometryCount = GeometryDesc.BoxCount + GeometryDesc.TriangleCount; - const Uint32 BeginIndex = InstanceIndex + 0 * this->m_Desc.HitShadersPerInstance + RayOffsetInHitGroupIndex; - const Uint32 EndIndex = InstanceIndex + GeometryCount * this->m_Desc.HitShadersPerInstance + RayOffsetInHitGroupIndex; - PipelineStateImplType* pPSO = ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO); + const Uint32 InstanceIndex = Desc.ContributionToHitGroupIndex; + const auto& GeometryDesc = Desc.pBLAS->GetDesc(); + Uint32 GeometryCount = 0; + + switch (pTLAS->GetDesc().BindingMode) + { + // clang-format off + case SHADER_BINDING_MODE_PER_GEOMETRY: GeometryCount = GeometryDesc.BoxCount + GeometryDesc.TriangleCount; break; + case SHADER_BINDING_MODE_PER_INSTANCE: GeometryCount = 1; break; + default: UNEXPECTED("unknown binding mode"); + // clang-format on + } + + VERIFY_EXPR(Data == nullptr || (DataSize == this->m_ShaderRecordSize * GeometryCount)); - this->m_HitGroupsRecord.resize(std::max<size_t>(this->m_HitGroupsRecord.size(), EndIndex * this->m_ShaderRecordStride)); + const Uint32 BeginIndex = InstanceIndex + 0 * this->m_Desc.HitShadersPerInstance + RayOffsetInHitGroupIndex; + const Uint32 EndIndex = InstanceIndex + GeometryCount * this->m_Desc.HitShadersPerInstance + RayOffsetInHitGroupIndex; + const Uint32 GroupSize = this->m_pDevice->GetShaderGroupHandleSize(); + const auto* DataPtr = static_cast<const Uint8*>(Data); + + this->m_HitGroupsRecord.resize(std::max<size_t>(this->m_HitGroupsRecord.size(), EndIndex * this->m_ShaderRecordStride), EmptyElem); for (Uint32 i = 0; i < GeometryCount; ++i) { Uint32 Offset = (BeginIndex + i) * this->m_ShaderRecordStride; - pPSO->CopyShaderHandle(ShaderGroupName, this->m_HitGroupsRecord.data() + Offset, this->m_ShaderRecordStride); + this->m_pPSO->CopyShaderHandle(ShaderGroupName, this->m_HitGroupsRecord.data() + Offset, this->m_ShaderRecordStride); + + std::memcpy(this->m_HitGroupsRecord.data() + Offset + GroupSize, DataPtr, this->m_ShaderRecordSize); + DataPtr += this->m_ShaderRecordSize; } this->m_Changed = true; } - void BindCallableShader(const char* ShaderGroupName, - Uint32 CallableIndex, - const void* Data, - Uint32 DataSize) override final + void DILIGENT_CALL_TYPE BindCallableShader(const char* ShaderGroupName, + Uint32 CallableIndex, + const void* Data, + Uint32 DataSize) override final { - VERIFY(Data == nullptr && DataSize == 0, "not supported yet"); + VERIFY_EXPR((Data == nullptr) == (DataSize == 0)); + VERIFY_EXPR(Data == nullptr || (DataSize == this->m_ShaderRecordSize)); - const Uint32 Offset = CallableIndex * this->m_ShaderRecordStride; - this->m_CallableShadersRecord.resize(std::max<size_t>(this->m_CallableShadersRecord.size(), Offset + this->m_ShaderRecordStride)); + const Uint32 GroupSize = this->m_pDevice->GetShaderGroupHandleSize(); + const Uint32 Offset = CallableIndex * this->m_ShaderRecordStride; + this->m_CallableShadersRecord.resize(std::max<size_t>(this->m_CallableShadersRecord.size(), Offset + this->m_ShaderRecordStride), EmptyElem); - ValidatedCast<PipelineStateImplType>(this->m_Desc.pPSO)->CopyShaderHandle(ShaderGroupName, this->m_CallableShadersRecord.data() + Offset, this->m_ShaderRecordStride); + this->m_pPSO->CopyShaderHandle(ShaderGroupName, this->m_CallableShadersRecord.data() + Offset, this->m_ShaderRecordStride); + std::memcpy(this->m_CallableShadersRecord.data() + Offset + GroupSize, Data, DataSize); this->m_Changed = true; } + Bool DILIGENT_CALL_TYPE Verify() const override final + { + // AZ TODO + return true; + } + protected: - static void ValidateShaderBindingTableDesc(const ShaderBindingTableDesc& Desc) + void ValidateShaderBindingTableDesc(const ShaderBindingTableDesc& Desc) const { #define LOG_SBT_ERROR_AND_THROW(...) LOG_ERROR_AND_THROW("Description of Shader binding table '", (Desc.Name ? Desc.Name : ""), "' is invalid: ", ##__VA_ARGS__) @@ -180,6 +244,20 @@ protected: LOG_SBT_ERROR_AND_THROW("pPSO must be ray tracing pipeline"); } + const auto ShaderGroupHandleSize = this->m_pDevice->GetShaderGroupHandleSize(); + const auto MaxShaderRecordStride = this->m_pDevice->GetMaxShaderRecordStride(); + const auto ShaderRecordSize = Desc.pPSO->GetRayTracingPipelineDesc().ShaderRecordSize; + const auto ShaderRecordStride = ShaderRecordSize + ShaderGroupHandleSize; + + if (ShaderRecordStride > MaxShaderRecordStride) + { + LOG_SBT_ERROR_AND_THROW("ShaderRecordSize(", ShaderRecordSize, ") is too big, max size is: ", MaxShaderRecordStride - ShaderGroupHandleSize); + } + + if (ShaderRecordStride % ShaderGroupHandleSize != 0) + { + LOG_SBT_ERROR_AND_THROW("ShaderRecordSize(", ShaderRecordSize, ") plus ShaderGroupHandleSize(", ShaderGroupHandleSize, ") must be multiple of ", ShaderGroupHandleSize); + } #undef LOG_SBT_ERROR_AND_THROW } @@ -191,8 +269,13 @@ protected: std::vector<Uint8> m_CallableShadersRecord; std::vector<Uint8> m_HitGroupsRecord; + RefCntAutoPtr<PipelineStateImplType> m_pPSO; + + Uint32 m_ShaderRecordSize = 0; Uint32 m_ShaderRecordStride = 0; bool m_Changed = true; + + static const Uint8 EmptyElem = 0xA7; }; } // namespace Diligent diff --git a/Graphics/GraphicsEngine/include/TopLevelASBase.hpp b/Graphics/GraphicsEngine/include/TopLevelASBase.hpp index ab56636f..b03d9e5c 100644 --- a/Graphics/GraphicsEngine/include/TopLevelASBase.hpp +++ b/Graphics/GraphicsEngine/include/TopLevelASBase.hpp @@ -47,7 +47,7 @@ namespace Diligent /// (Diligent::ITopLevelASD3D12 or Diligent::ITopLevelASVk). /// \tparam RenderDeviceImplType - type of the render device implementation /// (Diligent::RenderDeviceD3D12Impl or Diligent::RenderDeviceVkImpl) -template <class BaseInterface, class RenderDeviceImplType> +template <class BaseInterface, class BottomLevelASType, class RenderDeviceImplType> class TopLevelASBase : public DeviceObjectBase<BaseInterface, RenderDeviceImplType, TopLevelASDesc> { public: @@ -73,8 +73,9 @@ public: void SetInstanceData(const TLASBuildInstanceData* pInstances, Uint32 InstanceCount, Uint32 HitShadersPerInstance) { - m_Instances.clear(); - m_StringPool.Release(); + this->m_Instances.clear(); + this->m_StringPool.Release(); + this->m_HitShadersPerInstance = HitShadersPerInstance; size_t StringPoolSize = 0; for (Uint32 i = 0; i < InstanceCount; ++i) @@ -82,30 +83,61 @@ public: StringPoolSize += strlen(pInstances[i].InstanceName) + 1; } - m_StringPool.Reserve(StringPoolSize, GetRawAllocator()); + this->m_StringPool.Reserve(StringPoolSize, GetRawAllocator()); Uint32 InstanceOffset = 0; for (Uint32 i = 0; i < InstanceCount; ++i) { auto& inst = pInstances[i]; - const char* NameCopy = m_StringPool.CopyString(inst.InstanceName); + const char* NameCopy = this->m_StringPool.CopyString(inst.InstanceName); InstanceDesc Desc = {}; Desc.ContributionToHitGroupIndex = inst.ContributionToHitGroupIndex; - Desc.pBLAS = inst.pBLAS; + Desc.pBLAS = ValidatedCast<BottomLevelASType>(inst.pBLAS); + +#ifdef DILIGENT_DEVELOPMENT + Desc.Version = Desc.pBLAS->GetVersion(); +#endif if (Desc.ContributionToHitGroupIndex == TLAS_INSTANCE_OFFSET_AUTO) { Desc.ContributionToHitGroupIndex = InstanceOffset; auto& BLASDesc = Desc.pBLAS->GetDesc(); - InstanceOffset += (BLASDesc.TriangleCount + BLASDesc.BoxCount) * HitShadersPerInstance; + switch (this->m_Desc.BindingMode) + { + // clang-format off + case SHADER_BINDING_MODE_PER_GEOMETRY: InstanceOffset += (BLASDesc.TriangleCount + BLASDesc.BoxCount) * HitShadersPerInstance; break; + case SHADER_BINDING_MODE_PER_INSTANCE: InstanceOffset += HitShadersPerInstance; break; + case SHADER_BINDING_USER_DEFINED: UNEXPECTED("TLAS_INSTANCE_OFFSET_AUTO is not compatible with SHADER_BINDING_USER_DEFINED"); break; + default: UNEXPECTED("unknown ray tracing shader binding mode"); + // clang-format on + } } - bool IsUniqueName = m_Instances.emplace(NameCopy, Desc).second; + bool IsUniqueName = this->m_Instances.emplace(NameCopy, Desc).second; if (!IsUniqueName) LOG_ERROR_AND_THROW("Instance name must be unique!"); } + + VERIFY_EXPR(this->m_StringPool.GetRemainingSize() == 0); + } + + void CopyInstancceData(const TopLevelASBase& Src) + { + this->m_Instances.clear(); + this->m_StringPool.Release(); + this->m_StringPool.Reserve(Src.m_StringPool.GetReservedSize(), GetRawAllocator()); + this->m_HitShadersPerInstance = Src.m_HitShadersPerInstance; + this->m_Desc.BindingMode = Src.m_Desc.BindingMode; + + for (auto& SrcInst : Src.m_Instances) + { + const char* NameCopy = this->m_StringPool.CopyString(SrcInst.first.GetStr()); + this->m_Instances.emplace(NameCopy, SrcInst.second); + } + + VERIFY_EXPR(this->m_StringPool.GetRemainingSize() == 0); } virtual TLASInstanceDesc DILIGENT_CALL_TYPE GetInstanceDesc(const char* Name) const override final @@ -114,11 +146,11 @@ public: TLASInstanceDesc Result = {}; - auto iter = m_Instances.find(Name); - if (iter != m_Instances.end()) + auto iter = this->m_Instances.find(Name); + if (iter != this->m_Instances.end()) { Result.ContributionToHitGroupIndex = iter->second.ContributionToHitGroupIndex; - Result.pBLAS = iter->second.pBLAS; + Result.pBLAS = iter->second.pBLAS.RawPtr<IBottomLevelAS>(); } else { @@ -150,6 +182,22 @@ public: return (this->m_State & State) == State; } +#ifdef DILIGENT_DEVELOPMENT + bool CheckBLASVersion() const + { + for (auto& NameAndInst : m_Instances) + { + auto& Inst = NameAndInst.second; + if (Inst.Version != Inst.pBLAS->GetVersion()) + { + LOG_ERROR_MESSAGE("Instance with name ('", NameAndInst.first.GetStr(), "') has BLAS that was changed after TLAS build, you must rebuild TLAS."); + return false; + } + } + return true; + } +#endif + protected: static void ValidateTopLevelASDesc(const TopLevelASDesc& Desc) { @@ -172,14 +220,19 @@ protected: IMPLEMENT_QUERY_INTERFACE_IN_PLACE(IID_TopLevelAS, TDeviceObjectBase) protected: - RESOURCE_STATE m_State = RESOURCE_STATE_UNKNOWN; + RESOURCE_STATE m_State = RESOURCE_STATE_UNKNOWN; + Uint32 m_HitShadersPerInstance = 0; StringPool m_StringPool; struct InstanceDesc { - Uint32 ContributionToHitGroupIndex = 0; - mutable RefCntAutoPtr<IBottomLevelAS> pBLAS; + Uint32 ContributionToHitGroupIndex = 0; + RefCntAutoPtr<BottomLevelASType> pBLAS; + +#ifdef DILIGENT_DEVELOPMENT + Uint32 Version = 0; +#endif }; std::unordered_map<HashMapStringKey, InstanceDesc, HashMapStringKey::Hasher> m_Instances; }; diff --git a/Graphics/GraphicsEngine/interface/DeviceContext.h b/Graphics/GraphicsEngine/interface/DeviceContext.h index 35e17e91..0cb7fe7d 100644 --- a/Graphics/GraphicsEngine/interface/DeviceContext.h +++ b/Graphics/GraphicsEngine/interface/DeviceContext.h @@ -741,7 +741,7 @@ DILIGENT_TYPED_ENUM(RAYTRACING_INSTANCE_FLAGS, Uint8) /// geometries referenced by this instance. This behavior can be overridden by the SPIR-V OpaqueKHR ray flag. RAYTRACING_INSTANCE_FORCE_NO_OPAQUE = 0x08, - RAYTRACING_INSTANCE_FLAGS_LAST = 0x08 + RAYTRACING_INSTANCE_FLAGS_LAST = RAYTRACING_INSTANCE_FORCE_NO_OPAQUE }; DEFINE_FLAG_ENUM_OPERATORS(RAYTRACING_INSTANCE_FLAGS) @@ -757,7 +757,7 @@ DILIGENT_TYPED_ENUM(COPY_AS_MODE, Uint8) // after the build of the acceleration structure specified by src. //COPY_AS_MODE_COMPACT, - COPY_AS_MODE_LAST = 0, + COPY_AS_MODE_LAST = COPY_AS_MODE_CLONE, }; /// Defines geometry flags for ray tracing. @@ -775,7 +775,7 @@ DILIGENT_TYPED_ENUM(RAYTRACING_GEOMETRY_FLAGS, Uint8) /// If this bit is absent an implementation may invoke the any-hit shader more than once for this geometry. RAYTRACING_GEOMETRY_NO_DUPLICATE_ANY_HIT_INVOCATION = 0x02, - RAYTRACING_GEOMETRY_FLAGS_LAST = 0x02 + RAYTRACING_GEOMETRY_FLAGS_LAST = RAYTRACING_GEOMETRY_NO_DUPLICATE_ANY_HIT_INVOCATION }; DEFINE_FLAG_ENUM_OPERATORS(RAYTRACING_GEOMETRY_FLAGS) @@ -910,6 +910,35 @@ static const Uint32 TLAS_INSTANCE_OFFSET_AUTO = ~0u; /// AZ TODO static const Uint32 TLAS_INSTANCE_DATA_SIZE = 64; +/// AZ TODO +struct InstanceMatrix +{ + /// rotation translation + /// (0 1 2) [ 3] + /// (4 5 6) [ 7] + /// (8 9 10) [11] + float data [3][4]; + +#if DILIGENT_CPP_INTERFACE + /// AZ TODO + InstanceMatrix() noexcept : + data{{1.0f, 0.0f, 0.0f, 0.0f}, + {0.0f, 1.0f, 0.0f, 0.0f}, + {0.0f, 0.0f, 1.0f, 0.0f}} + {} + + InstanceMatrix(const InstanceMatrix&) noexcept = default; + + InstanceMatrix& SetTranslation(float x, float y, float z) noexcept + { + data[0][3] = x; + data[1][3] = y; + data[2][3] = z; + return *this; + } +#endif +}; +typedef struct InstanceMatrix InstanceMatrix; /// AZ TODO struct TLASBuildInstanceData @@ -921,7 +950,7 @@ struct TLASBuildInstanceData IBottomLevelAS* pBLAS DEFAULT_INITIALIZER(nullptr); // can be null to deactive instance /// AZ TODO - float Transform[3][4] DEFAULT_INITIALIZER({}); + InstanceMatrix Transform; /// AZ TODO Uint32 CustomId DEFAULT_INITIALIZER(0); // 24 bits, in shader: gl_InstanceCustomIndexNV for GLSL, InstanceID() for HLSL diff --git a/Graphics/GraphicsEngine/interface/PipelineState.h b/Graphics/GraphicsEngine/interface/PipelineState.h index 4115a9d1..ebec8344 100644 --- a/Graphics/GraphicsEngine/interface/PipelineState.h +++ b/Graphics/GraphicsEngine/interface/PipelineState.h @@ -299,8 +299,11 @@ typedef struct RayTracingProceduralHitShaderGroup RayTracingProceduralHitShaderG /// AZ TODO struct RayTracingPipelineDesc { + // Size of the additional data passed to the shader. + Uint16 ShaderRecordSize DEFAULT_INITIALIZER(0); + /// AZ TODO - Uint8 MaxRecursionDepth DEFAULT_INITIALIZER(0); // must be 0..31 (check current device limits) + Uint8 MaxRecursionDepth DEFAULT_INITIALIZER(0); // must be 0..31 (check current device limits) }; typedef struct RayTracingPipelineDesc RayTracingPipelineDesc; @@ -438,7 +441,7 @@ typedef struct ComputePipelineStateCreateInfo ComputePipelineStateCreateInfo; struct RayTracingPipelineStateCreateInfo DILIGENT_DERIVE(PipelineStateCreateInfo) /// AZ TODO - RayTracingPipelineDesc RayTracingPipeline; + RayTracingPipelineDesc RayTracingPipeline; /// AZ TODO const RayTracingGeneralShaderGroup* pGeneralShaders DEFAULT_INITIALIZER(nullptr); @@ -457,6 +460,10 @@ struct RayTracingPipelineStateCreateInfo DILIGENT_DERIVE(PipelineStateCreateInfo /// AZ TODO Uint16 ProceduralHitShaderCount DEFAULT_INITIALIZER(0); + + /// Direct3D12 only: set name of constant buffer that will be used by local root signature. + /// Ignored if RayTracingPipelineDesc::ShaderRecordSize is zero. + const char* ShaderRecordName DEFAULT_INITIALIZER(nullptr); }; typedef struct RayTracingPipelineStateCreateInfo RayTracingPipelineStateCreateInfo; diff --git a/Graphics/GraphicsEngine/interface/ShaderBindingTable.h b/Graphics/GraphicsEngine/interface/ShaderBindingTable.h index 2a5e587a..d6f0740f 100644 --- a/Graphics/GraphicsEngine/interface/ShaderBindingTable.h +++ b/Graphics/GraphicsEngine/interface/ShaderBindingTable.h @@ -51,9 +51,6 @@ struct ShaderBindingTableDesc DILIGENT_DERIVE(DeviceObjectAttribs) /// AZ TODO IPipelineState* pPSO DEFAULT_INITIALIZER(nullptr); - - // Size of the additional data passed to the shader, maximum size is 4064 bytes. - Uint32 ShaderRecordSize DEFAULT_INITIALIZER(0); /// AZ TODO Uint32 HitShadersPerInstance DEFAULT_INITIALIZER(1); @@ -114,7 +111,7 @@ DILIGENT_BEGIN_INTERFACE(IShaderBindingTable, IDeviceObject) #endif /// AZ TODO - VIRTUAL void METHOD(Verify)(THIS) CONST PURE; + VIRTUAL Bool METHOD(Verify)(THIS) CONST PURE; /// AZ TODO VIRTUAL void METHOD(Reset)(THIS_ diff --git a/Graphics/GraphicsEngineD3D11/include/DeviceContextD3D11Impl.hpp b/Graphics/GraphicsEngineD3D11/include/DeviceContextD3D11Impl.hpp index 5e69da8e..5be505a1 100644 --- a/Graphics/GraphicsEngineD3D11/include/DeviceContextD3D11Impl.hpp +++ b/Graphics/GraphicsEngineD3D11/include/DeviceContextD3D11Impl.hpp @@ -61,7 +61,7 @@ struct DeviceContextD3D11ImplTraits using FramebufferType = FramebufferD3D11Impl; using RenderPassType = RenderPassD3D11Impl; using BottomLevelASType = BottomLevelASBase<IBottomLevelAS, RenderDeviceD3D11Impl>; - using TopLevelASType = TopLevelASBase<ITopLevelAS, RenderDeviceD3D11Impl>; + using TopLevelASType = TopLevelASBase<ITopLevelAS, BottomLevelASType, RenderDeviceD3D11Impl>; }; /// Device context implementation in Direct3D11 backend. diff --git a/Graphics/GraphicsEngineD3D12/include/PipelineStateD3D12Impl.hpp b/Graphics/GraphicsEngineD3D12/include/PipelineStateD3D12Impl.hpp index 6a27a029..f336c6ff 100644 --- a/Graphics/GraphicsEngineD3D12/include/PipelineStateD3D12Impl.hpp +++ b/Graphics/GraphicsEngineD3D12/include/PipelineStateD3D12Impl.hpp @@ -144,8 +144,11 @@ private: void Destruct(); - CComPtr<ID3D12DeviceChild> m_pd3d12PSO; - RootSignature m_RootSig; + void CreateLocalRootSignature(const RayTracingPipelineDesc& Desc); + + CComPtr<ID3D12DeviceChild> m_pd3d12PSO; + RootSignature m_RootSig; + CComPtr<ID3D12RootSignature> m_LocalRootSignature; // Must be defined before default SRB SRBMemoryAllocator m_SRBMemAllocator; diff --git a/Graphics/GraphicsEngineD3D12/include/RenderDeviceD3D12Impl.hpp b/Graphics/GraphicsEngineD3D12/include/RenderDeviceD3D12Impl.hpp index ab72e22a..636175d5 100644 --- a/Graphics/GraphicsEngineD3D12/include/RenderDeviceD3D12Impl.hpp +++ b/Graphics/GraphicsEngineD3D12/include/RenderDeviceD3D12Impl.hpp @@ -178,10 +178,8 @@ public: ShaderVersion GetMaxShaderModel() const; D3D_FEATURE_LEVEL GetD3DFeatureLevel() const; - static Uint32 GetShaderGroupHandleSize() - { - return D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES; - } + static Uint32 GetShaderGroupHandleSize() { return D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES; } + static Uint32 GetMaxShaderRecordStride() { return D3D12_RAYTRACING_MAX_SHADER_RECORD_STRIDE; } private: template <typename PSOCreateInfoType> diff --git a/Graphics/GraphicsEngineD3D12/include/ShaderBindingTableD3D12Impl.hpp b/Graphics/GraphicsEngineD3D12/include/ShaderBindingTableD3D12Impl.hpp index e866c6f2..daf84d14 100644 --- a/Graphics/GraphicsEngineD3D12/include/ShaderBindingTableD3D12Impl.hpp +++ b/Graphics/GraphicsEngineD3D12/include/ShaderBindingTableD3D12Impl.hpp @@ -54,10 +54,6 @@ public: virtual void DILIGENT_CALL_TYPE QueryInterface(const INTERFACE_ID& IID, IObject** ppInterface) override final; - virtual void DILIGENT_CALL_TYPE Verify() const override; - - virtual void DILIGENT_CALL_TYPE Reset(const ShaderBindingTableDesc& Desc) override; - virtual void DILIGENT_CALL_TYPE ResetHitGroups(Uint32 HitShadersPerInstance) override; virtual void DILIGENT_CALL_TYPE BindAll(const BindAllAttribs& Attribs) override; @@ -70,9 +66,6 @@ public: D3D12_GPU_VIRTUAL_ADDRESS_RANGE_AND_STRIDE& CallableShaderBindingTable) override; private: - void ValidateDesc(const ShaderBindingTableDesc& Desc) const; - -private: RefCntAutoPtr<IBuffer> m_pBuffer; }; diff --git a/Graphics/GraphicsEngineD3D12/include/TopLevelASD3D12Impl.hpp b/Graphics/GraphicsEngineD3D12/include/TopLevelASD3D12Impl.hpp index 8eccf530..c61ab368 100644 --- a/Graphics/GraphicsEngineD3D12/include/TopLevelASD3D12Impl.hpp +++ b/Graphics/GraphicsEngineD3D12/include/TopLevelASD3D12Impl.hpp @@ -33,6 +33,7 @@ #include "TopLevelASD3D12.h" #include "RenderDeviceD3D12.h" #include "TopLevelASBase.hpp" +#include "BottomLevelASD3D12Impl.hpp" #include "D3D12ResourceBase.hpp" #include "RenderDeviceD3D12Impl.hpp" @@ -40,10 +41,10 @@ namespace Diligent { /// Top-level acceleration structure object implementation in Direct3D12 backend. -class TopLevelASD3D12Impl final : public TopLevelASBase<ITopLevelASD3D12, RenderDeviceD3D12Impl>, public D3D12ResourceBase +class TopLevelASD3D12Impl final : public TopLevelASBase<ITopLevelASD3D12, BottomLevelASD3D12Impl, RenderDeviceD3D12Impl>, public D3D12ResourceBase { public: - using TTopLevelASBase = TopLevelASBase<ITopLevelASD3D12, RenderDeviceD3D12Impl>; + using TTopLevelASBase = TopLevelASBase<ITopLevelASD3D12, BottomLevelASD3D12Impl, RenderDeviceD3D12Impl>; TopLevelASD3D12Impl(IReferenceCounters* pRefCounters, class RenderDeviceD3D12Impl* pDeviceD3D12, diff --git a/Graphics/GraphicsEngineD3D12/interface/ShaderBindingTableD3D12.h b/Graphics/GraphicsEngineD3D12/interface/ShaderBindingTableD3D12.h index 4b0835cd..aa50251a 100644 --- a/Graphics/GraphicsEngineD3D12/interface/ShaderBindingTableD3D12.h +++ b/Graphics/GraphicsEngineD3D12/interface/ShaderBindingTableD3D12.h @@ -65,6 +65,7 @@ DILIGENT_END_INTERFACE #if DILIGENT_C_INTERFACE +# define IShaderBindingTableD3D12_GetD3D12AddressRangeAndStride(This, ...) CALL_IFACE_METHOD(ShaderBindingTableD3D12, GetD3D12AddressRangeAndStride, This, __VA_ARGS__) #endif diff --git a/Graphics/GraphicsEngineD3D12/src/DeviceContextD3D12Impl.cpp b/Graphics/GraphicsEngineD3D12/src/DeviceContextD3D12Impl.cpp index 3cd896ac..0adcaa59 100644 --- a/Graphics/GraphicsEngineD3D12/src/DeviceContextD3D12Impl.cpp +++ b/Graphics/GraphicsEngineD3D12/src/DeviceContextD3D12Impl.cpp @@ -2198,8 +2198,6 @@ void DeviceContextD3D12Impl::TransitionOrVerifyTLASState(CommandContext& RESOURCE_STATE RequiredState, const char* OperationName) { - // AZ TODO: transit BLAS state too? - if (TransitionMode == RESOURCE_STATE_TRANSITION_MODE_TRANSITION) { if (TLAS.IsInKnownState() && !TLAS.CheckState(RequiredState)) @@ -2210,6 +2208,11 @@ void DeviceContextD3D12Impl::TransitionOrVerifyTLASState(CommandContext& { DvpVerifyTLASState(TLAS, RequiredState, OperationName); } + + if (RequiredState & RESOURCE_STATE_RAY_TRACING) + { + TLAS.CheckBLASVersion(); + } #endif } @@ -2318,6 +2321,8 @@ void DeviceContextD3D12Impl::BuildBLAS(const BLASBuildAttribs& Attribs) d3d12Tris.VertexBuffer.StartAddress = pVB->GetGPUAddress() + SrcTris.VertexOffset; d3d12Tris.VertexBuffer.StrideInBytes = SrcTris.VertexStride; + TransitionOrVerifyBufferState(CmdCtx, *pVB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, OpName); + if (SrcTris.pIndexBuffer) { auto* const pIB = ValidatedCast<BufferD3D12Impl>(SrcTris.pIndexBuffer); @@ -2389,6 +2394,10 @@ void DeviceContextD3D12Impl::BuildBLAS(const BLASBuildAttribs& Attribs) CmdCtx.AsGraphicsContext4().BuildRaytracingAccelerationStructure(Desc, 0, nullptr); ++m_State.NumCommands; + +#ifdef DILIGENT_DEVELOPMENT + pBLASD12->UpdateVersion(); +#endif } void DeviceContextD3D12Impl::BuildTLAS(const TLASBuildAttribs& Attribs) @@ -2401,7 +2410,6 @@ void DeviceContextD3D12Impl::BuildTLAS(const TLASBuildAttribs& Attribs) auto* pTLASD12 = ValidatedCast<TopLevelASD3D12Impl>(Attribs.pTLAS); auto* pScratchD12 = ValidatedCast<BufferD3D12Impl>(Attribs.pScratchBuffer); auto* pInstancesD12 = ValidatedCast<BufferD3D12Impl>(Attribs.pInstanceBuffer); - //auto& TLASDesc = pTLASD12->GetDesc(); auto& CmdCtx = GetCmdContext(); const char* OpName = "Build TopLevelAS (DeviceContextD3D12Impl::BuildTLAS)"; @@ -2423,7 +2431,7 @@ void DeviceContextD3D12Impl::BuildTLAS(const TLASBuildAttribs& Attribs) auto* const pBLASD12 = ValidatedCast<BottomLevelASD3D12Impl>(Inst.pBLAS); static_assert(sizeof(d3d12Inst.Transform) == sizeof(Inst.Transform), "size mismatch"); - std::memcpy(&d3d12Inst.Transform, Inst.Transform, sizeof(d3d12Inst.Transform)); + std::memcpy(&d3d12Inst.Transform, Inst.Transform.data, sizeof(d3d12Inst.Transform)); d3d12Inst.InstanceID = Inst.CustomId; d3d12Inst.InstanceContributionToHitGroupIndex = pTLASD12->GetInstanceDesc(Inst.InstanceName).ContributionToHitGroupIndex; // AZ TODO: optimize @@ -2458,7 +2466,20 @@ void DeviceContextD3D12Impl::CopyBLAS(const CopyBLASAttribs& Attribs) if (!TDeviceContextBase::CopyBLAS(Attribs, 0)) return; - // AZ TODO + auto* pSrcD3D12 = ValidatedCast<BottomLevelASD3D12Impl>(Attribs.pSrc); + auto* pDstD3D12 = ValidatedCast<BottomLevelASD3D12Impl>(Attribs.pDst); + auto& CmdCtx = GetCmdContext(); + + const char* OpName = "Copy BottomLevelAS (DeviceContextD3D12Impl::CopyBLAS)"; + TransitionOrVerifyBLASState(CmdCtx, *pSrcD3D12, Attribs.TransitionMode, RESOURCE_STATE_BUILD_AS_READ, OpName); + TransitionOrVerifyBLASState(CmdCtx, *pDstD3D12, Attribs.TransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, OpName); + + CmdCtx.AsGraphicsContext4().CopyRaytracingAccelerationStructure(pSrcD3D12->GetGPUAddress(), pDstD3D12->GetGPUAddress(), D3D12_RAYTRACING_ACCELERATION_STRUCTURE_COPY_MODE_CLONE); + ++m_State.NumCommands; + +#ifdef DILIGENT_DEVELOPMENT + pDstD3D12->UpdateVersion(); +#endif } void DeviceContextD3D12Impl::CopyTLAS(const CopyTLASAttribs& Attribs) @@ -2466,7 +2487,18 @@ void DeviceContextD3D12Impl::CopyTLAS(const CopyTLASAttribs& Attribs) if (!TDeviceContextBase::CopyTLAS(Attribs, 0)) return; - // AZ TODO + auto* pSrcD3D12 = ValidatedCast<TopLevelASD3D12Impl>(Attribs.pSrc); + auto* pDstD3D12 = ValidatedCast<TopLevelASD3D12Impl>(Attribs.pDst); + auto& CmdCtx = GetCmdContext(); + + pDstD3D12->CopyInstancceData(*pSrcD3D12); + + const char* OpName = "Copy BottomLevelAS (DeviceContextD3D12Impl::CopyTLAS)"; + TransitionOrVerifyTLASState(CmdCtx, *pSrcD3D12, Attribs.TransitionMode, RESOURCE_STATE_BUILD_AS_READ, OpName); + TransitionOrVerifyTLASState(CmdCtx, *pDstD3D12, Attribs.TransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, OpName); + + CmdCtx.AsGraphicsContext4().CopyRaytracingAccelerationStructure(pSrcD3D12->GetGPUAddress(), pDstD3D12->GetGPUAddress(), D3D12_RAYTRACING_ACCELERATION_STRUCTURE_COPY_MODE_CLONE); + ++m_State.NumCommands; } void DeviceContextD3D12Impl::TraceRays(const TraceRaysAttribs& Attribs) diff --git a/Graphics/GraphicsEngineD3D12/src/PipelineStateD3D12Impl.cpp b/Graphics/GraphicsEngineD3D12/src/PipelineStateD3D12Impl.cpp index 405f6d4b..b1fe19fe 100644 --- a/Graphics/GraphicsEngineD3D12/src/PipelineStateD3D12Impl.cpp +++ b/Graphics/GraphicsEngineD3D12/src/PipelineStateD3D12Impl.cpp @@ -224,7 +224,7 @@ void BuildRTPipelineDescription(const RayTracingPipelineStateCreateInfo& CreateI } template <typename TNameToGroupIndexMap> -void GetShaderIdentifiers(ID3D12StateObject* pSO, +void GetShaderIdentifiers(ID3D12DeviceChild* pSO, const RayTracingPipelineStateCreateInfo& CreateInfo, const TNameToGroupIndexMap& NameToGroupIndex, Uint8* ShaderData) @@ -625,6 +625,8 @@ PipelineStateD3D12Impl::PipelineStateD3D12Impl(IReferenceCounters* { try { + CreateLocalRootSignature(CreateInfo.RayTracingPipeline); + TShaderStages ShaderStages; std::vector<D3D12_STATE_SUBOBJECT> Subobjects; DynamicLinearAllocator TempPool{GetRawAllocator(), 4 << 10}; @@ -640,21 +642,21 @@ PipelineStateD3D12Impl::PipelineStateD3D12Impl(IReferenceCounters* D3D12_GLOBAL_ROOT_SIGNATURE GlobalRoot = {m_RootSig.GetD3D12RootSignature()}; Subobjects.push_back({D3D12_STATE_SUBOBJECT_TYPE_GLOBAL_ROOT_SIGNATURE, &GlobalRoot}); + D3D12_LOCAL_ROOT_SIGNATURE LocalRoot = {m_LocalRootSignature}; + if (m_LocalRootSignature) + Subobjects.push_back({D3D12_STATE_SUBOBJECT_TYPE_LOCAL_ROOT_SIGNATURE, &LocalRoot}); + D3D12_STATE_OBJECT_DESC RTPipelineDesc = {}; RTPipelineDesc.Type = D3D12_STATE_OBJECT_TYPE_RAYTRACING_PIPELINE; RTPipelineDesc.NumSubobjects = static_cast<UINT>(Subobjects.size()); RTPipelineDesc.pSubobjects = Subobjects.data(); - CComPtr<ID3D12StateObject> pSO; - auto pd3d12Device = pDeviceD3D12->GetD3D12Device5(); - HRESULT hr = pd3d12Device->CreateStateObject(&RTPipelineDesc, IID_PPV_ARGS(&pSO)); + HRESULT hr = pd3d12Device->CreateStateObject(&RTPipelineDesc, IID_PPV_ARGS(&m_pd3d12PSO)); if (FAILED(hr)) LOG_ERROR_AND_THROW("Failed to create ray tracing state object"); - m_pd3d12PSO = pSO; - - GetShaderIdentifiers(pSO, CreateInfo, m_pRayTracingPipelineData->NameToGroupIndex, m_pRayTracingPipelineData->Shaders); + GetShaderIdentifiers(m_pd3d12PSO, CreateInfo, m_pRayTracingPipelineData->NameToGroupIndex, m_pRayTracingPipelineData->Shaders); if (*m_Desc.Name != 0) { @@ -672,6 +674,35 @@ PipelineStateD3D12Impl::PipelineStateD3D12Impl(IReferenceCounters* } } +void PipelineStateD3D12Impl::CreateLocalRootSignature(const RayTracingPipelineDesc& Desc) +{ + // AZ TODO + /*if (Desc.ShaderRecordSize == 0) + return; + + D3D12_ROOT_SIGNATURE_DESC d3d12RootSignatureDesc = {}; + D3D12_ROOT_PARAMETER d3d12Params = {}; + + d3d12Params.ParameterType = D3D12_ROOT_PARAMETER_TYPE_32BIT_CONSTANTS; + d3d12Params.ShaderVisibility = D3D12_SHADER_VISIBILITY_ALL; + d3d12Params.Constants.Num32BitValues = Desc.ShaderRecordSize / 4; + d3d12Params.Constants.RegisterSpace = Desc.LocalRootRegisterSpace; + d3d12Params.Constants.ShaderRegister = 0; + + d3d12RootSignatureDesc.Flags = D3D12_ROOT_SIGNATURE_FLAG_LOCAL_ROOT_SIGNATURE; + d3d12RootSignatureDesc.NumParameters = 1; + d3d12RootSignatureDesc.pParameters = &d3d12Params; + + CComPtr<ID3DBlob> signature; + auto hr = D3D12SerializeRootSignature(&d3d12RootSignatureDesc, D3D_ROOT_SIGNATURE_VERSION_1, &signature, nullptr); + CHECK_D3D_RESULT_THROW(hr, "Failed to serialize root signature"); + + auto pd3d12Device = GetDevice()->GetD3D12Device(); + + hr = pd3d12Device->CreateRootSignature(0, signature->GetBufferPointer(), signature->GetBufferSize(), IID_PPV_ARGS(&m_LocalRootSignature)); + CHECK_D3D_RESULT_THROW(hr, "Failed to create root signature");*/ +} + PipelineStateD3D12Impl::~PipelineStateD3D12Impl() { Destruct(); diff --git a/Graphics/GraphicsEngineD3D12/src/RootSignature.cpp b/Graphics/GraphicsEngineD3D12/src/RootSignature.cpp index d0db2201..e6af12c8 100644 --- a/Graphics/GraphicsEngineD3D12/src/RootSignature.cpp +++ b/Graphics/GraphicsEngineD3D12/src/RootSignature.cpp @@ -705,7 +705,7 @@ __forceinline void TransitionResource(CommandContext& Ctx, { VERIFY(RangeType == D3D12_DESCRIPTOR_RANGE_TYPE_SRV, "Unexpected descriptor range type"); auto* pTLASD3D12 = Res.pObject.RawPtr<TopLevelASD3D12Impl>(); - if (pTLASD3D12->IsInKnownState() && !pTLASD3D12->CheckState(RESOURCE_STATE_RAY_TRACING)) + if (pTLASD3D12->IsInKnownState()) Ctx.TransitionResource(pTLASD3D12, RESOURCE_STATE_RAY_TRACING); } break; diff --git a/Graphics/GraphicsEngineD3D12/src/ShaderBindingTableD3D12Impl.cpp b/Graphics/GraphicsEngineD3D12/src/ShaderBindingTableD3D12Impl.cpp index d352caf0..d71d98b7 100644 --- a/Graphics/GraphicsEngineD3D12/src/ShaderBindingTableD3D12Impl.cpp +++ b/Graphics/GraphicsEngineD3D12/src/ShaderBindingTableD3D12Impl.cpp @@ -44,9 +44,6 @@ ShaderBindingTableD3D12Impl::ShaderBindingTableD3D12Impl(IReferenceCounters* bool bIsDeviceInternal) : TShaderBindingTableBase{pRefCounters, pDeviceD3D12, Desc, bIsDeviceInternal} { - ValidateDesc(Desc); - - m_ShaderRecordStride = m_Desc.ShaderRecordSize + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES; } ShaderBindingTableD3D12Impl::~ShaderBindingTableD3D12Impl() @@ -55,43 +52,6 @@ ShaderBindingTableD3D12Impl::~ShaderBindingTableD3D12Impl() IMPLEMENT_QUERY_INTERFACE(ShaderBindingTableD3D12Impl, IID_ShaderBindingTableD3D12, TShaderBindingTableBase) -void ShaderBindingTableD3D12Impl::ValidateDesc(const ShaderBindingTableDesc& Desc) const -{ - if (Desc.ShaderRecordSize + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES > D3D12_RAYTRACING_MAX_SHADER_RECORD_STRIDE) - { - LOG_ERROR_AND_THROW("Description of Shader binding table '", (Desc.Name ? Desc.Name : ""), - "' is invalid: ShaderRecordSize is too big, max size is: ", D3D12_RAYTRACING_MAX_SHADER_RECORD_STRIDE - D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES); - } -} - -void ShaderBindingTableD3D12Impl::Verify() const -{ - // AZ TODO -} - -void ShaderBindingTableD3D12Impl::Reset(const ShaderBindingTableDesc& Desc) -{ - m_RayGenShaderRecord.clear(); - m_MissShadersRecord.clear(); - m_CallableShadersRecord.clear(); - m_HitGroupsRecord.clear(); - m_Changed = true; - - try - { - ValidateShaderBindingTableDesc(Desc); - ValidateDesc(Desc); - } - catch (const std::runtime_error&) - { - // AZ TODO - return; - } - - m_Desc = Desc; - m_ShaderRecordStride = m_Desc.ShaderRecordSize + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES; -} - void ShaderBindingTableD3D12Impl::ResetHitGroups(Uint32 HitShadersPerInstance) { // AZ TODO diff --git a/Graphics/GraphicsEngineOpenGL/include/DeviceContextGLImpl.hpp b/Graphics/GraphicsEngineOpenGL/include/DeviceContextGLImpl.hpp index 0b2154e6..32055367 100644 --- a/Graphics/GraphicsEngineOpenGL/include/DeviceContextGLImpl.hpp +++ b/Graphics/GraphicsEngineOpenGL/include/DeviceContextGLImpl.hpp @@ -56,7 +56,7 @@ struct DeviceContextGLImplTraits using FramebufferType = FramebufferGLImpl; using RenderPassType = RenderPassGLImpl; using BottomLevelASType = BottomLevelASBase<IBottomLevelAS, RenderDeviceGLImpl>; - using TopLevelASType = TopLevelASBase<ITopLevelAS, RenderDeviceGLImpl>; + using TopLevelASType = TopLevelASBase<ITopLevelAS, BottomLevelASType, RenderDeviceGLImpl>; }; /// Device context implementation in OpenGL backend. diff --git a/Graphics/GraphicsEngineVulkan/include/RenderDeviceVkImpl.hpp b/Graphics/GraphicsEngineVulkan/include/RenderDeviceVkImpl.hpp index e13f1a76..5440a6c8 100644 --- a/Graphics/GraphicsEngineVulkan/include/RenderDeviceVkImpl.hpp +++ b/Graphics/GraphicsEngineVulkan/include/RenderDeviceVkImpl.hpp @@ -201,6 +201,10 @@ public: { return GetPhysicalDevice().GetExtProperties().RayTracing.shaderGroupHandleSize; } + Uint32 GetMaxShaderRecordStride() const + { + return GetPhysicalDevice().GetExtProperties().RayTracing.maxShaderGroupStride; + } private: template <typename PSOCreateInfoType> diff --git a/Graphics/GraphicsEngineVulkan/include/ShaderBindingTableVkImpl.hpp b/Graphics/GraphicsEngineVulkan/include/ShaderBindingTableVkImpl.hpp index 1b2db950..cef50a4e 100644 --- a/Graphics/GraphicsEngineVulkan/include/ShaderBindingTableVkImpl.hpp +++ b/Graphics/GraphicsEngineVulkan/include/ShaderBindingTableVkImpl.hpp @@ -51,10 +51,6 @@ public: bool bIsDeviceInternal = false); ~ShaderBindingTableVkImpl(); - virtual void DILIGENT_CALL_TYPE Verify() const override; - - virtual void DILIGENT_CALL_TYPE Reset(const ShaderBindingTableDesc& Desc) override; - virtual void DILIGENT_CALL_TYPE ResetHitGroups(Uint32 HitShadersPerInstance) override; virtual void DILIGENT_CALL_TYPE BindAll(const BindAllAttribs& Attribs) override; @@ -68,9 +64,6 @@ public: IMPLEMENT_QUERY_INTERFACE_IN_PLACE(IID_ShaderBindingTableVk, TShaderBindingTableBase); private: - void ValidateDesc(const ShaderBindingTableDesc& Desc) const; - -private: RefCntAutoPtr<IBuffer> m_pBuffer; }; diff --git a/Graphics/GraphicsEngineVulkan/include/TopLevelASVkImpl.hpp b/Graphics/GraphicsEngineVulkan/include/TopLevelASVkImpl.hpp index b2801eca..0f8a94c8 100644 --- a/Graphics/GraphicsEngineVulkan/include/TopLevelASVkImpl.hpp +++ b/Graphics/GraphicsEngineVulkan/include/TopLevelASVkImpl.hpp @@ -34,15 +34,16 @@ #include "RenderDeviceVkImpl.hpp" #include "TopLevelASVk.h" #include "TopLevelASBase.hpp" +#include "BottomLevelASVkImpl.hpp" #include "VulkanUtilities/VulkanObjectWrappers.hpp" namespace Diligent { -class TopLevelASVkImpl final : public TopLevelASBase<ITopLevelASVk, RenderDeviceVkImpl> +class TopLevelASVkImpl final : public TopLevelASBase<ITopLevelASVk, BottomLevelASVkImpl, RenderDeviceVkImpl> { public: - using TTopLevelASBase = TopLevelASBase<ITopLevelASVk, RenderDeviceVkImpl>; + using TTopLevelASBase = TopLevelASBase<ITopLevelASVk, BottomLevelASVkImpl, RenderDeviceVkImpl>; TopLevelASVkImpl(IReferenceCounters* pRefCounters, RenderDeviceVkImpl* pRenderDeviceVk, diff --git a/Graphics/GraphicsEngineVulkan/src/DeviceContextVkImpl.cpp b/Graphics/GraphicsEngineVulkan/src/DeviceContextVkImpl.cpp index 69a3a1ba..a77fb96d 100644 --- a/Graphics/GraphicsEngineVulkan/src/DeviceContextVkImpl.cpp +++ b/Graphics/GraphicsEngineVulkan/src/DeviceContextVkImpl.cpp @@ -2334,6 +2334,22 @@ void DeviceContextVkImpl::TransitionImageLayout(ITexture* pTexture, VkImageLayou } } +namespace +{ +NODISCARD inline bool ResourceStateHasWriteAccess(RESOURCE_STATE State) +{ + static_assert(RESOURCE_STATE_MAX_BIT == RESOURCE_STATE_RAY_TRACING, "This function must be updated to handle new resource state flag"); + constexpr RESOURCE_STATE WriteAccessStates = + RESOURCE_STATE_RENDER_TARGET | + RESOURCE_STATE_UNORDERED_ACCESS | + RESOURCE_STATE_COPY_DEST | + RESOURCE_STATE_RESOLVE_DEST | + RESOURCE_STATE_BUILD_AS_WRITE; + + return State & WriteAccessStates; +} +} // namespace + void DeviceContextVkImpl::TransitionTextureState(TextureVkImpl& TextureVk, RESOURCE_STATE OldState, RESOURCE_STATE NewState, @@ -2396,17 +2412,22 @@ void DeviceContextVkImpl::TransitionTextureState(TextureVkImpl& Textur pSubresRange->aspectMask = VK_IMAGE_ASPECT_COLOR_BIT; } - // Note that when both old and new states are RESOURCE_STATE_UNORDERED_ACCESS, we need to execute UAV barrier - // to make sure that all UAV writes are complete and visible. + // Always add barrier after writes. + const bool AfterWrite = ResourceStateHasWriteAccess(OldState); + auto OldLayout = ResourceStateToVkImageLayout(OldState); auto NewLayout = ResourceStateToVkImageLayout(NewState); auto OldStages = ResourceStateFlagsToVkPipelineStageFlags(OldState, m_CommandBuffer.GetEnabledShaderStages()); auto NewStages = ResourceStateFlagsToVkPipelineStageFlags(NewState, m_CommandBuffer.GetEnabledShaderStages()); - m_CommandBuffer.TransitionImageLayout(vkImg, OldLayout, NewLayout, *pSubresRange, OldStages, NewStages); - if (UpdateTextureState) + + if (((OldState & NewState) != NewState) || OldLayout != NewLayout || AfterWrite) { - TextureVk.SetState(NewState); - VERIFY_EXPR(TextureVk.GetLayout() == NewLayout); + m_CommandBuffer.TransitionImageLayout(vkImg, OldLayout, NewLayout, *pSubresRange, OldStages, NewStages); + if (UpdateTextureState) + { + TextureVk.SetState(NewState); + VERIFY_EXPR(TextureVk.GetLayout() == NewLayout); + } } } @@ -2421,10 +2442,7 @@ void DeviceContextVkImpl::TransitionOrVerifyTextureState(TextureVkImpl& VERIFY(m_pActiveRenderPass == nullptr, "State transitions are not allowed inside a render pass"); if (Texture.IsInKnownState()) { - if (!Texture.CheckState(RequiredState)) - { - TransitionTextureState(Texture, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + TransitionTextureState(Texture, RESOURCE_STATE_UNKNOWN, RequiredState, true); VERIFY_EXPR(Texture.GetLayout() == ExpectedLayout); } } @@ -2489,9 +2507,10 @@ void DeviceContextVkImpl::TransitionBufferState(BufferVkImpl& BufferVk, RESOURCE } } - // When both old and new states are RESOURCE_STATE_UNORDERED_ACCESS, we need to execute UAV barrier - // to make sure that all UAV writes are complete and visible. - if (((OldState & NewState) != NewState) || NewState == RESOURCE_STATE_UNORDERED_ACCESS || NewState == RESOURCE_STATE_BUILD_AS_WRITE) + // Always add barrier after writes. + const bool AfterWrite = ResourceStateHasWriteAccess(OldState); + + if (((OldState & NewState) != NewState) || AfterWrite) { DEV_CHECK_ERR(BufferVk.m_VulkanBuffer != VK_NULL_HANDLE, "Cannot transition suballocated buffer"); VERIFY_EXPR(BufferVk.GetDynamicOffset(m_ContextId, this) == 0); @@ -2521,10 +2540,7 @@ void DeviceContextVkImpl::TransitionOrVerifyBufferState(BufferVkImpl& VERIFY(m_pActiveRenderPass == nullptr, "State transitions are not allowed inside a render pass"); if (Buffer.IsInKnownState()) { - if (!Buffer.CheckState(RequiredState)) - { - TransitionBufferState(Buffer, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + TransitionBufferState(Buffer, RESOURCE_STATE_UNKNOWN, RequiredState, true); VERIFY_EXPR(Buffer.CheckAccessFlags(ExpectedAccessFlags)); } } @@ -2564,7 +2580,10 @@ void DeviceContextVkImpl::TransitionBLASState(BottomLevelASVkImpl& BLAS, } } - if ((OldState & NewState) != NewState) + // Always add barrier after writes. + const bool AfterWrite = ResourceStateHasWriteAccess(OldState); + + if ((OldState & NewState) != NewState || AfterWrite) { EnsureVkCmdBuffer(); auto OldAccessFlags = ResourceStateFlagsToVkAccessFlags(OldState); @@ -2584,8 +2603,6 @@ void DeviceContextVkImpl::TransitionTLASState(TopLevelASVkImpl& TLAS, RESOURCE_STATE NewState, bool UpdateInternalState) { - // AZ TODO: transit BLAS state too? - VERIFY(m_pActiveRenderPass == nullptr, "State transitions are not allowed inside a render pass"); if (OldState == RESOURCE_STATE_UNKNOWN) { @@ -2609,7 +2626,10 @@ void DeviceContextVkImpl::TransitionTLASState(TopLevelASVkImpl& TLAS, } } - if ((OldState & NewState) != NewState) + // Always add barrier after writes. + const bool AfterWrite = ResourceStateHasWriteAccess(OldState); + + if ((OldState & NewState) != NewState || AfterWrite) { EnsureVkCmdBuffer(); auto OldAccessFlags = ResourceStateFlagsToVkAccessFlags(OldState); @@ -2634,10 +2654,7 @@ void DeviceContextVkImpl::TransitionOrVerifyBLASState(BottomLevelASVkImpl& VERIFY(m_pActiveRenderPass == nullptr, "State transitions are not allowed inside a render pass"); if (BLAS.IsInKnownState()) { - if (!BLAS.CheckState(RequiredState)) - { - TransitionBLASState(BLAS, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + TransitionBLASState(BLAS, RESOURCE_STATE_UNKNOWN, RequiredState, true); } } #ifdef DILIGENT_DEVELOPMENT @@ -2658,10 +2675,7 @@ void DeviceContextVkImpl::TransitionOrVerifyTLASState(TopLevelASVkImpl& VERIFY(m_pActiveRenderPass == nullptr, "State transitions are not allowed inside a render pass"); if (TLAS.IsInKnownState()) { - if (!TLAS.CheckState(RequiredState)) - { - TransitionTLASState(TLAS, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + TransitionTLASState(TLAS, RESOURCE_STATE_UNKNOWN, RequiredState, true); } } #ifdef DILIGENT_DEVELOPMENT @@ -2669,6 +2683,11 @@ void DeviceContextVkImpl::TransitionOrVerifyTLASState(TopLevelASVkImpl& { DvpVerifyTLASState(TLAS, RequiredState, OperationName); } + + if (RequiredState & RESOURCE_STATE_RAY_TRACING) + { + TLAS.CheckBLASVersion(); + } #endif } @@ -2718,13 +2737,13 @@ void DeviceContextVkImpl::TransitionResourceStates(Uint32 BarrierCount, StateTra { TransitionBufferState(*pBuffer, Barrier.OldState, Barrier.NewState, Barrier.UpdateResourceState); } - else if (RefCntAutoPtr<BottomLevelASVkImpl> pBLAS{Barrier.pResource, IID_BottomLevelAS}) + else if (RefCntAutoPtr<BottomLevelASVkImpl> pBottomLevelAS{Barrier.pResource, IID_BottomLevelAS}) { - TransitionBLASState(*pBLAS, Barrier.OldState, Barrier.NewState, Barrier.UpdateResourceState); + TransitionBLASState(*pBottomLevelAS, Barrier.OldState, Barrier.NewState, Barrier.UpdateResourceState); } - else if (RefCntAutoPtr<TopLevelASVkImpl> pTLAS{Barrier.pResource, IID_TopLevelAS}) + else if (RefCntAutoPtr<TopLevelASVkImpl> pTopLevelAS{Barrier.pResource, IID_TopLevelAS}) { - TransitionTLASState(*pTLAS, Barrier.OldState, Barrier.NewState, Barrier.UpdateResourceState); + TransitionTLASState(*pTopLevelAS, Barrier.OldState, Barrier.NewState, Barrier.UpdateResourceState); } else { @@ -2808,7 +2827,7 @@ void DeviceContextVkImpl::BuildBLAS(const BLASBuildAttribs& Attribs) const char* OpName = "Build BottomLevelAS (DeviceContextVkImpl::BuildBLAS)"; TransitionOrVerifyBLASState(*pBLASVk, Attribs.BLASTransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, OpName); - TransitionOrVerifyBufferState(*pScratchVk, Attribs.ScratchBufferTransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, VkAccessFlagBits(0), OpName); + TransitionOrVerifyBufferState(*pScratchVk, Attribs.ScratchBufferTransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, VK_ACCESS_ACCELERATION_STRUCTURE_WRITE_BIT_KHR, OpName); VkAccelerationStructureBuildGeometryInfoKHR Info = {}; std::vector<VkAccelerationStructureBuildOffsetInfoKHR> Offsets; @@ -2845,7 +2864,7 @@ void DeviceContextVkImpl::BuildBLAS(const BLASBuildAttribs& Attribs) vkTris.vertexStride = SrcTris.VertexStride; vkTris.vertexData.deviceAddress = pVB->GetVkDeviceAddress() + SrcTris.VertexOffset; - TransitionOrVerifyBufferState(*pVB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, static_cast<VkAccessFlagBits>(0), OpName); + TransitionOrVerifyBufferState(*pVB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VK_ACCESS_ACCELERATION_STRUCTURE_READ_BIT_KHR, OpName); if (SrcTris.pIndexBuffer) { @@ -2854,7 +2873,7 @@ void DeviceContextVkImpl::BuildBLAS(const BLASBuildAttribs& Attribs) vkTris.indexData.deviceAddress = pIB->GetVkDeviceAddress() + SrcTris.IndexOffset; off.primitiveCount = SrcTris.IndexCount / 3; - TransitionOrVerifyBufferState(*pIB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, static_cast<VkAccessFlagBits>(0), OpName); + TransitionOrVerifyBufferState(*pIB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VK_ACCESS_ACCELERATION_STRUCTURE_READ_BIT_KHR, OpName); } else { @@ -2870,7 +2889,7 @@ void DeviceContextVkImpl::BuildBLAS(const BLASBuildAttribs& Attribs) auto* const pTB = ValidatedCast<BufferVkImpl>(SrcTris.pTransformBuffer); vkTris.transformData.deviceAddress = pTB->GetVkDeviceAddress() + SrcTris.TransformBufferOffset; - TransitionOrVerifyBufferState(*pTB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VkAccessFlagBits(0), OpName); + TransitionOrVerifyBufferState(*pTB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VK_ACCESS_ACCELERATION_STRUCTURE_READ_BIT_KHR, OpName); } else { @@ -2913,7 +2932,7 @@ void DeviceContextVkImpl::BuildBLAS(const BLASBuildAttribs& Attribs) vkAABBs.stride = SrcBoxes.BoxStride; vkAABBs.data.deviceAddress = pBB->GetVkDeviceAddress() + SrcBoxes.BoxOffset; - TransitionOrVerifyBufferState(*pBB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VkAccessFlagBits(0), OpName); + TransitionOrVerifyBufferState(*pBB, Attribs.GeometryTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VK_ACCESS_ACCELERATION_STRUCTURE_READ_BIT_KHR, OpName); off.firstVertex = 0; off.transformOffset = 0; @@ -2939,6 +2958,10 @@ void DeviceContextVkImpl::BuildBLAS(const BLASBuildAttribs& Attribs) EnsureVkCmdBuffer(); m_CommandBuffer.BuildAccelerationStructure(1, &Info, &OffsetsPtr); ++m_State.NumCommands; + +#ifdef DILIGENT_DEVELOPMENT + pBLASVk->UpdateVersion(); +#endif } void DeviceContextVkImpl::BuildTLAS(const TLASBuildAttribs& Attribs) @@ -2964,7 +2987,7 @@ void DeviceContextVkImpl::BuildTLAS(const TLASBuildAttribs& Attribs) const char* OpName = "Build TopLevelAS (DeviceContextVkImpl::BuildTLAS)"; TransitionOrVerifyTLASState(*pTLASVk, Attribs.TLASTransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, OpName); - TransitionOrVerifyBufferState(*pScratchVk, Attribs.ScratchBufferTransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, VkAccessFlagBits(0), OpName); + TransitionOrVerifyBufferState(*pScratchVk, Attribs.ScratchBufferTransitionMode, RESOURCE_STATE_BUILD_AS_WRITE, VK_ACCESS_ACCELERATION_STRUCTURE_WRITE_BIT_KHR, OpName); pTLASVk->SetInstanceData(Attribs.pInstances, Attribs.InstanceCount, Attribs.HitShadersPerInstance); @@ -2981,7 +3004,7 @@ void DeviceContextVkImpl::BuildTLAS(const TLASBuildAttribs& Attribs) auto* const pBLASVk = ValidatedCast<BottomLevelASVkImpl>(Inst.pBLAS); static_assert(sizeof(vkASInst.transform) == sizeof(Inst.Transform), "size mismatch"); - std::memcpy(&vkASInst.transform, Inst.Transform, sizeof(vkASInst.transform)); + std::memcpy(&vkASInst.transform, Inst.Transform.data, sizeof(vkASInst.transform)); vkASInst.instanceCustomIndex = Inst.CustomId; vkASInst.instanceShaderBindingTableRecordOffset = pTLASVk->GetInstanceDesc(Inst.InstanceName).ContributionToHitGroupIndex; // AZ TODO: optimize @@ -2994,7 +3017,7 @@ void DeviceContextVkImpl::BuildTLAS(const TLASBuildAttribs& Attribs) UpdateBufferRegion(pInstancesVk, Attribs.InstanceBufferOffset, Size, TmpSpace.vkBuffer, TmpSpace.AlignedOffset, Attribs.InstanceBufferTransitionMode); } - TransitionOrVerifyBufferState(*pInstancesVk, Attribs.InstanceBufferTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VkAccessFlagBits(0), OpName); + TransitionOrVerifyBufferState(*pInstancesVk, Attribs.InstanceBufferTransitionMode, RESOURCE_STATE_BUILD_AS_READ, VK_ACCESS_ACCELERATION_STRUCTURE_READ_BIT_KHR, OpName); VkAccelerationStructureBuildGeometryInfoKHR vkASBuildInfo = {}; VkAccelerationStructureBuildOffsetInfoKHR vkASBuildOffset = {}; @@ -3060,6 +3083,10 @@ void DeviceContextVkImpl::CopyBLAS(const CopyBLASAttribs& Attribs) m_CommandBuffer.CopyAccelerationStructure(Info); ++m_State.NumCommands; + +#ifdef DILIGENT_DEVELOPMENT + pDstVk->UpdateVersion(); +#endif } void DeviceContextVkImpl::CopyTLAS(const CopyTLASAttribs& Attribs) @@ -3077,6 +3104,8 @@ void DeviceContextVkImpl::CopyTLAS(const CopyTLASAttribs& Attribs) auto* pSrcVk = ValidatedCast<TopLevelASVkImpl>(Attribs.pSrc); auto* pDstVk = ValidatedCast<TopLevelASVkImpl>(Attribs.pDst); + pDstVk->CopyInstancceData(*pSrcVk); + VkCopyAccelerationStructureInfoKHR Info = {}; Info.sType = VK_STRUCTURE_TYPE_COPY_ACCELERATION_STRUCTURE_INFO_KHR; diff --git a/Graphics/GraphicsEngineVulkan/src/PipelineStateVkImpl.cpp b/Graphics/GraphicsEngineVulkan/src/PipelineStateVkImpl.cpp index 36ca3a42..1cac2b6f 100644 --- a/Graphics/GraphicsEngineVulkan/src/PipelineStateVkImpl.cpp +++ b/Graphics/GraphicsEngineVulkan/src/PipelineStateVkImpl.cpp @@ -758,10 +758,16 @@ PipelineStateVkImpl::PipelineStateVkImpl(IReferenceCounters* { try { + const auto& LogicalDevice = GetDevice()->GetLogicalDevice(); + const auto ShaderGroupHandleSize = pDeviceVk->GetShaderGroupHandleSize(); + + if (LogicalDevice.GetEnabledExtFeatures().RayTracing.rayTracing == VK_FALSE) + LOG_ERROR_AND_THROW("Ray tracing is not supported by this device"); + std::vector<VkPipelineShaderStageCreateInfo> vkShaderStages; std::vector<VulkanUtilities::ShaderModuleWrapper> ShaderModules; - std::vector<VkRayTracingShaderGroupCreateInfoKHR> ShaderGroups; + InitInternalObjects(CreateInfo, vkShaderStages, ShaderModules, [&](const RayTracingPipelineStateCreateInfo& CreateInfo, LinearAllocator& MemPool, TShaderStages& ShaderStages) // { @@ -773,9 +779,6 @@ PipelineStateVkImpl::PipelineStateVkImpl(IReferenceCounters* CreateRayTracingPipeline(pDeviceVk, vkShaderStages, ShaderGroups, m_PipelineLayout, m_Desc, GetRayTracingPipelineDesc(), m_Pipeline); - const auto& LogicalDevice = GetDevice()->GetLogicalDevice(); - const auto ShaderGroupHandleSize = pDeviceVk->GetShaderGroupHandleSize(); - auto err = LogicalDevice.GetRayTracingShaderGroupHandles(m_Pipeline, 0, static_cast<uint32_t>(ShaderGroups.size()), ShaderGroupHandleSize, &m_pRayTracingPipelineData->Shaders[0]); VERIFY(err == VK_SUCCESS, "Failed to get shader group handles"); (void)err; diff --git a/Graphics/GraphicsEngineVulkan/src/ShaderBindingTableVkImpl.cpp b/Graphics/GraphicsEngineVulkan/src/ShaderBindingTableVkImpl.cpp index 6f5091e0..3940769f 100644 --- a/Graphics/GraphicsEngineVulkan/src/ShaderBindingTableVkImpl.cpp +++ b/Graphics/GraphicsEngineVulkan/src/ShaderBindingTableVkImpl.cpp @@ -39,57 +39,12 @@ ShaderBindingTableVkImpl::ShaderBindingTableVkImpl(IReferenceCounters* bool bIsDeviceInternal) : TShaderBindingTableBase{pRefCounters, pRenderDeviceVk, Desc, bIsDeviceInternal} { - ValidateDesc(Desc); - - const auto& RTLimits = GetDevice()->GetPhysicalDevice().GetExtProperties().RayTracing; - m_ShaderRecordStride = m_Desc.ShaderRecordSize + RTLimits.shaderGroupHandleSize; } ShaderBindingTableVkImpl::~ShaderBindingTableVkImpl() { } -void ShaderBindingTableVkImpl::ValidateDesc(const ShaderBindingTableDesc& Desc) const -{ - const auto& RTLimits = GetDevice()->GetPhysicalDevice().GetExtProperties().RayTracing; - - if (Desc.ShaderRecordSize + RTLimits.shaderGroupHandleSize > RTLimits.maxShaderGroupStride) - { - LOG_ERROR_AND_THROW("Description of Shader binding table '", (Desc.Name ? Desc.Name : ""), - "' is invalid: ShaderRecordSize is too big, max size is: ", RTLimits.maxShaderGroupStride - RTLimits.shaderGroupHandleSize); - } -} - -void ShaderBindingTableVkImpl::Verify() const -{ - // AZ TODO -} - -void ShaderBindingTableVkImpl::Reset(const ShaderBindingTableDesc& Desc) -{ - m_RayGenShaderRecord.clear(); - m_MissShadersRecord.clear(); - m_CallableShadersRecord.clear(); - m_HitGroupsRecord.clear(); - m_Changed = true; - - try - { - ValidateShaderBindingTableDesc(Desc); - ValidateDesc(Desc); - } - catch (const std::runtime_error&) - { - // AZ TODO - return; - } - - m_Desc = Desc; - - const auto& RTLimits = GetDevice()->GetPhysicalDevice().GetExtProperties().RayTracing; - m_ShaderRecordStride = m_Desc.ShaderRecordSize + RTLimits.shaderGroupHandleSize; -} - void ShaderBindingTableVkImpl::ResetHitGroups(Uint32 HitShadersPerInstance) { // AZ TODO diff --git a/Graphics/GraphicsEngineVulkan/src/ShaderResourceCacheVk.cpp b/Graphics/GraphicsEngineVulkan/src/ShaderResourceCacheVk.cpp index 8101fefc..27e5ee72 100644 --- a/Graphics/GraphicsEngineVulkan/src/ShaderResourceCacheVk.cpp +++ b/Graphics/GraphicsEngineVulkan/src/ShaderResourceCacheVk.cpp @@ -166,10 +166,9 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) { constexpr RESOURCE_STATE RequiredState = RESOURCE_STATE_CONSTANT_BUFFER; VERIFY_EXPR((ResourceStateFlagsToVkAccessFlags(RequiredState) & VK_ACCESS_UNIFORM_READ_BIT) == VK_ACCESS_UNIFORM_READ_BIT); - const bool IsInRequiredState = pBufferVk->CheckState(RequiredState); if (VerifyOnly) { - if (!IsInRequiredState) + if (!pBufferVk->CheckState(RequiredState)) { LOG_ERROR_MESSAGE("State of buffer '", pBufferVk->GetDesc().Name, "' is incorrect. Required state: ", GetResourceStateString(RequiredState), ". Actual state: ", @@ -181,10 +180,7 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) } else { - if (!IsInRequiredState) - { - pCtxVkImpl->TransitionBufferState(*pBufferVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + pCtxVkImpl->TransitionBufferState(*pBufferVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); VERIFY_EXPR(pBufferVk->CheckAccessFlags(VK_ACCESS_UNIFORM_READ_BIT)); } } @@ -211,11 +207,10 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) (VK_ACCESS_SHADER_READ_BIT | VK_ACCESS_SHADER_WRITE_BIT); VERIFY_EXPR((ResourceStateFlagsToVkAccessFlags(RequiredState) & RequiredAccessFlags) == RequiredAccessFlags); #endif - const bool IsInRequiredState = pBufferVk->CheckState(RequiredState); if (VerifyOnly) { - if (!IsInRequiredState) + if (!pBufferVk->CheckState(RequiredState)) { LOG_ERROR_MESSAGE("State of buffer '", pBufferVk->GetDesc().Name, "' is incorrect. Required state: ", GetResourceStateString(RequiredState), ". Actual state: ", @@ -227,12 +222,7 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) } else { - // When both old and new states are RESOURCE_STATE_UNORDERED_ACCESS, we need to execute UAV barrier - // to make sure that all UAV writes are complete and visible. - if (!IsInRequiredState || RequiredState == RESOURCE_STATE_UNORDERED_ACCESS) - { - pCtxVkImpl->TransitionBufferState(*pBufferVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + pCtxVkImpl->TransitionBufferState(*pBufferVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); VERIFY_EXPR(pBufferVk->CheckAccessFlags(RequiredAccessFlags)); } } @@ -275,11 +265,10 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) VERIFY_EXPR(ResourceStateToVkImageLayout(RequiredState) == VK_IMAGE_LAYOUT_SHADER_READ_ONLY_OPTIMAL); } } - const bool IsInRequiredState = pTextureVk->CheckState(RequiredState); if (VerifyOnly) { - if (!IsInRequiredState) + if (!pTextureVk->CheckState(RequiredState)) { LOG_ERROR_MESSAGE("State of texture '", pTextureVk->GetDesc().Name, "' is incorrect. Required state: ", GetResourceStateString(RequiredState), ". Actual state: ", @@ -291,12 +280,7 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) } else { - // When both old and new states are RESOURCE_STATE_UNORDERED_ACCESS, we need to execute UAV barrier - // to make sure that all UAV writes are complete and visible. - if (!IsInRequiredState || RequiredState == RESOURCE_STATE_UNORDERED_ACCESS) - { - pCtxVkImpl->TransitionTextureState(*pTextureVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + pCtxVkImpl->TransitionTextureState(*pTextureVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); } } } @@ -327,11 +311,10 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) auto* pTLASVk = Res.pObject.RawPtr<TopLevelASVkImpl>(); if (pTLASVk != nullptr && pTLASVk->IsInKnownState()) { - constexpr RESOURCE_STATE RequiredState = RESOURCE_STATE_RAY_TRACING; - const bool IsInRequiredState = pTLASVk->CheckState(RequiredState); + constexpr RESOURCE_STATE RequiredState = RESOURCE_STATE_RAY_TRACING; if (VerifyOnly) { - if (!IsInRequiredState) + if (!pTLASVk->CheckState(RequiredState)) { LOG_ERROR_MESSAGE("State of TLAS '", pTLASVk->GetDesc().Name, "' is incorrect. Required state: ", GetResourceStateString(RequiredState), ". Actual state: ", @@ -340,13 +323,12 @@ void ShaderResourceCacheVk::TransitionResources(DeviceContextVkImpl* pCtxVkImpl) "when calling IDeviceContext::CommitShaderResources() or explicitly transition the TLAS state " "with IDeviceContext::TransitionResourceStates()."); } + + pTLASVk->CheckBLASVersion(); } else { - if (!IsInRequiredState) - { - pCtxVkImpl->TransitionTLASState(*pTLASVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); - } + pCtxVkImpl->TransitionTLASState(*pTLASVk, RESOURCE_STATE_UNKNOWN, RequiredState, true); } } } diff --git a/Graphics/GraphicsEngineVulkan/src/VulkanTypeConversions.cpp b/Graphics/GraphicsEngineVulkan/src/VulkanTypeConversions.cpp index 143042b8..e039311b 100644 --- a/Graphics/GraphicsEngineVulkan/src/VulkanTypeConversions.cpp +++ b/Graphics/GraphicsEngineVulkan/src/VulkanTypeConversions.cpp @@ -1613,12 +1613,10 @@ VkBuildAccelerationStructureFlagsKHR BuildASFlagsToVkBuildAccelerationStructureF "Please update the switch below to handle the new ray tracing build flag"); VkBuildAccelerationStructureFlagsKHR Result = 0; - for (Uint32 Bit = 1; Bit <= Flags; Bit <<= 1) + while (Flags != RAYTRACING_BUILD_AS_NONE) { - if ((Flags & Bit) != Bit) - continue; - - switch (static_cast<RAYTRACING_BUILD_AS_FLAGS>(Bit)) + auto FlagBit = static_cast<RAYTRACING_BUILD_AS_FLAGS>(1 << PlatformMisc::GetLSB(Uint32{Flags})); + switch (FlagBit) { // clang-format off case RAYTRACING_BUILD_AS_ALLOW_UPDATE: Result |= VK_BUILD_ACCELERATION_STRUCTURE_ALLOW_UPDATE_BIT_KHR; break; @@ -1629,6 +1627,7 @@ VkBuildAccelerationStructureFlagsKHR BuildASFlagsToVkBuildAccelerationStructureF // clang-format on default: UNEXPECTED("unknown build AS flag"); } + Flags = Flags & ~FlagBit; } return Result; } @@ -1639,12 +1638,10 @@ VkGeometryFlagsKHR GeometryFlagsToVkGeometryFlags(RAYTRACING_GEOMETRY_FLAGS Flag "Please update the switch below to handle the new ray tracing geometry flag"); VkGeometryFlagsKHR Result = 0; - for (Uint32 Bit = 1; Bit <= Flags; Bit <<= 1) + while (Flags != RAYTRACING_GEOMETRY_NONE) { - if ((Flags & Bit) != Bit) - continue; - - switch (static_cast<RAYTRACING_GEOMETRY_FLAGS>(Bit)) + auto FlagBit = static_cast<RAYTRACING_GEOMETRY_FLAGS>(1 << PlatformMisc::GetLSB(Uint32{Flags})); + switch (FlagBit) { // clang-format off case RAYTRACING_GEOMETRY_OPAQUE: Result |= VK_GEOMETRY_OPAQUE_BIT_KHR; break; @@ -1652,6 +1649,7 @@ VkGeometryFlagsKHR GeometryFlagsToVkGeometryFlags(RAYTRACING_GEOMETRY_FLAGS Flag // clang-format on default: UNEXPECTED("unknown geometry flag"); } + Flags = Flags & ~FlagBit; } return Result; } @@ -1662,12 +1660,10 @@ VkGeometryInstanceFlagsKHR InstanceFlagsToVkGeometryInstanceFlags(RAYTRACING_INS "Please update the switch below to handle the new ray tracing instance flag"); VkGeometryInstanceFlagsKHR Result = 0; - for (Uint32 Bit = 1; Bit <= Flags; Bit <<= 1) + while (Flags != RAYTRACING_INSTANCE_NONE) { - if ((Flags & Bit) != Bit) - continue; - - switch (static_cast<RAYTRACING_INSTANCE_FLAGS>(Bit)) + auto FlagBit = static_cast<RAYTRACING_INSTANCE_FLAGS>(1 << PlatformMisc::GetLSB(Uint32{Flags})); + switch (FlagBit) { // clang-format off case RAYTRACING_INSTANCE_TRIANGLE_FACING_CULL_DISABLE: Result |= VK_GEOMETRY_INSTANCE_TRIANGLE_FACING_CULL_DISABLE_BIT_KHR; break; @@ -1677,6 +1673,7 @@ VkGeometryInstanceFlagsKHR InstanceFlagsToVkGeometryInstanceFlags(RAYTRACING_INS // clang-format on default: UNEXPECTED("unknown instance flag"); } + Flags = Flags & ~FlagBit; } return Result; } @@ -1693,7 +1690,7 @@ VkCopyAccelerationStructureModeKHR CopyASModeToVkCopyAccelerationStructureMode(C // clang-format on default: UNEXPECTED("unknown AS copy mode"); - return static_cast<VkCopyAccelerationStructureModeKHR>(0); + return VK_COPY_ACCELERATION_STRUCTURE_MODE_MAX_ENUM_KHR; } } |
