hf2q 0.1.23

Pure Rust CLI for converting HuggingFace models to hardware-optimized formats and serving them over an OpenAI-compatible API on Apple Silicon
// glp_project_mhc.metal — GLP projection for multi-hyper-connection (mHC)
// state: `[rows, hc, hidden]` F32. One projection per (row, stream) slice.
//
// Launch contract (S13 fix): one FULL 256-lane threadgroup per (row,
// stream) pair — the host dispatches grid = rows*hc*256 threads with
// threads_per_threadgroup = 256. Threadgroup id maps to (row, stream);
// columns are covered strided. Per-stream norms arrive in buffer(3): the
// shared-direction form repeats one norm; the per-stream form
// (per_stream_directions) carries each stream's own norm, so unequal
// per-stack norms scale correctly.

#include <metal_stdlib>
using namespace metal;

struct GlpProjectMhcParams {
    uint rows;
    uint hc;
    uint hidden;
    uint per_stream_directions;
    float alpha;
};

kernel void glp_project_mhc_f32(
    constant GlpProjectMhcParams& params [[buffer(0)]],
    device const float*           direction [[buffer(1)]],
    device float*                 state [[buffer(2)]],
    device const float*           norms    [[buffer(3)]],
    uint                          tid_in_tg [[thread_position_in_threadgroup]],
    uint                          tg_size [[threads_per_threadgroup]],
    uint                          tg_id [[threadgroup_position_in_grid]]
) {
    const uint row = tg_id / params.hc;
    const uint stream = tg_id - row * params.hc;
    if (row >= params.rows) {
        return;
    }
    device float* slice = state + (row * params.hc + stream) * params.hidden;
    device const float* d = direction + (params.per_stream_directions ? stream * params.hidden : 0);
    const float norm_sq = norms[params.per_stream_directions ? stream : 0];

    // accumulate strided elements
    float local_dot = 0.0f;
    for (uint col = tid_in_tg; col < params.hidden; col += tg_size) {
        local_dot += slice[col] * d[col];
    }

    // threadgroup tree reduction on the 256-thread window
    threadgroup float partial[256];
    partial[tid_in_tg] = local_dot;
    threadgroup_barrier(mem_flags::mem_threadgroup);
    for (uint stride = 128; stride > 0; stride >>= 1) {
        if (tid_in_tg < stride) {
            partial[tid_in_tg] += partial[tid_in_tg + stride];
        }
        threadgroup_barrier(mem_flags::mem_threadgroup);
    }
    float dot = partial[0];
    const float scale = params.alpha * dot / norm_sq;
    for (uint col = tid_in_tg; col < params.hidden; col += tg_size) {
        slice[col] -= scale * d[col];
    }
}