yum-mirror/slang

Making it easier to work with shaders

git clone https://git.yummers.dev/yum-mirror/slang

Jay KwakFix intermittent failure of slang-unit-test-tool/ReplayRecord (#6981)1b539d890

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