yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
1b539d890
master
1// d3d12-shader-table.cpp 2#include "d3d12-shader-table.h" 3 4#include "d3d12-device.h" 5#include "d3d12-pipeline-state.h" 6#include "d3d12-transient-heap.h" 7 8namespace gfx 9{ 10namespace d3d12 11{ 12 13using namespace Slang ; 14 15RefPtr < BufferResource > ShaderTableImpl ::createDeviceBuffer ( 16PipelineStateBase * pipeline , 17TransientResourceHeapBase * transientHeap , 18IResourceCommandEncoder * encoder ) 19{ 20uint32_t raygenTableSize = m_rayGenShaderCount * kRayGenRecordSize ; 21uint32_t missTableSize = m_missShaderCount * D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES ; 22uint32_t hitgroupTableSize = m_hitGroupCount * D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES ; 23uint32_t callableTableSize = m_callableShaderCount * D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES ; 24m_rayGenTableOffset = 0 ; 25m_missTableOffset = raygenTableSize ; 26m_hitGroupTableOffset = (uint32_t )D3DUtil ::calcAligned ( 27m_missTableOffset + missTableSize , 28D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT ); 29m_callableTableOffset = (uint32_t )D3DUtil ::calcAligned ( 30m_hitGroupTableOffset + hitgroupTableSize , 31D3D12_RAYTRACING_SHADER_TABLE_BYTE_ALIGNMENT ); 32uint32_t tableSize = m_callableTableOffset + callableTableSize ; 33 34auto pipelineImpl = static_cast < RayTracingPipelineStateImpl *> (pipeline ); 35ComPtr < IBufferResource > bufferResource ; 36IBufferResource ::Desc bufferDesc = {}; 37bufferDesc .memoryType = gfx::MemoryType ::DeviceLocal ; 38bufferDesc .defaultState = ResourceState ::General ; 39bufferDesc .allowedStates .add (ResourceState ::NonPixelShaderResource ); 40bufferDesc .type = IResource ::Type ::Buffer ; 41bufferDesc .sizeInBytes = tableSize ; 42m_device -> createBufferResource (bufferDesc ,nullptr ,bufferResource .writeRef ()); 43 44ComPtr < ID3D12StateObjectProperties > stateObjectProperties ; 45pipelineImpl -> m_stateObject -> QueryInterface (stateObjectProperties .writeRef ()); 46 47TransientResourceHeapImpl * transientHeapImpl = 48static_cast < TransientResourceHeapImpl *> (transientHeap ); 49 50IBufferResource * stagingBuffer = nullptr ; 51Offset stagingBufferOffset = 0 ; 52transientHeapImpl 53-> allocateStagingBuffer (tableSize ,stagingBuffer ,stagingBufferOffset ,MemoryType ::Upload ); 54 55assert (stagingBuffer ); 56void * stagingPtr = nullptr ; 57stagingBuffer -> map (nullptr ,& stagingPtr ); 58 59auto copyShaderIdInto = [& ](void * dest ,String & name ,const ShaderRecordOverwrite & overwrite ) 60 { 61if (name .getLength ()) 62 { 63void * shaderId = stateObjectProperties -> GetShaderIdentifier (name .toWString ().begin ()); 64if (nullptr == shaderId ) 65throw Exception (String ("Failed to get shader identifier for '" )+ name + "'" ); 66memcpy (dest ,shaderId ,D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES ); 67 } 68if (overwrite .size ) 69 { 70memcpy ((uint8_t * )dest + overwrite .offset ,overwrite .data ,overwrite .size ); 71 } 72 }; 73 74uint8_t * stagingBufferPtr = (uint8_t * )stagingPtr + stagingBufferOffset ; 75memset (stagingBufferPtr ,0 ,tableSize ); 76 77for (uint32_t i = 0 ;i < m_rayGenShaderCount ;i ++ ) 78 { 79copyShaderIdInto ( 80stagingBufferPtr + m_rayGenTableOffset + kRayGenRecordSize * i , 81m_shaderGroupNames [i ], 82m_recordOverwrites [i ]); 83 } 84for (uint32_t i = 0 ;i < m_missShaderCount ;i ++ ) 85 { 86copyShaderIdInto ( 87stagingBufferPtr + m_missTableOffset + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES * i , 88m_shaderGroupNames [m_rayGenShaderCount + i ], 89m_recordOverwrites [m_rayGenShaderCount + i ]); 90 } 91for (uint32_t i = 0 ;i < m_hitGroupCount ;i ++ ) 92 { 93copyShaderIdInto ( 94stagingBufferPtr + m_hitGroupTableOffset + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES * i , 95m_shaderGroupNames [m_rayGenShaderCount + m_missShaderCount + i ], 96m_recordOverwrites [m_rayGenShaderCount + m_missShaderCount + i ]); 97 } 98for (uint32_t i = 0 ;i < m_callableShaderCount ;i ++ ) 99 { 100copyShaderIdInto ( 101stagingBufferPtr + m_callableTableOffset + D3D12_SHADER_IDENTIFIER_SIZE_IN_BYTES * i , 102m_shaderGroupNames [m_rayGenShaderCount + m_missShaderCount + m_hitGroupCount + i ], 103m_recordOverwrites [m_rayGenShaderCount + m_missShaderCount + m_hitGroupCount + i ]); 104 } 105 106stagingBuffer -> unmap (nullptr ); 107encoder -> copyBuffer (bufferResource ,0 ,stagingBuffer ,stagingBufferOffset ,tableSize ); 108encoder -> bufferBarrier ( 1091 , 110bufferResource .readRef (), 111 gfx::ResourceState ::CopyDestination , 112 gfx::ResourceState ::NonPixelShaderResource ); 113RefPtr < BufferResource > resultPtr = static_cast < BufferResource *> (bufferResource .get ()); 114return _Move (resultPtr ); 115} 116 117}// namespace d3d12 118}// namespace gfx