yum-mirror/slang

Making it easier to work with shaders

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

Jay KwakAdd command-line arguments to examples (#7835)13dd01489

master
26.8 KiB690 linesraw
1// main.cpp
2
3// This file implements an example of hardware ray-tracing using
4// Slang shaders and the `slang-rhi` graphics API.
5
6#include "core/slang-basic.h"
7#include "examples/example-base/example-base.h"
8#include "platform/vector-math.h"
9#include "platform/window.h"
10#include "slang-com-ptr.h"
11#include "slang-rhi.h"
12#include "slang-rhi/acceleration-structure-utils.h"
13#include "slang-rhi/shader-cursor.h"
14#include "slang.h"
15
16using namespace rhi;
17using namespace Slang;
18
19static const ExampleResources resourceBase("ray-tracing-pipeline");
20
21struct Uniforms
22{
23    float screenWidth, screenHeight;
24    float focalLength = 24.0f, frameHeight = 24.0f;
25    float cameraDir[4];
26    float cameraUp[4];
27    float cameraRight[4];
28    float cameraPosition[4];
29    float lightDir[4];
30};
31
32struct Vertex
33{
34    float position[3];
35};
36
37// Define geometry data for our test scene.
38// The scene contains a floor plane, and a cube placed on top of it at the center.
39static const int kVertexCount = 24;
40static const Vertex kVertexData[kVertexCount] = {
41    // Floor plane
42    {{-100.0f, 0, 100.0f}},
43    {{100.0f, 0, 100.0f}},
44    {{100.0f, 0, -100.0f}},
45    {{-100.0f, 0, -100.0f}},
46    // Cube face (+y).
47    {{-1.0f, 2.0, 1.0f}},
48    {{1.0f, 2.0, 1.0f}},
49    {{1.0f, 2.0, -1.0f}},
50    {{-1.0f, 2.0, -1.0f}},
51    // Cube face (+z).
52    {{-1.0f, 0.0, 1.0f}},
53    {{1.0f, 0.0, 1.0f}},
54    {{1.0f, 2.0, 1.0f}},
55    {{-1.0f, 2.0, 1.0f}},
56    // Cube face (-z).
57    {{-1.0f, 0.0, -1.0f}},
58    {{-1.0f, 2.0, -1.0f}},
59    {{1.0f, 2.0, -1.0f}},
60    {{1.0f, 0.0, -1.0f}},
61    // Cube face (-x).
62    {{-1.0f, 0.0, -1.0f}},
63    {{-1.0f, 0.0, 1.0f}},
64    {{-1.0f, 2.0, 1.0f}},
65    {{-1.0f, 2.0, -1.0f}},
66    // Cube face (+x).
67    {{1.0f, 2.0, -1.0f}},
68    {{1.0f, 2.0, 1.0f}},
69    {{1.0f, 0.0, 1.0f}},
70    {{1.0f, 0.0, -1.0f}},
71};
72static const int kIndexCount = 36;
73static const int kIndexData[kIndexCount] = {0,  1,  2,  0,  2,  3,  4,  5,  6,  4,  6,  7,
74                                            8,  9,  10, 8,  10, 11, 12, 13, 14, 12, 14, 15,
75                                            16, 17, 18, 16, 18, 19, 20, 21, 22, 20, 22, 23};
76
77struct Primitive
78{
79    float data[4];
80    float color[4];
81};
82static const int kPrimitiveCount = 12;
83static const Primitive kPrimitiveData[kPrimitiveCount] = {
84    {{0.0f, 1.0f, 0.0f, 0.0f}, {0.75f, 0.8f, 0.85f, 1.0f}},
85    {{0.0f, 1.0f, 0.0f, 0.0f}, {0.75f, 0.8f, 0.85f, 1.0f}},
86    {{0.0f, 1.0f, 0.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
87    {{0.0f, 1.0f, 0.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
88    {{0.0f, 0.0f, 1.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
89    {{0.0f, 0.0f, 1.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
90    {{0.0f, 0.0f, -1.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
91    {{0.0f, 0.0f, -1.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
92    {{-1.0f, 0.0f, 0.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
93    {{-1.0f, 0.0f, 0.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
94    {{1.0f, 0.0f, 0.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
95    {{1.0f, 0.0f, 0.0f, 0.0f}, {0.95f, 0.85f, 0.05f, 1.0f}},
96};
97
98
99// We need to use a rasterization pipeline to copy the ray-traced image
100// to the swapchain. To do so we need to render a full-screen triangle.
101// We will define a small helper type that defines the data for such a triangle.
102//
103struct FullScreenTriangle
104{
105    struct Vertex
106    {
107        float position[2];
108    };
109
110    enum
111    {
112        kVertexCount = 3
113    };
114
115    static const Vertex kVertices[kVertexCount];
116};
117const FullScreenTriangle::Vertex FullScreenTriangle::kVertices[FullScreenTriangle::kVertexCount] = {
118    {{-1, -1}},
119    {{-1, 3}},
120    {{3, -1}},
121};
122
123// The example application will be implemented as a `struct`, so that
124// we can scope the resources it allocates without using global variables.
125//
126struct RayTracing : public WindowedAppBase
127{
128
129
130    Uniforms gUniforms = {};
131
132
133    // Many Slang API functions return detailed diagnostic information
134    // (error messages, warnings, etc.) as a "blob" of data, or return
135    // a null blob pointer instead if there were no issues.
136    //
137    // For convenience, we define a subroutine that will dump the information
138    // in a diagnostic blob if one is produced, and skip it otherwise.
139    //
140    void diagnoseIfNeeded(slang::IBlob* diagnosticsBlob)
141    {
142        if (diagnosticsBlob != nullptr)
143        {
144            printf("%s", (const char*)diagnosticsBlob->getBufferPointer());
145#ifdef _WIN32
146            _Win32OutputDebugString((const char*)diagnosticsBlob->getBufferPointer());
147#endif
148        }
149    }
150
151    // Load and compile shader code from souce.
152    Result loadShaderProgram(
153        IDevice* device,
154        bool isRayTracingPipeline,
155        IShaderProgram** outProgram)
156    {
157        ComPtr<slang::ISession> slangSession;
158        slangSession = device->getSlangSession();
159
160        ComPtr<slang::IBlob> diagnosticsBlob;
161        Slang::String path = resourceBase.resolveResource("shaders.slang");
162        slang::IModule* module =
163            slangSession->loadModule(path.getBuffer(), diagnosticsBlob.writeRef());
164        diagnoseIfNeeded(diagnosticsBlob);
165        if (!module)
166            return SLANG_FAIL;
167
168        Slang::List<slang::IComponentType*> componentTypes;
169        componentTypes.add(module);
170        if (isRayTracingPipeline)
171        {
172            ComPtr<slang::IEntryPoint> entryPoint;
173            SLANG_RETURN_ON_FAIL(
174                module->findEntryPointByName("rayGenShader", entryPoint.writeRef()));
175            componentTypes.add(entryPoint);
176            SLANG_RETURN_ON_FAIL(module->findEntryPointByName("missShader", entryPoint.writeRef()));
177            componentTypes.add(entryPoint);
178            SLANG_RETURN_ON_FAIL(
179                module->findEntryPointByName("closestHitShader", entryPoint.writeRef()));
180            componentTypes.add(entryPoint);
181            SLANG_RETURN_ON_FAIL(
182                module->findEntryPointByName("shadowRayHitShader", entryPoint.writeRef()));
183            componentTypes.add(entryPoint);
184        }
185        else
186        {
187            ComPtr<slang::IEntryPoint> entryPoint;
188            SLANG_RETURN_ON_FAIL(module->findEntryPointByName("vertexMain", entryPoint.writeRef()));
189            componentTypes.add(entryPoint);
190            SLANG_RETURN_ON_FAIL(
191                module->findEntryPointByName("fragmentMain", entryPoint.writeRef()));
192            componentTypes.add(entryPoint);
193        }
194
195        ComPtr<slang::IComponentType> linkedProgram;
196        SlangResult result = slangSession->createCompositeComponentType(
197            componentTypes.getBuffer(),
198            componentTypes.getCount(),
199            linkedProgram.writeRef(),
200            diagnosticsBlob.writeRef());
201        diagnoseIfNeeded(diagnosticsBlob);
202        SLANG_RETURN_ON_FAIL(result);
203
204        if (isTestMode())
205        {
206            printEntrypointHashes(componentTypes.getCount() - 1, 1, linkedProgram);
207        }
208
209        ShaderProgramDesc programDesc = {};
210        programDesc.slangGlobalScope = linkedProgram;
211        SLANG_RETURN_ON_FAIL(device->createShaderProgram(programDesc, outProgram));
212
213        return SLANG_OK;
214    }
215
216    ComPtr<IRenderPipeline> gPresentPipeline;
217    ComPtr<IRayTracingPipeline> gRenderPipeline;
218    ComPtr<IBuffer> gFullScreenVertexBuffer;
219    ComPtr<IBuffer> gVertexBuffer;
220    ComPtr<IBuffer> gIndexBuffer;
221    ComPtr<IBuffer> gPrimitiveBuffer;
222    ComPtr<IBuffer> gTransformBuffer;
223    ComPtr<IBuffer> gInstanceBuffer;
224    ComPtr<IAccelerationStructure> gBLAS;
225    ComPtr<IAccelerationStructure> gTLAS;
226    ComPtr<ITexture> gResultTexture;
227    ComPtr<IShaderTable> gShaderTable;
228
229    uint64_t lastTime = 0;
230
231    // glm::vec3 lightDir = normalize(glm::vec3(10, 10, 10));
232    // glm::vec3 lightColor = glm::vec3(1, 1, 1);
233
234    glm::vec3 cameraPosition = glm::vec3(-2.53f, 2.72f, 4.3f);
235    float cameraOrientationAngles[2] = {-0.475f, -0.35f}; // Spherical angles (theta, phi).
236
237    float translationScale = 0.5f;
238    float rotationScale = 0.01f;
239
240    // In order to control camera movement, we will
241    // use good old WASD
242    bool wPressed = false;
243    bool aPressed = false;
244    bool sPressed = false;
245    bool dPressed = false;
246
247    bool isMouseDown = false;
248    float lastMouseX = 0.0f;
249    float lastMouseY = 0.0f;
250
251    void setKeyState(platform::KeyCode key, bool state)
252    {
253        switch (key)
254        {
255        default:
256            break;
257        case platform::KeyCode::W:
258            wPressed = state;
259            break;
260        case platform::KeyCode::A:
261            aPressed = state;
262            break;
263        case platform::KeyCode::S:
264            sPressed = state;
265            break;
266        case platform::KeyCode::D:
267            dPressed = state;
268            break;
269        }
270    }
271    void onKeyDown(platform::KeyEventArgs args) { setKeyState(args.key, true); }
272    void onKeyUp(platform::KeyEventArgs args) { setKeyState(args.key, false); }
273
274    void onMouseDown(platform::MouseEventArgs args)
275    {
276        isMouseDown = true;
277        lastMouseX = (float)args.x;
278        lastMouseY = (float)args.y;
279    }
280
281    void onMouseMove(platform::MouseEventArgs args)
282    {
283        if (isMouseDown)
284        {
285            float deltaX = args.x - lastMouseX;
286            float deltaY = args.y - lastMouseY;
287
288            cameraOrientationAngles[0] += -deltaX * rotationScale;
289            cameraOrientationAngles[1] += -deltaY * rotationScale;
290            lastMouseX = (float)args.x;
291            lastMouseY = (float)args.y;
292        }
293    }
294    void onMouseUp(platform::MouseEventArgs args) { isMouseDown = false; }
295
296    Slang::Result initialize()
297    {
298        SLANG_RETURN_ON_FAIL(initializeBase("Ray Tracing Pipeline", 1024, 768, getDeviceType()));
299        if (!isTestMode())
300        {
301            gWindow->events.mouseMove = [this](const platform::MouseEventArgs& e)
302            { onMouseMove(e); };
303            gWindow->events.mouseUp = [this](const platform::MouseEventArgs& e) { onMouseUp(e); };
304            gWindow->events.mouseDown = [this](const platform::MouseEventArgs& e)
305            { onMouseDown(e); };
306            gWindow->events.keyDown = [this](const platform::KeyEventArgs& e) { onKeyDown(e); };
307            gWindow->events.keyUp = [this](const platform::KeyEventArgs& e) { onKeyUp(e); };
308        }
309
310        BufferDesc vertexBufferDesc;
311        vertexBufferDesc.size = kVertexCount * sizeof(Vertex);
312        vertexBufferDesc.usage = BufferUsage::AccelerationStructureBuildInput;
313        vertexBufferDesc.defaultState = ResourceState::AccelerationStructureBuildInput;
314        gVertexBuffer = gDevice->createBuffer(vertexBufferDesc, &kVertexData[0]);
315        if (!gVertexBuffer)
316            return SLANG_FAIL;
317
318        BufferDesc indexBufferDesc;
319        indexBufferDesc.size = kIndexCount * sizeof(int32_t);
320        indexBufferDesc.usage = BufferUsage::AccelerationStructureBuildInput;
321        indexBufferDesc.defaultState = ResourceState::AccelerationStructureBuildInput;
322        gIndexBuffer = gDevice->createBuffer(indexBufferDesc, &kIndexData[0]);
323        if (!gIndexBuffer)
324            return SLANG_FAIL;
325
326        BufferDesc primitiveBufferDesc;
327        primitiveBufferDesc.size = kPrimitiveCount * sizeof(Primitive);
328        primitiveBufferDesc.elementSize = sizeof(Primitive);
329        primitiveBufferDesc.usage = BufferUsage::ShaderResource;
330        primitiveBufferDesc.defaultState = ResourceState::ShaderResource;
331        gPrimitiveBuffer = gDevice->createBuffer(primitiveBufferDesc, &kPrimitiveData[0]);
332        if (!gPrimitiveBuffer)
333            return SLANG_FAIL;
334
335        BufferDesc transformBufferDesc;
336        transformBufferDesc.size = sizeof(float) * 12;
337        transformBufferDesc.usage = BufferUsage::AccelerationStructureBuildInput;
338        transformBufferDesc.defaultState = ResourceState::AccelerationStructureBuildInput;
339        float transformData[12] =
340            {1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f};
341        gTransformBuffer = gDevice->createBuffer(transformBufferDesc, &transformData);
342        if (!gTransformBuffer)
343            return SLANG_FAIL;
344        // Build bottom level acceleration structure.
345        {
346            AccelerationStructureBuildInput buildInput = {};
347            buildInput.type = AccelerationStructureBuildInputType::Triangles;
348            buildInput.triangles.vertexBuffers[0] = gVertexBuffer;
349            buildInput.triangles.vertexBufferCount = 1;
350            buildInput.triangles.vertexFormat = Format::RGB32Float;
351            buildInput.triangles.vertexCount = kVertexCount;
352            buildInput.triangles.vertexStride = sizeof(Vertex);
353            buildInput.triangles.indexBuffer = gIndexBuffer;
354            buildInput.triangles.indexFormat = IndexFormat::Uint32;
355            buildInput.triangles.indexCount = kIndexCount;
356            buildInput.triangles.preTransformBuffer = gTransformBuffer;
357            buildInput.triangles.flags = AccelerationStructureGeometryFlags::Opaque;
358
359            AccelerationStructureBuildDesc buildDesc = {};
360            buildDesc.inputs = &buildInput;
361            buildDesc.inputCount = 1;
362            buildDesc.flags = AccelerationStructureBuildFlags::AllowCompaction;
363
364            // Query buffer size for acceleration structure build.
365            AccelerationStructureSizes sizes;
366            SLANG_RETURN_ON_FAIL(gDevice->getAccelerationStructureSizes(buildDesc, &sizes));
367
368            // Allocate buffers for acceleration structure.
369            BufferDesc scratchBufferDesc;
370            scratchBufferDesc.usage = BufferUsage::UnorderedAccess;
371            scratchBufferDesc.defaultState = ResourceState::UnorderedAccess;
372            scratchBufferDesc.size = sizes.scratchSize;
373            ComPtr<IBuffer> scratchBuffer = gDevice->createBuffer(scratchBufferDesc);
374            if (!scratchBuffer)
375                return SLANG_FAIL;
376
377            // Build acceleration structure.
378            ComPtr<IQueryPool> compactedSizeQuery;
379            QueryPoolDesc queryPoolDesc;
380            queryPoolDesc.count = 1;
381            queryPoolDesc.type = QueryType::AccelerationStructureCompactedSize;
382            SLANG_RETURN_ON_FAIL(
383                gDevice->createQueryPool(queryPoolDesc, compactedSizeQuery.writeRef()));
384
385            ComPtr<IAccelerationStructure> draftAS;
386            AccelerationStructureDesc draftCreateDesc;
387            draftCreateDesc.size = sizes.accelerationStructureSize;
388            SLANG_RETURN_ON_FAIL(
389                gDevice->createAccelerationStructure(draftCreateDesc, draftAS.writeRef()));
390
391            compactedSizeQuery->reset();
392
393            auto commandEncoder = gQueue->createCommandEncoder();
394            AccelerationStructureQueryDesc compactedSizeQueryDesc = {};
395            compactedSizeQueryDesc.queryPool = compactedSizeQuery;
396            compactedSizeQueryDesc.queryType = QueryType::AccelerationStructureCompactedSize;
397            commandEncoder->buildAccelerationStructure(
398                buildDesc,
399                draftAS,
400                nullptr,
401                scratchBuffer,
402                1,
403                &compactedSizeQueryDesc);
404            gQueue->submit(commandEncoder->finish());
405            gQueue->waitOnHost();
406
407            uint64_t compactedSize = 0;
408            compactedSizeQuery->getResult(0, 1, &compactedSize);
409            AccelerationStructureDesc createDesc;
410            createDesc.size = compactedSize;
411            gDevice->createAccelerationStructure(createDesc, gBLAS.writeRef());
412
413            commandEncoder = gQueue->createCommandEncoder();
414            commandEncoder->copyAccelerationStructure(
415                gBLAS,
416                draftAS,
417                AccelerationStructureCopyMode::Compact);
418            gQueue->submit(commandEncoder->finish());
419            gQueue->waitOnHost();
420        }
421
422        // Build top level acceleration structure.
423        {
424            AccelerationStructureInstanceDescType nativeInstanceDescType =
425                getAccelerationStructureInstanceDescType(gDevice);
426            Size nativeInstanceDescSize =
427                getAccelerationStructureInstanceDescSize(nativeInstanceDescType);
428
429            std::vector<AccelerationStructureInstanceDescGeneric> instanceDescs;
430            instanceDescs.resize(1);
431            float transformMatrix[] =
432                {1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f};
433            memcpy(&instanceDescs[0].transform[0][0], transformMatrix, sizeof(float) * 12);
434
435            instanceDescs[0].instanceID = 0;
436            instanceDescs[0].instanceMask = 0xFF;
437            instanceDescs[0].instanceContributionToHitGroupIndex = 0;
438            instanceDescs[0].flags = AccelerationStructureInstanceFlags::TriangleFacingCullDisable;
439            instanceDescs[0].accelerationStructure = gBLAS->getHandle();
440
441            std::vector<uint8_t> nativeInstanceDescs(instanceDescs.size() * nativeInstanceDescSize);
442            convertAccelerationStructureInstanceDescs(
443                instanceDescs.size(),
444                nativeInstanceDescType,
445                nativeInstanceDescs.data(),
446                nativeInstanceDescSize,
447                instanceDescs.data(),
448                sizeof(AccelerationStructureInstanceDescGeneric));
449
450            BufferDesc instanceBufferDesc;
451            instanceBufferDesc.size =
452                instanceDescs.size() * sizeof(AccelerationStructureInstanceDescGeneric);
453            instanceBufferDesc.usage = BufferUsage::ShaderResource;
454            instanceBufferDesc.defaultState = ResourceState::ShaderResource;
455            gInstanceBuffer = gDevice->createBuffer(instanceBufferDesc, nativeInstanceDescs.data());
456            if (!gInstanceBuffer)
457                return SLANG_FAIL;
458
459            AccelerationStructureBuildInput buildInput = {};
460            buildInput.type = AccelerationStructureBuildInputType::Instances;
461            buildInput.instances.instanceBuffer = gInstanceBuffer;
462            buildInput.instances.instanceCount = 1;
463            buildInput.instances.instanceStride = nativeInstanceDescSize;
464
465            AccelerationStructureBuildDesc buildDesc = {};
466            buildDesc.inputs = &buildInput;
467            buildDesc.inputCount = 1;
468
469            // Query buffer size for acceleration structure build.
470            AccelerationStructureSizes sizes;
471            SLANG_RETURN_ON_FAIL(gDevice->getAccelerationStructureSizes(buildDesc, &sizes));
472
473            BufferDesc scratchBufferDesc;
474            scratchBufferDesc.usage = BufferUsage::UnorderedAccess;
475            scratchBufferDesc.defaultState = ResourceState::UnorderedAccess;
476            scratchBufferDesc.size = sizes.scratchSize;
477            ComPtr<IBuffer> scratchBuffer = gDevice->createBuffer(scratchBufferDesc);
478
479            AccelerationStructureDesc createDesc;
480            createDesc.size = sizes.accelerationStructureSize;
481            SLANG_RETURN_ON_FAIL(
482                gDevice->createAccelerationStructure(createDesc, gTLAS.writeRef()));
483
484            auto commandEncoder = gQueue->createCommandEncoder();
485            commandEncoder
486                ->buildAccelerationStructure(buildDesc, gTLAS, nullptr, scratchBuffer, 0, nullptr);
487            gQueue->submit(commandEncoder->finish());
488            gQueue->waitOnHost();
489        }
490
491        BufferDesc fullScreenVertexBufferDesc;
492        fullScreenVertexBufferDesc.size =
493            FullScreenTriangle::kVertexCount * sizeof(FullScreenTriangle::Vertex);
494        fullScreenVertexBufferDesc.usage = BufferUsage::VertexBuffer;
495        fullScreenVertexBufferDesc.defaultState = ResourceState::VertexBuffer;
496        gFullScreenVertexBuffer =
497            gDevice->createBuffer(fullScreenVertexBufferDesc, &FullScreenTriangle::kVertices[0]);
498        if (!gFullScreenVertexBuffer)
499            return SLANG_FAIL;
500
501        InputElementDesc inputElements[] = {
502            {"POSITION", 0, Format::RG32Float, offsetof(FullScreenTriangle::Vertex, position)},
503        };
504        auto inputLayout = gDevice->createInputLayout(
505            sizeof(FullScreenTriangle::Vertex),
506            &inputElements[0],
507            SLANG_COUNT_OF(inputElements));
508        if (!inputLayout)
509            return SLANG_FAIL;
510
511        ComPtr<IShaderProgram> shaderProgram;
512        SLANG_RETURN_ON_FAIL(loadShaderProgram(gDevice, false, shaderProgram.writeRef()));
513        ColorTargetDesc colorTarget;
514        colorTarget.format = Format::RGBA16Float;
515        RenderPipelineDesc desc;
516        desc.inputLayout = inputLayout;
517        desc.program = shaderProgram;
518        desc.targetCount = 1;
519        desc.targets = &colorTarget;
520        desc.depthStencil.depthTestEnable = false;
521        desc.depthStencil.depthWriteEnable = false;
522        desc.primitiveTopology = PrimitiveTopology::TriangleList;
523        gPresentPipeline = gDevice->createRenderPipeline(desc);
524        if (!gPresentPipeline)
525            return SLANG_FAIL;
526
527        const char* hitgroupNames[] = {"hitgroup0", "hitgroup1"};
528
529        ComPtr<IShaderProgram> rayTracingProgram;
530        SLANG_RETURN_ON_FAIL(loadShaderProgram(gDevice, true, rayTracingProgram.writeRef()));
531        RayTracingPipelineDesc rtpDesc = {};
532        rtpDesc.program = rayTracingProgram;
533        rtpDesc.hitGroupCount = 2;
534        HitGroupDesc hitGroups[2];
535        hitGroups[0].closestHitEntryPoint = "closestHitShader";
536        hitGroups[0].hitGroupName = hitgroupNames[0];
537        hitGroups[1].closestHitEntryPoint = "shadowRayHitShader";
538        hitGroups[1].hitGroupName = hitgroupNames[1];
539        rtpDesc.hitGroups = hitGroups;
540        rtpDesc.maxRayPayloadSize = 64;
541        rtpDesc.maxRecursion = 2;
542        SLANG_RETURN_ON_FAIL(
543            gDevice->createRayTracingPipeline(rtpDesc, gRenderPipeline.writeRef()));
544        if (!gRenderPipeline)
545            return SLANG_FAIL;
546
547        ShaderTableDesc shaderTableDesc = {};
548        const char* raygenName = "rayGenShader";
549        const char* missName = "missShader";
550        shaderTableDesc.program = rayTracingProgram;
551        shaderTableDesc.hitGroupCount = 2;
552        shaderTableDesc.hitGroupNames = hitgroupNames;
553        shaderTableDesc.rayGenShaderCount = 1;
554        shaderTableDesc.rayGenShaderEntryPointNames = &raygenName;
555        shaderTableDesc.missShaderCount = 1;
556        shaderTableDesc.missShaderEntryPointNames = &missName;
557        SLANG_RETURN_ON_FAIL(gDevice->createShaderTable(shaderTableDesc, gShaderTable.writeRef()));
558
559        createResultTexture();
560        return SLANG_OK;
561    }
562
563    void createResultTexture()
564    {
565        TextureDesc resultTextureDesc = {};
566        resultTextureDesc.type = TextureType::Texture2D;
567        resultTextureDesc.mipCount = 1;
568        resultTextureDesc.size.width = windowWidth;
569        resultTextureDesc.size.height = windowHeight;
570        resultTextureDesc.size.depth = 1;
571        resultTextureDesc.usage = TextureUsage::UnorderedAccess | TextureUsage::ShaderResource;
572        resultTextureDesc.defaultState = ResourceState::UnorderedAccess;
573        resultTextureDesc.format = Format::RGBA16Float;
574        gResultTexture = gDevice->createTexture(resultTextureDesc);
575    }
576
577    virtual void windowSizeChanged() override
578    {
579        WindowedAppBase::windowSizeChanged();
580        createResultTexture();
581    }
582
583    glm::vec3 getVectorFromSphericalAngles(float theta, float phi)
584    {
585        auto sinTheta = sin(theta);
586        auto cosTheta = cos(theta);
587        auto sinPhi = sin(phi);
588        auto cosPhi = cos(phi);
589        return glm::vec3(-sinTheta * cosPhi, sinPhi, -cosTheta * cosPhi);
590    }
591    void updateUniforms()
592    {
593        gUniforms.screenWidth = (float)windowWidth;
594        gUniforms.screenHeight = (float)windowHeight;
595        if (!lastTime)
596            lastTime = getCurrentTime();
597        uint64_t currentTime = getCurrentTime();
598        float deltaTime = float(double(currentTime - lastTime) / double(getTimerFrequency()));
599        lastTime = currentTime;
600
601        auto camDir =
602            getVectorFromSphericalAngles(cameraOrientationAngles[0], cameraOrientationAngles[1]);
603        auto camUp = getVectorFromSphericalAngles(
604            cameraOrientationAngles[0],
605            cameraOrientationAngles[1] + glm::pi<float>() * 0.5f);
606        auto camRight = glm::cross(camDir, camUp);
607
608        glm::vec3 movement = glm::vec3(0);
609        if (wPressed)
610            movement += camDir;
611        if (sPressed)
612            movement -= camDir;
613        if (aPressed)
614            movement -= camRight;
615        if (dPressed)
616            movement += camRight;
617
618        cameraPosition += deltaTime * translationScale * movement;
619
620        memcpy(gUniforms.cameraDir, &camDir, sizeof(float) * 3);
621        memcpy(gUniforms.cameraUp, &camUp, sizeof(float) * 3);
622        memcpy(gUniforms.cameraRight, &camRight, sizeof(float) * 3);
623        memcpy(gUniforms.cameraPosition, &cameraPosition, sizeof(float) * 3);
624        auto lightDir = glm::normalize(glm::vec3(1.0f, 3.0f, 2.0f));
625        memcpy(gUniforms.lightDir, &lightDir, sizeof(float) * 3);
626    }
627
628    virtual void renderFrame(ITexture* texture) override
629    {
630        updateUniforms();
631        {
632            auto commandEncoder = gQueue->createCommandEncoder();
633            auto rayTracingPassEncoder = commandEncoder->beginRayTracingPass();
634            auto rootObject = rayTracingPassEncoder->bindPipeline(gRenderPipeline, gShaderTable);
635            auto cursor = ShaderCursor(rootObject);
636            cursor["resultTexture"].setBinding(gResultTexture);
637            cursor["uniforms"].setData(&gUniforms, sizeof(Uniforms));
638            cursor["sceneBVH"].setBinding(gTLAS);
639            cursor["primitiveBuffer"].setBinding(gPrimitiveBuffer);
640            rayTracingPassEncoder->dispatchRays(0, windowWidth, windowHeight, 1);
641            rayTracingPassEncoder->end();
642            gQueue->submit(commandEncoder->finish());
643        }
644
645        {
646            auto commandEncoder = gQueue->createCommandEncoder();
647
648            ComPtr<ITextureView> textureView = gDevice->createTextureView(texture, {});
649            RenderPassColorAttachment colorAttachment = {};
650            colorAttachment.view = textureView;
651            colorAttachment.loadOp = LoadOp::Clear;
652
653            RenderPassDesc renderPassDesc = {};
654            renderPassDesc.colorAttachments = &colorAttachment;
655            renderPassDesc.colorAttachmentCount = 1;
656
657            auto renderPassEncoder = commandEncoder->beginRenderPass(renderPassDesc);
658
659            RenderState renderState = {};
660            renderState.viewports[0] = Viewport::fromSize(windowWidth, windowHeight);
661            renderState.viewportCount = 1;
662            renderState.scissorRects[0] = ScissorRect::fromSize(windowWidth, windowHeight);
663            renderState.scissorRectCount = 1;
664            renderState.vertexBuffers[0] = gFullScreenVertexBuffer;
665            renderState.vertexBufferCount = 1;
666            renderPassEncoder->setRenderState(renderState);
667
668            auto rootObject = renderPassEncoder->bindPipeline(gPresentPipeline);
669            auto cursor = ShaderCursor(rootObject);
670            cursor["t"].setBinding(gResultTexture);
671
672            DrawArguments drawArgs = {};
673            drawArgs.vertexCount = 3;
674            renderPassEncoder->draw(drawArgs);
675            renderPassEncoder->end();
676            gQueue->submit(commandEncoder->finish());
677        }
678
679        if (!isTestMode())
680        {
681            // With that, we are done drawing for one frame, and ready for the next.
682            //
683            gSurface->present();
684        }
685    }
686};
687
688// This macro instantiates an appropriate main function to
689// run the application defined above.
690EXAMPLE_MAIN(innerMain<RayTracing>);