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
25.2 KiB658 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");
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(IDevice* device, bool isComputePipeline, IShaderProgram** outProgram)
153    {
154        ComPtr<slang::ISession> slangSession;
155        slangSession = device->getSlangSession();
156
157        ComPtr<slang::IBlob> diagnosticsBlob;
158        Slang::String path = resourceBase.resolveResource("shaders.slang");
159        slang::IModule* module =
160            slangSession->loadModule(path.getBuffer(), diagnosticsBlob.writeRef());
161        diagnoseIfNeeded(diagnosticsBlob);
162        if (!module)
163            return SLANG_FAIL;
164
165        Slang::List<slang::IComponentType*> componentTypes;
166        componentTypes.add(module);
167        if (isComputePipeline)
168        {
169            ComPtr<slang::IEntryPoint> computeEntryPoint;
170            SLANG_RETURN_ON_FAIL(
171                module->findEntryPointByName("computeMain", computeEntryPoint.writeRef()));
172            componentTypes.add(computeEntryPoint);
173        }
174        else
175        {
176            ComPtr<slang::IEntryPoint> entryPoint;
177            SLANG_RETURN_ON_FAIL(module->findEntryPointByName("vertexMain", entryPoint.writeRef()));
178            componentTypes.add(entryPoint);
179            SLANG_RETURN_ON_FAIL(
180                module->findEntryPointByName("fragmentMain", entryPoint.writeRef()));
181            componentTypes.add(entryPoint);
182        }
183
184        ComPtr<slang::IComponentType> linkedProgram;
185        SlangResult result = slangSession->createCompositeComponentType(
186            componentTypes.getBuffer(),
187            componentTypes.getCount(),
188            linkedProgram.writeRef(),
189            diagnosticsBlob.writeRef());
190        diagnoseIfNeeded(diagnosticsBlob);
191        SLANG_RETURN_ON_FAIL(result);
192
193        if (isTestMode())
194        {
195            printEntrypointHashes(componentTypes.getCount() - 1, 1, linkedProgram);
196        }
197
198        ShaderProgramDesc programDesc = {};
199        programDesc.slangGlobalScope = linkedProgram;
200        SLANG_RETURN_ON_FAIL(device->createShaderProgram(programDesc, outProgram));
201
202        return SLANG_OK;
203    }
204
205    ComPtr<IRenderPipeline> gPresentPipeline;
206    ComPtr<IComputePipeline> gRenderPipeline;
207    ComPtr<IBuffer> gFullScreenVertexBuffer;
208    ComPtr<IBuffer> gVertexBuffer;
209    ComPtr<IBuffer> gIndexBuffer;
210    ComPtr<IBuffer> gPrimitiveBuffer;
211    ComPtr<IBuffer> gTransformBuffer;
212    ComPtr<IBuffer> gInstanceBuffer;
213    ComPtr<IAccelerationStructure> gBLAS;
214    ComPtr<IAccelerationStructure> gTLAS;
215    ComPtr<ITexture> gResultTexture;
216
217    uint64_t lastTime = 0;
218
219    // glm::vec3 lightDir = normalize(glm::vec3(10, 10, 10));
220    // glm::vec3 lightColor = glm::vec3(1, 1, 1);
221
222    glm::vec3 cameraPosition = glm::vec3(-2.53f, 2.72f, 4.3f);
223    float cameraOrientationAngles[2] = {-0.475f, -0.35f}; // Spherical angles (theta, phi).
224
225    float translationScale = 0.5f;
226    float rotationScale = 0.01f;
227
228    // In order to control camera movement, we will
229    // use good old WASD
230    bool wPressed = false;
231    bool aPressed = false;
232    bool sPressed = false;
233    bool dPressed = false;
234
235    bool isMouseDown = false;
236    float lastMouseX = 0.0f;
237    float lastMouseY = 0.0f;
238
239    void setKeyState(platform::KeyCode key, bool state)
240    {
241        switch (key)
242        {
243        default:
244            break;
245        case platform::KeyCode::W:
246            wPressed = state;
247            break;
248        case platform::KeyCode::A:
249            aPressed = state;
250            break;
251        case platform::KeyCode::S:
252            sPressed = state;
253            break;
254        case platform::KeyCode::D:
255            dPressed = state;
256            break;
257        }
258    }
259    void onKeyDown(platform::KeyEventArgs args) { setKeyState(args.key, true); }
260    void onKeyUp(platform::KeyEventArgs args) { setKeyState(args.key, false); }
261
262    void onMouseDown(platform::MouseEventArgs args)
263    {
264        isMouseDown = true;
265        lastMouseX = (float)args.x;
266        lastMouseY = (float)args.y;
267    }
268
269    void onMouseMove(platform::MouseEventArgs args)
270    {
271        if (isMouseDown)
272        {
273            float deltaX = args.x - lastMouseX;
274            float deltaY = args.y - lastMouseY;
275
276            cameraOrientationAngles[0] += -deltaX * rotationScale;
277            cameraOrientationAngles[1] += -deltaY * rotationScale;
278            lastMouseX = (float)args.x;
279            lastMouseY = (float)args.y;
280        }
281    }
282    void onMouseUp(platform::MouseEventArgs args) { isMouseDown = false; }
283
284    Slang::Result initialize()
285    {
286        SLANG_RETURN_ON_FAIL(initializeBase("Ray Tracing", 1024, 768, getDeviceType()));
287
288        if (!isTestMode())
289        {
290            gWindow->events.mouseMove = [this](const platform::MouseEventArgs& e)
291            { onMouseMove(e); };
292            gWindow->events.mouseUp = [this](const platform::MouseEventArgs& e) { onMouseUp(e); };
293            gWindow->events.mouseDown = [this](const platform::MouseEventArgs& e)
294            { onMouseDown(e); };
295            gWindow->events.keyDown = [this](const platform::KeyEventArgs& e) { onKeyDown(e); };
296            gWindow->events.keyUp = [this](const platform::KeyEventArgs& e) { onKeyUp(e); };
297        }
298
299        BufferDesc vertexBufferDesc;
300        vertexBufferDesc.size = kVertexCount * sizeof(Vertex);
301        vertexBufferDesc.usage = BufferUsage::AccelerationStructureBuildInput;
302        vertexBufferDesc.defaultState = ResourceState::AccelerationStructureBuildInput;
303        gVertexBuffer = gDevice->createBuffer(vertexBufferDesc, &kVertexData[0]);
304        if (!gVertexBuffer)
305            return SLANG_FAIL;
306
307        BufferDesc indexBufferDesc;
308        indexBufferDesc.size = kIndexCount * sizeof(int32_t);
309        indexBufferDesc.usage = BufferUsage::AccelerationStructureBuildInput;
310        indexBufferDesc.defaultState = ResourceState::AccelerationStructureBuildInput;
311        gIndexBuffer = gDevice->createBuffer(indexBufferDesc, &kIndexData[0]);
312        if (!gIndexBuffer)
313            return SLANG_FAIL;
314
315        BufferDesc primitiveBufferDesc;
316        primitiveBufferDesc.size = kPrimitiveCount * sizeof(Primitive);
317        primitiveBufferDesc.elementSize = sizeof(Primitive);
318        primitiveBufferDesc.usage = BufferUsage::ShaderResource;
319        primitiveBufferDesc.defaultState = ResourceState::ShaderResource;
320        gPrimitiveBuffer = gDevice->createBuffer(primitiveBufferDesc, &kPrimitiveData[0]);
321        if (!gPrimitiveBuffer)
322            return SLANG_FAIL;
323
324        BufferDesc transformBufferDesc;
325        transformBufferDesc.size = sizeof(float) * 12;
326        transformBufferDesc.usage = BufferUsage::AccelerationStructureBuildInput;
327        transformBufferDesc.defaultState = ResourceState::AccelerationStructureBuildInput;
328        float transformData[12] =
329            {1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f};
330        gTransformBuffer = gDevice->createBuffer(transformBufferDesc, &transformData);
331        if (!gTransformBuffer)
332            return SLANG_FAIL;
333        // Build bottom level acceleration structure.
334        {
335            AccelerationStructureBuildInput buildInput = {};
336            buildInput.type = AccelerationStructureBuildInputType::Triangles;
337            buildInput.triangles.vertexBuffers[0] = gVertexBuffer;
338            buildInput.triangles.vertexBufferCount = 1;
339            buildInput.triangles.vertexFormat = Format::RGB32Float;
340            buildInput.triangles.vertexCount = kVertexCount;
341            buildInput.triangles.vertexStride = sizeof(Vertex);
342            buildInput.triangles.indexBuffer = gIndexBuffer;
343            buildInput.triangles.indexFormat = IndexFormat::Uint32;
344            buildInput.triangles.indexCount = kIndexCount;
345            buildInput.triangles.preTransformBuffer = gTransformBuffer;
346            buildInput.triangles.flags = AccelerationStructureGeometryFlags::Opaque;
347
348            AccelerationStructureBuildDesc buildDesc = {};
349            buildDesc.inputs = &buildInput;
350            buildDesc.inputCount = 1;
351            buildDesc.flags = AccelerationStructureBuildFlags::AllowCompaction;
352
353            // Query buffer size for acceleration structure build.
354            AccelerationStructureSizes sizes;
355            SLANG_RETURN_ON_FAIL(gDevice->getAccelerationStructureSizes(buildDesc, &sizes));
356
357            // Allocate buffers for acceleration structure.
358            BufferDesc scratchBufferDesc;
359            scratchBufferDesc.usage = BufferUsage::UnorderedAccess;
360            scratchBufferDesc.defaultState = ResourceState::UnorderedAccess;
361            scratchBufferDesc.size = sizes.scratchSize;
362            ComPtr<IBuffer> scratchBuffer = gDevice->createBuffer(scratchBufferDesc);
363            if (!scratchBuffer)
364                return SLANG_FAIL;
365
366            // Build acceleration structure.
367            ComPtr<IQueryPool> compactedSizeQuery;
368            QueryPoolDesc queryPoolDesc;
369            queryPoolDesc.count = 1;
370            queryPoolDesc.type = QueryType::AccelerationStructureCompactedSize;
371            SLANG_RETURN_ON_FAIL(
372                gDevice->createQueryPool(queryPoolDesc, compactedSizeQuery.writeRef()));
373
374            ComPtr<IAccelerationStructure> draftAS;
375            AccelerationStructureDesc draftCreateDesc;
376            draftCreateDesc.size = sizes.accelerationStructureSize;
377            SLANG_RETURN_ON_FAIL(
378                gDevice->createAccelerationStructure(draftCreateDesc, draftAS.writeRef()));
379
380            compactedSizeQuery->reset();
381
382            auto commandEncoder = gQueue->createCommandEncoder();
383            AccelerationStructureQueryDesc compactedSizeQueryDesc = {};
384            compactedSizeQueryDesc.queryPool = compactedSizeQuery;
385            compactedSizeQueryDesc.queryType = QueryType::AccelerationStructureCompactedSize;
386            commandEncoder->buildAccelerationStructure(
387                buildDesc,
388                draftAS,
389                nullptr,
390                scratchBuffer,
391                1,
392                &compactedSizeQueryDesc);
393            gQueue->submit(commandEncoder->finish());
394            gQueue->waitOnHost();
395
396            uint64_t compactedSize = 0;
397            compactedSizeQuery->getResult(0, 1, &compactedSize);
398            AccelerationStructureDesc createDesc;
399            createDesc.size = compactedSize;
400            gDevice->createAccelerationStructure(createDesc, gBLAS.writeRef());
401
402            commandEncoder = gQueue->createCommandEncoder();
403            commandEncoder->copyAccelerationStructure(
404                gBLAS,
405                draftAS,
406                AccelerationStructureCopyMode::Compact);
407            gQueue->submit(commandEncoder->finish());
408            gQueue->waitOnHost();
409        }
410
411        // Build top level acceleration structure.
412        {
413            AccelerationStructureInstanceDescType nativeInstanceDescType =
414                getAccelerationStructureInstanceDescType(gDevice);
415            Size nativeInstanceDescSize =
416                getAccelerationStructureInstanceDescSize(nativeInstanceDescType);
417
418            std::vector<AccelerationStructureInstanceDescGeneric> instanceDescs;
419            instanceDescs.resize(1);
420            float transformMatrix[] =
421                {1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f};
422            memcpy(&instanceDescs[0].transform[0][0], transformMatrix, sizeof(float) * 12);
423
424            instanceDescs[0].instanceID = 0;
425            instanceDescs[0].instanceMask = 0xFF;
426            instanceDescs[0].instanceContributionToHitGroupIndex = 0;
427            instanceDescs[0].flags = AccelerationStructureInstanceFlags::TriangleFacingCullDisable;
428            instanceDescs[0].accelerationStructure = gBLAS->getHandle();
429
430            std::vector<uint8_t> nativeInstanceDescs(instanceDescs.size() * nativeInstanceDescSize);
431            convertAccelerationStructureInstanceDescs(
432                instanceDescs.size(),
433                nativeInstanceDescType,
434                nativeInstanceDescs.data(),
435                nativeInstanceDescSize,
436                instanceDescs.data(),
437                sizeof(AccelerationStructureInstanceDescGeneric));
438
439            BufferDesc instanceBufferDesc;
440            instanceBufferDesc.size =
441                instanceDescs.size() * sizeof(AccelerationStructureInstanceDescGeneric);
442            instanceBufferDesc.usage = BufferUsage::ShaderResource;
443            instanceBufferDesc.defaultState = ResourceState::ShaderResource;
444            gInstanceBuffer = gDevice->createBuffer(instanceBufferDesc, nativeInstanceDescs.data());
445            if (!gInstanceBuffer)
446                return SLANG_FAIL;
447
448            AccelerationStructureBuildInput buildInput = {};
449            buildInput.type = AccelerationStructureBuildInputType::Instances;
450            buildInput.instances.instanceBuffer = gInstanceBuffer;
451            buildInput.instances.instanceCount = 1;
452            buildInput.instances.instanceStride = nativeInstanceDescSize;
453
454            AccelerationStructureBuildDesc buildDesc = {};
455            buildDesc.inputs = &buildInput;
456            buildDesc.inputCount = 1;
457
458            // Query buffer size for acceleration structure build.
459            AccelerationStructureSizes sizes;
460            SLANG_RETURN_ON_FAIL(gDevice->getAccelerationStructureSizes(buildDesc, &sizes));
461
462            BufferDesc scratchBufferDesc;
463            scratchBufferDesc.usage = BufferUsage::UnorderedAccess;
464            scratchBufferDesc.defaultState = ResourceState::UnorderedAccess;
465            scratchBufferDesc.size = sizes.scratchSize;
466            ComPtr<IBuffer> scratchBuffer = gDevice->createBuffer(scratchBufferDesc);
467
468            AccelerationStructureDesc createDesc;
469            createDesc.size = sizes.accelerationStructureSize;
470            SLANG_RETURN_ON_FAIL(
471                gDevice->createAccelerationStructure(createDesc, gTLAS.writeRef()));
472
473            auto commandEncoder = gQueue->createCommandEncoder();
474            commandEncoder
475                ->buildAccelerationStructure(buildDesc, gTLAS, nullptr, scratchBuffer, 0, nullptr);
476            gQueue->submit(commandEncoder->finish());
477            gQueue->waitOnHost();
478        }
479
480        BufferDesc fullScreenVertexBufferDesc;
481        fullScreenVertexBufferDesc.size =
482            FullScreenTriangle::kVertexCount * sizeof(FullScreenTriangle::Vertex);
483        fullScreenVertexBufferDesc.usage = BufferUsage::VertexBuffer;
484        fullScreenVertexBufferDesc.defaultState = ResourceState::VertexBuffer;
485        gFullScreenVertexBuffer =
486            gDevice->createBuffer(fullScreenVertexBufferDesc, &FullScreenTriangle::kVertices[0]);
487        if (!gFullScreenVertexBuffer)
488            return SLANG_FAIL;
489
490        InputElementDesc inputElements[] = {
491            {"POSITION", 0, Format::RG32Float, offsetof(FullScreenTriangle::Vertex, position)},
492        };
493        auto inputLayout = gDevice->createInputLayout(
494            sizeof(FullScreenTriangle::Vertex),
495            &inputElements[0],
496            SLANG_COUNT_OF(inputElements));
497        if (!inputLayout)
498            return SLANG_FAIL;
499
500        ComPtr<IShaderProgram> shaderProgram;
501        SLANG_RETURN_ON_FAIL(loadShaderProgram(gDevice, false, shaderProgram.writeRef()));
502        ColorTargetDesc colorTarget;
503        colorTarget.format = Format::RGBA16Float;
504        RenderPipelineDesc desc;
505        desc.inputLayout = inputLayout;
506        desc.program = shaderProgram;
507        desc.targetCount = 1;
508        desc.targets = &colorTarget;
509        desc.depthStencil.depthTestEnable = false;
510        desc.depthStencil.depthWriteEnable = false;
511        desc.primitiveTopology = PrimitiveTopology::TriangleList;
512        gPresentPipeline = gDevice->createRenderPipeline(desc);
513        if (!gPresentPipeline)
514            return SLANG_FAIL;
515
516        ComPtr<IShaderProgram> computeProgram;
517        SLANG_RETURN_ON_FAIL(loadShaderProgram(gDevice, true, computeProgram.writeRef()));
518        ComputePipelineDesc computeDesc;
519        computeDesc.program = computeProgram;
520        gRenderPipeline = gDevice->createComputePipeline(computeDesc);
521        if (!gRenderPipeline)
522            return SLANG_FAIL;
523
524        createResultTexture();
525        return SLANG_OK;
526    }
527
528    void createResultTexture()
529    {
530        TextureDesc resultTextureDesc = {};
531        resultTextureDesc.type = TextureType::Texture2D;
532        resultTextureDesc.mipCount = 1;
533        resultTextureDesc.size.width = windowWidth;
534        resultTextureDesc.size.height = windowHeight;
535        resultTextureDesc.size.depth = 1;
536        resultTextureDesc.usage = TextureUsage::UnorderedAccess | TextureUsage::ShaderResource;
537        resultTextureDesc.defaultState = ResourceState::UnorderedAccess;
538        resultTextureDesc.format = Format::RGBA16Float;
539        gResultTexture = gDevice->createTexture(resultTextureDesc);
540    }
541
542    virtual void windowSizeChanged() override
543    {
544        WindowedAppBase::windowSizeChanged();
545        createResultTexture();
546    }
547
548    glm::vec3 getVectorFromSphericalAngles(float theta, float phi)
549    {
550        auto sinTheta = sin(theta);
551        auto cosTheta = cos(theta);
552        auto sinPhi = sin(phi);
553        auto cosPhi = cos(phi);
554        return glm::vec3(-sinTheta * cosPhi, sinPhi, -cosTheta * cosPhi);
555    }
556    void updateUniforms()
557    {
558        gUniforms.screenWidth = (float)windowWidth;
559        gUniforms.screenHeight = (float)windowHeight;
560        if (!lastTime)
561            lastTime = getCurrentTime();
562        uint64_t currentTime = getCurrentTime();
563        float deltaTime = float(double(currentTime - lastTime) / double(getTimerFrequency()));
564        lastTime = currentTime;
565
566        auto camDir =
567            getVectorFromSphericalAngles(cameraOrientationAngles[0], cameraOrientationAngles[1]);
568        auto camUp = getVectorFromSphericalAngles(
569            cameraOrientationAngles[0],
570            cameraOrientationAngles[1] + glm::pi<float>() * 0.5f);
571        auto camRight = glm::cross(camDir, camUp);
572
573        glm::vec3 movement = glm::vec3(0);
574        if (wPressed)
575            movement += camDir;
576        if (sPressed)
577            movement -= camDir;
578        if (aPressed)
579            movement -= camRight;
580        if (dPressed)
581            movement += camRight;
582
583        cameraPosition += deltaTime * translationScale * movement;
584
585        memcpy(gUniforms.cameraDir, &camDir, sizeof(float) * 3);
586        memcpy(gUniforms.cameraUp, &camUp, sizeof(float) * 3);
587        memcpy(gUniforms.cameraRight, &camRight, sizeof(float) * 3);
588        memcpy(gUniforms.cameraPosition, &cameraPosition, sizeof(float) * 3);
589        auto lightDir = glm::normalize(glm::vec3(1.0f, 3.0f, 2.0f));
590        memcpy(gUniforms.lightDir, &lightDir, sizeof(float) * 3);
591    }
592
593    virtual void renderFrame(ITexture* texture) override
594    {
595        updateUniforms();
596        {
597            auto commandEncoder = gQueue->createCommandEncoder();
598            auto computePassEncoder = commandEncoder->beginComputePass();
599            auto rootObject = computePassEncoder->bindPipeline(gRenderPipeline);
600            auto cursor = ShaderCursor(rootObject);
601            cursor["resultTexture"].setBinding(gResultTexture);
602            cursor["uniforms"].setData(&gUniforms, sizeof(Uniforms));
603            cursor["sceneBVH"].setBinding(gTLAS);
604            cursor["primitiveBuffer"].setBinding(gPrimitiveBuffer);
605            computePassEncoder->dispatchCompute(
606                (windowWidth + 15) / 16,
607                (windowHeight + 15) / 16,
608                1);
609            computePassEncoder->end();
610            gQueue->submit(commandEncoder->finish());
611        }
612
613        {
614            auto commandEncoder = gQueue->createCommandEncoder();
615
616            ComPtr<ITextureView> textureView = gDevice->createTextureView(texture, {});
617            RenderPassColorAttachment colorAttachment = {};
618            colorAttachment.view = textureView;
619            colorAttachment.loadOp = LoadOp::Clear;
620
621            RenderPassDesc renderPassDesc = {};
622            renderPassDesc.colorAttachments = &colorAttachment;
623            renderPassDesc.colorAttachmentCount = 1;
624
625            auto renderPassEncoder = commandEncoder->beginRenderPass(renderPassDesc);
626
627            RenderState renderState = {};
628            renderState.viewports[0] = Viewport::fromSize(windowWidth, windowHeight);
629            renderState.viewportCount = 1;
630            renderState.scissorRects[0] = ScissorRect::fromSize(windowWidth, windowHeight);
631            renderState.scissorRectCount = 1;
632            renderState.vertexBuffers[0] = gFullScreenVertexBuffer;
633            renderState.vertexBufferCount = 1;
634            renderPassEncoder->setRenderState(renderState);
635
636            auto rootObject = renderPassEncoder->bindPipeline(gPresentPipeline);
637            auto cursor = ShaderCursor(rootObject);
638            cursor["t"].setBinding(gResultTexture);
639
640            DrawArguments drawArgs = {};
641            drawArgs.vertexCount = 3;
642            renderPassEncoder->draw(drawArgs);
643            renderPassEncoder->end();
644            gQueue->submit(commandEncoder->finish());
645        }
646
647        if (!isTestMode())
648        {
649            // With that, we are done drawing for one frame, and ready for the next.
650            //
651            gSurface->present();
652        }
653    }
654};
655
656// This macro instantiates an appropriate main function to
657// run the application defined above.
658EXAMPLE_MAIN(innerMain<RayTracing>);