fg_utils.h

cross platform rendering playground

src/passes/fg_utils.h

4.19 KB
#pragma once

#include "passes/fg.h"

struct PassInspectorEntry {
    const char *name = nullptr;
    f32 gpu_ms = 0.0f;
    bool timing_valid = false;
};

class PassTimings {
  public:
    PassTimings() = default;
    ~PassTimings() = default;

    void enable(FrameGraph &fg, rhi::Device &device, u32 max_frames_in_flight) {
        destroy(device);

        fg_ = &fg;
        timestamp_period_ns_ = static_cast<f64>(device.timestamp_period);

        u32 steps = static_cast<u32>(fg.passes.size());
        if (steps == 0) {
            return;
        }

        rhi::QueryPoolDesc desc;
        desc.queryType = rhi::QueryType::TIMESTAMP;
        desc.capacity = steps * 2 * max_frames_in_flight;
        desc.name = "queries/pass-timings";
        rhi::create_query_pool(device, desc, pool_);

        max_frames_ = max_frames_in_flight;
        query_count_per_frame_ = steps * 2;
        entries_.resize(steps);
        slot_written_.assign(max_frames_in_flight, false);
        for (u32 i = 0; i < steps; ++i) {
            entries_[i].name = fg.passes[i].name;
        }

        fg.set_timings(this);
    }

    void destroy(rhi::Device &device) {
        if (rhi::valid(pool_)) {
            rhi::destroy_query_pool(device, pool_);
            pool_ = {};
        }
        if (fg_) {
            fg_->set_timings(nullptr);
        }
        fg_ = nullptr;
        entries_.clear();
        slot_written_.clear();
        max_frames_ = 0;
        query_count_per_frame_ = 0;
    }

    void reset_frame(rhi::CmdBuffer &cmd, u32 frame_slot) {
        if (!rhi::valid(pool_) || frame_slot >= max_frames_) {
            return;
        }
        rhi::reset_query_pool(cmd, pool_, frame_slot * query_count_per_frame_, query_count_per_frame_);
    }

    void begin_pass(rhi::CmdBuffer &cmd, u32 pass_index, u32 frame_slot) {
        if (!rhi::valid(pool_) || frame_slot >= max_frames_) {
            return;
        }
        rhi::write_timestamp(
            cmd, pool_, frame_slot * query_count_per_frame_ + pass_index * 2, rhi::PipelineStages::TOP_OF_PIPE
        );
    }

    void end_pass(rhi::CmdBuffer &cmd, u32 pass_index, u32 frame_slot) {
        if (!rhi::valid(pool_) || frame_slot >= max_frames_) {
            return;
        }
        rhi::write_timestamp(
            cmd, pool_, frame_slot * query_count_per_frame_ + pass_index * 2 + 1, rhi::PipelineStages::BOTTOM_OF_PIPE
        );
    }

    void readback(rhi::Device &device, u32 frame_slot) {
        if (!rhi::valid(pool_) || frame_slot >= slot_written_.size()) {
            return;
        }
        u32 query_count = query_count_per_frame_;
        if (!slot_written_[frame_slot]) {
            slot_written_[frame_slot] = true;
            return;
        }

        std::vector<u64> results(query_count * 2);
        bool ok = rhi::query_pool_results(
            device,
            pool_,
            frame_slot * query_count,
            query_count,
            results.data(),
            results.size() * sizeof(u64),
            sizeof(u64) * 2,
            true
        );
        if (!ok) {
            return;
        }

        u32 steps = static_cast<u32>(entries_.size());
        for (u32 pass_idx = 0; pass_idx < steps; ++pass_idx) {
            u64 begin_timestamp = results[(pass_idx * 2 + 0) * 2 + 0];
            u64 begin_available = results[(pass_idx * 2 + 0) * 2 + 1];
            u64 end_timestamp = results[(pass_idx * 2 + 1) * 2 + 0];
            u64 end_available = results[(pass_idx * 2 + 1) * 2 + 1];
            if (begin_available != 0 && end_available != 0 && end_timestamp >= begin_timestamp) {
                entries_[pass_idx].gpu_ms = static_cast<f32>(
                    static_cast<f64>(end_timestamp - begin_timestamp) * timestamp_period_ns_ / 1000000.0
                );
                entries_[pass_idx].timing_valid = true;
            }
        }
    }

    const PassInspectorEntry *entries(u32 &count) const {
        count = static_cast<u32>(entries_.size());
        return entries_.data();
    }

  private:
    rhi::QueryPool pool_ = {};
    FrameGraph *fg_ = nullptr;
    f64 timestamp_period_ns_ = 0.0;
    std::vector<PassInspectorEntry> entries_;
    std::vector<bool> slot_written_;
    u32 max_frames_ = 0;
    u32 query_count_per_frame_ = 0;
};