Skip to content

Instantly share code, notes, and snippets.

@leok7v
Last active July 4, 2026 19:12
Show Gist options
  • Select an option

  • Save leok7v/b200acf63e3f42c3d9f87f2bf3a84e0c to your computer and use it in GitHub Desktop.

Select an option

Save leok7v/b200acf63e3f42c3d9f87f2bf3a84e0c to your computer and use it in GitHub Desktop.
bitnet.metal
#ifndef BN_BACKEND_H
#define BN_BACKEND_H
#ifdef __cplusplus
extern "C" {
#endif
#include <stddef.h>
#include <stdint.h>
#include <stdbool.h>
typedef struct bn_backend bn_backend;
typedef struct bn_buf bn_buf; // device buffer (Metal) or heap buffer (SIMD)
typedef struct bn_session bn_session; // one command buffer (Metal) or sync batch (SIMD)
typedef struct bn_pipes bn_pipes; // pipeline cache (Metal) or kernel table (SIMD)
enum bn_storage {
BN_BUF_SHARED = 0, // host + device visible. SIMD: heap.
BN_BUF_PRIVATE = 1, // device-only. SIMD: same as SHARED (no host/device split).
};
// Kernel ids. Add new ops at the end; SIMD's per-pipe adapter table is
// indexed by this enum so order must stay stable across both backends.
// When hybrid SSM lands this gets new entries (conv1d_state, ssm_scan,
// ssm_state_update). The backend layer is pure addition for new pipes.
enum bn_pipe_id {
BPI_ACTIVATION,
BPI_ATTN_SCORES,
BPI_ATTN_V,
BPI_ELEMENTWISE,
BPI_EMBEDDING,
BPI_F32_MATMUL,
BPI_QUANTIZE,
BPI_RMSNORM,
BPI_ROPE,
BPI_SOFTMAX,
BPI_SPLIT,
BPI_TERNARY_GEMM,
BPI_TERNARY_GEMV,
BPI_TQ2_GEMV, // GGML TQ2_0 (Gemma 3 / Qwen 3.5): 256-elem blocks,
// 64 packed bytes + fp16 d, no per-tensor scale.
BPI_Q4_GEMV, // GGML Q4_0 (Gemma 3 attn_q/k/v + ffn_up/gate):
// 32-elem blocks, fp16 d + 16 packed nibble bytes.
// Input-row block quantizers. Produce the (x_qs, x_d) pair the block-
// quant gemvs read so the engine reproduces llama.cpp's
// `ggml_vec_dot_q4_0_q8_0` / `ggml_vec_dot_tq2_0_q8_K` arithmetic.
BPI_QUANT_Q8_0, // per-32-block i8 + fp16 d (pairs with Q4_0 gemv)
BPI_QUANT_Q8_K, // per-256-block i8 + fp16 d (pairs with TQ2_0 gemv)
BPI__COUNT
};
// Max storage bindings per dispatch. 12 covers our current 5 plus headroom
// for SSM scan (8) and any future hybrid op. Wider would just waste pointer
// slots on the stack per job; no runtime cost beyond that.
#define BN_DISPATCH_MAX_BINDINGS 12
// One kernel invocation. Storage bindings sit in [[buffer(N)]] order; when
// params_bytes is set it slots in at params_binding via setBytes and any
// binding[i] with i >= params_binding gets pushed to slot (i+1).
// params_binding == UINT32_MAX means "no inline params at all".
struct bn_dispatch_job {
enum bn_pipe_id pipe;
const char * label;
bn_buf * bindings[BN_DISPATCH_MAX_BINDINGS];
uint32_t n_bindings;
uint32_t params_binding;
const void * params_bytes;
size_t params_size;
uint32_t gx, gy, gz; // grid in threadgroup counts (Metal) /
// total work units (SIMD interprets as
// total threadgroups × per-pipe tg size)
// Optional GPU timestamp sampling. Metal backend honours when non-NULL;
// SIMD ignores. The counter_buf is the backend-specific opaque returned
// by bn_backend_metal_counter_buf_new(). UINT32_MAX in idx fields = skip.
void * counter_buf;
uint32_t timing_start_idx;
uint32_t timing_end_idx;
};
// Backend vtable. Each backend exposes ONE instance via its factory.
struct bn_backend_ops {
const char * name; // "metal" / "simd" — for trace labels
// ===== lifecycle =====
void (*destroy)(bn_backend *);
size_t (*max_buffer_length)(bn_backend *);
// ===== pipeline cache =====
// resource_path is backend-defined (metallib path for Metal; ignored on SIMD).
bn_pipes * (*pipes_create)(bn_backend *, const char * resource_path);
void (*pipes_destroy)(bn_pipes *);
// ===== buffers =====
bn_buf * (*buf_new)(bn_backend *, size_t bytes, enum bn_storage);
bn_buf * (*buf_new_with_data)(bn_backend *, const void * src,
size_t bytes, enum bn_storage);
void (*buf_release)(bn_buf *);
size_t (*buf_length)(bn_buf *);
void * (*buf_contents)(bn_buf *); // NULL on Metal PRIVATE
void (*buf_set_label)(bn_buf *, const char *); // optional / debug-only
void (*buf_set_user_data)(bn_buf *, uintptr_t);
uintptr_t (*buf_get_user_data)(bn_buf *);
// ===== command buffer =====
bn_session * (*session_begin)(bn_backend *, bn_pipes *);
void (*session_dispatch)(bn_session *, const struct bn_dispatch_job *);
void (*session_blit)(bn_session *,
bn_buf * src, size_t src_off,
bn_buf * dst, size_t dst_off, size_t n);
void (*session_commit_wait)(bn_session *); // also releases the session
// ===== GPU timestamp profiling (Metal-specific; stubbed on SIMD) =====
// counter_buf is opaque; treat as void*. Metal backend resolves it via
// bn_backend_metal_counter_buf_resolve. SIMD impl returns NULL for new.
bool (*supports_timestamps)(bn_backend *);
void * (*counter_buf_new)(bn_backend *, uint32_t n_samples);
void (*counter_buf_release)(bn_backend *, void * cb);
void (*counter_buf_resolve)(bn_backend *, void * cb,
uint32_t first, uint32_t count,
uint64_t * out_ticks);
void (*sample_cpu_gpu_ticks)(bn_backend *,
uint64_t * out_cpu_ts, uint64_t * out_gpu_ts);
};
// Generic backend header — embedded as the first field of every per-backend
// struct, so callers can read be->ops without knowing the concrete type.
struct bn_backend {
const struct bn_backend_ops * ops;
};
// Factories. Each returns a heap-owned backend; caller releases via
// be->ops->destroy(be).
bn_backend * bn_backend_metal_create(void);
bn_backend * bn_backend_simd_create(void); // Phase 2
#ifdef __cplusplus
}
#endif
#endif // BN_BACKEND_H
// Metal backend — vtable instance that forwards every op to mtl.h's pure-C
// surface. The C++ that talks to metal-cpp lives in mtl.cc; this file is
// plain C and links against the prebuilt mtl.o
#include "backend.h"
#ifndef __APPLE__
bn_backend * bn_backend_metal_create(void) { return NULL; }
#else
#include "mtl.h"
#include <stdlib.h>
#include <string.h>
// ============================================================================
// Pipeline specs — kernel entry names + threadgroup sizes baked into the
// .metallib at compile time. Order MUST match enum bn_pipe_id.
// ============================================================================
struct bn_metal_pipe_spec {
const char * entry;
uint32_t tx, ty, tz;
};
static const struct bn_metal_pipe_spec PIPE_SPECS[BPI__COUNT] = {
[BPI_ACTIVATION] = { "activation_main", 256, 1, 1 },
[BPI_ATTN_SCORES] = { "attention_compute_scores", 16, 16, 1 },
[BPI_ATTN_V] = { "attention_attn_v", 256, 1, 1 },
[BPI_ELEMENTWISE] = { "elementwise_main", 256, 1, 1 },
[BPI_EMBEDDING] = { "embedding_main", 256, 1, 1 },
[BPI_F32_MATMUL] = { "f32_matmul_main", 256, 1, 1 },
[BPI_QUANTIZE] = { "quantize_main", 256, 1, 1 },
[BPI_RMSNORM] = { "rmsnorm_main", 256, 1, 1 },
[BPI_ROPE] = { "rope_main", 256, 1, 1 },
[BPI_SOFTMAX] = { "softmax_main", 256, 1, 1 },
[BPI_SPLIT] = { "split_main", 256, 1, 1 },
[BPI_TERNARY_GEMM] = { "ternary_gemm_main", 8, 8, 1 },
[BPI_TERNARY_GEMV] = { "ternary_gemv_main", 32, 1, 1 },
[BPI_TQ2_GEMV] = { "ternary_tq2_gemv_main", 32, 1, 1 },
[BPI_Q4_GEMV] = { "ternary_q4_gemv_main", 32, 1, 1 },
[BPI_QUANT_Q8_0] = { "quantize_q8_0_main", 1, 1, 1 },
[BPI_QUANT_Q8_K] = { "quantize_q8_K_main", 1, 1, 1 },
};
// ============================================================================
// Concrete struct definitions (opaque to callers via forward decl in backend.h)
// ============================================================================
struct bn_backend_metal {
struct bn_backend base; // MUST be first — vtable lookup via casts
mtl_device * device;
mtl_queue * queue;
};
struct bn_buf {
mtl_buffer * mb;
};
struct bn_session {
mtl_cmdbuf * cb;
bn_pipes * pipes;
mtl_queue * queue;
};
struct bn_pipes {
mtl_library * library;
mtl_pipeline * pipes[BPI__COUNT];
};
static inline struct bn_backend_metal * as_metal(bn_backend * be) {
return (struct bn_backend_metal *)be;
}
// ============================================================================
// vtable implementations
// ============================================================================
static void m_destroy(bn_backend * be) {
struct bn_backend_metal * m = as_metal(be);
if (m->queue) { mtl_queue_release(m->queue); m->queue = NULL; }
if (m->device) { mtl_device_release(m->device); m->device = NULL; }
free(m);
}
static size_t m_max_buffer_length(bn_backend * be) {
return mtl_device_max_buffer_length(as_metal(be)->device);
}
static bn_pipes * m_pipes_create(bn_backend * be, const char * resource_path) {
struct bn_backend_metal * m = as_metal(be);
bn_pipes * p = (bn_pipes *)calloc(1, sizeof *p);
p->library = mtl_library_load(m->device, resource_path);
if (p->library) {
int i = 0;
while (i < BPI__COUNT && (i == 0 || p->pipes[i - 1] != NULL)) {
p->pipes[i] = mtl_pipeline_create(m->device, p->library,
PIPE_SPECS[i].entry,
PIPE_SPECS[i].tx,
PIPE_SPECS[i].ty,
PIPE_SPECS[i].tz);
i++;
}
}
if (!p->library || !p->pipes[BPI__COUNT - 1]) {
// Partial init — clean up what's there and return NULL.
for (int i = 0; i < BPI__COUNT; i++) {
if (p->pipes[i]) mtl_pipeline_release(p->pipes[i]);
}
if (p->library) mtl_library_release(p->library);
free(p);
p = NULL;
}
return p;
}
static void m_pipes_destroy(bn_pipes * p) {
if (p) {
for (int i = 0; i < BPI__COUNT; i++) {
if (p->pipes[i]) mtl_pipeline_release(p->pipes[i]);
}
if (p->library) mtl_library_release(p->library);
free(p);
}
}
static bn_buf * m_buf_new(bn_backend * be, size_t bytes, enum bn_storage s) {
mtl_buffer * mb = mtl_buffer_new(as_metal(be)->device, bytes, (mtl_storage)s);
if (!mb) return NULL;
bn_buf * b = (bn_buf *)malloc(sizeof *b);
b->mb = mb;
return b;
}
static bn_buf * m_buf_new_with_data(bn_backend * be, const void * src,
size_t bytes, enum bn_storage s) {
mtl_buffer * mb = mtl_buffer_new_with_data(as_metal(be)->device, src,
bytes, (mtl_storage)s);
if (!mb) return NULL;
bn_buf * b = (bn_buf *)malloc(sizeof *b);
b->mb = mb;
return b;
}
static void m_buf_release(bn_buf * b) {
if (b) {
if (b->mb) mtl_buffer_release(b->mb);
free(b);
}
}
static size_t m_buf_length (bn_buf * b) { return mtl_buffer_length(b->mb); }
static void * m_buf_contents (bn_buf * b) { return mtl_buffer_contents(b->mb); }
static void m_buf_set_label (bn_buf * b, const char * s) { mtl_buffer_set_label(b->mb, s); }
static void m_buf_set_user_data (bn_buf * b, uintptr_t v) { mtl_buffer_set_user_data(b->mb, v); }
static uintptr_t m_buf_get_user_data (bn_buf * b) { return mtl_buffer_get_user_data(b->mb); }
static bn_session * m_session_begin(bn_backend * be, bn_pipes * pipes) {
struct bn_backend_metal * m = as_metal(be);
bn_session * s = (bn_session *)malloc(sizeof *s);
s->cb = mtl_cmdbuf_new(m->queue);
s->pipes = pipes;
s->queue = m->queue;
return s;
}
static void m_session_dispatch(bn_session * s, const struct bn_dispatch_job * job) {
bool has_params = (job->params_bytes != NULL && job->params_size > 0
&& job->params_binding != UINT32_MAX);
bool timed = (job->counter_buf != NULL
&& job->timing_start_idx != UINT32_MAX
&& job->timing_end_idx != UINT32_MAX);
mtl_enc * enc = timed
? mtl_enc_compute_timed(s->cb, (mtl_counter_buf *)job->counter_buf,
job->timing_start_idx, job->timing_end_idx)
: mtl_enc_compute(s->cb);
mtl_enc_set_pipeline(enc, s->pipes->pipes[job->pipe]);
for (uint32_t i = 0; i < job->n_bindings; i++) {
uint32_t slot = (has_params && i >= job->params_binding) ? (i + 1) : i;
mtl_enc_set_buffer(enc, slot, job->bindings[i]->mb, 0);
}
if (has_params) {
mtl_enc_set_bytes(enc, job->params_binding,
job->params_bytes, job->params_size);
}
mtl_enc_dispatch(enc, job->gx, job->gy, job->gz);
mtl_enc_end(enc);
mtl_enc_release(enc);
}
static void m_session_blit(bn_session * s,
bn_buf * src, size_t src_off,
bn_buf * dst, size_t dst_off, size_t n) {
mtl_blit * bl = mtl_enc_blit(s->cb);
mtl_blit_copy(bl, src->mb, src_off, dst->mb, dst_off, n);
mtl_blit_end(bl);
mtl_blit_release(bl);
}
static void m_session_commit_wait(bn_session * s) {
mtl_cmdbuf_commit_wait(s->cb);
mtl_cmdbuf_release(s->cb);
free(s);
}
// ===== GPU timestamps =====
static bool m_supports_timestamps(bn_backend * be) {
return mtl_device_supports_timestamps(as_metal(be)->device);
}
static void * m_counter_buf_new(bn_backend * be, uint32_t n_samples) {
return mtl_counter_buf_new(as_metal(be)->device, n_samples);
}
static void m_counter_buf_release(bn_backend * be, void * cb) {
(void)be;
if (cb) mtl_counter_buf_release((mtl_counter_buf *)cb);
}
static void m_counter_buf_resolve(bn_backend * be, void * cb,
uint32_t first, uint32_t count,
uint64_t * out_ticks) {
(void)be;
mtl_counter_buf_resolve((mtl_counter_buf *)cb, first, count, out_ticks);
}
static void m_sample_cpu_gpu_ticks(bn_backend * be,
uint64_t * out_cpu_ts, uint64_t * out_gpu_ts) {
mtl_device_sample_timestamps(as_metal(be)->device, out_cpu_ts, out_gpu_ts);
}
// ============================================================================
// vtable instance + factory
// ============================================================================
static const struct bn_backend_ops METAL_OPS = {
.name = "metal",
.destroy = m_destroy,
.max_buffer_length = m_max_buffer_length,
.pipes_create = m_pipes_create,
.pipes_destroy = m_pipes_destroy,
.buf_new = m_buf_new,
.buf_new_with_data = m_buf_new_with_data,
.buf_release = m_buf_release,
.buf_length = m_buf_length,
.buf_contents = m_buf_contents,
.buf_set_label = m_buf_set_label,
.buf_set_user_data = m_buf_set_user_data,
.buf_get_user_data = m_buf_get_user_data,
.session_begin = m_session_begin,
.session_dispatch = m_session_dispatch,
.session_blit = m_session_blit,
.session_commit_wait = m_session_commit_wait,
.supports_timestamps = m_supports_timestamps,
.counter_buf_new = m_counter_buf_new,
.counter_buf_release = m_counter_buf_release,
.counter_buf_resolve = m_counter_buf_resolve,
.sample_cpu_gpu_ticks = m_sample_cpu_gpu_ticks,
};
bn_backend * bn_backend_metal_create(void) {
struct bn_backend_metal * m =
(struct bn_backend_metal *)calloc(1, sizeof *m);
m->base.ops = &METAL_OPS;
m->device = mtl_device_create();
if (m->device) {
m->queue = mtl_queue_create(m->device);
}
if (!m->device || !m->queue) {
m_destroy(&m->base);
return NULL;
}
return &m->base;
}
#endif // __APPLE__
// shaders.metal
// @group(0) @binding(N) var<storage,read> -> device const T* x [[buffer(N)]]
// @group(0) @binding(N) var<storage,read_write> -> device T* x [[buffer(N)]]
// @group(0) @binding(N) var<uniform> -> constant T& p [[buffer(N)]]
// @builtin(global_invocation_id) -> uint3 gid [[thread_position_in_grid]]
// @builtin(workgroup_id) -> uint3 wg_id [[threadgroup_position_in_grid]]
// @builtin(local_invocation_id) -> uint3 lid [[thread_position_in_threadgroup]]
// var<workgroup> shared: array<T,N> -> threadgroup T shared[N]
// workgroupBarrier() -> threadgroup_barrier(mem_flags::mem_threadgroup)
// inverseSqrt(x) -> rsqrt(x)
// unpack2x16float(packed) -> float2(as_type<half2>(packed))
// select(a, b, cond) -> select(a, b, cond) // same arg order
#include <metal_stdlib>
using namespace metal;
// ===========================================================================
// activation
// ===========================================================================
struct ActivationParams {
uint N;
uint activation_type; // 0 = ReLU^2, 1 = SiLU, 2 = GELU (tanh approx)
};
kernel void activation_main(
device const float * input [[buffer(0)]],
device float * output [[buffer(1)]],
constant ActivationParams& params [[buffer(2)]],
uint3 gid [[thread_position_in_grid]])
{
uint idx = gid.x;
if (idx >= params.N) { return; }
float x = input[idx];
if (params.activation_type == 0u) {
// ReLU^2 (Falcon-E / bitnet-25)
float r = max(0.0f, x);
output[idx] = r * r;
} else if (params.activation_type == 1u) {
// SiLU / Swish (BitNet / Falcon-E / Qwen 3.5)
output[idx] = x / (1.0f + exp(-x));
} else {
// GELU, tanh-based approximation (Gemma 3). Matches the
// 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))) form
// used by Gemma's reference implementation.
float x3 = x * x * x;
float inner = 0.7978845608028654f * (x + 0.044715f * x3);
output[idx] = 0.5f * x * (1.0f + tanh(inner));
}
}
// ===========================================================================
// attention — two kernels: compute_scores + attn_v
// ===========================================================================
struct ScoreParams {
uint N;
uint S;
uint num_heads;
uint num_kv_heads;
uint head_dim;
uint window_size; // 0 = no sliding window (full causal, BitNet path)
float scale;
};
kernel void attention_compute_scores(
device const float * Q [[buffer(0)]],
device const float * K [[buffer(1)]],
device float * scores [[buffer(2)]],
constant ScoreParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]])
{
uint q_pos = gid.x;
uint k_pos = gid.y;
uint head = gid.z;
if (q_pos >= params.N || k_pos >= params.S || head >= params.num_heads) { return; }
uint kv_head = head / (params.num_heads / params.num_kv_heads);
uint q_offset = (q_pos * params.num_heads + head) * params.head_dim;
uint k_offset = (k_pos * params.num_kv_heads + kv_head) * params.head_dim;
float dot_v = 0.0f;
for (uint d = 0u; d < params.head_dim; d++) {
dot_v += Q[q_offset + d] * K[k_offset + d];
}
// q_abs = absolute sequence index of this query token. S - N is the
// number of prefix tokens already in the KV cache; the n new tokens
// occupy positions [S - N, S).
uint q_abs = q_pos + (params.S - params.N);
bool is_causal = k_pos > q_abs;
bool is_outside = (params.window_size > 0u)
&& (k_pos + params.window_size <= q_abs);
bool masked_b = is_causal || is_outside;
float masked = select(dot_v * params.scale, -3.402823e+38f, masked_b);
uint idx = (head * params.N + q_pos) * params.S + k_pos;
scores[idx] = masked;
}
struct AttnVParams {
uint N;
uint S;
uint num_heads;
uint num_kv_heads;
uint head_dim;
};
kernel void attention_attn_v(
device const float * attn [[buffer(0)]],
device const float * V [[buffer(1)]],
device float * attn_output [[buffer(2)]],
constant AttnVParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]])
{
uint total = params.N * params.num_heads * params.head_dim;
uint idx = gid.x;
if (idx >= total) { return; }
uint d = idx % params.head_dim;
uint remainder = idx / params.head_dim;
uint head = remainder % params.num_heads;
uint q_pos = remainder / params.num_heads;
uint kv_head = head / (params.num_heads / params.num_kv_heads);
float sum = 0.0f;
for (uint s = 0u; s < params.S; s++) {
float a = attn[(head * params.N + q_pos) * params.S + s];
float v = V[(s * params.num_kv_heads + kv_head) * params.head_dim + d];
sum += a * v;
}
uint out_idx = (q_pos * params.num_heads + head) * params.head_dim + d;
attn_output[out_idx] = sum;
}
// ===========================================================================
// elementwise
// ===========================================================================
struct ElemParams {
uint N;
uint op; // 0 = add, 1 = multiply
};
kernel void elementwise_main(
device const float * a [[buffer(0)]],
device const float * b [[buffer(1)]],
device float * output [[buffer(2)]],
constant ElemParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]])
{
uint idx = gid.x;
if (idx >= params.N) { return; }
output[idx] = (params.op == 0u) ? (a[idx] + b[idx]) : (a[idx] * b[idx]);
}
// ===========================================================================
// embedding
// ===========================================================================
struct EmbedParams {
uint N;
uint D;
uint V;
float scale; // post-lookup multiplier; 1.0 = identity (BitNet/Qwen/Falcon-E),
// sqrt(hidden) on Gemma 3 where the reference impl scales
// the hidden state right after the embed gather in fp32 —
// baking into the f16 dequant at upload time would saturate
// some entries to inf (e.g. d * 127 * sqrt(640) overflows
// the fp16 max for sufficiently large d).
};
kernel void embedding_main(
device const uint * token_ids [[buffer(0)]],
device const uint * embed_table [[buffer(1)]],
device float * output [[buffer(2)]],
constant EmbedParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]])
{
uint idx = gid.x;
uint total = params.N * params.D;
if (idx >= total) { return; }
uint token = idx / params.D;
uint dim = idx % params.D;
uint token_id = token_ids[token];
if (token_id < params.V) {
uint flat = token_id * params.D + dim;
uint packed = embed_table[flat / 2u];
float2 pair = float2(as_type<half2>(packed));
float v = ((flat & 1u) == 1u) ? pair.y : pair.x;
output[idx] = v * params.scale;
} else {
output[idx] = 0.0f;
}
}
// ===========================================================================
// f32_matmul — F32 GEMV against F16-packed embedding (LM head)
// ===========================================================================
struct F32MMParams {
uint N;
uint V;
uint D;
};
constant uint F32MM_WG = 256u;
kernel void f32_matmul_main(
device const float * hidden [[buffer(0)]],
device const uint * embed [[buffer(1)]],
device float * output [[buffer(2)]],
constant F32MMParams& params [[buffer(3)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]])
{
threadgroup float shared_sums[256];
uint flat_id = wg_id.x + wg_id.y * 65535u;
uint n = flat_id / params.V;
uint v = flat_id % params.V;
if (n >= params.N || v >= params.V) { return; }
uint tid = lid.x;
float acc = 0.0f;
uint hidden_base = n * params.D;
uint embed_base = v * params.D;
uint D_half = params.D / 2u;
for (uint dh = tid; dh < D_half; dh += F32MM_WG) {
uint d = dh * 2u;
uint packed = embed[embed_base / 2u + dh];
float2 pair = float2(as_type<half2>(packed));
acc += hidden[hidden_base + d] * pair.x;
acc += hidden[hidden_base + d + 1u] * pair.y;
}
shared_sums[tid] = acc;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = F32MM_WG / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shared_sums[tid] += shared_sums[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (tid == 0u) {
output[n * params.V + v] = shared_sums[0];
}
}
// ===========================================================================
// quantize — per-token absmax f32 -> int8
// ===========================================================================
struct QuantParams {
uint N;
uint D;
};
constant uint QUANT_WG = 256u;
constant uint QUANT_N_SIMDS = 8u; // QUANT_WG / 32
kernel void quantize_main(
device const float * input [[buffer(0)]],
device int * output [[buffer(1)]],
device float * scales [[buffer(2)]],
constant QuantParams& params [[buffer(3)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]],
uint simd_lane [[thread_index_in_simdgroup]],
uint simd_id [[simdgroup_index_in_threadgroup]])
{
// Slot 0 doubles as the broadcast point for absmax after reduction,
// so phase-2 reads it in every thread without another reduce.
threadgroup float simd_partials[QUANT_N_SIMDS];
uint row = wg_id.x;
if (row >= params.N) { return; }
uint tid = lid.x;
uint row_offset = row * params.D;
float local_max = 0.0f;
for (uint col = tid; col < params.D; col += QUANT_WG) {
local_max = max(local_max, fabs(input[row_offset + col]));
}
// Two-step max reduction: simd_max within each simdgroup, then one
// simdgroup folds the 8 partials. Cuts ~7 threadgroup barriers vs
// the tree reduction.
local_max = simd_max(local_max);
if (simd_lane == 0u) simd_partials[simd_id] = local_max;
threadgroup_barrier(mem_flags::mem_threadgroup);
if (simd_id == 0u) {
float p = (simd_lane < QUANT_N_SIMDS) ? simd_partials[simd_lane] : 0.0f;
p = simd_max(p);
if (simd_lane == 0u) {
simd_partials[0] = p; // broadcast slot
scales[row] = (p == 0.0f) ? 1.0f : (p / 127.0f);
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float absmax = simd_partials[0];
float inv_scale = (absmax == 0.0f) ? 0.0f : (127.0f / absmax);
for (uint col = tid; col < params.D; col += QUANT_WG) {
float val = input[row_offset + col];
int q = clamp(int(round(val * inv_scale)), -127, 127);
output[row_offset + col] = q;
}
}
// ===========================================================================
// rmsnorm
// ===========================================================================
struct RMSNormParams {
uint N;
uint D;
float eps;
};
constant uint RMS_WG = 256u;
// NOTE: simd_sum reduction was tried here (the natural counterpart to
// quantize's simd_max), but float sum-of-squares is non-associative — the
// reordered accumulation produced bit-different rms values, which cascade
// through the transformer and pick slightly different tokens after a few
// dozen output bytes. Reverted to the original tree reduction so the
// metal port stays exact-match with the standalone WebGPU baseline. Cost:
// ~3% decode TPS that we'd otherwise have. Kept quantize (max is
// associative) and ternary_gemv (int32 sum is associative) at simd_sum.
kernel void rmsnorm_main(
device const float * input [[buffer(0)]],
device const float * weight [[buffer(1)]],
device float * output [[buffer(2)]],
constant RMSNormParams& params [[buffer(3)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]])
{
threadgroup float shared_sum[256];
uint row = wg_id.x;
if (row >= params.N) { return; }
uint tid = lid.x;
uint row_offset = row * params.D;
float local_sum = 0.0f;
for (uint col = tid; col < params.D; col += RMS_WG) {
float v = input[row_offset + col];
local_sum += v * v;
}
shared_sum[tid] = local_sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = RMS_WG / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shared_sum[tid] += shared_sum[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float rms = rsqrt(shared_sum[0] / float(params.D) + params.eps);
for (uint col = tid; col < params.D; col += RMS_WG) {
output[row_offset + col] = input[row_offset + col] * rms * weight[col];
}
}
// ===========================================================================
// rope
// ===========================================================================
struct RopeParams {
uint N;
uint num_heads;
uint head_dim;
uint pos_offset;
uint neox; // 1 -> split-half pairs (Gemma3/Qwen3/llama-BitNet); 0 -> consecutive
float theta_base;
};
kernel void rope_main(
device const float * input [[buffer(0)]],
device float * output [[buffer(1)]],
constant RopeParams& params [[buffer(2)]],
uint3 gid [[thread_position_in_grid]])
{
uint half_dim = params.head_dim / 2u;
uint total_pairs = params.N * params.num_heads * half_dim;
uint pair_idx = gid.x;
if (pair_idx >= total_pairs) { return; }
uint dim_pair = pair_idx % half_dim;
uint remainder = pair_idx / half_dim;
uint head = remainder % params.num_heads;
uint token = remainder / params.num_heads;
float pos = float(token + params.pos_offset);
float freq_exp = -2.0f * float(dim_pair) / float(params.head_dim);
float theta = pos * pow(params.theta_base, freq_exp);
float cos_theta = cos(theta);
float sin_theta = sin(theta);
uint base = (token * params.num_heads + head) * params.head_dim;
uint i0, i1;
if (params.neox != 0u) {
i0 = base + dim_pair;
i1 = base + dim_pair + half_dim;
} else {
i0 = base + dim_pair * 2u;
i1 = i0 + 1u;
}
float x0 = input[i0];
float x1 = input[i1];
output[i0] = x0 * cos_theta - x1 * sin_theta;
output[i1] = x0 * sin_theta + x1 * cos_theta;
}
// ===========================================================================
// softmax
// ===========================================================================
struct SoftmaxParams {
uint N;
uint D;
};
constant uint SM_WG = 256u;
kernel void softmax_main(
device const float * input [[buffer(0)]],
device float * output [[buffer(1)]],
constant SoftmaxParams& params [[buffer(2)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]])
{
threadgroup float shared_val[256];
uint row = wg_id.x;
if (row >= params.N) { return; }
uint tid = lid.x;
uint row_offset = row * params.D;
// Pass 1: row max
float local_max = -3.402823e+38f;
for (uint col = tid; col < params.D; col += SM_WG) {
local_max = max(local_max, input[row_offset + col]);
}
shared_val[tid] = local_max;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = SM_WG / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shared_val[tid] = max(shared_val[tid], shared_val[tid + stride]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float row_max = shared_val[0];
threadgroup_barrier(mem_flags::mem_threadgroup);
// Pass 2: sum of exp(x - max)
float local_sum = 0.0f;
for (uint col = tid; col < params.D; col += SM_WG) {
local_sum += exp(input[row_offset + col] - row_max);
}
shared_val[tid] = local_sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = SM_WG / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shared_val[tid] += shared_val[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float inv_sum = 1.0f / shared_val[0];
threadgroup_barrier(mem_flags::mem_threadgroup);
// Pass 3: normalize
for (uint col = tid; col < params.D; col += SM_WG) {
output[row_offset + col] =
exp(input[row_offset + col] - row_max) * inv_sum;
}
}
// ===========================================================================
// split
// ===========================================================================
struct SplitParams {
uint N;
uint src_stride;
uint src_offset;
uint dst_size;
};
kernel void split_main(
device const float * src [[buffer(0)]],
device float * dst [[buffer(1)]],
constant SplitParams& params [[buffer(2)]],
uint3 gid [[thread_position_in_grid]])
{
uint total = params.N * params.dst_size;
uint idx = gid.x;
if (idx >= total) { return; }
uint token = idx / params.dst_size;
uint i = idx % params.dst_size;
dst[idx] = src[token * params.src_stride + params.src_offset + i];
}
// ===========================================================================
// ternary_gemm — 32×32 tile, 4×4 per-thread, 64 threads/WG.
// ===========================================================================
// Host dispatch grid (GEMM_TILE_M, GEMM_TILE_N in llm.c) MUST match the
// TILE_M / TILE_N constants below or output rows go uncomputed.
//
// Tile sizing (perf log so a future maintainer doesn't redo the search):
// 64×64 / 256 -> 32×64 / 128: +3-5% prefill (M=2048 doubles gx).
// 32×64 -> 32×32 / 64: +3% short prefill (more concurrent WGs), +4-5%
// long prefill (N-direction waste 14% -> 4%).
// Tried and rejected:
// TILE_K=64: shared mem 24 KB halved concurrent WGs/SM (regression).
// TILE_N=16: WG drops to one simdgroup (not enough work to amortize).
// THREAD_TILE_M=8: -10% prefill (register spills with 32 i32 accumulators).
struct GemmParams {
uint M;
uint N;
uint K;
uint K_packed;
};
constant uint TILE_M = 32u;
constant uint TILE_N = 32u;
constant uint TILE_K = 32u;
constant uint THREADS_M = 8u;
constant uint THREADS_N = 8u;
constant uint THREAD_TILE_M = 4u;
constant uint THREAD_TILE_N = 4u;
kernel void ternary_gemm_main(
device const uint * weights [[buffer(0)]],
device const int * input [[buffer(1)]],
device const float * scales [[buffer(2)]],
constant GemmParams& params [[buffer(3)]],
device const float * input_scales [[buffer(4)]],
device float * output [[buffer(5)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]])
{
threadgroup int shared_w[1024]; // TILE_M(32) * TILE_K(32)
threadgroup int shared_x[1024]; // TILE_K(32) * TILE_N(32)
uint wg_row = wg_id.x * TILE_M;
uint wg_col = wg_id.y * TILE_N;
uint tid_m = lid.x;
uint tid_n = lid.y;
int acc[16];
for (uint i = 0u; i < 16u; i++) { acc[i] = 0; }
uint k_tiles = (params.K + TILE_K - 1u) / TILE_K;
for (uint kt = 0u; kt < k_tiles; kt++) {
uint k_base = kt * TILE_K;
// ---- cooperatively load weights tile ----
uint linear_id = tid_m * THREADS_N + tid_n;
uint block = k_base / 128u;
uint group = (k_base % 128u) / 32u;
uint group_shift = 6u - 2u * group;
// 32 rows × 8 u32/row = 256 u32s, 64 threads -> 4 loads each
for (uint ld = 0u; ld < 4u; ld++) {
uint flat_idx = linear_id + ld * 64u;
uint local_row = flat_idx / 8u;
uint u32_in_row = flat_idx % 8u;
uint base_gp = u32_in_row * 4u;
uint global_row = wg_row + local_row;
int w0 = 0, w1 = 0, w2 = 0, w3 = 0;
if (global_row < params.M && k_base < params.K) {
uint packed = weights[global_row * params.K_packed + block * 8u + u32_in_row];
w0 = int((packed >> group_shift) & 3u) - 1;
w1 = int((packed >> (8u + group_shift)) & 3u) - 1;
w2 = int((packed >> (16u + group_shift)) & 3u) - 1;
w3 = int((packed >> (24u + group_shift)) & 3u) - 1;
}
uint sm_base = local_row * TILE_K + base_gp;
shared_w[sm_base] = w0;
shared_w[sm_base + 1u] = w1;
shared_w[sm_base + 2u] = w2;
shared_w[sm_base + 3u] = w3;
}
// ---- cooperatively load input tile ----
uint load_count_x = (TILE_K * TILE_N) / (THREADS_M * THREADS_N);
for (uint ld = 0u; ld < load_count_x; ld++) {
uint idx = linear_id + ld * (THREADS_M * THREADS_N);
uint local_k = idx / TILE_N;
uint local_col = idx % TILE_N;
uint global_k = k_base + local_k;
uint global_col = wg_col + local_col;
int x_val = 0;
if (global_k < params.K && global_col < params.N) {
x_val = input[global_col * params.K + global_k];
}
shared_x[local_k * TILE_N + local_col] = x_val;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
// ---- per-thread 4×4 accumulate (explicit CSE; same as wgsl OPT) ----
uint row0 = tid_m * THREAD_TILE_M;
uint col0 = tid_n * THREAD_TILE_N;
for (uint k = 0u; k < TILE_K; k++) {
int x0 = shared_x[k * TILE_N + col0 ];
int x1 = shared_x[k * TILE_N + col0 + 1u];
int x2 = shared_x[k * TILE_N + col0 + 2u];
int x3 = shared_x[k * TILE_N + col0 + 3u];
int w0 = shared_w[(row0 ) * TILE_K + k];
int w1 = shared_w[(row0 + 1u) * TILE_K + k];
int w2 = shared_w[(row0 + 2u) * TILE_K + k];
int w3 = shared_w[(row0 + 3u) * TILE_K + k];
acc[ 0] += w0 * x0; acc[ 1] += w0 * x1; acc[ 2] += w0 * x2; acc[ 3] += w0 * x3;
acc[ 4] += w1 * x0; acc[ 5] += w1 * x1; acc[ 6] += w1 * x2; acc[ 7] += w1 * x3;
acc[ 8] += w2 * x0; acc[ 9] += w2 * x1; acc[10] += w2 * x2; acc[11] += w2 * x3;
acc[12] += w3 * x0; acc[13] += w3 * x1; acc[14] += w3 * x2; acc[15] += w3 * x3;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
// ---- write results with dequantization ----
for (uint tm = 0u; tm < THREAD_TILE_M; tm++) {
uint global_row = wg_row + tid_m * THREAD_TILE_M + tm;
if (global_row >= params.M) { continue; }
float w_scale = scales[global_row];
for (uint tn = 0u; tn < THREAD_TILE_N; tn++) {
uint global_col = wg_col + tid_n * THREAD_TILE_N + tn;
if (global_col >= params.N) { continue; }
float scale = w_scale * input_scales[global_col];
output[global_col * params.M + global_row] =
float(acc[tm * THREAD_TILE_N + tn]) * scale;
}
}
}
// ===========================================================================
// ternary_gemv — multi-row per workgroup.
// ===========================================================================
// One WG produces GEMV_ROWS_PER_WG output rows (was 1); the reduction
// tail runs in parallel across the rows and the inputs hit Apple's L1/L2
// without an explicit threadgroup stage, so per-row launch + reduction
// overhead is amortized.
//
// Host: dispatch DIV_CEIL(M, GEMV_ROWS_PER_WG) workgroups; host's
// GEMV_ROWS_PER_WG MUST match the shader constant below.
//
// GEMV_WG=32 = exactly one Apple GPU simdgroup, so the reduction
// collapses to a single simd_sum (no inter-simd fold, no shared mem,
// no barriers).
struct GemvParams {
uint M;
uint K;
uint K_packed;
};
constant uint GEMV_WG = 32u;
constant uint GEMV_ROWS_PER_WG = 4u;
kernel void ternary_gemv_main(
device const uint * weights [[buffer(0)]],
device const int * input [[buffer(1)]],
device const float * scales [[buffer(2)]],
constant GemvParams& params [[buffer(3)]],
constant float& input_scale [[buffer(4)]],
device float * output [[buffer(5)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]],
uint simd_lane [[thread_index_in_simdgroup]],
uint simd_id [[simdgroup_index_in_threadgroup]])
{
// With GEMV_WG=32 each WG is exactly one simdgroup, so the entire
// reduction collapses to a single simd_sum() call — no threadgroup
// memory, no barriers. (Previous design used WG=128 + 4 simdgroups
// with an inter-simd shared-mem fold; the smaller WG fits more in
// flight per SM and amortizes launch overhead better for the K
// values we actually see in Falcon-E-class models.)
uint row_base = wg_id.x * GEMV_ROWS_PER_WG;
if (row_base >= params.M) { return; }
uint tid = lid.x;
int acc[GEMV_ROWS_PER_WG];
for (uint r = 0u; r < GEMV_ROWS_PER_WG; r++) { acc[r] = 0; }
for (uint col = tid; col < params.K_packed; col += GEMV_WG) {
uint block = col / 8u;
uint base_gp = (col % 8u) * 4u;
// Pre-load the 16 inputs this column block consumes — they're the
// same for every row in the block, so we read them once into
// registers and reuse across the row loop.
int x0_0 = input[block * 128u + base_gp + 0u + 0u];
int x0_1 = input[block * 128u + base_gp + 0u + 32u];
int x0_2 = input[block * 128u + base_gp + 0u + 64u];
int x0_3 = input[block * 128u + base_gp + 0u + 96u];
int x1_0 = input[block * 128u + base_gp + 1u + 0u];
int x1_1 = input[block * 128u + base_gp + 1u + 32u];
int x1_2 = input[block * 128u + base_gp + 1u + 64u];
int x1_3 = input[block * 128u + base_gp + 1u + 96u];
int x2_0 = input[block * 128u + base_gp + 2u + 0u];
int x2_1 = input[block * 128u + base_gp + 2u + 32u];
int x2_2 = input[block * 128u + base_gp + 2u + 64u];
int x2_3 = input[block * 128u + base_gp + 2u + 96u];
int x3_0 = input[block * 128u + base_gp + 3u + 0u];
int x3_1 = input[block * 128u + base_gp + 3u + 32u];
int x3_2 = input[block * 128u + base_gp + 3u + 64u];
int x3_3 = input[block * 128u + base_gp + 3u + 96u];
for (uint r = 0u; r < GEMV_ROWS_PER_WG; r++) {
uint row = row_base + r;
if (row >= params.M) { break; }
uint packed = weights[row * params.K_packed + col];
uint b0 = (packed >> 0u) & 0xFFu;
uint b1 = (packed >> 8u) & 0xFFu;
uint b2 = (packed >> 16u) & 0xFFu;
uint b3 = (packed >> 24u) & 0xFFu;
acc[r] += (int((b0 >> 6u) & 3u) - 1) * x0_0
+ (int((b0 >> 4u) & 3u) - 1) * x0_1
+ (int((b0 >> 2u) & 3u) - 1) * x0_2
+ (int( b0 & 3u) - 1) * x0_3
+ (int((b1 >> 6u) & 3u) - 1) * x1_0
+ (int((b1 >> 4u) & 3u) - 1) * x1_1
+ (int((b1 >> 2u) & 3u) - 1) * x1_2
+ (int( b1 & 3u) - 1) * x1_3
+ (int((b2 >> 6u) & 3u) - 1) * x2_0
+ (int((b2 >> 4u) & 3u) - 1) * x2_1
+ (int((b2 >> 2u) & 3u) - 1) * x2_2
+ (int( b2 & 3u) - 1) * x2_3
+ (int((b3 >> 6u) & 3u) - 1) * x3_0
+ (int((b3 >> 4u) & 3u) - 1) * x3_1
+ (int((b3 >> 2u) & 3u) - 1) * x3_2
+ (int( b3 & 3u) - 1) * x3_3;
}
}
// simd_sum: 32-lane reduction in one op, broadcast to all lanes.
// With one simdgroup per WG there's no inter-simd fold — every lane
// already holds the sum. Lanes 0..3 each write one row's result.
for (uint r = 0u; r < GEMV_ROWS_PER_WG; r++) {
acc[r] = simd_sum(acc[r]);
}
if (simd_lane < GEMV_ROWS_PER_WG) {
uint r = simd_lane;
uint row = row_base + r;
if (row < params.M) {
output[row] = float(acc[r]) * scales[row] * input_scale;
}
}
}
// ===========================================================================
// ternary_tq2_gemv — GGML TQ2_0 matvec, on-disk layout, n==1 token.
// ===========================================================================
// Block geometry (66 bytes / 256 elements / row × n_blocks blocks):
// qs[0..63]: packed 2-bit trits, layout described in
// vendor/simd/bn_simd_kernels.h::bn_simd_tq2_gemv. In half
// h in {0,1}, byte m in [0,32) covers four trits, where
// bit-pair p in [0,4) maps to element h*128 + p*32 + m.
// d [64..65]: fp16 per-block scale (little-endian, 2-aligned because
// 66*k + 64 is always even).
//
// One workgroup = one 32-thread simdgroup = 4 output rows × 8 threads/row.
// Each row's 8 threads cooperatively walk all n_blocks blocks; per block
// they split the 64 packed bytes 8-byte-per-thread, decoding 32 trits each.
// A simdgroup XOR reduction folds the 8 per-thread sums into one row sum;
// lane 0 of each row writes the result. The XOR pattern (4, 2, 1) reduces
// within each 8-lane chunk independently, which is exactly the 8-thread
// row group layout (tid 0..7, 8..15, 16..23, 24..31).
//
// No threadgroup memory, no barriers. Mirrors the design of the existing
// ternary_gemv_main shader above so the same dispatch grid (M/4 WGs)
// applies.
struct TQ2GemvParams {
uint M;
uint K;
float tensor_scale;
};
constant uint TQ2_GEMV_ROWS_PER_WG = 4u;
constant uint TQ2_GEMV_THREADS_PER_ROW = 8u; // 4 rows × 8 = 32 = one simdgroup
constant uint TQ2_GEMV_BLOCK_BYTES = 66u;
constant uint TQ2_GEMV_BYTES_PER_THR = 8u; // 64 qs bytes / 8 threads
// Dispatch grid: (ceil(M/4), n_tokens, 1). For n==1 the y-axis collapses
// to a single WG slice and token=0, preserving the original behaviour.
// llama.cpp's TQ2_0_Q8_K dot product, bit-for-bit:
// sumi(b) = sum_{j={0,32}} sum_{l=0..3} sum_{m=0..31}
// ((W.qs[j+m] >> (l*2)) & 3 - 1) * X.qs[j*4 + l*32 + m]
// sumf += sumi * d_w * d_x (per block)
// Input is supplied as (x_qs: int8 per token, x_d: fp16 per block).
kernel void ternary_tq2_gemv_main(
device const uchar * weights [[buffer(0)]],
device const char * x_qs [[buffer(1)]],
device const ushort * x_d [[buffer(2)]],
device float * output [[buffer(3)]],
constant TQ2GemvParams& params [[buffer(4)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]])
{
uint token = wg_id.y;
uint tid = lid.x;
uint local_row = tid / TQ2_GEMV_THREADS_PER_ROW; // 0..3
uint local_elem = tid % TQ2_GEMV_THREADS_PER_ROW; // 0..7
uint global_row = wg_id.x * TQ2_GEMV_ROWS_PER_WG + local_row;
uint n_blocks = params.K / 256u;
uint row_bytes = n_blocks * TQ2_GEMV_BLOCK_BYTES;
uint x_qs_base = token * params.K;
uint x_d_base = token * n_blocks;
float row_sum = 0.0f;
if (global_row < params.M) {
uint row_off = global_row * row_bytes;
uint byte_start = local_elem * TQ2_GEMV_BYTES_PER_THR; // 0,8,16,...,56
for (uint b = 0u; b < n_blocks; b++) {
uint block_off = row_off + b * TQ2_GEMV_BLOCK_BYTES;
ushort d_w_bits = (ushort)weights[block_off + 64u]
| ((ushort)weights[block_off + 65u] << 8);
float d_w = float(as_type<half>(d_w_bits));
float d_x = float(as_type<half>(x_d[x_d_base + b]));
uint x_block_off = x_qs_base + b * 256u;
int sumi = 0;
for (uint k = 0u; k < TQ2_GEMV_BYTES_PER_THR; k++) {
uint byte_idx = byte_start + k;
uint half_idx = byte_idx / 32u;
uint m_idx = byte_idx % 32u;
uchar byte_v = weights[block_off + byte_idx];
uint x_half_off = x_block_off + half_idx * 128u;
int q0 = (int)((byte_v >> 0) & 3u) - 1;
int q1 = (int)((byte_v >> 2) & 3u) - 1;
int q2 = (int)((byte_v >> 4) & 3u) - 1;
int q3 = (int)((byte_v >> 6) & 3u) - 1;
sumi += q0 * (int)x_qs[x_half_off + 0u + m_idx];
sumi += q1 * (int)x_qs[x_half_off + 32u + m_idx];
sumi += q2 * (int)x_qs[x_half_off + 64u + m_idx];
sumi += q3 * (int)x_qs[x_half_off + 96u + m_idx];
}
row_sum += (float)sumi * d_w * d_x;
}
}
// Reduce within the row's 8-thread chunk via simdgroup XOR shuffles.
// Threads outside the row group (global_row >= M) contributed 0, so
// their reduction is a no-op.
row_sum += simd_shuffle_xor(row_sum, 4);
row_sum += simd_shuffle_xor(row_sum, 2);
row_sum += simd_shuffle_xor(row_sum, 1);
if (local_elem == 0u && global_row < params.M) {
output[token * params.M + global_row] = row_sum * params.tensor_scale;
}
}
// ===========================================================================
// ternary_q4_gemv — GGML Q4_0 matvec, on-disk layout, n==1 token.
// ===========================================================================
// Block geometry (18 bytes / 32 elements / row * n_blocks blocks):
// d [0..1]: little-endian fp16 per-block scale.
// qs [2..17]: 16 packed bytes. Byte i holds two 4-bit nibbles:
// low nibble (qs[i] & 0xF) - 8 -> element i
// high nibble (qs[i] >> 4) - 8 -> element i + 16
// Total: 32 signed-4-bit elements per block, signed range [-8, +7].
//
// Same WG geometry as ternary_tq2_gemv_main: 32 threads = 1 simdgroup =
// 4 rows * 8 threads/row. Each row's 8 threads split the 16 qs bytes
// 2-per-thread, decoding 4 trits each (2 nibbles per byte, 2 bytes per
// thread). The simd_shuffle_xor reduction folds the 8 per-thread sums
// into one row sum on lane 0.
struct Q4GemvParams {
uint M;
uint K;
};
constant uint Q4_GEMV_ROWS_PER_WG = 4u;
constant uint Q4_GEMV_THREADS_PER_ROW = 8u; // 4 rows * 8 = 32 = simdgroup
constant uint Q4_GEMV_BLOCK_BYTES = 18u;
constant uint Q4_GEMV_BYTES_PER_THR = 2u; // 16 qs bytes / 8 threads
// Dispatch grid: (ceil(M/4), n_tokens, 1). See ternary_tq2_gemv_main for
// the multi-token shape rationale.
// llama.cpp's Q4_0_Q8_0 dot product, bit-for-bit:
// sumi(b) = sum_{i=0..15} ((W.qs[i]&0xF)-8)*X.qs[i] + ((W.qs[i]>>4)-8)*X.qs[i+16]
// sumf += sumi * d_w * d_x (per block)
// Input is supplied as (x_qs: int8 per token, x_d: fp16 per block).
kernel void ternary_q4_gemv_main(
device const uchar * weights [[buffer(0)]],
device const char * x_qs [[buffer(1)]],
device const ushort * x_d [[buffer(2)]],
device float * output [[buffer(3)]],
constant Q4GemvParams& params [[buffer(4)]],
uint3 wg_id [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]])
{
uint token = wg_id.y;
uint tid = lid.x;
uint local_row = tid / Q4_GEMV_THREADS_PER_ROW; // 0..3
uint local_elem = tid % Q4_GEMV_THREADS_PER_ROW; // 0..7
uint global_row = wg_id.x * Q4_GEMV_ROWS_PER_WG + local_row;
uint n_blocks = params.K / 32u;
uint row_bytes = n_blocks * Q4_GEMV_BLOCK_BYTES;
uint x_qs_base = token * params.K;
uint x_d_base = token * n_blocks;
float row_sum = 0.0f;
if (global_row < params.M) {
uint row_off = global_row * row_bytes;
for (uint b = 0u; b < n_blocks; b++) {
uint block_off = row_off + b * Q4_GEMV_BLOCK_BYTES;
ushort d_w_bits = (ushort)weights[block_off + 0u]
| ((ushort)weights[block_off + 1u] << 8);
float d_w = float(as_type<half>(d_w_bits));
float d_x = float(as_type<half>(x_d[x_d_base + b]));
uint x_block_off = x_qs_base + b * 32u;
int sumi = 0;
for (uint k = 0u; k < Q4_GEMV_BYTES_PER_THR; k++) {
uint i = local_elem * Q4_GEMV_BYTES_PER_THR + k; // 0..15
uchar byte_v = weights[block_off + 2u + i];
int q_lo = (int)(byte_v & 0x0Fu) - 8;
int q_hi = (int)(byte_v >> 4) - 8;
sumi += q_lo * (int)x_qs[x_block_off + i];
sumi += q_hi * (int)x_qs[x_block_off + i + 16u];
}
row_sum += (float)sumi * d_w * d_x;
}
}
row_sum += simd_shuffle_xor(row_sum, 4);
row_sum += simd_shuffle_xor(row_sum, 2);
row_sum += simd_shuffle_xor(row_sum, 1);
if (local_elem == 0u && global_row < params.M) {
output[token * params.M + global_row] = row_sum;
}
}
// ===========================================================================
// Input-row block quantizers — produce (x_qs, x_d) for the block-quant
// gemvs above. One workgroup per token: each thread handles one block, the
// workgroup is sized to cover the largest token's block count we expect.
// Single-threaded per block (block widths 32 / 256 are tiny — the matmul
// itself dominates runtime).
// ===========================================================================
struct QuantBlkParams { uint K; uint n_tokens; };
// f32 -> f16 round-to-nearest-even (flush subnormals to ±0). The Metal
// builtin half(...) handles this correctly via the IEEE conversion.
kernel void quantize_q8_0_main(
device const float * input [[buffer(0)]],
device char * x_qs [[buffer(1)]],
device ushort * x_d [[buffer(2)]],
constant QuantBlkParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]])
{
uint t = gid.x;
if (t >= params.n_tokens) { return; }
uint n_blocks = params.K / 32u;
device const float * in = input + (size_t)t * params.K;
device char * qs = x_qs + (size_t)t * params.K;
device ushort* d = x_d + (size_t)t * n_blocks;
for (uint b = 0u; b < n_blocks; b++) {
float amax = 0.0f;
uint off = b * 32u;
for (uint j = 0u; j < 32u; j++) {
float a = fabs(in[off + j]);
if (a > amax) amax = a;
}
float dv = amax / 127.0f;
float id = (dv > 0.0f) ? 1.0f / dv : 0.0f;
d[b] = as_type<ushort>(half(dv));
for (uint j = 0u; j < 32u; j++) {
float v = in[off + j] * id;
int q = (int)round(v);
if (q > 127) q = 127;
if (q < -128) q = -128;
qs[off + j] = (char)q;
}
}
}
kernel void quantize_q8_K_main(
device const float * input [[buffer(0)]],
device char * x_qs [[buffer(1)]],
device ushort * x_d [[buffer(2)]],
constant QuantBlkParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]])
{
uint t = gid.x;
if (t >= params.n_tokens) { return; }
uint n_blocks = params.K / 256u;
device const float * in = input + (size_t)t * params.K;
device char * qs = x_qs + (size_t)t * params.K;
device ushort* d = x_d + (size_t)t * n_blocks;
for (uint b = 0u; b < n_blocks; b++) {
float amax = 0.0f;
uint off = b * 256u;
for (uint j = 0u; j < 256u; j++) {
float a = fabs(in[off + j]);
if (a > amax) amax = a;
}
float dv = amax / 127.0f;
float id = (dv > 0.0f) ? 1.0f / dv : 0.0f;
d[b] = as_type<ushort>(half(dv));
for (uint j = 0u; j < 256u; j++) {
float v = in[off + j] * id;
int q = (int)round(v);
if (q > 127) q = 127;
if (q < -128) q = -128;
qs[off + j] = (char)q;
}
}
}
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment