hiz.cpp

cross platform rendering playground

src/passes/hiz/hiz.cpp

4.28 KB
#include "hiz.h"

#include "passes/fg.h"
#include "pipelines.h"
#include "shaders/renderer_types.h"

namespace hiz {

Resources create(FrameGraph &fg) {

    Resources out;
    out.hiz = fg.add_image({
        .desc = {.mip_levels = 0, .format = rhi::ImageFormat::R32_FLOAT},
        .name = "fg/hiz/hiz",
        .lifetime = Lifetime::Transient,
        .size_class = SizeOp::Swapchain,
        .scale = {0.5f, 0.5f},
        // Full-chain sampled view (mip_count 0 = full chain, substituted at
        // compile). Per-mip storage views stay on the lazy push path.
        .view = {.type = rhi::ImageViewType::Sampled, .mip_count = 0},
    });
    return out;
};

void record(
    rhi::Device &device,
    Pipelines &pipelines,
    rhi::ShaderCompiler &sc,
    FrameGraph &fg,
    Resources &res,
    ImageHandle depth
) {

    rhi::ComputePipelineDesc pd;
    pd.device = &device;
    pd.set_shader(sc, "src/passes/hiz/hiz_downsample.slang", "csMain");
    rhi::ComputePipeline &pipeline = pipelines.add_compute("HiZ", pd);

    auto &builder = fg.add_pass("HiZ");
    builder.read(depth, rhi::ResourceState::TextureSampleNonFragment);
    builder.write(res.hiz, rhi::ResourceState::StorageReadWrite);
    builder.read(res.hiz, rhi::ResourceState::TextureSampleNonFragment);

    // Per-mip storage views resolved once per compile(); execute only reads.
    res.storage_views = fg.request_extra_view(res.hiz, true);

    builder.execute([=, &pipeline](PassContext &ctx) {
        auto &img = ctx.get_image(res.hiz);
        u32 mip_count = img.desc.mip_levels;
        u32 w = img.desc.width;
        u32 h = img.desc.height;

        auto &depth_img = ctx.get_image(depth);
        u32 depth_w = depth_img.desc.width;
        u32 depth_h = depth_img.desc.height;

        for (u32 mip = 0; mip < mip_count; ++mip) {
            u32 src_w = (mip == 0) ? depth_w : std::max(w >> (mip - 1), 1u);
            u32 src_h = (mip == 0) ? depth_h : std::max(h >> (mip - 1), 1u);
            u32 dst_w = std::max(w >> mip, 1u);
            u32 dst_h = std::max(h >> mip, 1u);

            rhi::set_pipeline(ctx.cmd, pipeline);

            HiZPushConstants pc{};
            pc.depth = (u32)ctx.view(depth);
            pc.hiz_src = (u32)ctx.view(res.hiz);
            pc.hiz_dst = (u32)ctx.view_extra(res.storage_views, 0, mip);
            pc.mip_level = mip;
            pc.src_width = src_w;
            pc.src_height = src_h;
            {
                static_assert(sizeof(pc) <= 128, "push constants exceed global 128B range");
                char pc128[128] = {};
                std::memcpy(pc128, &pc, sizeof(pc));
                rhi::set_constants(ctx.cmd, pipeline, rhi::ShaderStage::ALL, sizeof(pc128), pc128);
            }

            u32 gx = (dst_w + 15u) / 16u;
            u32 gy = (dst_h + 15u) / 16u;
            if (gx > 0 && gy > 0) {
                rhi::dispatch(ctx.cmd, gx, gy);
            }

            if (mip < mip_count - 1) {
                rhi::ImageRange range{};
                range.mip = mip;
                range.mip_count = 1;
                rhi::barrier(
                    ctx.cmd,
                    ctx.get_image(res.hiz),
                    range,
                    rhi::ResourceState::StorageReadWrite,
                    rhi::ResourceState::TextureSampleNonFragment
                );
            }
        }

        // Restore: mips 0..n-2 were left in SHADER_READ_ONLY_OPTIMAL for
        // sampling, but the FG tracks this image as one whole-image
        // StorageReadWrite (GENERAL) state. Without per-mip state tracking,
        // the next frame's dispatches (STORAGE_IMAGE descriptors expecting
        // GENERAL) would hit the SRO leftovers. Transition every non-last
        // mip back to GENERAL so pass boundaries always leave a uniform
        // GENERAL image matching the tracked state. Per-mip ranges never
        // touch image.state, so no tracking divergence.
        for (u32 mip = 0; mip + 1 < mip_count; ++mip) {
            rhi::ImageRange range{};
            range.mip = mip;
            range.mip_count = 1;
            rhi::barrier(
                ctx.cmd,
                ctx.get_image(res.hiz),
                range,
                rhi::ResourceState::TextureSampleNonFragment,
                rhi::ResourceState::StorageReadWrite
            );
        }
    });
}

} // namespace hiz