#include "CustomDX12API.h" #include #include #include #include #include "CustomChunkAnyHitShader.h" #include "CustomChunkClosestHitShader.h" #include "CustomChunkIntersectionShader.h" #include "CustomCustomMissShader.h" #include "CustomCustomRayGenShader.h" #include "CustomDefaultAnyHitShader.h" #include "CustomDefaultClosestHitShader.h" #include "CustomSimpleBlocksAnyHitShader.h" #include "CustomSimpleBlocksClosestHitShader.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( CustomCustomRayGenShader, sizeof(CustomCustomRayGenShader)); // 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( CustomCustomMissShader, sizeof(CustomCustomMissShader)); // 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( CustomDefaultAnyHitShader, sizeof(CustomDefaultAnyHitShader)); 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(CustomDefaultClosestHitShader, sizeof(CustomDefaultClosestHitShader)); // 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(28 * DEFAULT_MAX_TRANSPARENT_HITS + 4); defaultHitGroup->setAnyHitShaderFunction(anyHitFunction); defaultHitGroup->setClosestHitShaderFunction(closestHitFunction); pipeline->addHitGroup(defaultHitGroup); // Chunk ray traversal shaders 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); sbtChunkLightBufferOffset = chunkHitSignature->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 4); 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(20); chunkHitGroup->setPayloadSize(28 * DEFAULT_MAX_TRANSPARENT_HITS + 4); pipeline->addHitGroup(chunkHitGroup); // simple blocks hit shaders DX12ShaderSignature* simpleBlocksHitSignature = new DX12ShaderSignature(); simpleBlocksHitSignature->addRegisterUsageLinkedToDescriptorHeap( 0, DX12_SHADER_REGISTER_S_SAMPLER, 0, 0, SAMPLER_DESCRIPTOR_HEAP); simpleBlocksHitSignature->addRegisterUsageLinkedToDescriptorHeap(0, DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 0, 1, TEXTURE_DESCRIPTOR_HEAP, 1); sbtSimpleBlocksTextureIdBufferOffset = simpleBlocksHitSignature ->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 2); sbtSimpleBlocksIndexBufferOffset = simpleBlocksHitSignature ->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 3); sbtSimpleBlocksVertexDataBufferOffset = simpleBlocksHitSignature ->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 4); sbtSimpleBlocksIndexOffsetBufferOffset = simpleBlocksHitSignature ->addRegisterUsageLinkedToShaderBindingTable( DX12_SHADER_REGISTER_T_SHADER_RESOURCE, 1, 5); DX12Shader* simpleBlocksAnyHitShader = new DX12Shader(CustomSimpleBlocksAnyHitShader, sizeof(CustomSimpleBlocksAnyHitShader)); DX12ShaderFunction* simpleBlocksAnyHitFunction = new DX12ShaderFunction("SimpleBlocksAnyHit", simpleBlocksHitSignature, DX12_SHADER_FUNCTION_TYPE_ANY_HIT); simpleBlocksAnyHitShader->addFunction(simpleBlocksAnyHitFunction); pipeline->addShader(simpleBlocksAnyHitShader); DX12Shader* simpleBlocksClosestHitShader = new DX12Shader(CustomSimpleBlocksClosestHitShader, sizeof(CustomSimpleBlocksClosestHitShader)); DX12ShaderFunction* simpleBlocksClosestHitFunction = new DX12ShaderFunction("SimpleBlocksClosestHit", dynamic_cast( simpleBlocksHitSignature->getThis()), DX12_SHADER_FUNCTION_TYPE_CLOSEST_HIT); simpleBlocksClosestHitShader->addFunction( simpleBlocksClosestHitFunction); pipeline->addShader(simpleBlocksClosestHitShader); simpleBlocksHitGroup = new DX12ShaderHitGroup("SimpleBlocksHitGroup"); simpleBlocksHitGroup->setAnyHitShaderFunction( simpleBlocksAnyHitFunction); simpleBlocksHitGroup->setClosestHitShaderFunction( simpleBlocksClosestHitFunction); simpleBlocksHitGroup->setAttributeSize(8); simpleBlocksHitGroup->setPayloadSize( 28 * DEFAULT_MAX_TRANSPARENT_HITS + 4); pipeline->addHitGroup(simpleBlocksHitGroup); 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); if (chunkData->hasCustomBlocks()) { bool changed = chunkData->wasCustomBufferChanged() || chunkData->getLastCustomObjectIndex() != objectIndex; if (changed) { chunkData->setLastCustomObjectIndex(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->zCustomBlocksBlasBuffer() ->zBuffer() ->GetGPUVirtualAddress(); Framework::Mat4 transform = Framework::Mat4::translation( Vec3((float)chunk->getCenter().x, (float)chunk->getCenter().y, 0.f)); memcpy(desc->Transform, &transform, sizeof(float) * 12); } else { zTLAS->nextInstanceDesc(); } int hitGroupOffset = zSBT->addHitGroup(simpleBlocksHitGroup, chunkData->getLastCustomHitGroupIndex()); if (chunkData->getLastCustomHitGroupIndex() != hitGroupOffset || chunkData->wasCustomBufferChanged()) { chunkData->setLastCustomHitGroupIndex(hitGroupOffset); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtSimpleBlocksTextureIdBufferOffset, chunkData->zCustomBlocksTextureBuffer() ->zBuffer() ->GetGPUVirtualAddress()); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtSimpleBlocksIndexBufferOffset, chunkData->zCustomBlocksCombinedIndexBuffer() ->zBuffer() ->GetGPUVirtualAddress()); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtSimpleBlocksVertexDataBufferOffset, chunkData->zCustomBlocksVertexDataBuffer() ->zBuffer() ->GetGPUVirtualAddress()); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtSimpleBlocksIndexOffsetBufferOffset, chunkData->zCustomBlocksIndexOffsetBuffer() ->zBuffer() ->GetGPUVirtualAddress()); } chunkData->setCustomBufferChanged(0); ++objectIndex; } 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->wasBufferChanged()) { 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()); zSBT->setHitGroupShaderInput(hitGroupOffset, sbtChunkLightBufferOffset, chunkData->zLightBuffer() ->zBuffer() ->GetGPUVirtualAddress()); } chunkData->setBufferChanged(0); ++objectIndex; } } } DirectX12::renderWorld(zWorld, zTLAS, zSBT, objectIndex); }