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
27.4 KiB724 linesraw
1#include "core/slang-basic.h"
2#include "examples/example-base/example-base.h"
3#include "platform/vector-math.h"
4#include "platform/window.h"
5#include "slang-com-ptr.h"
6#include "slang-rhi.h"
7#include "slang-rhi/shader-cursor.h"
8#include "slang.h"
9
10using namespace rhi;
11using namespace Slang;
12
13static const ExampleResources resourceBase("autodiff-texture");
14
15struct Vertex
16{
17    float position[3];
18};
19
20static const int kVertexCount = 4;
21static const Vertex kVertexData[kVertexCount] = {
22    {{0, 0, 0}},
23    {{0, 1, 0}},
24    {{1, 0, 0}},
25    {{1, 1, 0}},
26};
27float clearValue[] = {0.0f, 0.0f, 0.0f, 0.0f};
28
29
30struct AutoDiffTexture : public WindowedAppBase
31{
32
33    List<uint32_t> mipMapOffset;
34    int textureWidth;
35    int textureHeight;
36
37    void diagnoseIfNeeded(slang::IBlob* diagnosticsBlob)
38    {
39        if (diagnosticsBlob != nullptr)
40        {
41            printf("%s", (const char*)diagnosticsBlob->getBufferPointer());
42        }
43    }
44
45    Result loadRenderProgram(
46        IDevice* device,
47        const char* fileName,
48        const char* fragmentShader,
49        IShaderProgram** outProgram)
50    {
51        ComPtr<slang::ISession> slangSession;
52        slangSession = device->getSlangSession();
53
54        ComPtr<slang::IBlob> diagnosticsBlob;
55        Slang::String path = resourceBase.resolveResource(fileName);
56        slang::IModule* module =
57            slangSession->loadModule(path.getBuffer(), diagnosticsBlob.writeRef());
58        diagnoseIfNeeded(diagnosticsBlob);
59        if (!module)
60            return SLANG_FAIL;
61
62        ComPtr<slang::IEntryPoint> vertexEntryPoint;
63        SLANG_RETURN_ON_FAIL(
64            module->findEntryPointByName("vertexMain", vertexEntryPoint.writeRef()));
65        ComPtr<slang::IEntryPoint> fragmentEntryPoint;
66        SLANG_RETURN_ON_FAIL(
67            module->findEntryPointByName(fragmentShader, fragmentEntryPoint.writeRef()));
68
69        Slang::List<slang::IComponentType*> componentTypes;
70        componentTypes.add(module);
71        int entryPointCount = 0;
72        int vertexEntryPointIndex = entryPointCount++;
73        componentTypes.add(vertexEntryPoint);
74
75        int fragmentEntryPointIndex = entryPointCount++;
76        componentTypes.add(fragmentEntryPoint);
77
78        ComPtr<slang::IComponentType> linkedProgram;
79        SlangResult result = slangSession->createCompositeComponentType(
80            componentTypes.getBuffer(),
81            componentTypes.getCount(),
82            linkedProgram.writeRef(),
83            diagnosticsBlob.writeRef());
84        diagnoseIfNeeded(diagnosticsBlob);
85        SLANG_RETURN_ON_FAIL(result);
86
87        if (isTestMode())
88        {
89            printEntrypointHashes(componentTypes.getCount() - 1, 1, linkedProgram);
90        }
91
92        ShaderProgramDesc programDesc = {};
93        programDesc.slangGlobalScope = linkedProgram;
94        SLANG_RETURN_ON_FAIL(device->createShaderProgram(programDesc, outProgram));
95
96        return SLANG_OK;
97    }
98
99    Result loadComputeProgram(IDevice* device, const char* fileName, IShaderProgram** outProgram)
100    {
101        ComPtr<slang::ISession> slangSession;
102        slangSession = device->getSlangSession();
103
104        ComPtr<slang::IBlob> diagnosticsBlob;
105        Slang::String path = resourceBase.resolveResource(fileName);
106        slang::IModule* module =
107            slangSession->loadModule(path.getBuffer(), diagnosticsBlob.writeRef());
108        diagnoseIfNeeded(diagnosticsBlob);
109        if (!module)
110            return SLANG_FAIL;
111
112        Slang::List<slang::IComponentType*> componentTypes;
113        componentTypes.add(module);
114        ComPtr<slang::IEntryPoint> computeEntryPoint;
115        SLANG_RETURN_ON_FAIL(
116            module->findEntryPointByName("computeMain", computeEntryPoint.writeRef()));
117        componentTypes.add(computeEntryPoint);
118
119        ComPtr<slang::IComponentType> linkedProgram;
120        SlangResult result = slangSession->createCompositeComponentType(
121            componentTypes.getBuffer(),
122            componentTypes.getCount(),
123            linkedProgram.writeRef(),
124            diagnosticsBlob.writeRef());
125        diagnoseIfNeeded(diagnosticsBlob);
126        SLANG_RETURN_ON_FAIL(result);
127
128        if (isTestMode())
129        {
130            printEntrypointHashes(componentTypes.getCount() - 1, 1, linkedProgram);
131        }
132
133        ShaderProgramDesc programDesc = {};
134        programDesc.slangGlobalScope = linkedProgram;
135        SLANG_RETURN_ON_FAIL(device->createShaderProgram(programDesc, outProgram));
136
137        return SLANG_OK;
138    }
139
140    ComPtr<IRenderPipeline> gRefPipeline;
141    ComPtr<IRenderPipeline> gIterPipeline;
142    ComPtr<IComputePipeline> gReconstructPipeline;
143    ComPtr<IComputePipeline> gConvertPipeline;
144    ComPtr<IComputePipeline> gBuildMipPipeline;
145    ComPtr<IComputePipeline> gLearnMipPipeline;
146    ComPtr<IRenderPipeline> gDrawQuadPipeline;
147
148    ComPtr<ITexture> gLearningTexture;
149    ComPtr<ITextureView> gLearningTextureSRV;
150    List<ComPtr<ITextureView>> gLearningTextureUAVs;
151
152    ComPtr<ITexture> gDiffTexture;
153    ComPtr<ITextureView> gDiffTextureSRV;
154    List<ComPtr<ITextureView>> gDiffTextureUAVs;
155
156    ComPtr<IBuffer> gVertexBuffer;
157    ComPtr<ITextureView> gTexView;
158    ComPtr<ISampler> gSampler;
159
160    ComPtr<ITexture> gDepthTexture;
161    ComPtr<ITextureView> gDepthTextureView;
162
163    ComPtr<ITexture> gIterImage;
164    ComPtr<ITextureView> gIterImageSRV;
165
166    ComPtr<ITexture> gRefImage;
167    ComPtr<ITextureView> gRefImageSRV;
168
169    ComPtr<IBuffer> gAccumulateBuffer;
170    ComPtr<ITextureView> gAccumulateBufferView;
171
172    ComPtr<IBuffer> gReconstructBuffer;
173    ComPtr<ITextureView> gReconstructBufferView;
174
175
176    bool resetLearntTexture = false;
177
178    ComPtr<ITexture> createRenderTargetTexture(Format format, int w, int h, int levels)
179    {
180        TextureDesc textureDesc = {};
181        textureDesc.format = format;
182        textureDesc.size.width = w;
183        textureDesc.size.height = h;
184        textureDesc.size.depth = 1;
185        textureDesc.mipCount = levels;
186        textureDesc.usage = TextureUsage::ShaderResource | TextureUsage::UnorderedAccess |
187                            TextureUsage::RenderTarget;
188        textureDesc.defaultState = ResourceState::RenderTarget;
189        return gDevice->createTexture(textureDesc);
190    }
191    ComPtr<ITexture> createDepthTexture()
192    {
193        TextureDesc textureDesc = {};
194        textureDesc.format = Format::D32Float;
195        textureDesc.size.width = windowWidth;
196        textureDesc.size.height = windowHeight;
197        textureDesc.size.depth = 1;
198        textureDesc.mipCount = 1;
199        textureDesc.usage = TextureUsage::DepthStencil;
200        textureDesc.defaultState = ResourceState::DepthWrite;
201        return gDevice->createTexture(textureDesc);
202    }
203    ComPtr<ITextureView> createRTV(ITexture* tex, Format f)
204    {
205        TextureViewDesc rtvDesc = {};
206        rtvDesc.format = f;
207        rtvDesc.subresourceRange.mipCount = 1;
208        return gDevice->createTextureView(tex, rtvDesc);
209    }
210    ComPtr<ITextureView> createDSV(ITexture* tex)
211    {
212        TextureViewDesc dsvDesc = {};
213        dsvDesc.format = Format::D32Float;
214        dsvDesc.subresourceRange.mipCount = 1;
215        return gDevice->createTextureView(tex, dsvDesc);
216    }
217    ComPtr<ITextureView> createSRV(ITexture* tex)
218    {
219        TextureViewDesc srvDesc = {};
220        return gDevice->createTextureView(tex, srvDesc);
221    }
222    ComPtr<IRenderPipeline> createRenderPipeline(IInputLayout* inputLayout, IShaderProgram* program)
223    {
224        ColorTargetDesc colorTarget;
225        colorTarget.format = Format::RGBA8Unorm;
226        RenderPipelineDesc desc;
227        desc.inputLayout = inputLayout;
228        desc.program = program;
229        desc.targetCount = 1;
230        desc.targets = &colorTarget;
231        desc.depthStencil.depthTestEnable = true;
232        desc.depthStencil.depthWriteEnable = true;
233        desc.depthStencil.format = Format::D32Float;
234        desc.rasterizer.cullMode = CullMode::None;
235        desc.primitiveTopology = PrimitiveTopology::TriangleStrip;
236        return gDevice->createRenderPipeline(desc);
237    }
238    ComPtr<IComputePipeline> createComputePipeline(IShaderProgram* program)
239    {
240        ComputePipelineDesc desc = {};
241        desc.program = program;
242        return gDevice->createComputePipeline(desc);
243    }
244    ComPtr<ITextureView> createUAV(ITexture* texture, int level)
245    {
246        TextureViewDesc desc = {};
247        SubresourceRange textureViewRange = {};
248        textureViewRange.mipCount = 1;
249        textureViewRange.mip = level;    // Fixed: should be level, not 0
250        textureViewRange.layerCount = 1; // Fixed: should be 1, not level
251        textureViewRange.layer = 0;
252        desc.subresourceRange = textureViewRange;
253        return gDevice->createTextureView(texture, desc);
254    }
255    Slang::Result initialize()
256    {
257        SLANG_RETURN_ON_FAIL(initializeBase("autodiff-texture", 1024, 768, DeviceType::Default));
258        srand(20421);
259
260        if (!isTestMode())
261        {
262            gWindow->events.keyPress = [this](platform::KeyEventArgs& e)
263            {
264                if (e.keyChar == 'R' || e.keyChar == 'r')
265                    resetLearntTexture = true;
266            };
267        }
268
269        platform::Rect clientRect{};
270        if (isTestMode())
271        {
272            clientRect.width = 1024;
273            clientRect.height = 768;
274        }
275        else
276        {
277            clientRect = getWindow()->getClientRect();
278        }
279
280        windowWidth = clientRect.width;
281        windowHeight = clientRect.height;
282
283        InputElementDesc inputElements[] = {
284            {"POSITION", 0, Format::RGB32Float, offsetof(Vertex, position)}};
285        auto inputLayout = gDevice->createInputLayout(sizeof(Vertex), &inputElements[0], 1);
286        if (!inputLayout)
287            return SLANG_FAIL;
288
289        BufferDesc vertexBufferDesc;
290        vertexBufferDesc.size = kVertexCount * sizeof(Vertex);
291        vertexBufferDesc.elementSize = sizeof(Vertex);
292        vertexBufferDesc.usage = BufferUsage::VertexBuffer;
293        gVertexBuffer = gDevice->createBuffer(vertexBufferDesc, &kVertexData[0]);
294        if (!gVertexBuffer)
295            return SLANG_FAIL;
296
297        {
298            ComPtr<IShaderProgram> shaderProgram;
299            SLANG_RETURN_ON_FAIL(loadRenderProgram(
300                gDevice,
301                "train.slang",
302                "fragmentMain",
303                shaderProgram.writeRef()));
304            gRefPipeline = createRenderPipeline(inputLayout, shaderProgram);
305        }
306        {
307            ComPtr<IShaderProgram> shaderProgram;
308            SLANG_RETURN_ON_FAIL(loadRenderProgram(
309                gDevice,
310                "train.slang",
311                "diffFragmentMain",
312                shaderProgram.writeRef()));
313            gIterPipeline = createRenderPipeline(inputLayout, shaderProgram);
314        }
315        {
316            ComPtr<IShaderProgram> shaderProgram;
317            SLANG_RETURN_ON_FAIL(loadRenderProgram(
318                gDevice,
319                "draw-quad.slang",
320                "fragmentMain",
321                shaderProgram.writeRef()));
322            gDrawQuadPipeline = createRenderPipeline(inputLayout, shaderProgram);
323        }
324        {
325            ComPtr<IShaderProgram> shaderProgram;
326            SLANG_RETURN_ON_FAIL(
327                loadComputeProgram(gDevice, "reconstruct.slang", shaderProgram.writeRef()));
328            gReconstructPipeline = createComputePipeline(shaderProgram);
329        }
330        {
331            ComPtr<IShaderProgram> shaderProgram;
332            SLANG_RETURN_ON_FAIL(
333                loadComputeProgram(gDevice, "convert.slang", shaderProgram.writeRef()));
334            gConvertPipeline = createComputePipeline(shaderProgram);
335        }
336        {
337            ComPtr<IShaderProgram> shaderProgram;
338            SLANG_RETURN_ON_FAIL(
339                loadComputeProgram(gDevice, "buildmip.slang", shaderProgram.writeRef()));
340            gBuildMipPipeline = createComputePipeline(shaderProgram);
341        }
342        {
343            ComPtr<IShaderProgram> shaderProgram;
344            SLANG_RETURN_ON_FAIL(
345                loadComputeProgram(gDevice, "learnmip.slang", shaderProgram.writeRef()));
346            gLearnMipPipeline = createComputePipeline(shaderProgram);
347        }
348
349        // Load texture from file - this would need to be adapted to use slang-rhi texture loading
350        Slang::String imagePath = resourceBase.resolveResource("checkerboard.jpg");
351        gTexView = createTextureFromFile(imagePath.getBuffer(), textureWidth, textureHeight);
352        textureWidth = 512; // Placeholder values
353        textureHeight = 512;
354        initMipOffsets(textureWidth, textureHeight);
355
356        BufferDesc bufferDesc = {};
357        bufferDesc.size = mipMapOffset.getLast() * sizeof(uint32_t);
358        bufferDesc.usage = BufferUsage::ShaderResource | BufferUsage::UnorderedAccess;
359
360        gAccumulateBuffer = gDevice->createBuffer(bufferDesc);
361        if (!gAccumulateBuffer)
362        {
363            printf("ERROR: Failed to create accumulate buffer!\n");
364            return SLANG_FAIL;
365        }
366
367        gReconstructBuffer = gDevice->createBuffer(bufferDesc);
368        if (!gReconstructBuffer)
369        {
370            printf("ERROR: Failed to create reconstruct buffer!\n");
371            return SLANG_FAIL;
372        }
373
374        int mipCount = 1 + Math::Log2Ceil(Math::Max(textureWidth, textureHeight));
375        SubresourceData initialData = {};
376        initialData.data = gLearningTexture =
377            createRenderTargetTexture(Format::RGBA32Float, textureWidth, textureHeight, mipCount);
378        gLearningTextureSRV = createSRV(gLearningTexture);
379        for (int i = 0; i < mipCount; i++)
380            gLearningTextureUAVs.add(createUAV(gLearningTexture, i));
381
382        gDiffTexture =
383            createRenderTargetTexture(Format::RGBA32Float, textureWidth, textureHeight, mipCount);
384        gDiffTextureSRV = createSRV(gDiffTexture);
385        for (int i = 0; i < mipCount; i++)
386            gDiffTextureUAVs.add(createUAV(gDiffTexture, i));
387
388        SamplerDesc samplerDesc = {};
389        gSampler = gDevice->createSampler(samplerDesc);
390
391        gDepthTexture = createDepthTexture();
392        gDepthTextureView = createDSV(gDepthTexture);
393
394        gRefImage = createRenderTargetTexture(Format::RGBA8Unorm, windowWidth, windowHeight, 1);
395        gRefImageSRV = createSRV(gRefImage);
396
397        gIterImage = createRenderTargetTexture(Format::RGBA8Unorm, windowWidth, windowHeight, 1);
398        gIterImageSRV = createSRV(gIterImage);
399
400        // Initialize textures
401        {
402            auto commandEncoder = gQueue->createCommandEncoder();
403            // Clear learning and diff textures
404            commandEncoder->clearTextureFloat(gLearningTexture, kEntireTexture, clearValue);
405            commandEncoder->clearTextureFloat(gDiffTexture, kEntireTexture, clearValue);
406
407            gQueue->submit(commandEncoder->finish());
408        }
409
410        return SLANG_OK;
411    }
412
413    void initMipOffsets(int w, int h)
414    {
415        int layers = 1 + Math::Log2Ceil(Math::Max(w, h));
416        uint32_t offset = 0;
417        for (int i = 0; i < layers; i++)
418        {
419            auto lw = Math::Max(1, w >> i);
420            auto lh = Math::Max(1, h >> i);
421            mipMapOffset.add(offset);
422            offset += lw * lh * 4;
423        }
424        mipMapOffset.add(offset);
425    }
426
427    glm::mat4x4 getTransformMatrix()
428    {
429        float rotX = (rand() / (float)RAND_MAX) * 0.3f;
430        float rotY = (rand() / (float)RAND_MAX) * 0.2f;
431        glm::mat4x4 matProj = glm::perspectiveRH_ZO(
432            glm::radians(60.0f),
433            (float)windowWidth / (float)windowHeight,
434            0.1f,
435            1000.0f);
436        auto identity = glm::mat4(1.0f);
437        auto translate = glm::translate(
438            identity,
439            glm::vec3(
440                -0.6f + 0.2f * (rand() / (float)RAND_MAX),
441                -0.6f + 0.2f * (rand() / (float)RAND_MAX),
442                -1.0f));
443        auto rot = glm::rotate(translate, -glm::pi<float>() * rotX, glm::vec3(1.0f, 0.0f, 0.0f));
444        rot = glm::rotate(rot, -glm::pi<float>() * rotY, glm::vec3(0.0f, 1.0f, 0.0f));
445        auto transformMatrix = matProj * rot;
446        transformMatrix = glm::transpose(transformMatrix);
447        return transformMatrix;
448    }
449
450    template<typename SetupPipelineFunc>
451    void renderImage(ITexture* renderTarget, const SetupPipelineFunc& setupPipeline)
452    {
453        auto commandEncoder = gQueue->createCommandEncoder();
454
455        ComPtr<ITextureView> renderTargetView = createRTV(renderTarget, Format::RGBA8Unorm);
456        RenderPassColorAttachment colorAttachment = {};
457        colorAttachment.view = renderTargetView;
458        colorAttachment.loadOp = LoadOp::Clear;
459        colorAttachment.clearValue[0] = 0.3f;
460        colorAttachment.clearValue[1] = 0.5f;
461        colorAttachment.clearValue[2] = 0.7f;
462        colorAttachment.clearValue[3] = 1.0f;
463
464        RenderPassDepthStencilAttachment depthAttachment = {};
465        depthAttachment.view = gDepthTextureView;
466        depthAttachment.depthLoadOp = LoadOp::Clear;
467        depthAttachment.depthClearValue = 1.0f;
468
469        RenderPassDesc renderPass = {};
470        renderPass.colorAttachments = &colorAttachment;
471        renderPass.colorAttachmentCount = 1;
472        renderPass.depthStencilAttachment = &depthAttachment;
473
474        auto renderEncoder = commandEncoder->beginRenderPass(renderPass);
475
476        RenderState renderState = {};
477        renderState.viewports[0] = Viewport::fromSize(windowWidth, windowHeight);
478        renderState.viewportCount = 1;
479        renderState.scissorRects[0] = ScissorRect::fromSize(windowWidth, windowHeight);
480        renderState.scissorRectCount = 1;
481        renderState.vertexBuffers[0] = gVertexBuffer;
482        renderState.vertexBufferCount = 1;
483
484        setupPipeline(renderEncoder);
485
486
487        renderEncoder->setRenderState(renderState);
488
489        DrawArguments drawArgs = {};
490        drawArgs.vertexCount = 4;
491        renderEncoder->draw(drawArgs);
492        renderEncoder->end();
493        gQueue->submit(commandEncoder->finish());
494    }
495
496    void renderReferenceImage(glm::mat4x4 transformMatrix)
497    {
498        renderImage(
499            gRefImage,
500            [&](IRenderPassEncoder* encoder)
501            {
502                auto rootObject =
503                    encoder->bindPipeline(static_cast<IRenderPipeline*>(gRefPipeline.get()));
504                ShaderCursor rootCursor(rootObject);
505                rootCursor["Uniforms"]["modelViewProjection"].setData(
506                    &transformMatrix,
507                    sizeof(float) * 16);
508                rootCursor["Uniforms"]["bwdTexture"]["texture"].setBinding(gTexView);
509                rootCursor["Uniforms"]["sampler"].setBinding(gSampler);
510                rootCursor["Uniforms"]["mipOffset"].setData(
511                    mipMapOffset.getBuffer(),
512                    sizeof(uint32_t) * mipMapOffset.getCount());
513                rootCursor["Uniforms"]["texRef"].setBinding(gTexView);
514                rootCursor["Uniforms"]["bwdTexture"]["accumulateBuffer"].setBinding(
515                    gAccumulateBuffer);
516            });
517    }
518
519    virtual void renderFrame(ITexture* texture) override
520    {
521        static uint32_t frameCount = 0;
522        frameCount++;
523        auto transformMatrix = getTransformMatrix();
524        renderReferenceImage(transformMatrix);
525
526        // Clear buffers
527        {
528            auto commandEncoder = gQueue->createCommandEncoder();
529            commandEncoder->clearBuffer(gAccumulateBuffer, 0, gAccumulateBuffer->getDesc().size);
530            commandEncoder->clearBuffer(gReconstructBuffer, 0, gReconstructBuffer->getDesc().size);
531
532            if (resetLearntTexture)
533            {
534                commandEncoder->clearTextureFloat(gLearningTexture, kEntireTexture, clearValue);
535                resetLearntTexture = false;
536            }
537            gQueue->submit(commandEncoder->finish());
538        }
539
540        // Render image using backward propagate shader to obtain texture-space gradients.
541        renderImage(
542            gIterImage,
543            [&](IRenderPassEncoder* encoder)
544            {
545                auto rootObject = encoder->bindPipeline(gIterPipeline.get());
546                ShaderCursor rootCursor(rootObject);
547
548                rootCursor["Uniforms"]["modelViewProjection"].setData(
549                    &transformMatrix,
550                    sizeof(float) * 16);
551                rootCursor["Uniforms"]["bwdTexture"]["texture"].setBinding(gLearningTextureSRV);
552                rootCursor["Uniforms"]["sampler"].setBinding(gSampler);
553                rootCursor["Uniforms"]["mipOffset"].setData(
554                    mipMapOffset.getBuffer(),
555                    sizeof(uint32_t) * mipMapOffset.getCount());
556                rootCursor["Uniforms"]["texRef"].setBinding(gRefImageSRV);
557                rootCursor["Uniforms"]["bwdTexture"]["accumulateBuffer"].setBinding(
558                    gAccumulateBuffer);
559                rootCursor["Uniforms"]["bwdTexture"]["minLOD"].setData(5.0);
560            });
561
562        // Propagete gradients through mip map layers from top (lowest res) to bottom (highest res).
563        {
564            auto commandEncoder = gQueue->createCommandEncoder();
565            auto encoder = commandEncoder->beginComputePass();
566            auto rootObject = encoder->bindPipeline(gReconstructPipeline.get());
567            for (int i = (int)mipMapOffset.getCount() - 2; i >= 0; i--)
568            {
569                ShaderCursor rootCursor(rootObject);
570                rootCursor["Uniforms"]["mipOffset"].setData(
571                    mipMapOffset.getBuffer(),
572                    sizeof(uint32_t) * mipMapOffset.getCount());
573                rootCursor["Uniforms"]["dstLayer"].setData(i);
574                rootCursor["Uniforms"]["layerCount"].setData(mipMapOffset.getCount() - 1);
575                rootCursor["Uniforms"]["width"].setData(textureWidth);
576                rootCursor["Uniforms"]["height"].setData(textureHeight);
577                rootCursor["Uniforms"]["accumulateBuffer"].setBinding(gAccumulateBuffer);
578                rootCursor["Uniforms"]["dstBuffer"].setBinding(gReconstructBuffer);
579
580                encoder->dispatchCompute(
581                    ((textureWidth >> i) + 15) / 16,
582                    ((textureHeight >> i) + 15) / 16,
583                    1);
584            }
585            encoder->end();
586            gQueue->submit(commandEncoder->finish());
587
588            commandEncoder = gQueue->createCommandEncoder();
589            // Convert bottom layer mip from buffer to texture
590            {
591                auto encoder = commandEncoder->beginComputePass();
592                auto rootObject = encoder->bindPipeline(gConvertPipeline.get());
593                ShaderCursor rootCursor(rootObject);
594                rootCursor["Uniforms"]["mipOffset"].setData(
595                    mipMapOffset.getBuffer(),
596                    sizeof(uint32_t) * mipMapOffset.getCount());
597                rootCursor["Uniforms"]["dstLayer"].setData(0);
598                rootCursor["Uniforms"]["width"].setData(textureWidth);
599                rootCursor["Uniforms"]["height"].setData(textureHeight);
600                rootCursor["Uniforms"]["srcBuffer"].setBinding(gReconstructBuffer);
601                rootCursor["Uniforms"]["dstTexture"].setBinding(gDiffTextureUAVs[0]);
602                encoder->dispatchCompute((textureWidth + 15) / 16, (textureHeight + 15) / 16, 1);
603                encoder->end();
604            }
605
606            // Build higher level mip map layers
607            encoder = commandEncoder->beginComputePass();
608            rootObject = encoder->bindPipeline(gBuildMipPipeline.get());
609            for (int i = 1; i < (int)mipMapOffset.getCount() - 1; i++)
610            {
611
612                ShaderCursor rootCursor(rootObject);
613                rootCursor["Uniforms"]["dstWidth"].setData(textureWidth >> i);
614                rootCursor["Uniforms"]["dstHeight"].setData(textureHeight >> i);
615                rootCursor["Uniforms"]["srcTexture"].setBinding(gDiffTextureUAVs[i - 1]);
616                rootCursor["Uniforms"]["dstTexture"].setBinding(gDiffTextureUAVs[i]);
617                encoder->dispatchCompute(
618                    ((textureWidth >> i) + 15) / 16,
619                    ((textureHeight >> i) + 15) / 16,
620                    1);
621            }
622            encoder->end();
623
624            // Accumulate gradients to learnt texture
625            encoder = commandEncoder->beginComputePass();
626            rootObject = encoder->bindPipeline(gLearnMipPipeline.get());
627            for (int i = 0; i < (int)mipMapOffset.getCount() - 1; i++)
628            {
629                ShaderCursor rootCursor(rootObject);
630                rootCursor["Uniforms"]["dstWidth"].setData(textureWidth >> i);
631                rootCursor["Uniforms"]["dstHeight"].setData(textureHeight >> i);
632                rootCursor["Uniforms"]["learningRate"].setData(0.1f);
633                rootCursor["Uniforms"]["srcTexture"].setBinding(gDiffTextureUAVs[i]);
634                rootCursor["Uniforms"]["dstTexture"].setBinding(gLearningTextureUAVs[i]);
635                encoder->dispatchCompute(
636                    ((textureWidth >> i) + 15) / 16,
637                    ((textureHeight >> i) + 15) / 16,
638                    1);
639            }
640            encoder->end();
641
642            gQueue->submit(commandEncoder->finish());
643        }
644
645        // Draw currently learnt texture
646        {
647            auto commandEncoder = gQueue->createCommandEncoder();
648
649            ComPtr<ITextureView> textureView = gDevice->createTextureView(texture, {});
650            RenderPassColorAttachment colorAttachment = {};
651            colorAttachment.view = textureView;
652            colorAttachment.loadOp = LoadOp::Clear;
653
654            RenderPassDesc renderPass = {};
655            renderPass.colorAttachments = &colorAttachment;
656            renderPass.colorAttachmentCount = 1;
657
658            auto renderEncoder = commandEncoder->beginRenderPass(renderPass);
659
660            drawTexturedQuad(renderEncoder, 0, 0, textureWidth, textureHeight, gLearningTextureSRV);
661
662            int refImageWidth = windowWidth - textureWidth - 10;
663            int refImageHeight = refImageWidth * windowHeight / windowWidth;
664            drawTexturedQuad(
665                renderEncoder,
666                textureWidth + 10,
667                0,
668                refImageWidth,
669                refImageHeight,
670                gRefImageSRV);
671
672            drawTexturedQuad(
673                renderEncoder,
674                textureWidth + 10,
675                refImageHeight + 10,
676                refImageWidth,
677                refImageHeight,
678                gIterImageSRV);
679            renderEncoder->end();
680            gQueue->submit(commandEncoder->finish());
681        }
682
683        if (!isTestMode())
684        {
685            gSurface->present();
686        }
687    }
688
689    void drawTexturedQuad(
690        IRenderPassEncoder* renderEncoder,
691        int x,
692        int y,
693        int w,
694        int h,
695        ITextureView* srv)
696    {
697        RenderState renderState = {};
698        renderState.viewports[0] = Viewport::fromSize(windowWidth, windowHeight);
699        renderState.viewportCount = 1;
700        renderState.scissorRects[0] = ScissorRect::fromSize(windowWidth, windowHeight);
701        renderState.scissorRectCount = 1;
702        renderState.vertexBuffers[0] = gVertexBuffer;
703        renderState.vertexBufferCount = 1;
704        renderEncoder->setRenderState(renderState);
705
706        auto root =
707            renderEncoder->bindPipeline(static_cast<IRenderPipeline*>(gDrawQuadPipeline.get()));
708        ShaderCursor rootCursor(root);
709        rootCursor["Uniforms"]["x"].setData(x);
710        rootCursor["Uniforms"]["y"].setData(y);
711        rootCursor["Uniforms"]["width"].setData(w);
712        rootCursor["Uniforms"]["height"].setData(h);
713        rootCursor["Uniforms"]["viewWidth"].setData(windowWidth);
714        rootCursor["Uniforms"]["viewHeight"].setData(windowHeight);
715        rootCursor["Uniforms"]["texture"].setBinding(srv);
716        rootCursor["Uniforms"]["sampler"].setBinding(gSampler);
717
718        DrawArguments drawArgs = {};
719        drawArgs.vertexCount = 4;
720        renderEncoder->draw(drawArgs);
721    }
722};
723
724EXAMPLE_MAIN(innerMain<AutoDiffTexture>);