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
18.3 KiB463 linesraw
1#if 0
2// Duplcated: This is ported to slang-rhi\tests\test-ray-tracing.cpp
3
4#include "core/slang-basic.h"
5#include "gfx-test-texture-util.h"
6#include "gfx-test-util.h"
7#include "platform/vector-math.h"
8#include "unit-test/slang-unit-test.h"
9
10#include <chrono>
11#include <slang-rhi.h>
12#include <slang-rhi/shader-cursor.h>
13
14using namespace rhi;
15using namespace Slang;
16
17namespace gfx_test
18{
19struct Vertex
20{
21    float position[3];
22};
23
24static const int kVertexCount = 9;
25static const Vertex kVertexData[kVertexCount] = {
26    // Triangle 1
27    {0, 0, 1},
28    {4, 0, 1},
29    {0, 4, 1},
30
31    // Triangle 2
32    {-4, 0, 1},
33    {0, 0, 1},
34    {0, 4, 1},
35
36    // Triangle 3
37    {0, 0, 1},
38    {4, 0, 1},
39    {0, -4, 1},
40};
41static const int kIndexCount = 9;
42static const uint32_t kIndexData[kIndexCount] = {
43    0,
44    1,
45    2,
46    3,
47    4,
48    5,
49    6,
50    7,
51    8,
52};
53
54struct BaseRayTracingTest
55{
56    IDevice* device;
57    UnitTestContext* context;
58
59    ComPtr<ICommandQueue> queue;
60
61    ComPtr<IRayTracingPipeline> renderPipelineState;
62    ComPtr<IBuffer> vertexBuffer;
63    ComPtr<IBuffer> indexBuffer;
64    ComPtr<IBuffer> transformBuffer;
65    ComPtr<IBuffer> instanceBuffer;
66    ComPtr<IBuffer> BLASBuffer;
67    ComPtr<IAccelerationStructure> BLAS;
68    ComPtr<IBuffer> TLASBuffer;
69    ComPtr<IAccelerationStructure> TLAS;
70    ComPtr<ITexture> resultTexture;
71    ComPtr<ITextureView> resultTextureUAV;
72    ComPtr<IShaderTable> shaderTable;
73
74    uint32_t width = 2;
75    uint32_t height = 2;
76
77    void init(IDevice* device, UnitTestContext* context)
78    {
79        if (!device->hasFeature("ray-tracing"))
80        {
81            SLANG_IGNORE_TEST;
82        }
83
84        this->device = device;
85        this->context = context;
86    }
87
88    // Load and compile shader code from source.
89    Result loadShaderProgram(IDevice* device, IShaderProgram** outProgram)
90    {
91        ComPtr<slang::ISession> slangSession;
92        slangSession = device->getSlangSession();
93
94        ComPtr<slang::IBlob> diagnosticsBlob;
95        slang::IModule* module =
96            slangSession->loadModule("ray-tracing-test-shaders", diagnosticsBlob.writeRef());
97        if (!module)
98            return SLANG_FAIL;
99
100        Slang::List<slang::IComponentType*> componentTypes;
101        componentTypes.add(module);
102        ComPtr<slang::IEntryPoint> entryPoint;
103        SLANG_RETURN_ON_FAIL(module->findEntryPointByName("rayGenShaderA", entryPoint.writeRef()));
104        componentTypes.add(entryPoint);
105        SLANG_RETURN_ON_FAIL(module->findEntryPointByName("rayGenShaderB", entryPoint.writeRef()));
106        componentTypes.add(entryPoint);
107        SLANG_RETURN_ON_FAIL(module->findEntryPointByName("missShaderA", entryPoint.writeRef()));
108        componentTypes.add(entryPoint);
109        SLANG_RETURN_ON_FAIL(module->findEntryPointByName("missShaderB", entryPoint.writeRef()));
110        componentTypes.add(entryPoint);
111        SLANG_RETURN_ON_FAIL(
112            module->findEntryPointByName("closestHitShaderA", entryPoint.writeRef()));
113        componentTypes.add(entryPoint);
114        SLANG_RETURN_ON_FAIL(
115            module->findEntryPointByName("closestHitShaderB", entryPoint.writeRef()));
116        componentTypes.add(entryPoint);
117
118        ComPtr<slang::IComponentType> linkedProgram;
119        SlangResult result = slangSession->createCompositeComponentType(
120            componentTypes.getBuffer(),
121            componentTypes.getCount(),
122            linkedProgram.writeRef(),
123            diagnosticsBlob.writeRef());
124        SLANG_RETURN_ON_FAIL(result);
125
126        ShaderProgramDesc programDesc = {};
127        programDesc.slangGlobalScope = linkedProgram;
128        SLANG_RETURN_ON_FAIL(device->createShaderProgram(programDesc, outProgram));
129
130        return SLANG_OK;
131    }
132
133    void createResultTexture()
134    {
135        TextureDesc resultTextureDesc = {};
136        resultTextureDesc.type = TextureType::Texture2D;
137        resultTextureDesc.mipCount = 1;
138        resultTextureDesc.size.width = width;
139        resultTextureDesc.size.height = height;
140        resultTextureDesc.size.depth = 1;
141        resultTextureDesc.defaultState = ResourceState::UnorderedAccess;
142        resultTextureDesc.format = Format::RGBA32Float;
143        resultTextureDesc.usage = TextureUsage::UnorderedAccess | TextureUsage::CopySource;
144        resultTexture = device->createTexture(resultTextureDesc);
145        
146        TextureViewDesc resultUAVDesc = {};
147        resultUAVDesc.format = resultTextureDesc.format;
148        resultTextureUAV = resultTexture->createView(resultUAVDesc);
149    }
150
151    void createRequiredResources()
152    {
153        GFX_CHECK_CALL_ABORT(device->getQueue(QueueType::Graphics, queue.writeRef()));
154
155        BufferDesc vertexBufferDesc;
156        vertexBufferDesc.size = kVertexCount * sizeof(Vertex);
157        vertexBufferDesc.defaultState = ResourceState::ShaderResource;
158        vertexBufferDesc.usage = BufferUsage::ShaderResource | BufferUsage::AccelerationStructureBuildInput;
159        vertexBuffer = device->createBuffer(vertexBufferDesc, &kVertexData[0]);
160        SLANG_CHECK_ABORT(vertexBuffer != nullptr);
161
162        BufferDesc indexBufferDesc;
163        indexBufferDesc.size = kIndexCount * sizeof(int32_t);
164        indexBufferDesc.defaultState = ResourceState::ShaderResource;
165        indexBufferDesc.usage = BufferUsage::ShaderResource | BufferUsage::AccelerationStructureBuildInput;
166        indexBuffer = device->createBuffer(indexBufferDesc, &kIndexData[0]);
167        SLANG_CHECK_ABORT(indexBuffer != nullptr);
168
169        BufferDesc transformBufferDesc;
170        transformBufferDesc.size = sizeof(float) * 12;
171        transformBufferDesc.defaultState = ResourceState::ShaderResource;
172        transformBufferDesc.usage = BufferUsage::ShaderResource | BufferUsage::AccelerationStructureBuildInput;
173        float transformData[12] =
174            {1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f};
175        transformBuffer = device->createBuffer(transformBufferDesc, &transformData);
176        SLANG_CHECK_ABORT(transformBuffer != nullptr);
177
178        createResultTexture();
179
180        // Build bottom level acceleration structure.
181        {
182            AccelerationStructureBuildInput geomInput = {};
183            geomInput.type = AccelerationStructureBuildInputType::Triangles;
184            geomInput.triangles.flags = AccelerationStructureGeometryFlags::Opaque;
185            geomInput.triangles.indexCount = kIndexCount;
186            geomInput.triangles.indexBuffer = BufferOffsetPair(indexBuffer, 0);
187            geomInput.triangles.indexFormat = IndexFormat::Uint32;
188            geomInput.triangles.vertexCount = kVertexCount;
189            geomInput.triangles.vertexBuffers[0] = BufferOffsetPair(vertexBuffer, 0);
190            geomInput.triangles.vertexBufferCount = 1;
191            geomInput.triangles.vertexFormat = Format::RGB32Float;
192            geomInput.triangles.vertexStride = sizeof(Vertex);
193            geomInput.triangles.preTransformBuffer = BufferOffsetPair(transformBuffer, 0);
194
195            AccelerationStructureBuildDesc buildInputs = {};
196            buildInputs.inputs = &geomInput;
197            buildInputs.inputCount = 1;
198            buildInputs.flags = AccelerationStructureBuildFlags::AllowCompaction;
199
200            // Query buffer size for acceleration structure build.
201            AccelerationStructureSizes sizes;
202            GFX_CHECK_CALL_ABORT(device->getAccelerationStructureSizes(buildInputs, &sizes));
203            
204            // Allocate buffers for acceleration structure.
205            BufferDesc asDraftBufferDesc;
206            asDraftBufferDesc.defaultState = ResourceState::AccelerationStructure;
207            asDraftBufferDesc.size = sizes.accelerationStructureSize;
208            asDraftBufferDesc.usage = BufferUsage::AccelerationStructure;
209            ComPtr<IBuffer> draftBuffer = device->createBuffer(asDraftBufferDesc);
210            
211            BufferDesc scratchBufferDesc;
212            scratchBufferDesc.defaultState = ResourceState::UnorderedAccess;
213            scratchBufferDesc.size = sizes.scratchSize;
214            scratchBufferDesc.usage = BufferUsage::UnorderedAccess;
215            ComPtr<IBuffer> scratchBuffer = device->createBuffer(scratchBufferDesc);
216
217            // Build acceleration structure.
218            ComPtr<IQueryPool> compactedSizeQuery;
219            QueryPoolDesc queryPoolDesc;
220            queryPoolDesc.count = 1;
221            queryPoolDesc.type = QueryType::AccelerationStructureCompactedSize;
222            GFX_CHECK_CALL_ABORT(
223                device->createQueryPool(queryPoolDesc, compactedSizeQuery.writeRef()));
224
225            ComPtr<IAccelerationStructure> draftAS;
226            AccelerationStructureDesc draftCreateDesc;
227            draftCreateDesc.size = sizes.accelerationStructureSize;
228            GFX_CHECK_CALL_ABORT(
229                device->createAccelerationStructure(draftCreateDesc, draftAS.writeRef()));
230
231            compactedSizeQuery->reset();
232
233            auto commandEncoder = queue->createCommandEncoder();
234            AccelerationStructureQueryDesc compactedSizeQueryDesc = {};
235            compactedSizeQueryDesc.queryPool = compactedSizeQuery;
236            compactedSizeQueryDesc.queryType = QueryType::AccelerationStructureCompactedSize;
237            commandEncoder->buildAccelerationStructure(buildInputs, draftAS, nullptr, BufferOffsetPair(scratchBuffer, 0), 1, &compactedSizeQueryDesc);
238            auto commandBuffer = commandEncoder->finish();
239            queue->submit(commandBuffer);
240            queue->waitOnHost();
241
242            uint64_t compactedSize = 0;
243            compactedSizeQuery->getResult(0, 1, &compactedSize);
244            
245            BufferDesc asBufferDesc;
246            asBufferDesc.defaultState = ResourceState::AccelerationStructure;
247            asBufferDesc.size = (size_t)compactedSize;
248            asBufferDesc.usage = BufferUsage::AccelerationStructure;
249            BLASBuffer = device->createBuffer(asBufferDesc);
250            
251            AccelerationStructureDesc createDesc;
252            createDesc.size = (size_t)compactedSize;
253            device->createAccelerationStructure(createDesc, BLAS.writeRef());
254
255            commandEncoder = queue->createCommandEncoder();
256            commandEncoder->copyAccelerationStructure(
257                BLAS,
258                draftAS,
259                AccelerationStructureCopyMode::Compact);
260            commandBuffer = commandEncoder->finish();
261            queue->submit(commandBuffer);
262            queue->waitOnHost();
263        }
264
265        // Build top level acceleration structure.
266        {
267            List<AccelerationStructureInstanceDescGeneric> instanceDescs;
268            instanceDescs.setCount(1);
269            instanceDescs[0].accelerationStructure.value = BLAS->getDeviceAddress();
270            instanceDescs[0].flags = AccelerationStructureInstanceFlags::TriangleFacingCullDisable;
271            instanceDescs[0].instanceContributionToHitGroupIndex = 0;
272            instanceDescs[0].instanceID = 0;
273            instanceDescs[0].instanceMask = 0xFF;
274            float transformMatrix[] =
275                {1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f};
276            memcpy(&instanceDescs[0].transform[0][0], transformMatrix, sizeof(float) * 12);
277
278            BufferDesc instanceBufferDesc;
279            instanceBufferDesc.size = instanceDescs.getCount() * sizeof(AccelerationStructureInstanceDescGeneric);
280            instanceBufferDesc.defaultState = ResourceState::ShaderResource;
281            instanceBufferDesc.usage = BufferUsage::ShaderResource | BufferUsage::AccelerationStructureBuildInput;
282            instanceBuffer = device->createBuffer(instanceBufferDesc, instanceDescs.getBuffer());
283            SLANG_CHECK_ABORT(instanceBuffer != nullptr);
284
285            AccelerationStructureBuildInput instanceInput = {};
286            instanceInput.type = AccelerationStructureBuildInputType::Instances;
287            instanceInput.instances.instanceBuffer = BufferOffsetPair(instanceBuffer, 0);
288            instanceInput.instances.instanceStride = sizeof(AccelerationStructureInstanceDescGeneric);
289            instanceInput.instances.instanceCount = instanceDescs.getCount();
290
291            AccelerationStructureBuildDesc buildInputs = {};
292            buildInputs.inputs = &instanceInput;
293            buildInputs.inputCount = 1;
294
295            // Query buffer size for acceleration structure build.
296            AccelerationStructureSizes sizes;
297            GFX_CHECK_CALL_ABORT(device->getAccelerationStructureSizes(buildInputs, &sizes));
298
299            BufferDesc asBufferDesc;
300            asBufferDesc.defaultState = ResourceState::AccelerationStructure;
301            asBufferDesc.size = sizes.accelerationStructureSize;
302            asBufferDesc.usage = BufferUsage::AccelerationStructure;
303            TLASBuffer = device->createBuffer(asBufferDesc);
304
305            BufferDesc scratchBufferDesc;
306            scratchBufferDesc.defaultState = ResourceState::UnorderedAccess;
307            scratchBufferDesc.size = sizes.scratchSize;
308            scratchBufferDesc.usage = BufferUsage::UnorderedAccess;
309            ComPtr<IBuffer> scratchBuffer = device->createBuffer(scratchBufferDesc);
310
311            AccelerationStructureDesc createDesc;
312            createDesc.size = sizes.accelerationStructureSize;
313            GFX_CHECK_CALL_ABORT(device->createAccelerationStructure(createDesc, TLAS.writeRef()));
314
315            auto commandEncoder = queue->createCommandEncoder();
316            commandEncoder->buildAccelerationStructure(buildInputs, TLAS, nullptr, BufferOffsetPair(scratchBuffer, 0), 0, nullptr);
317            auto commandBuffer = commandEncoder->finish();
318            queue->submit(commandBuffer);
319            queue->waitOnHost();
320        }
321
322        const char* hitgroupNames[] = {"hitgroupA", "hitgroupB"};
323
324        ComPtr<IShaderProgram> rayTracingProgram;
325        SLANG_CHECK_ABORT(loadShaderProgram(device, rayTracingProgram.writeRef()));
326        RayTracingPipelineDesc rtpDesc = {};
327        rtpDesc.program = rayTracingProgram;
328        rtpDesc.hitGroupCount = 2;
329        HitGroupDesc hitGroups[2];
330        hitGroups[0].closestHitEntryPoint = "closestHitShaderA";
331        hitGroups[0].hitGroupName = hitgroupNames[0];
332        hitGroups[1].closestHitEntryPoint = "closestHitShaderB";
333        hitGroups[1].hitGroupName = hitgroupNames[1];
334        rtpDesc.hitGroups = hitGroups;
335        rtpDesc.maxRayPayloadSize = 64;
336        rtpDesc.maxRecursion = 2;
337        GFX_CHECK_CALL_ABORT(
338            device->createRayTracingPipeline(rtpDesc, renderPipelineState.writeRef()));
339        SLANG_CHECK_ABORT(renderPipelineState != nullptr);
340
341        const char* raygenNames[] = {"rayGenShaderA", "rayGenShaderB"};
342        const char* missNames[] = {"missShaderA", "missShaderB"};
343
344        ShaderTableDesc shaderTableDesc = {};
345        shaderTableDesc.program = rayTracingProgram;
346        shaderTableDesc.hitGroupCount = 2;
347        shaderTableDesc.hitGroupNames = hitgroupNames;
348        shaderTableDesc.rayGenShaderCount = 2;
349        shaderTableDesc.rayGenShaderEntryPointNames = raygenNames;
350        shaderTableDesc.missShaderCount = 2;
351        shaderTableDesc.missShaderEntryPointNames = missNames;
352        GFX_CHECK_CALL_ABORT(device->createShaderTable(shaderTableDesc, shaderTable.writeRef()));
353    }
354
355    void checkTestResults(float* expectedResult, uint32_t count)
356    {
357        ComPtr<ISlangBlob> resultBlob;
358        auto commandEncoder = queue->createCommandEncoder();
359        commandEncoder->setTextureState(resultTexture, ResourceState::CopySource);
360        queue->submit(commandEncoder->finish());
361        queue->waitOnHost();
362
363        SubresourceLayout layout;
364        GFX_CHECK_CALL_ABORT(device->readTexture(
365            resultTexture,
366            0, 0,
367            resultBlob.writeRef(),
368            &layout));
369        size_t rowPitch = layout.rowPitch;
370        size_t pixelSize = 4 ;
371
372#if 0 // for debugging only
373            writeImage("test.hdr", resultBlob, width, height, (uint32_t)rowPitch, (uint32_t)pixelSize);
374#endif
375        auto buffer = removePadding(resultBlob, width, height, rowPitch, pixelSize);
376        auto actualData = (float*)buffer.data();
377        SLANG_CHECK_ABORT(memcmp(actualData, expectedResult, count * sizeof(float)) == 0)
378    }
379};
380
381struct RayTracingTestA : BaseRayTracingTest
382{
383    void renderFrame()
384    {
385        auto commandEncoder = queue->createCommandEncoder();
386        auto renderEncoder = commandEncoder->beginRayTracingPass();
387        auto rootObject = renderEncoder->bindPipeline(renderPipelineState, shaderTable);
388        auto cursor = ShaderCursor(rootObject);
389        cursor["resultTexture"].setBinding(Binding(resultTextureUAV));
390        cursor["sceneBVH"].setBinding(Binding(TLAS));
391        renderEncoder->dispatchRays(0, width, height, 1);
392        renderEncoder->end();
393        auto commandBuffer = commandEncoder->finish();
394        queue->submit(commandBuffer);
395        queue->waitOnHost();
396    }
397
398    void run()
399    {
400        createRequiredResources();
401        renderFrame();
402
403        float expectedResult[16] = {1, 1, 1, 1, 0, 0, 1, 1, 0, 1, 0, 1, 1, 0, 0, 1};
404        checkTestResults(expectedResult, 16);
405    }
406};
407
408struct RayTracingTestB : BaseRayTracingTest
409{
410    void renderFrame()
411    {
412        auto commandEncoder = queue->createCommandEncoder();
413        auto renderEncoder = commandEncoder->beginRayTracingPass();
414        auto rootObject = renderEncoder->bindPipeline(renderPipelineState, shaderTable);
415        auto cursor = ShaderCursor(rootObject);
416        cursor["resultTexture"].setBinding(Binding(resultTextureUAV));
417        cursor["sceneBVH"].setBinding(Binding(TLAS));
418        renderEncoder->dispatchRays(1, width, height, 1);
419        renderEncoder->end();
420        auto commandBuffer = commandEncoder->finish();
421        queue->submit(commandBuffer);
422        queue->waitOnHost();
423    }
424
425    void run()
426    {
427        createRequiredResources();
428        renderFrame();
429
430        float expectedResult[16] = {0, 0, 0, 1, 1, 1, 0, 1, 1, 0, 1, 1, 0, 1, 1, 1};
431        checkTestResults(expectedResult, 16);
432    }
433};
434
435template<typename T>
436void rayTracingTestImpl(IDevice* device, UnitTestContext* context)
437{
438    T test;
439    test.init(device, context);
440    test.run();
441}
442
443SLANG_UNIT_TEST(RayTracingTestAD3D12)
444{
445    runTestImpl(rayTracingTestImpl<RayTracingTestA>, unitTestContext, DeviceType::D3D12);
446}
447
448SLANG_UNIT_TEST(RayTracingTestAVulkan)
449{
450    runTestImpl(rayTracingTestImpl<RayTracingTestA>, unitTestContext, DeviceType::Vulkan);
451}
452
453SLANG_UNIT_TEST(RayTracingTestBD3D12)
454{
455    runTestImpl(rayTracingTestImpl<RayTracingTestB>, unitTestContext, DeviceType::D3D12);
456}
457
458SLANG_UNIT_TEST(RayTracingTestBVulkan)
459{
460    runTestImpl(rayTracingTestImpl<RayTracingTestB>, unitTestContext, DeviceType::Vulkan);
461}
462} // namespace gfx_test
463#endif