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
4.1 KiB136 linesraw
1#include "core/slang-basic.h"
2#include "gfx-test-util.h"
3#include "slang-rhi.h"
4#include "unit-test/slang-unit-test.h"
5
6#include <slang-rhi/shader-cursor.h>
7
8using namespace rhi;
9
10namespace gfx_test
11{
12struct Shader
13{
14    ComPtr<IShaderProgram> program;
15    slang::ProgramLayout* reflection = nullptr;
16    ComputePipelineDesc pipelineDesc = {};
17    ComPtr<IComputePipeline> pipeline;
18};
19
20struct Buffer
21{
22    BufferDesc desc;
23    ComPtr<IBuffer> buffer;
24    ComPtr<ITextureView> view;
25};
26
27ComPtr<IBuffer> createFloatBuffer(
28    IDevice* device,
29    bool unorderedAccess,
30    size_t elementCount,
31    float* initialData = nullptr)
32{
33    BufferDesc desc = {};
34    desc.size = elementCount * sizeof(float);
35    desc.elementSize = sizeof(float);
36    desc.format = Format::Undefined;
37    desc.memoryType = MemoryType::DeviceLocal;
38    desc.usage =
39        BufferUsage::ShaderResource | BufferUsage::CopyDestination | BufferUsage::CopySource;
40    if (unorderedAccess)
41        desc.usage |= BufferUsage::UnorderedAccess;
42
43    ComPtr<IBuffer> buffer;
44    GFX_CHECK_CALL_ABORT(device->createBuffer(desc, (void*)initialData, buffer.writeRef()));
45    return buffer;
46}
47
48void barrierTestImpl(IDevice* device, UnitTestContext* context)
49{
50    Shader programA;
51    Shader programB;
52    GFX_CHECK_CALL_ABORT(loadComputeProgram(
53        device,
54        programA.program,
55        "buffer-barrier-test",
56        "computeA",
57        programA.reflection));
58    GFX_CHECK_CALL_ABORT(loadComputeProgram(
59        device,
60        programB.program,
61        "buffer-barrier-test",
62        "computeB",
63        programB.reflection));
64    programA.pipelineDesc.program = programA.program.get();
65    programB.pipelineDesc.program = programB.program.get();
66    GFX_CHECK_CALL_ABORT(
67        device->createComputePipeline(programA.pipelineDesc, programA.pipeline.writeRef()));
68
69    GFX_CHECK_CALL_ABORT(
70        device->createComputePipeline(programB.pipelineDesc, programB.pipeline.writeRef()));
71
72    float initialData[] = {1.0f, 2.0f, 3.0f, 4.0f};
73    ComPtr<IBuffer> inputBuffer = createFloatBuffer(device, false, 4, initialData);
74    ComPtr<IBuffer> intermediateBuffer = createFloatBuffer(device, true, 4, nullptr);
75    ComPtr<IBuffer> outputBuffer = createFloatBuffer(device, true, 4, nullptr);
76
77    // We have done all the set up work, now it is time to start recording a command buffer for
78    // GPU execution.
79    {
80        auto queue = device->getQueue(QueueType::Graphics);
81        auto commandEncoder = queue->createCommandEncoder();
82
83        // Write inputBuffer data to intermediateBuffer
84        {
85            auto passEncoder = commandEncoder->beginComputePass();
86            auto rootObject = passEncoder->bindPipeline(programA.pipeline);
87
88            ShaderCursor cursor(rootObject->getEntryPoint(0));
89            cursor["inBuffer"].setBinding(inputBuffer);
90            cursor["outBuffer"].setBinding(intermediateBuffer);
91            passEncoder->dispatchCompute(1, 1, 1);
92            passEncoder->end();
93        }
94
95        // Resource transition is automatically handled.
96
97        // Write intermediateBuffer data to outputBuffer
98
99        {
100            auto passEncoder = commandEncoder->beginComputePass();
101            auto rootObject = passEncoder->bindPipeline(programB.pipeline);
102            ShaderCursor cursor(rootObject->getEntryPoint(0));
103            cursor["inBuffer"].setBinding(intermediateBuffer);
104            cursor["outBuffer"].setBinding(outputBuffer);
105            passEncoder->dispatchCompute(1, 1, 1);
106            passEncoder->end();
107        }
108
109
110        queue->submit(commandEncoder->finish());
111        queue->waitOnHost();
112    }
113
114
115    compareComputeResult(device, outputBuffer, makeArray<float>(11.0f, 12.0f, 13.0f, 14.0f));
116}
117
118void barrierTestAPI(UnitTestContext* context, DeviceType deviceType)
119{
120    Slang::List<const char*> searchPaths = {"", "../../tools/gfx-unit-test", "tools/gfx-unit-test"};
121    auto device = createTestingDevice(context, deviceType, searchPaths);
122
123    if (!device)
124    {
125        SLANG_IGNORE_TEST
126    }
127
128    barrierTestImpl(device.get(), context);
129}
130
131SLANG_UNIT_TEST(bufferBarrierVulkan)
132{
133    barrierTestAPI(unitTestContext, DeviceType::Vulkan);
134}
135
136} // namespace gfx_test