#include "CustomDX12API.h" #include #include #include #include #include #include #include #include #include #include "CustomChunkAnyHitShader.h" #include "CustomChunkClosestHitShader.h" #include "CustomChunkIntersectionShader.h" #include "Dimension.h" #include "DX12ChunkData.h" using namespace Framework; CustomDX12API::CustomDX12API() : DirectX12() {} CustomDX12API::~CustomDX12API() {} void CustomDX12API::initializePipeline() { if (pipeline->getShaders().getEntryCount() == 0) { // add default shaders DX12Shader* rayGenShader = new DX12Shader( DX12DefaultRayGenShaderBytes, sizeof(DX12DefaultRayGenShaderBytes)); // RayGen from RayGen.hlsl DX12ShaderSignature* rayGenSignature = new DX12ShaderSignature(); rayGenSignature->addRegisterUsageLinkedToDescriptorHeap( 0, DX12_SHADER_REGISTER_U_UNORDERED_ACCESS, 0); rayGenSignature->addRegisterUsageLinkedToDescriptorHeap( 1, DX12_SHADER_REGISTER_U_UNORDERED_ACCESS, 1); rayGenSignature->addRegisterUsageLinkedToDescriptorHeap( 2, DX12_SHADER_REGISTER_B_CONST_BUFFER, 0); rayGenSignature->addRegisterUsageLinkedToDescriptorHeap( 3, DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 0); defaultRayGenerationShaderFunction = new DX12ShaderFunction( "RayGen", rayGenSignature, DX12_SHADER_FUNCTION_TYPE_RAY_GEN); rayGenShader->addFunction(defaultRayGenerationShaderFunction); pipeline->addShader(rayGenShader); DX12Shader* missShader = new DX12Shader( DX12DefaultMissShaderBytes, sizeof(DX12DefaultMissShaderBytes)); // Miss from Miss.hlsl missShader->addFunction(new DX12ShaderFunction( "Miss", new DX12ShaderSignature(), DX12_SHADER_FUNCTION_TYPE_MISS)); pipeline->addShader(missShader); DX12ShaderSignature* hitSignature = new DX12ShaderSignature(); hitSignature->addRegisterUsageLinkedToDescriptorHeap( 0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP); hitSignature->addRegisterUsageLinkedToDescriptorHeap(0, DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 0, 1, TEXTURE_DESCRIPTOR_HEAP, 1); sbtTextureIdBufferOffset = hitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2); sbtIndexBufferOffset = hitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 3); sbtVertexDataBufferOffset = hitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 4); sbtPolygonSizeBufferOffset = hitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 5); DX12Shader* anyHitShader = new DX12Shader( DX12DefaultAnyHitShaderBytes, sizeof(DX12DefaultAnyHitShaderBytes)); DX12ShaderFunction* anyHitFunction = new DX12ShaderFunction( "AnyHit", hitSignature, DX12_SHADER_FUNCTION_TYPE_ANY_HIT); // AnyHit from AnyHit.hlsl anyHitShader->addFunction(anyHitFunction); pipeline->addShader(anyHitShader); DX12Shader* hitShader = new DX12Shader( DX12DefaultHitShaderBytes, sizeof(DX12DefaultHitShaderBytes)); // ClosestHit from Hit.hlsl DX12ShaderFunction* closestHitFunction = new DX12ShaderFunction("ClosestHit", dynamic_cast(hitSignature->getThis()), DX12_SHADER_FUNCTION_TYPE_CLOSEST_HIT); hitShader->addFunction(closestHitFunction); pipeline->addShader(hitShader); defaultHitGroup = new DX12ShaderHitGroup("HitGroup"); defaultHitGroup->setAttributeSize(8); defaultHitGroup->setPayloadSize(20 * DEFAULT_MAX_TRANSPARENT_HITS + 4); defaultHitGroup->setAnyHitShaderFunction(anyHitFunction); defaultHitGroup->setClosestHitShaderFunction(closestHitFunction); pipeline->addHitGroup(defaultHitGroup); DX12ShaderSignature* chunkHitSignature = new DX12ShaderSignature(); chunkHitSignature->addRegisterUsageLinkedToDescriptorHeap( 0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP); chunkHitSignature->addRegisterUsageLinkedToDescriptorHeap(0, DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 0, 1, TEXTURE_DESCRIPTOR_HEAP, 1); sbtChunkTextureIdBufferOffset = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2); sbtChunkIndexBufferOffset = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 3); sbtChunkDataBufferOffset = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_B_CONST_BUFFER, 0, 2); DX12Shader* chunkIntersectionShader = new DX12Shader(CustomChunkIntersectionShader, sizeof(CustomChunkIntersectionShader)); DX12ShaderFunction* chunkIntersectionFunction = new DX12ShaderFunction("ChunkIntersection", chunkHitSignature, DX12_SHADER_FUNCTION_TYPE_INTERSECTION); chunkIntersectionShader->addFunction(chunkIntersectionFunction); pipeline->addShader(chunkIntersectionShader); DX12Shader* chunkAnyHitShader = new DX12Shader( CustomChunkAnyHitShader, sizeof(CustomChunkAnyHitShader)); DX12ShaderFunction* chunkAnyHitFunction = new DX12ShaderFunction( "ChunkAnyHit", dynamic_cast(chunkHitSignature->getThis()), DX12_SHADER_FUNCTION_TYPE_ANY_HIT); chunkAnyHitShader->addFunction(chunkAnyHitFunction); pipeline->addShader(chunkAnyHitShader); DX12Shader* chunkClosestHitShader = new DX12Shader( CustomChunkClosestHitShader, sizeof(CustomChunkClosestHitShader)); DX12ShaderFunction* chunkClosestHitFunction = new DX12ShaderFunction( "ChunkClosestHit", dynamic_cast(chunkHitSignature->getThis()), DX12_SHADER_FUNCTION_TYPE_CLOSEST_HIT); chunkClosestHitShader->addFunction(chunkClosestHitFunction); pipeline->addShader(chunkClosestHitShader); chunkHitGroup = new DX12ShaderHitGroup("ChunkHitGroup"); chunkHitGroup->setIntersectionShaderFunction(chunkIntersectionFunction); chunkHitGroup->setAnyHitShaderFunction(chunkAnyHitFunction); chunkHitGroup->setClosestHitShaderFunction(chunkClosestHitFunction); chunkHitGroup->setAttributeSize(12); chunkHitGroup->setPayloadSize(20 * DEFAULT_MAX_TRANSPARENT_HITS + 4); pipeline->addHitGroup(chunkHitGroup); pipeline->setMaxRecursionDepth(10); } pipeline->createPipelineState(device, pfnD3D12SerializeRootSignature); } void CustomDX12API::initializeGlobalDescriptorHeap() { DirectX12::initializeGlobalDescriptorHeap(); } void CustomDX12API::fillShaderBindingTable( Framework::DX12ShaderBindingTable* zShaderBindingTable, Framework::Model3D* zModel, int objectIndex, const Framework::DX12BLAS* zBLAS, int& lastHitGroupIndex) { DirectX12::fillShaderBindingTable( zShaderBindingTable, zModel, objectIndex, zBLAS, lastHitGroupIndex); } void CustomDX12API::renderWorld(Framework::World3D* zWorld, DX12TLAS* zTLAS, DX12ShaderBindingTable* zSBT, int& objectIndex) { Mat4 identity = Mat4::identity(); for (const Model3DCollection* collection : zWorld->getCollections()) { const Dimension* dim = dynamic_cast(collection); if (dim) { for (const Chunk* chunk : dim->getChunks()) { DX12ChunkData* chunkData = chunk->zDX12Data(); chunkData->updateBuffers(device, directCommandQueue); bool changed = chunkData->wasBufferChanged() || chunkData->getLastObjectIndex() != objectIndex; if (changed) { chunkData->setLastObjectIndex(objectIndex); D3D12_RAYTRACING_INSTANCE_DESC* desc = zTLAS->nextInstanceDesc(); desc->InstanceID = objectIndex; desc->InstanceContributionToHitGroupIndex = objectIndex; desc->Flags = D3D12_RAYTRACING_INSTANCE_FLAG_NONE; desc->InstanceMask = 0xFF; desc->AccelerationStructure = chunkData->zBlasBuffer() ->zBuffer() ->GetGPUVirtualAddress(); memcpy(desc->Transform, &identity, sizeof(float) * 12); } else { zTLAS->nextInstanceDesc(); } int hitGroupOffset = zSBT->addHitGroup( chunkHitGroup, chunkData->getLastHitGroupIndex()); if (chunkData->getLastHitGroupIndex() != hitGroupOffset) { chunkData->setLastHitGroupIndex(hitGroupOffset); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtChunkIndexBufferOffset, chunkData->zBlockIndexBuffer() ->zBuffer() ->GetGPUVirtualAddress()); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtChunkTextureIdBufferOffset, chunkData->zTextureIdBuffer() ->zBuffer() ->GetGPUVirtualAddress()); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtChunkDataBufferOffset, chunkData->zChunkInfoBuffer() ->zBuffer() ->GetGPUVirtualAddress()); } ++objectIndex; } } } DirectX12::renderWorld(zWorld, zTLAS, zSBT, objectIndex); }