|
// 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; |
|
} |
|
} |
|
} |