yum-mirror/slang

Making it easier to work with shaders

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

Ellie Hermaszewskaformatf65d756bf

master
13.6 KiB392 linesraw
1#pragma once
2
3#include "core/slang-basic.h"
4#include "core/slang-com-object.h"
5#include "renderer-shared.h"
6#include "slang-gfx.h"
7
8namespace gfx
9{
10class ShaderObjectLayoutBase;
11
12template<typename T>
13class VersionedObjectPool
14{
15public:
16    struct ObjectVersion
17    {
18        Slang::RefPtr<T> object;
19        Slang::RefPtr<TransientResourceHeapBase> transientHeap;
20        uint64_t transientHeapVersion;
21        bool canRecycle() { return (transientHeap->getVersion() != transientHeapVersion); }
22    };
23    Slang::List<ObjectVersion> objects;
24    SlangInt lastAllocationIndex = -1;
25    ObjectVersion& allocate(TransientResourceHeapBase* currentTransientHeap)
26    {
27        for (SlangInt i = 0; i < objects.getCount(); i++)
28        {
29            auto& object = objects[i];
30            if (object.canRecycle())
31            {
32                object.transientHeap = currentTransientHeap;
33                object.transientHeapVersion = currentTransientHeap->getVersion();
34                lastAllocationIndex = i;
35                return object;
36            }
37        }
38        ObjectVersion v;
39        v.transientHeap = currentTransientHeap;
40        v.transientHeapVersion = currentTransientHeap->getVersion();
41        objects.add(v);
42        lastAllocationIndex = objects.getCount() - 1;
43        return objects.getLast();
44    }
45    ObjectVersion& getLastAllocation() { return objects[lastAllocationIndex]; }
46};
47
48class MutableShaderObjectData
49{
50public:
51    // Any "ordinary" / uniform data for this object
52    Slang::List<char> m_ordinaryData;
53
54    bool m_dirty = true;
55
56    Slang::Index getCount() { return m_ordinaryData.getCount(); }
57    void setCount(Slang::Index count) { m_ordinaryData.setCount(count); }
58    char* getBuffer() { return m_ordinaryData.getBuffer(); }
59    void markDirty() { m_dirty = true; }
60
61    // We don't actually create any GPU buffers here, since they will be handled
62    // by the immutable shader objects once the user calls `getCurrentVersion`.
63    ResourceViewBase* getResourceView(
64        RendererBase* device,
65        slang::TypeLayoutReflection* elementLayout,
66        slang::BindingType bindingType)
67    {
68        return nullptr;
69    }
70};
71
72template<typename TShaderObject, typename TShaderObjectLayoutImpl>
73class MutableShaderObject
74    : public ShaderObjectBaseImpl<TShaderObject, TShaderObjectLayoutImpl, MutableShaderObjectData>
75{
76    typedef ShaderObjectBaseImpl<TShaderObject, TShaderObjectLayoutImpl, MutableShaderObjectData>
77        Super;
78
79protected:
80    Slang::OrderedDictionary<ShaderOffset, Slang::RefPtr<ResourceViewBase>> m_resources;
81    Slang::OrderedDictionary<ShaderOffset, Slang::RefPtr<SamplerStateBase>> m_samplers;
82    Slang::OrderedHashSet<ShaderOffset> m_objectOffsets;
83    VersionedObjectPool<ShaderObjectBase> m_shaderObjectVersions;
84    bool m_dirty = true;
85    bool isDirty()
86    {
87        if (m_dirty)
88            return true;
89        if (this->m_data.m_dirty)
90            return true;
91        for (auto& object : this->m_objects)
92        {
93            if (object && object->isDirty())
94                return true;
95        }
96        return false;
97    }
98
99    void markDirty() { m_dirty = true; }
100
101public:
102    Result init(RendererBase* device, ShaderObjectLayoutBase* layout)
103    {
104        this->m_device = device;
105        auto layoutImpl = static_cast<TShaderObjectLayoutImpl*>(layout);
106        this->m_layout = layoutImpl;
107        Slang::Index subObjectCount = layoutImpl->getSubObjectCount();
108        this->m_objects.setCount(subObjectCount);
109        auto dataSize = layoutImpl->getElementTypeLayout()->getSize();
110        assert(dataSize >= 0);
111        this->m_data.setCount(dataSize);
112        memset(this->m_data.getBuffer(), 0, dataSize);
113        return SLANG_OK;
114    }
115
116public:
117    virtual SLANG_NO_THROW const void* SLANG_MCALL getRawData() override
118    {
119        return this->m_data.getBuffer();
120    }
121    virtual SLANG_NO_THROW size_t SLANG_MCALL getSize() override { return this->m_data.getCount(); }
122    virtual SLANG_NO_THROW Result SLANG_MCALL
123    setData(ShaderOffset const& offset, void const* data, size_t size) override
124    {
125        if (!size)
126            return SLANG_OK;
127        if (SlangInt(offset.uniformOffset + size) > this->m_data.getCount())
128            this->m_data.setCount(offset.uniformOffset + size);
129        memcpy(this->m_data.getBuffer() + offset.uniformOffset, data, size);
130        this->m_data.markDirty();
131        markDirty();
132        return SLANG_OK;
133    }
134
135    virtual SLANG_NO_THROW Result SLANG_MCALL
136    setObject(ShaderOffset const& offset, IShaderObject* object) override
137    {
138        Super::setObject(offset, object);
139        m_objectOffsets.add(offset);
140        markDirty();
141        return SLANG_OK;
142    }
143
144    virtual SLANG_NO_THROW Result SLANG_MCALL
145    setResource(ShaderOffset const& offset, IResourceView* resourceView) override
146    {
147        m_resources[offset] = static_cast<ResourceViewBase*>(resourceView);
148        markDirty();
149        return SLANG_OK;
150    }
151
152    virtual SLANG_NO_THROW Result SLANG_MCALL
153    setSampler(ShaderOffset const& offset, ISamplerState* sampler) override
154    {
155        m_samplers[offset] = static_cast<SamplerStateBase*>(sampler);
156        markDirty();
157        return SLANG_OK;
158    }
159
160    virtual SLANG_NO_THROW Result SLANG_MCALL setCombinedTextureSampler(
161        ShaderOffset const& offset,
162        IResourceView* textureView,
163        ISamplerState* sampler) override
164    {
165        m_samplers[offset] = static_cast<SamplerStateBase*>(sampler);
166        m_resources[offset] = static_cast<ResourceViewBase*>(textureView);
167        markDirty();
168        return SLANG_OK;
169    }
170
171    virtual SLANG_NO_THROW Result SLANG_MCALL
172    getCurrentVersion(ITransientResourceHeap* transientHeap, IShaderObject** outObject) override
173    {
174        if (!isDirty())
175        {
176            returnComPtr(outObject, getLastAllocatedShaderObject());
177            return SLANG_OK;
178        }
179
180        Slang::RefPtr<ShaderObjectBase> object =
181            allocateShaderObject(static_cast<TransientResourceHeapBase*>(transientHeap));
182        SLANG_RETURN_ON_FAIL(
183            object->setData(ShaderOffset(), this->m_data.getBuffer(), this->m_data.getCount()));
184        for (auto res : m_resources)
185            SLANG_RETURN_ON_FAIL(object->setResource(res.key, res.value));
186        for (auto sampler : m_samplers)
187            SLANG_RETURN_ON_FAIL(object->setSampler(sampler.key, sampler.value));
188        for (auto offset : m_objectOffsets)
189        {
190            if (offset.bindingRangeIndex < 0)
191                return SLANG_E_INVALID_ARG;
192            auto layout = this->getLayout();
193            if (offset.bindingRangeIndex >= layout->getBindingRangeCount())
194                return SLANG_E_INVALID_ARG;
195            auto bindingRange = layout->getBindingRange(offset.bindingRangeIndex);
196
197            auto subObject =
198                this->m_objects[bindingRange.subObjectIndex + offset.bindingArrayIndex];
199            if (subObject)
200            {
201                ComPtr<IShaderObject> subObjectVersion;
202                SLANG_RETURN_ON_FAIL(
203                    subObject->getCurrentVersion(transientHeap, subObjectVersion.writeRef()));
204                SLANG_RETURN_ON_FAIL(object->setObject(offset, subObjectVersion));
205            }
206        }
207        m_dirty = false;
208        this->m_data.m_dirty = false;
209        returnComPtr(outObject, object);
210        return SLANG_OK;
211    }
212
213public:
214    Slang::RefPtr<ShaderObjectBase> allocateShaderObject(TransientResourceHeapBase* transientHeap)
215    {
216        auto& version = m_shaderObjectVersions.allocate(transientHeap);
217        if (!version.object)
218        {
219            ComPtr<IShaderObject> shaderObject;
220            SLANG_RETURN_NULL_ON_FAIL(
221                this->m_device->createShaderObject(this->m_layout, shaderObject.writeRef()));
222            version.object = static_cast<ShaderObjectBase*>(shaderObject.get());
223        }
224        return version.object;
225    }
226    Slang::RefPtr<ShaderObjectBase> getLastAllocatedShaderObject()
227    {
228        return m_shaderObjectVersions.getLastAllocation().object;
229    }
230};
231
232// A proxy shader object to hold mutable shader parameters for global scope and entry-points.
233class MutableRootShaderObject : public ShaderObjectBase
234{
235public:
236    Slang::List<uint8_t> m_data;
237    Slang::OrderedDictionary<ShaderOffset, Slang::RefPtr<ResourceViewBase>> m_resources;
238    Slang::OrderedDictionary<ShaderOffset, Slang::RefPtr<SamplerStateBase>> m_samplers;
239    Slang::OrderedDictionary<ShaderOffset, Slang::RefPtr<ShaderObjectBase>> m_objects;
240    Slang::OrderedDictionary<ShaderOffset, Slang::List<slang::SpecializationArg>>
241        m_specializationArgs;
242    Slang::List<Slang::RefPtr<MutableRootShaderObject>> m_entryPoints;
243    Slang::RefPtr<BufferResource> m_constantBufferOverride;
244    slang::TypeLayoutReflection* m_elementTypeLayout;
245
246    MutableRootShaderObject(RendererBase* device, slang::TypeLayoutReflection* entryPointLayout)
247    {
248        this->m_device = device;
249        m_elementTypeLayout = entryPointLayout;
250        m_data.setCount(entryPointLayout->getSize());
251        memset(m_data.begin(), 0, m_data.getCount());
252    }
253
254    MutableRootShaderObject(RendererBase* device, Slang::RefPtr<ShaderProgramBase> program)
255    {
256        this->m_device = device;
257        auto programLayout = program->slangGlobalScope->getLayout();
258        SlangInt entryPointCount = programLayout->getEntryPointCount();
259        for (SlangInt e = 0; e < entryPointCount; ++e)
260        {
261            auto slangEntryPoint = programLayout->getEntryPointByIndex(e);
262            Slang::RefPtr<MutableRootShaderObject> entryPointObject = new MutableRootShaderObject(
263                device,
264                slangEntryPoint->getTypeLayout()->getElementTypeLayout());
265
266            m_entryPoints.add(entryPointObject);
267        }
268        m_data.setCount(programLayout->getGlobalParamsTypeLayout()->getSize());
269        memset(m_data.begin(), 0, m_data.getCount());
270        m_elementTypeLayout = programLayout->getGlobalParamsTypeLayout();
271    }
272
273
274    virtual SLANG_NO_THROW slang::TypeLayoutReflection* SLANG_MCALL getElementTypeLayout() override
275    {
276        return m_elementTypeLayout;
277    }
278
279    virtual SLANG_NO_THROW ShaderObjectContainerType SLANG_MCALL getContainerType() override
280    {
281        return ShaderObjectContainerType::None;
282    }
283
284    virtual SLANG_NO_THROW GfxCount SLANG_MCALL getEntryPointCount() override
285    {
286        return (GfxCount)m_entryPoints.getCount();
287    }
288
289    virtual SLANG_NO_THROW Result SLANG_MCALL
290    getEntryPoint(GfxIndex index, IShaderObject** entryPoint) override
291    {
292        returnComPtr(entryPoint, m_entryPoints[index]);
293        return SLANG_OK;
294    }
295
296    virtual SLANG_NO_THROW Result SLANG_MCALL
297    setData(ShaderOffset const& offset, void const* data, Size size) override
298    {
299        auto newSize = Slang::Index(size + offset.uniformOffset);
300        if (newSize > m_data.getCount())
301            m_data.setCount((Slang::Index)newSize);
302        memcpy(m_data.begin() + offset.uniformOffset, data, size);
303        return SLANG_OK;
304    }
305
306    virtual SLANG_NO_THROW Result SLANG_MCALL
307    getObject(ShaderOffset const& offset, IShaderObject** object) override
308    {
309        *object = nullptr;
310
311        Slang::RefPtr<ShaderObjectBase> subObject;
312        if (m_objects.tryGetValue(offset, subObject))
313        {
314            returnComPtr(object, subObject);
315        }
316        return SLANG_OK;
317    }
318
319    virtual SLANG_NO_THROW Result SLANG_MCALL
320    setObject(ShaderOffset const& offset, IShaderObject* object) override
321    {
322        m_objects[offset] = static_cast<ShaderObjectBase*>(object);
323        return SLANG_OK;
324    }
325
326    virtual SLANG_NO_THROW Result SLANG_MCALL
327    setResource(ShaderOffset const& offset, IResourceView* resourceView) override
328    {
329        m_resources[offset] = static_cast<ResourceViewBase*>(resourceView);
330        return SLANG_OK;
331    }
332
333    virtual SLANG_NO_THROW Result SLANG_MCALL
334    setSampler(ShaderOffset const& offset, ISamplerState* sampler) override
335    {
336        m_samplers[offset] = static_cast<SamplerStateBase*>(sampler);
337        return SLANG_OK;
338    }
339    virtual SLANG_NO_THROW Result SLANG_MCALL setCombinedTextureSampler(
340        ShaderOffset const& offset,
341        IResourceView* textureView,
342        ISamplerState* sampler) override
343    {
344        m_resources[offset] = static_cast<ResourceViewBase*>(textureView);
345        m_samplers[offset] = static_cast<SamplerStateBase*>(sampler);
346        return SLANG_OK;
347    }
348
349    virtual SLANG_NO_THROW Result SLANG_MCALL setSpecializationArgs(
350        ShaderOffset const& offset,
351        const slang::SpecializationArg* args,
352        GfxCount count) override
353    {
354        Slang::List<slang::SpecializationArg> specArgs;
355        specArgs.addRange(args, count);
356        m_specializationArgs[offset] = specArgs;
357        return SLANG_OK;
358    }
359
360    virtual SLANG_NO_THROW Result SLANG_MCALL
361    getCurrentVersion(ITransientResourceHeap* transientHeap, IShaderObject** outObject) override
362    {
363        return SLANG_FAIL;
364    }
365
366    virtual SLANG_NO_THROW Result SLANG_MCALL
367    copyFrom(IShaderObject* other, ITransientResourceHeap* transientHeap) override
368    {
369        auto otherObject = static_cast<MutableRootShaderObject*>(other);
370        *this = *otherObject;
371        return SLANG_OK;
372    }
373
374    virtual SLANG_NO_THROW const void* SLANG_MCALL getRawData() override { return m_data.begin(); }
375
376    virtual SLANG_NO_THROW Size SLANG_MCALL getSize() override { return (Size)m_data.getCount(); }
377
378    virtual SLANG_NO_THROW Result SLANG_MCALL
379    setConstantBufferOverride(IBufferResource* constantBuffer) override
380    {
381        m_constantBufferOverride = static_cast<BufferResource*>(constantBuffer);
382        return SLANG_OK;
383    }
384
385    virtual Result collectSpecializationArgs(ExtendedShaderObjectTypeList& args) override
386    {
387        SLANG_UNUSED(args);
388        return SLANG_OK;
389    }
390};
391
392} // namespace gfx