pipeline.cpp

cross platform rendering playground

src/backend/metal/pipeline.cpp

12.75 KB
#include <math.h>
#ifndef INFINITY
#define INFINITY __builtin_huge_valf()
#endif
#include <Foundation/Foundation.hpp>
#include <Metal/Metal.hpp>

#include "backend/metal/mt_api.h"
#include "core/logger.h"

namespace {

// has texture/sampler?
static bool struct_has_opaque_members(MTL::StructType *st) {
    if (!st || !st->members()) {
        return false;
    }
    auto *members = st->members();
    for (NS::UInteger i = 0; i < members->count(); ++i) {
        MTL::StructMember *m = static_cast<MTL::StructMember *>(members->object(i));
        MTL::DataType dt = m->dataType();
        if (dt == MTL::DataTypeTexture || dt == MTL::DataTypeSampler) {
            return true;
        }
        if (dt == MTL::DataTypeStruct) {
            if (struct_has_opaque_members(m->structType())) {
                return true;
            }
        }
        if (dt == MTL::DataTypeArray) {
            MTL::ArrayType *at = m->arrayType();
            if (!at) {
                continue;
            }
            MTL::DataType et = at->elementType();
            if (et == MTL::DataTypeTexture || et == MTL::DataTypeSampler) {
                return true;
            }
            if (et == MTL::DataTypeStruct) {
                if (struct_has_opaque_members(at->elementStructType())) {
                    return true;
                }
            }
        }
    }
    return false;
}

static bool arg_struct_has_opaque_members(MTL::Argument *arg) {
    return struct_has_opaque_members(arg->bufferStructType());
}

template <typename TArray> static void record_stage_args(rhi::RenderPipeline &out, TArray *args, const char *stage) {
    for (NS::UInteger i = 0; i < args->count(); ++i) {
        MTL::Argument *a = static_cast<MTL::Argument *>(args->object(i));
        if (a->type() != MTL::ArgumentTypeBuffer) {
            continue;
        }
        MTL::StructType *st = a->bufferStructType();
        if (!st) {
            continue;
        }
        if (arg_struct_has_opaque_members(a)) {
            continue;
        }
        rhi::RenderPipeline::ConstBlock cb{};
        cb.byte_size = (u32)a->bufferDataSize();
        cb.table_index = (u32)a->index();
        bool known = false;
        for (auto &e : out.const_blocks) {
            if (e.table_index == cb.table_index) {
                known = true;
                break;
            }
        }
        if (!known) {
            out.const_blocks.push_back(cb);
            VEL_INFO("metal const: {} bytes at table index {} ({} stage)", cb.byte_size, cb.table_index, stage);
        }
    }
}

} // namespace

namespace rhi {

static NS::String *str(const std::string &s) {
    return NS::String::string(s.c_str(), NS::UTF8StringEncoding);
}

static MTL::VertexFormat to_vertex_format(rhi::ImageFormat format) {
    switch (format) {
    case rhi::ImageFormat::R32_FLOAT:
        return MTL::VertexFormatFloat;
    case rhi::ImageFormat::RG32_FLOAT:
        return MTL::VertexFormatFloat2;
    case rhi::ImageFormat::RGB32_FLOAT:
        return MTL::VertexFormatFloat3;
    case rhi::ImageFormat::RGBA32_FLOAT:
        return MTL::VertexFormatFloat4;
    default:
        return MTL::VertexFormatFloat2;
    }
}

static MTL::PixelFormat to_pixel_format(rhi::ImageFormat format) {
    switch (format) {
    case rhi::ImageFormat::RGBA8_UNORM:
        return MTL::PixelFormatRGBA8Unorm;
    case rhi::ImageFormat::RGBA8_SRGB:
        return MTL::PixelFormatRGBA8Unorm_sRGB;
    case rhi::ImageFormat::BGRA8_UNORM:
        return MTL::PixelFormatBGRA8Unorm;
    default:
        return MTL::PixelFormatBGRA8Unorm;
    }
}

static MTL::CullMode to_cull_mode(rhi::PipelineCullMode mode) {
    switch (mode) {
    case rhi::PipelineCullMode::NONE:
        return MTL::CullModeNone;
    case rhi::PipelineCullMode::FRONT:
        return MTL::CullModeFront;
    case rhi::PipelineCullMode::BACK:
    default:
        return MTL::CullModeBack;
    }
}

static MTL::Winding to_winding(rhi::PipelineTriangleWindingOrder winding) {
    switch (winding) {
    case rhi::PipelineTriangleWindingOrder::CW:
        return MTL::WindingClockwise;
    case rhi::PipelineTriangleWindingOrder::CCW:
    default:
        return MTL::WindingCounterClockwise;
    }
}

static MTL::PixelFormat to_depth_format(rhi::ImageFormat format) {
    switch (format) {
    case rhi::ImageFormat::D32_FLOAT:
        return MTL::PixelFormatDepth32Float;
    default:
        return MTL::PixelFormatInvalid;
    }
}

static MTL::CompareFunction to_compare(rhi::PipelineCompareOp op) {
    switch (op) {
    case rhi::PipelineCompareOp::NEVER:
        return MTL::CompareFunctionNever;
    case rhi::PipelineCompareOp::EQUAL:
        return MTL::CompareFunctionEqual;
    case rhi::PipelineCompareOp::LESS_EQUAL:
        return MTL::CompareFunctionLessEqual;
    case rhi::PipelineCompareOp::GREATER:
        return MTL::CompareFunctionGreater;
    case rhi::PipelineCompareOp::NOT_EQUAL:
        return MTL::CompareFunctionNotEqual;
    case rhi::PipelineCompareOp::GREATER_EQUAL:
        return MTL::CompareFunctionGreaterEqual;
    case rhi::PipelineCompareOp::ALWAYS:
        return MTL::CompareFunctionAlways;
    case rhi::PipelineCompareOp::LESS:
    default:
        return MTL::CompareFunctionLess;
    }
}

static rhi::ImageFormat first_color_format(const RenderPipelineDesc &desc) {
    if (!desc.color_attachment_formats.empty()) {
        return desc.color_attachment_formats[0];
    }
    return rhi::ImageFormat::BGRA8_UNORM;
}

RenderPipelineDesc &
RenderPipelineDesc::set_shaders(ShaderCompiler &sc, const char *path, const char *vs, const char *fs) {
    shader_path = path ? path : "";
    vs_entry = vs ? vs : "";
    fs_entry = fs ? fs : "";
    const char *entries[] = {vs, fs};
    msl_source = mt::compile_msl_linked(sc, path, entries, 2);
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_input_topology(rhi::PipelineTopology topo) {
    topology = topo;
    return *this;
}

RenderPipelineDesc &
RenderPipelineDesc::add_vertex_attribute(u32 location, u32 binding, rhi::ImageFormat format, u32 offset) {
    vertex_attributes.push_back({location, binding, format, offset});
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::add_vertex_binding(u32 binding, u32 stride, bool is_instanced) {
    if (binding >= mt::kMaxVertexBindings) {
        VEL_ERROR("add_vertex_binding: binding {} exceeds max {}", binding, mt::kMaxVertexBindings);
        return *this;
    }
    vertex_bindings.push_back({binding, stride, is_instanced});
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_polygon_mode(rhi::PipelineFillMode mode) {
    fill_mode = mode;
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_color_format(rhi::ImageFormat format) {
    color_attachment_formats.clear();
    color_attachment_formats.push_back(format);
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_depth_format(rhi::ImageFormat format) {
    depth_attachment_format = format;
    return *this;
}

RenderPipelineDesc &
RenderPipelineDesc::set_cull_mode(rhi::PipelineCullMode mode, rhi::PipelineTriangleWindingOrder wind) {
    cull_mode = mode;
    winding = wind;
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_multisampling(rhi::SampleCount count) {
    samples = count;
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_depth_testing(bool do_test, bool do_write, rhi::PipelineCompareOp op) {
    depth_test_enable = do_test;
    depth_write_enable = do_write;
    depth_compare_op = op;
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_depth_bias(bool enabled, f32 constant, f32 clamp, f32 slope) {
    depth_bias_enable = enabled;
    depth_bias_constant = constant;
    depth_bias_clamp = clamp;
    depth_bias_slope = slope;
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_depth_clamp(bool enabled) {
    depth_clamp_enable = enabled;
    return *this;
}

RenderPipelineDesc &RenderPipelineDesc::set_blending(
    rhi::PipelineBlendFactor src_color,
    rhi::PipelineBlendFactor dst_color,
    rhi::PipelineBlendOp color_op,
    rhi::PipelineBlendFactor src_alpha,
    rhi::PipelineBlendFactor dst_alpha,
    rhi::PipelineBlendOp alpha_op
) {
    blend_enable = true;
    src_color_blend_factor = src_color;
    dst_color_blend_factor = dst_color;
    color_blend_op = color_op;
    src_alpha_blend_factor = src_alpha;
    dst_alpha_blend_factor = dst_alpha;
    alpha_blend_op = alpha_op;
    return *this;
}

ComputePipelineDesc &ComputePipelineDesc::set_shader(ShaderCompiler &sc, const char *path, const char *cs) {
    (void)sc;
    shader_path = path ? path : "";
    cs_entry = cs ? cs : "";
    VEL_ERROR("ComputePipeline is not implemented on Metal");
    return *this;
}

ComputePipeline ComputePipelineDesc::build() {
    VEL_ERROR("ComputePipeline::build is not implemented on Metal");
    ComputePipeline out{};
    out.device = device;
    return out;
}

RenderPipeline RenderPipelineDesc::build() {
    RenderPipeline out{};
    out.device = device;
    MTL::Device *dev = device ? device->handle : nullptr;
    if (!dev) {
        VEL_ERROR("RenderPipeline::build: null device");
        return out;
    }
    NS::Error *error = nullptr;

    MTL::Library *lib = dev->newLibrary(str(msl_source), nullptr, &error);
    if (!lib) {
        VEL_ERROR(
            "Failed to compile Metal library: {}", error ? error->localizedDescription()->utf8String() : "unknown error"
        );
        return out;
    }

    MTL::Function *vs = lib->newFunction(str(vs_entry));
    MTL::Function *fs = lib->newFunction(str(fs_entry));
    if (!vs || !fs) {
        VEL_ERROR("Failed to find entry points '{}'/'{}' in Metal library", vs_entry, fs_entry);
        lib->release();
        return out;
    }

    MTL::RenderPipelineDescriptor *rpd = MTL::RenderPipelineDescriptor::alloc()->init();
    rpd->setVertexFunction(vs);
    rpd->setFragmentFunction(fs);
    rpd->colorAttachments()->object(0)->setPixelFormat(to_pixel_format(first_color_format(*this)));
    // pipeline depth must match pass.
    if (depth_attachment_format != rhi::ImageFormat::UNDEFINED) {
        rpd->setDepthAttachmentPixelFormat(to_depth_format(depth_attachment_format));
    }

    if (!vertex_attributes.empty()) {
        // fetch at 2+binding. matches cmd.
        MTL::VertexDescriptor *vd = MTL::VertexDescriptor::vertexDescriptor();
        for (const auto &attr : vertex_attributes) {
            MTL::VertexAttributeDescriptor *a = vd->attributes()->object(attr.location);
            a->setFormat(to_vertex_format(attr.format));
            a->setOffset(attr.offset);
            a->setBufferIndex(mt::kVertexBaseIndex + attr.binding);
        }
        for (const auto &binding : vertex_bindings) {
            MTL::VertexBufferLayoutDescriptor *l = vd->layouts()->object(mt::kVertexBaseIndex + binding.binding);
            l->setStride(binding.stride);
            l->setStepFunction(
                binding.is_instanced ? MTL::VertexStepFunctionPerInstance : MTL::VertexStepFunctionPerVertex
            );
        }
        rpd->setVertexDescriptor(vd);
    }

    (void)to_cull_mode(cull_mode);
    (void)to_winding(winding);

    MTL::RenderPipelineState *ps = nullptr;
    if (enable_abi_blocks) {
        MTL::RenderPipelineReflection *refl = nullptr;
        ps = dev->newRenderPipelineState(rpd, MTL::PipelineOptionArgumentInfo, &refl, &error);
        if (!ps || !refl) {
            VEL_ERROR(
                "Failed to create render pipeline state: {}",
                error ? error->localizedDescription()->utf8String() : "unknown error"
            );
            vs->release();
            fs->release();
            rpd->release();
            lib->release();
            return out;
        }
        record_stage_args(out, refl->vertexArguments(), "vertex");
        record_stage_args(out, refl->fragmentArguments(), "fragment");
        refl->release();
    } else {
        ps = dev->newRenderPipelineState(rpd, &error);
        if (!ps) {
            VEL_ERROR(
                "Failed to create render pipeline state: {}",
                error ? error->localizedDescription()->utf8String() : "unknown error"
            );
            vs->release();
            fs->release();
            rpd->release();
            lib->release();
            return out;
        }
    }

    vs->release();
    fs->release();
    rpd->release();
    lib->release();

    out.handle = ps;
    if (depth_attachment_format != rhi::ImageFormat::UNDEFINED && (depth_test_enable || depth_write_enable)) {
        MTL::DepthStencilDescriptor *dd = MTL::DepthStencilDescriptor::alloc()->init();
        dd->setDepthCompareFunction(to_compare(depth_compare_op));
        dd->setDepthWriteEnabled(depth_write_enable);
        out.depth_stencil = dev->newDepthStencilState(dd);
        dd->release();
        if (!out.depth_stencil) {
            VEL_ERROR("RenderPipeline::build: newDepthStencilState failed (depth runs untested)");
        }
    }
    return out;
}

} // namespace rhi