yum-mirror/slang

Making it easier to work with shaders

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

Gangzheng TongConvert gfx unit tests and examples to use slang-rhi (#7577)43d0c2100

master
7.1 KiB238 linesraw
1#pragma once
2
3#include "core/slang-basic.h"
4#include "core/slang-blob.h"
5#include "core/slang-render-api-util.h"
6#include "core/slang-test-tool-util.h"
7#include "slang-rhi.h"
8#include "span.h"
9#include "unit-test/slang-unit-test.h"
10
11// GFX_CHECK_CALL and GFX_CHECK_CALL_ABORT are used to check SlangResult
12#define GFX_CHECK_CALL(x) SLANG_CHECK(!SLANG_FAILED(x))
13#define GFX_CHECK_CALL_ABORT(x) SLANG_CHECK_ABORT(!SLANG_FAILED(x))
14
15using namespace rhi;
16
17namespace gfx_test
18{
19enum class PrecompilationMode
20{
21    None,
22    SlangIR,
23    InternalLink,
24    ExternalLink,
25};
26/// Helper function for print out diagnostic messages output by Slang compiler.
27void diagnoseIfNeeded(slang::IBlob* diagnosticsBlob);
28
29/// Loads a compute shader module and produces a `rhi::IShaderProgram`.
30Slang::Result loadComputeProgram(
31    rhi::IDevice* device,
32    Slang::ComPtr<rhi::IShaderProgram>& outShaderProgram,
33    const char* shaderModuleName,
34    const char* entryPointName,
35    slang::ProgramLayout*& slangReflection);
36
37Slang::Result loadComputeProgram(
38    rhi::IDevice* device,
39    slang::ISession* slangSession,
40    Slang::ComPtr<rhi::IShaderProgram>& outShaderProgram,
41    const char* shaderModuleName,
42    const char* entryPointName,
43    slang::ProgramLayout*& slangReflection);
44
45Slang::Result loadComputeProgramFromSource(
46    rhi::IDevice* device,
47    Slang::ComPtr<rhi::IShaderProgram>& outShaderProgram,
48    std::string_view source);
49
50Slang::Result loadGraphicsProgram(
51    rhi::IDevice* device,
52    Slang::ComPtr<rhi::IShaderProgram>& outShaderProgram,
53    const char* shaderModuleName,
54    const char* vertexEntryPointName,
55    const char* fragmentEntryPointName,
56    slang::ProgramLayout*& slangReflection);
57
58template<typename T>
59void compareResultFuzzy(const T* result, const T* expectedResult, size_t count)
60{
61    for (size_t i = 0; i < count; ++i)
62    {
63        SLANG_CHECK(abs(result[i] - expectedResult[i]) < 0.01f);
64    }
65}
66
67template<typename T>
68void compareResult(const T* result, const T* expectedResult, size_t count)
69{
70    for (size_t i = 0; i < count; i++)
71    {
72        SLANG_CHECK(result[i] == expectedResult[i]);
73    }
74}
75
76template<typename T>
77void compareComputeResult(rhi::IDevice* device, rhi::IBuffer* buffer, span<T> expectedResult)
78{
79    size_t bufferSize = expectedResult.size() * sizeof(T);
80    // Read back the results.`
81    ComPtr<ISlangBlob> bufferData;
82    SLANG_CHECK(SLANG_SUCCEEDED(device->readBuffer(buffer, 0, bufferSize, bufferData.writeRef())));
83    SLANG_CHECK(bufferData->getBufferSize() == bufferSize);
84    const T* result = reinterpret_cast<const T*>(bufferData->getBufferPointer());
85
86    if constexpr (std::is_same<T, float>::value || std::is_same<T, double>::value)
87        compareResultFuzzy(result, expectedResult.data(), expectedResult.size());
88    else
89        compareResult<T>(result, expectedResult.data(), expectedResult.size());
90}
91
92template<typename T, size_t Count>
93void compareComputeResult(
94    rhi::IDevice* device,
95    rhi::IBuffer* buffer,
96    std::array<T, Count> expectedResult)
97{
98    compareComputeResult(device, buffer, span<T>(expectedResult.data(), Count));
99}
100
101template<typename T>
102void compareComputeResult(
103    rhi::IDevice* device,
104    rhi::ITexture* texture,
105    uint32_t layer,
106    uint32_t mip,
107    span<T> expectedResult)
108{
109    size_t bufferSize = expectedResult.size() * sizeof(T);
110    // Read back the results.
111    ComPtr<ISlangBlob> textureData;
112    rhi::SubresourceLayout layout;
113    SLANG_CHECK(
114        SLANG_SUCCEEDED(device->readTexture(texture, layer, mip, textureData.writeRef(), &layout)));
115    SLANG_CHECK(textureData->getBufferSize() >= bufferSize);
116
117    uint8_t* buffer = (uint8_t*)textureData->getBufferPointer();
118    for (uint32_t z = 0; z < layout.size.depth; z++)
119    {
120        for (uint32_t y = 0; y < layout.size.height; y++)
121        {
122            for (uint32_t x = 0; x < layout.size.width; x++)
123            {
124                const uint8_t* src = reinterpret_cast<const uint8_t*>(
125                    buffer + z * layout.slicePitch + y * layout.rowPitch + x * layout.colPitch);
126                uint8_t* dst = reinterpret_cast<uint8_t*>(
127                    buffer +
128                    (((z * layout.size.depth + y) * layout.size.width) + x) * layout.colPitch);
129                ::memcpy(dst, src, layout.colPitch);
130            }
131        }
132    }
133
134    const T* result = reinterpret_cast<const T*>(textureData->getBufferPointer());
135
136    if constexpr (std::is_same<T, float>::value)
137        compareResultFuzzy(result, expectedResult.data(), expectedResult.size());
138    else
139        compareResult<T>(result, expectedResult.data(), expectedResult.size());
140}
141
142template<typename T, size_t Count>
143void compareComputeResult(
144    rhi::IDevice* device,
145    rhi::ITexture* texture,
146    uint32_t layer,
147    uint32_t mip,
148    std::array<T, Count> expectedResult)
149{
150    compareComputeResult(device, texture, layer, mip, span<T>(expectedResult.data(), Count));
151}
152
153Slang::ComPtr<rhi::IDevice> createTestingDevice(
154    UnitTestContext* context,
155    rhi::DeviceType deviceType,
156    Slang::List<const char*> additionalSearchPaths = {});
157
158Slang::List<const char*> getSlangSearchPaths();
159
160void initializeRenderDoc();
161void renderDocBeginFrame();
162void renderDocEndFrame();
163
164template<typename T, typename... Args>
165auto makeArray(Args... args)
166{
167    return std::array<T, sizeof...(Args)>{static_cast<T>(args)...};
168}
169
170inline bool deviceTypeInEnabledApis(rhi::DeviceType deviceType, Slang::RenderApiFlags enabledApis)
171{
172    switch (deviceType)
173    {
174    case rhi::DeviceType::Default:
175        return true;
176    case rhi::DeviceType::CPU:
177        return enabledApis & Slang::RenderApiFlag::CPU;
178    case rhi::DeviceType::CUDA:
179        return enabledApis & Slang::RenderApiFlag::CUDA;
180    case rhi::DeviceType::Metal:
181        return enabledApis & Slang::RenderApiFlag::Metal;
182    case rhi::DeviceType::WGPU:
183        return enabledApis & Slang::RenderApiFlag::WebGPU;
184    case rhi::DeviceType::Vulkan:
185        return enabledApis & Slang::RenderApiFlag::Vulkan;
186    case rhi::DeviceType::D3D11:
187        return enabledApis & Slang::RenderApiFlag::D3D11;
188    case rhi::DeviceType::D3D12:
189        return enabledApis & Slang::RenderApiFlag::D3D12;
190    }
191    return true;
192}
193
194
195template<typename ImplFunc>
196void runTestImpl(
197    const ImplFunc& f,
198    UnitTestContext* context,
199    rhi::DeviceType deviceType,
200    Slang::List<const char*> searchPaths = {})
201{
202    if (!deviceTypeInEnabledApis(deviceType, context->enabledApis))
203    {
204        SLANG_IGNORE_TEST
205    }
206
207    auto device = createTestingDevice(context, deviceType, searchPaths);
208    if (!device)
209    {
210        SLANG_IGNORE_TEST
211    }
212#if SLANG_WIN32
213    // Skip d3d12 tests on x86 now since dxc doesn't function correctly there on Windows 11.
214    if (rhi::DeviceType == rhi::DeviceType::D3D12)
215    {
216        SLANG_IGNORE_TEST
217    }
218#endif
219    // Skip d3d11 tests when we don't have DXBC support as they're bound to
220    // fail without a backend compiler
221    if (deviceType == rhi::DeviceType::D3D11 && !SLANG_ENABLE_DXBC_SUPPORT)
222    {
223        SLANG_IGNORE_TEST
224    }
225    try
226    {
227        renderDocBeginFrame();
228        f(device, context);
229    }
230    catch (AbortTestException& e)
231    {
232        renderDocEndFrame();
233        throw e;
234    }
235    renderDocEndFrame();
236}
237
238} // namespace gfx_test