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.2 KiB123 linesraw
1#include "core/slang-basic.h"
2#include "gfx-test-util.h"
3#include "unit-test/slang-unit-test.h"
4
5#include <slang-rhi.h>
6#include <slang-rhi/shader-cursor.h>
7
8using namespace rhi;
9
10namespace gfx_test
11{
12void mutableShaderObjectTestImpl(IDevice* device, UnitTestContext* context)
13{
14    ComPtr<IShaderProgram> shaderProgram;
15    slang::ProgramLayout* slangReflection;
16    GFX_CHECK_CALL_ABORT(loadComputeProgram(
17        device,
18        shaderProgram,
19        "mutable-shader-object",
20        "computeMain",
21        slangReflection));
22
23    ComputePipelineDesc pipelineDesc = {};
24    pipelineDesc.program = shaderProgram.get();
25    ComPtr<IComputePipeline> pipelineState;
26    GFX_CHECK_CALL_ABORT(device->createComputePipeline(pipelineDesc, pipelineState.writeRef()));
27
28    float initialData[] = {0.0f, 1.0f, 2.0f, 3.0f};
29    const int numberCount = SLANG_COUNT_OF(initialData);
30    BufferDesc bufferDesc = {};
31    bufferDesc.size = sizeof(initialData);
32    bufferDesc.format = Format::Undefined;
33    bufferDesc.elementSize = sizeof(float);
34    bufferDesc.usage = BufferUsage::ShaderResource | BufferUsage::UnorderedAccess |
35                       BufferUsage::CopyDestination | BufferUsage::CopySource;
36    bufferDesc.defaultState = ResourceState::UnorderedAccess;
37    bufferDesc.memoryType = MemoryType::DeviceLocal;
38
39    ComPtr<IBuffer> numbersBuffer;
40    GFX_CHECK_CALL_ABORT(
41        device->createBuffer(bufferDesc, (void*)initialData, numbersBuffer.writeRef()));
42
43    {
44        slang::TypeReflection* addTransformerType =
45            slangReflection->findTypeByName("AddTransformer");
46
47        ComPtr<IShaderObject> transformer;
48        GFX_CHECK_CALL_ABORT(device->createShaderObject(
49            addTransformerType,
50            ShaderObjectContainerType::None,
51            transformer.writeRef()));
52
53        // Set the `c` field of the `AddTransformer`.
54        float c = 1.0f;
55        ShaderCursor(transformer).getPath("c").setData(&c, sizeof(float));
56
57        ComPtr<ICommandQueue> queue;
58        GFX_CHECK_CALL_ABORT(device->getQueue(QueueType::Graphics, queue.writeRef()));
59
60        // Create root shader object
61        ComPtr<IShaderObject> rootObject;
62        GFX_CHECK_CALL_ABORT(device->createRootShaderObject(shaderProgram, rootObject.writeRef()));
63
64        auto commandEncoder = queue->createCommandEncoder();
65        auto computeEncoder = commandEncoder->beginComputePass();
66
67        // Bind pipeline with our root object
68        computeEncoder->bindPipeline(pipelineState, rootObject);
69
70        auto entryPointCursor = ShaderCursor(rootObject->getEntryPoint(0));
71
72        entryPointCursor.getPath("buffer").setBinding(Binding(numbersBuffer));
73
74        // Bind the transformer object to root object.
75        entryPointCursor.getPath("transformer").setObject(transformer);
76
77        computeEncoder->dispatchCompute(1, 1, 1);
78        computeEncoder->end();
79
80        // Set buffer state to ensure writes are visible
81        commandEncoder->setBufferState(numbersBuffer, ResourceState::UnorderedAccess);
82
83        computeEncoder = commandEncoder->beginComputePass();
84
85        // Bind pipeline with our root object again
86        computeEncoder->bindPipeline(pipelineState, rootObject);
87
88        // Mutate `transformer` object and run again.
89        c = 2.0f;
90        ShaderCursor(transformer).getPath("c").setData(&c, sizeof(float));
91        entryPointCursor.getPath("buffer").setBinding(Binding(numbersBuffer));
92        entryPointCursor.getPath("transformer").setObject(transformer);
93        computeEncoder->dispatchCompute(1, 1, 1);
94        computeEncoder->end();
95
96        auto commandBuffer = commandEncoder->finish();
97        queue->submit(commandBuffer);
98        queue->waitOnHost();
99    }
100
101    compareComputeResult(device, numbersBuffer, std::array{3.0f, 4.0f, 5.0f, 6.0f});
102}
103
104// SLANG_UNIT_TEST(mutableShaderObjectCPU)
105//{
106//     runTestImpl(mutableShaderObjectTestImpl, unitTestContext, Slang::RenderApiFlag::CPU);
107// }
108
109SLANG_UNIT_TEST(mutableShaderObjectD3D11)
110{
111    runTestImpl(mutableShaderObjectTestImpl, unitTestContext, DeviceType::D3D11);
112}
113
114SLANG_UNIT_TEST(mutableShaderObjectD3D12)
115{
116    runTestImpl(mutableShaderObjectTestImpl, unitTestContext, DeviceType::D3D12);
117}
118
119SLANG_UNIT_TEST(mutableShaderObjectVulkan)
120{
121    runTestImpl(mutableShaderObjectTestImpl, unitTestContext, DeviceType::Vulkan);
122}
123} // namespace gfx_test