chore: import upstream snapshot with attribution

This commit is contained in:
wehub-resource-sync
2026-07-13 12:34:46 +08:00
commit f4e68ed970
84 changed files with 14896 additions and 0 deletions
+337
View File
@@ -0,0 +1,337 @@
// cider_sdpa_vector.h — Cider v9 optimized SDPA vector kernels
//
// Improvements over stock MLX:
// 1. Contiguous chunk layout (cache-friendly sequential access)
// 2. FlashInfer-style register tiling (TILE=4 unroll)
// 3. Adaptive blocks selection per GQA ratio
//
// Template parameters:
// T — data type (float, half, bfloat)
// D — head dimension (64, 96, 128, 256)
// BLOCKS — number of blocks for 2-pass (32, 64, 128)
//
// Three kernels:
// cider_sdpa_vector<T, D> — 1-pass (short N)
// cider_sdpa_vector_2pass_1<T, D, BLOCKS> — 2-pass pass1 (per-block partials)
// cider_sdpa_vector_2pass_2<T, D, BLOCKS> — 2-pass pass2 (reduce)
#pragma once
#include <metal_simdgroup>
using namespace metal;
// ═══════════════════════════════════════════════════════════════
// 1-pass kernel (N <= threshold, no intermediate buffer)
// ═══════════════════════════════════════════════════════════════
template <typename T, int D>
[[kernel]] void cider_sdpa_vector(
const device T* queries [[buffer(0)]],
const device T* keys [[buffer(1)]],
const device T* values [[buffer(2)]],
device T* out [[buffer(3)]],
const constant int& gqa_factor [[buffer(4)]],
const constant int& N [[buffer(5)]],
const constant size_t& k_head_stride [[buffer(6)]],
const constant size_t& k_seq_stride [[buffer(7)]],
const constant size_t& v_head_stride [[buffer(8)]],
const constant size_t& v_seq_stride [[buffer(9)]],
const constant float& scale [[buffer(10)]],
uint3 tid [[threadgroup_position_in_grid]],
uint simd_gid [[simdgroup_index_in_threadgroup]],
uint simd_lid [[thread_index_in_simdgroup]]) {
constexpr int BN = 32;
constexpr int BD = 32;
constexpr int qk_per_thread = D / BD;
constexpr int v_per_thread = D / BD;
typedef float U;
const int inner_k_stride = BN * int(k_seq_stride);
const int inner_v_stride = BN * int(v_seq_stride);
thread U q_reg[qk_per_thread];
thread U k_reg[qk_per_thread];
thread U o_reg[v_per_thread];
threadgroup U tg_outputs[BN * BD];
threadgroup U tg_max_scores[BN];
threadgroup U tg_sum_exp_scores[BN];
const int q_batch_head_idx = tid.x;
const int kv_head_idx = q_batch_head_idx / gqa_factor;
const int o_offset = q_batch_head_idx;
const device T* q_ptr = queries + o_offset * D + simd_lid * qk_per_thread;
const device T* k_ptr = keys + kv_head_idx * int(k_head_stride)
+ simd_gid * int(k_seq_stride) + simd_lid * qk_per_thread;
const device T* v_ptr = values + kv_head_idx * int(v_head_stride)
+ simd_gid * int(v_seq_stride) + simd_lid * v_per_thread;
for (int i = 0; i < qk_per_thread; i++) {
q_reg[i] = static_cast<U>(scale) * static_cast<U>(q_ptr[i]);
}
for (int i = 0; i < v_per_thread; i++) {
o_reg[i] = 0;
}
U max_score = -1e38f;
U sum_exp_score = 0;
for (int i = simd_gid; i < N; i += BN) {
for (int j = 0; j < qk_per_thread; j++) {
k_reg[j] = static_cast<U>(k_ptr[j]);
}
U score = 0;
for (int j = 0; j < qk_per_thread; j++) {
score += q_reg[j] * k_reg[j];
}
score = simd_sum(score);
U new_max = max(max_score, score);
U factor = fast::exp(max_score - new_max);
U exp_score = fast::exp(score - new_max);
max_score = new_max;
sum_exp_score = sum_exp_score * factor + exp_score;
for (int j = 0; j < v_per_thread; j++) {
o_reg[j] = o_reg[j] * factor + exp_score * static_cast<U>(v_ptr[j]);
}
k_ptr += inner_k_stride;
v_ptr += inner_v_stride;
}
if (simd_lid == 0) {
tg_max_scores[simd_gid] = max_score;
tg_sum_exp_scores[simd_gid] = sum_exp_score;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
max_score = tg_max_scores[simd_lid];
U new_max_final = simd_max(max_score);
U factor = fast::exp(max_score - new_max_final);
sum_exp_score = simd_sum(tg_sum_exp_scores[simd_lid] * factor);
for (int i = 0; i < v_per_thread; i++) {
tg_outputs[simd_lid * BD + simd_gid] = o_reg[i];
threadgroup_barrier(mem_flags::mem_threadgroup);
o_reg[i] = simd_sum(tg_outputs[simd_gid * BD + simd_lid] * factor);
o_reg[i] = sum_exp_score == 0 ? o_reg[i] : (o_reg[i] / sum_exp_score);
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (simd_lid == 0) {
device T* o_ptr = out + o_offset * D + simd_gid * v_per_thread;
for (int i = 0; i < v_per_thread; i++) {
o_ptr[i] = static_cast<T>(o_reg[i]);
}
}
}
// ═══════════════════════════════════════════════════════════════
// 2-pass pass1: per-block partial results
// v9: contiguous chunks + TILE=4 register tiling
// ═══════════════════════════════════════════════════════════════
template <typename T, int D, int BLOCKS>
[[kernel]] void cider_sdpa_vector_2pass_1(
const device T* queries [[buffer(0)]],
const device T* keys [[buffer(1)]],
const device T* values [[buffer(2)]],
device T* partials [[buffer(3)]],
device float* sums [[buffer(4)]],
device float* maxs [[buffer(5)]],
const constant int& N [[buffer(7)]],
const constant size_t& k_head_stride [[buffer(8)]],
const constant size_t& k_seq_stride [[buffer(9)]],
const constant size_t& v_head_stride [[buffer(10)]],
const constant size_t& v_seq_stride [[buffer(11)]],
const constant float& scale [[buffer(12)]],
uint3 tptg [[threads_per_threadgroup]],
uint3 tidtg [[thread_position_in_threadgroup]],
uint3 tid [[threadgroup_position_in_grid]],
uint3 tpg [[threadgroups_per_grid]],
uint simd_lid [[thread_index_in_simdgroup]]) {
constexpr int BD = 32;
constexpr int qk_per_thread = D / BD;
constexpr int v_per_thread = D / BD;
constexpr int TILE = 4;
typedef float U;
thread U q_reg[qk_per_thread];
thread U o_reg[v_per_thread] = {0};
const int kv_head_idx = tid.x;
const int batch_idx = tid.y;
const int block_idx = tid.z;
const int gqa_factor = tptg.y;
const int q_head_idx = gqa_factor * kv_head_idx + tidtg.y;
const int num_kv_heads = tpg.x;
const int num_q_heads = num_kv_heads * gqa_factor;
const int q_batch_head_idx = batch_idx * num_q_heads + q_head_idx;
const int o_offset = q_batch_head_idx;
queries += o_offset * D + simd_lid * qk_per_thread;
const int kv_batch_head_idx = batch_idx * num_kv_heads + kv_head_idx;
// v9: contiguous chunk layout (not interleaved)
const int chunk_size = (N + BLOCKS - 1) / BLOCKS;
const int kv_start = block_idx * chunk_size;
const int kv_end = min(kv_start + chunk_size, N);
const device T* k_ptr = keys + kv_batch_head_idx * int(k_head_stride)
+ kv_start * int(k_seq_stride) + simd_lid * qk_per_thread;
const device T* v_ptr = values + kv_batch_head_idx * int(v_head_stride)
+ kv_start * int(v_seq_stride) + simd_lid * v_per_thread;
device T* o_ptr = partials + o_offset * BLOCKS * D + block_idx * D
+ simd_lid * v_per_thread;
// Read query
for (int i = 0; i < qk_per_thread; i++) {
q_reg[i] = static_cast<U>(scale) * static_cast<U>(queries[i]);
}
U max_score = -1e38f;
U sum_exp_score = 0;
const int kss = int(k_seq_stride);
const int vss = int(v_seq_stride);
// Main loop with TILE=4 unrolling
int pos = kv_start;
const int tiled_end = kv_start + ((kv_end - kv_start) / TILE) * TILE;
for (; pos < tiled_end; pos += TILE) {
U scores[TILE];
for (int t = 0; t < TILE; t++) {
U score = 0;
const device T* kt = k_ptr + t * kss;
for (int j = 0; j < qk_per_thread; j++) {
score += q_reg[j] * static_cast<U>(kt[j]);
}
scores[t] = simd_sum(score);
}
for (int t = 0; t < TILE; t++) {
U new_max = max(max_score, scores[t]);
U factor = fast::exp(max_score - new_max);
U exp_score = fast::exp(scores[t] - new_max);
max_score = new_max;
sum_exp_score = sum_exp_score * factor + exp_score;
const device T* vt = v_ptr + t * vss;
for (int j = 0; j < v_per_thread; j++) {
o_reg[j] = o_reg[j] * factor + exp_score * static_cast<U>(vt[j]);
}
}
k_ptr += TILE * kss;
v_ptr += TILE * vss;
}
// Remainder
for (; pos < kv_end; pos++) {
U score = 0;
for (int j = 0; j < qk_per_thread; j++) {
score += q_reg[j] * static_cast<U>(k_ptr[j]);
}
score = simd_sum(score);
U new_max = max(max_score, score);
U factor = fast::exp(max_score - new_max);
U exp_score = fast::exp(score - new_max);
max_score = new_max;
sum_exp_score = sum_exp_score * factor + exp_score;
for (int j = 0; j < v_per_thread; j++) {
o_reg[j] = o_reg[j] * factor + exp_score * static_cast<U>(v_ptr[j]);
}
k_ptr += kss;
v_ptr += vss;
}
// Write partial results
if (simd_lid == 0) {
sums[o_offset * BLOCKS + block_idx] = sum_exp_score;
maxs[o_offset * BLOCKS + block_idx] = max_score;
}
for (int i = 0; i < v_per_thread; i++) {
o_ptr[i] = static_cast<T>(o_reg[i]);
}
}
// ═══════════════════════════════════════════════════════════════
// 2-pass pass2: reduce partial results across blocks
// ═══════════════════════════════════════════════════════════════
template <typename T, int D, int BLOCKS>
[[kernel]] void cider_sdpa_vector_2pass_2(
const device T* partials [[buffer(0)]],
const device float* sums [[buffer(1)]],
const device float* maxs [[buffer(2)]],
device T* out [[buffer(3)]],
uint3 tid [[threadgroup_position_in_grid]],
uint simd_gid [[simdgroup_index_in_threadgroup]],
uint simd_lid [[thread_index_in_simdgroup]]) {
constexpr int BN = 32;
constexpr int BD = 32;
constexpr int elem_per_thread = D / BD;
typedef float U;
thread U o_reg[elem_per_thread] = {0};
threadgroup U tg_outputs[BN * BD];
const int head_idx = tid.x;
const device T* p_ptr = partials + head_idx * BLOCKS * D
+ simd_gid * D + simd_lid * elem_per_thread;
const device float* s_ptr = sums + head_idx * BLOCKS;
const device float* m_ptr = maxs + head_idx * BLOCKS;
// Reduce max
U max_score = -1e38f;
for (int b = 0; b < BLOCKS / BN; ++b) {
max_score = max(max_score, m_ptr[simd_lid + BN * b]);
}
max_score = simd_max(max_score);
// Reduce sum_exp
U sum_exp_score = 0;
for (int b = 0; b < BLOCKS / BN; ++b) {
U factor = fast::exp(m_ptr[simd_lid + BN * b] - max_score);
sum_exp_score += factor * s_ptr[simd_lid + BN * b];
}
sum_exp_score = simd_sum(sum_exp_score);
// Reduce partials
const device float* m_walk = m_ptr;
for (int b = 0; b < BLOCKS / BN; ++b) {
U factor = fast::exp(m_walk[simd_gid] - max_score);
for (int i = 0; i < elem_per_thread; i++) {
o_reg[i] += factor * static_cast<U>(p_ptr[i]);
}
m_walk += BN;
p_ptr += BN * D;
}
// Transpose + reduce via shared memory
for (int i = 0; i < elem_per_thread; i++) {
tg_outputs[simd_lid * BD + simd_gid] = o_reg[i];
threadgroup_barrier(mem_flags::mem_threadgroup);
o_reg[i] = simd_sum(tg_outputs[simd_gid * BD + simd_lid]);
o_reg[i] = sum_exp_score == 0 ? o_reg[i] : (o_reg[i] / sum_exp_score);
threadgroup_barrier(mem_flags::mem_threadgroup);
}
if (simd_lid == 0) {
device T* o_ptr = out + head_idx * D + simd_gid * elem_per_thread;
for (int i = 0; i < elem_per_thread; i++) {
o_ptr[i] = static_cast<T>(o_reg[i]);
}
}
}
+42
View File
@@ -0,0 +1,42 @@
// cider_sdpa_vector.metal — Template instantiations for Cider v9 SDPA kernels
//
// Naming convention:
// 1-pass: cider_sdpa_vector_{type}_{D}
// 2-pass1: cider_sdpa_vector_2pass_1_{type}_{D}_b{BLOCKS}
// 2-pass2: cider_sdpa_vector_2pass_2_{type}_{D}_b{BLOCKS}
#include <metal_stdlib>
#include "cider_sdpa_vector.h"
using namespace metal;
// ── Helper macros ────────────────────────────────────────────────
#define instantiate_1pass(type, type_name, D) \
template [[host_name("cider_sdpa_vector_" #type_name "_" #D)]] \
[[kernel]] decltype(cider_sdpa_vector<type, D>) cider_sdpa_vector<type, D>;
#define instantiate_2pass(type, type_name, D, B) \
template [[host_name("cider_sdpa_vector_2pass_1_" #type_name "_" #D "_b" #B)]] \
[[kernel]] decltype(cider_sdpa_vector_2pass_1<type, D, B>) cider_sdpa_vector_2pass_1<type, D, B>; \
template [[host_name("cider_sdpa_vector_2pass_2_" #type_name "_" #D "_b" #B)]] \
[[kernel]] decltype(cider_sdpa_vector_2pass_2<type, D, B>) cider_sdpa_vector_2pass_2<type, D, B>;
#define instantiate_all_blocks(type, type_name, D) \
instantiate_2pass(type, type_name, D, 32) \
instantiate_2pass(type, type_name, D, 64) \
instantiate_2pass(type, type_name, D, 128)
#define instantiate_heads(type, type_name) \
instantiate_1pass(type, type_name, 64) \
instantiate_1pass(type, type_name, 96) \
instantiate_1pass(type, type_name, 128) \
instantiate_1pass(type, type_name, 256) \
instantiate_all_blocks(type, type_name, 64) \
instantiate_all_blocks(type, type_name, 96) \
instantiate_all_blocks(type, type_name, 128) \
instantiate_all_blocks(type, type_name, 256)
// ── Instantiate for all types ────────────────────────────────────
instantiate_heads(float, float32)
instantiate_heads(half, float16)
instantiate_heads(bfloat, bfloat16)
+223
View File
@@ -0,0 +1,223 @@
// ============================================================
// Per-group INT8 TensorOps GEMM — symmetric quantization (bias=0)
// Target: Apple M5 (G17G), Metal 4
//
// Weight layout: B is [N, K] int8 (per-group symmetric quantized)
// - scales_w: [num_groups, N] float32 (TRANSPOSED for coalesced access)
// - scales_a: [M] float32 (per-token activation scales)
//
// Computes: C[m,n] = scale_a[m] * Sigma_g { float(dot_int32[m,n,g]) * scale_w[g,n] }
// where dot_int32[m,n,g] = Sigma_{k in group g} A_int8[m,k] * B_int8[n,k]
//
// Supported group_size: 64, 128, 256
// V2: scale_w transposed [num_groups, N] for coalesced SIMD access
// ============================================================
#include <MetalPerformancePrimitives/MetalPerformancePrimitives.h>
#include <metal_stdlib>
using namespace metal;
// -- NAXFrag layout constants
constant constexpr short kEPF = 8;
constant constexpr short kEC = 4;
constant constexpr short kERJ = 8;
// -- NAXFrag coordinate mapping
inline short2 nax_coord(ushort lid) {
short qid = short(lid >> 2);
short fm = ((qid & 4) | ((short(lid) >> 1) & 3));
short fn = ((qid & 2) | (short(lid) & 1)) * 4;
return short2{fn, fm};
}
// -- Fragment load: device -> register
template <typename T>
inline void frag_load(thread T *dst, const device T *src, int ld, short2 sc,
short off_m = 0, short off_n = 0) {
src += (sc.y + off_m) * ld + (sc.x + off_n);
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kEC; j++) {
dst[i * kEC + j] = src[(i * kERJ) * ld + j];
}
}
}
// -- Per-group GEMM implementation (scale_w transposed: [num_groups, N])
template <int BM, int BN, int BK, int SK, int WM, int WN>
void pergroup_gemm_impl(const device int8_t *A, const device int8_t *B,
device half *C, uint M, uint N, uint K,
const device float *scale_a,
const device float *scale_w, // [num_groups, N] transposed
const device half *bias, // [N] half
uint swizzle_log, uint tiles_m, uint tiles_n,
uint2 tgid, uint sgid, uint lid) {
constexpr int SM = BM / WM;
constexpr int SN = BN / WN;
constexpr short TM = SM / 16;
constexpr short TN = SN / 16;
constexpr short TK = SK / 16;
uint tid_y = (tgid.y << swizzle_log) + (tgid.x & ((1u << swizzle_log) - 1u));
uint tid_x = tgid.x >> swizzle_log;
if (tid_x >= tiles_n || tid_y >= tiles_m) {
return;
}
short2 sc = nax_coord(ushort(lid));
uint sg_row = sgid / WN;
uint sg_col = sgid % WN;
uint m_base = tid_y * BM + sg_row * SM;
uint n_base = tid_x * BN + sg_col * SN;
const device int8_t *sg_A = A + m_base * K;
const device int8_t *sg_B = B + n_base * K;
uint num_groups = K / BK;
constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor(
16, 32, 16, false, true, true,
mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate);
mpp::tensor_ops::matmul2d<desc, metal::execution_simdgroup> gemm_op;
auto ct_a =
gemm_op.get_left_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_b =
gemm_op.get_right_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_c =
gemm_op.get_destination_cooperative_tensor<decltype(ct_a), decltype(ct_b),
int32_t>();
// Float accumulator (across all groups)
float acc[TM * TN][kEPF];
for (int f = 0; f < TM * TN; f++) {
for (int i = 0; i < kEPF; i++) {
acc[f][i] = 0.0f;
}
}
// -- Main K loop: one iteration per group
for (uint g = 0; g < num_groups; g++) {
// INT32 accumulator for this group
int32_t c_frags[TM * TN][kEPF];
for (int f = 0; f < TM * TN; f++) {
for (int i = 0; i < kEPF; i++) {
c_frags[f][i] = 0;
}
}
// Inner loop within group
for (int kk1 = 0; kk1 < BK; kk1 += SK) {
int8_t a_frags[TM][TK][kEPF];
int8_t b_frags[TN][TK][kEPF];
for (short mm = 0; mm < TM; mm++) {
for (short kk = 0; kk < TK; kk++) {
frag_load(a_frags[mm][kk], sg_A + kk1, int(K), sc, short(mm * 16),
short(kk * 16));
}
}
for (short nn = 0; nn < TN; nn++) {
for (short kk = 0; kk < TK; kk++) {
frag_load(b_frags[nn][kk], sg_B + kk1, int(K), sc, short(nn * 16),
short(kk * 16));
}
}
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn += 2) {
for (short kk = 0; kk < TK; kk++) {
for (short i = 0; i < kEPF; i++) {
ct_a[i] = a_frags[mm][kk][i];
}
for (short i = 0; i < kEPF; i++) {
ct_b[i] = b_frags[nn][kk][i];
ct_b[kEPF + i] = b_frags[nn + 1][kk][i];
}
short c0 = mm * TN + nn, c1 = c0 + 1;
for (short i = 0; i < kEPF; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kEPF + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kEPF; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kEPF + i];
}
}
}
}
}
// -- Flush: int32 * scale_w[g, n] -> accumulate
// scale_w is [num_groups, N]: scale_w[g * N + n_idx] is coalesced for adjacent n
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn++) {
short fidx = mm * TN + nn;
float sw[kEC];
for (short j = 0; j < kEC; j++) {
uint n_idx = n_base + uint(sc.x) + uint(nn * 16) + uint(j);
sw[j] = (n_idx < N) ? scale_w[g * N + n_idx] : 0.0f;
}
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kEC; j++) {
acc[fidx][i * kEC + j] += float(c_frags[fidx][i * kEC + j]) * sw[j];
}
}
}
}
sg_A += BK;
sg_B += BK;
}
// -- Store: acc * scale_a + bias -> half
device half *D = C + m_base * N + n_base;
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn++) {
short fidx = mm * TN + nn;
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kEC; j++) {
uint mi = m_base + uint(sc.y) + uint(mm * 16 + i * kERJ);
uint ni = n_base + uint(sc.x) + uint(nn * 16 + j);
if (mi < M && ni < N) {
float val = acc[fidx][i * kEC + j] * scale_a[mi] + float(bias[ni]);
D[(sc.y + mm * 16 + i * kERJ) * int(N) + (sc.x + nn * 16 + j)] =
half(val);
}
}
}
}
}
}
// ============================================================
// Kernel entry points
// ============================================================
#define GEMM_ENTRY(SUFFIX, BM_V, BN_V, BK_V, WM_V, WN_V) \
kernel void pergroup_int8_gemm_##SUFFIX( \
const device int8_t *A [[buffer(0)]], \
const device int8_t *B [[buffer(1)]], device half *C [[buffer(2)]], \
constant uint &M [[buffer(3)]], constant uint &N [[buffer(4)]], \
constant uint &K [[buffer(5)]], \
const device float *scale_a [[buffer(6)]], \
const device float *scale_w [[buffer(7)]], \
constant uint &swizzle_log [[buffer(8)]], \
constant uint &tiles_m [[buffer(9)]], \
constant uint &tiles_n [[buffer(10)]], \
const device half *bias [[buffer(11)]], \
uint2 tgid [[threadgroup_position_in_grid]], \
uint sgid [[simdgroup_index_in_threadgroup]], \
uint lid [[thread_index_in_simdgroup]]) { \
pergroup_gemm_impl<BM_V, BN_V, BK_V, 32, WM_V, WN_V>( \
A, B, C, M, N, K, scale_a, scale_w, bias, swizzle_log, tiles_m, \
tiles_n, tgid, sgid, lid); \
}
GEMM_ENTRY(g64, 128, 128, 64, 4, 4)
GEMM_ENTRY(g64_small, 32, 128, 64, 1, 4)
GEMM_ENTRY(g128, 128, 128, 128, 4, 4)
GEMM_ENTRY(g128_small, 32, 128, 128, 1, 4)
GEMM_ENTRY(g256, 128, 128, 256, 4, 4)
GEMM_ENTRY(g256_small, 32, 128, 256, 1, 4)
+158
View File
@@ -0,0 +1,158 @@
// ============================================================
// Per-group INT8 MV kernel V5 — symmetric-only fast path
//
// Match MLX qmv_fast: 2 SG, VPT=8, block=256
// Symmetric quantization: new_bias is always zero, skip correction.
//
// Formula:
// y[n] = Sigma_g { scale_w[g,n] * dot_g(w[n], x) } + bias[n]
//
// V2: scale_w transposed [num_groups, N] for coalesced access
// ============================================================
#include <metal_stdlib>
using namespace metal;
constant constexpr int SIMD_SIZE = 32;
constant constexpr int NUM_SIMDGROUPS = 2;
constant constexpr int RESULTS_PER_SG = 4;
constant constexpr int VPT = 8;
constant constexpr int BLOCK_K = VPT * SIMD_SIZE; // 256
template <int GROUP_SIZE>
inline void pergroup_mv_v5_impl(
const device half *x,
const device int8_t *W,
device half *y,
const device float *scale_w, // [num_groups, N] transposed
const device float *new_bias [[maybe_unused]],
constant uint &N, constant uint &K,
const device half *bias,
uint tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]])
{
constexpr int rows_per_tg = NUM_SIMDGROUPS * RESULTS_PER_SG; // 8
const uint out_row = tgid * rows_per_tg + sgid * RESULTS_PER_SG;
const uint num_groups = K / GROUP_SIZE;
float result[RESULTS_PER_SG] = {0.0f, 0.0f, 0.0f, 0.0f};
// Pointers
const device half *xp = x + lid * VPT;
const device int8_t *wrows[RESULTS_PER_SG];
for (int r = 0; r < RESULTS_PER_SG; r++) {
uint n_idx = out_row + r;
wrows[r] = (n_idx < N) ? (W + n_idx * K + lid * VPT) : W;
}
for (uint k = 0; k < K; k += BLOCK_K) {
// Load x tile
float xv[VPT];
for (int i = 0; i < VPT; i++) {
xv[i] = float(xp[i]);
}
// Group index for this thread data
uint g = (k + lid * VPT) / GROUP_SIZE;
for (int r = 0; r < RESULTS_PER_SG; r++) {
uint n_idx = out_row + r;
if (n_idx >= N) continue;
// Load 8 int8 weights as 2x uint32
uint32_t packed0 = *reinterpret_cast<const device uint32_t *>(wrows[r]);
uint32_t packed1 = *reinterpret_cast<const device uint32_t *>(wrows[r] + 4);
float b0 = float(int8_t(packed0 & 0xFF));
float b1 = float(int8_t((packed0 >> 8) & 0xFF));
float b2 = float(int8_t((packed0 >> 16) & 0xFF));
float b3 = float(int8_t((packed0 >> 24) & 0xFF));
float b4 = float(int8_t(packed1 & 0xFF));
float b5 = float(int8_t((packed1 >> 8) & 0xFF));
float b6 = float(int8_t((packed1 >> 16) & 0xFF));
float b7 = float(int8_t((packed1 >> 24) & 0xFF));
float dot = xv[0]*b0 + xv[1]*b1 + xv[2]*b2 + xv[3]*b3
+ xv[4]*b4 + xv[5]*b5 + xv[6]*b6 + xv[7]*b7;
// Transposed: scale_w[g * N + n_idx] (coalesced for adjacent n)
float sw = scale_w[g * N + (out_row + r)];
result[r] += sw * dot;
}
// Advance pointers
xp += BLOCK_K;
for (int r = 0; r < RESULTS_PER_SG; r++) {
wrows[r] += BLOCK_K;
}
}
// Reduce across simdgroup
for (int r = 0; r < RESULTS_PER_SG; r++) {
result[r] = simd_sum(result[r]);
}
// Write result
if (lid == 0) {
for (int r = 0; r < RESULTS_PER_SG; r++) {
uint n_idx = out_row + r;
if (n_idx < N) {
y[n_idx] = half(result[r] + float(bias[n_idx]));
}
}
}
}
// ============================================================
// Entry points
// ============================================================
kernel void pergroup_int8_mv_g64(
const device half *x [[buffer(0)]],
const device int8_t *W [[buffer(1)]],
device half *y [[buffer(2)]],
const device float *scale_w [[buffer(3)]],
constant uint &N [[buffer(4)]],
constant uint &K [[buffer(5)]],
const device half *bias [[buffer(6)]],
const device float *new_bias [[buffer(7)]],
uint tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]])
{
pergroup_mv_v5_impl<64>(x, W, y, scale_w, new_bias, N, K, bias, tgid, sgid, lid);
}
kernel void pergroup_int8_mv_g128(
const device half *x [[buffer(0)]],
const device int8_t *W [[buffer(1)]],
device half *y [[buffer(2)]],
const device float *scale_w [[buffer(3)]],
constant uint &N [[buffer(4)]],
constant uint &K [[buffer(5)]],
const device half *bias [[buffer(6)]],
const device float *new_bias [[buffer(7)]],
uint tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]])
{
pergroup_mv_v5_impl<128>(x, W, y, scale_w, new_bias, N, K, bias, tgid, sgid, lid);
}
kernel void pergroup_int8_mv_g256(
const device half *x [[buffer(0)]],
const device int8_t *W [[buffer(1)]],
device half *y [[buffer(2)]],
const device float *scale_w [[buffer(3)]],
constant uint &N [[buffer(4)]],
constant uint &K [[buffer(5)]],
const device half *bias [[buffer(6)]],
const device float *new_bias [[buffer(7)]],
uint tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]])
{
pergroup_mv_v5_impl<256>(x, W, y, scale_w, new_bias, N, K, bias, tgid, sgid, lid);
}
+288
View File
@@ -0,0 +1,288 @@
// ============================================================
// W4A8 INT4-weight × INT8-activation → FP16 TensorOps GEMM
// Target: Apple M5, Metal 4
//
// V3: Optimized inline unpack with precomputed base pointers.
// Key insight: for fragment's 8 elements (2 rows × 4 cols),
// the 2 rows are at k and k+8. For packed [K/2, N]:
// - Row 0 (k): byte at (k/2)*N + n, use k&1 to select nibble
// - Row 1 (k+8): byte at ((k+8)/2)*N + n, use (k+8)&1 to select nibble
// Since k+8 has same parity as k (8 is even), both rows use same nibble.
// => Can share nibble selection logic.
//
// Further: read 4 consecutive bytes at once per row (cols are contiguous
// in N dimension), then extract 4 nibbles. This maximizes memory bandwidth.
// ============================================================
#include <MetalPerformancePrimitives/MetalPerformancePrimitives.h>
#include <metal_stdlib>
using namespace metal;
constant constexpr short kElemsPerFrag = 8;
constant constexpr short kElemCols = 4;
constant constexpr short kElemRowsJump = 8;
inline short2 nax_get_coord(ushort lid) {
short qid = short(lid >> 2);
short fm = ((qid & 4) | ((short(lid) >> 1) & 3));
short fn = ((qid & 2) | (short(lid) & 1)) * 4;
return short2{fn, fm};
}
// ── Fragment load A from device memory ──────────────────────────
inline void frag_load_a(thread int8_t *dst, const device int8_t *src, int ld,
short2 sc, short off_m, short off_n) {
src += (sc.y + off_m) * ld + (sc.x + off_n);
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
dst[i * kElemCols + j] = src[(i * kElemRowsJump) * ld + j];
}
}
}
// ── Fragment load B with optimized W4 unpack ────────────────────
// Fragment maps to 2 rows (k, k+8) × 4 cols (n, n+1, n+2, n+3).
// packed_w layout: [K/2, N] uint8, high nibble = even k, low nibble = odd k.
// Pre-compute row pointers and read 4 consecutive bytes per row.
inline void frag_load_b_w4(thread int8_t *dst, const device uint8_t *packed_w,
uint N, uint k_base, uint n_base, short2 sc) {
uint k0 = k_base + uint(sc.y); // first row
uint k1 = k0 + kElemRowsJump; // second row (k+8)
uint n = n_base + uint(sc.x); // column start
// Both k0 and k1 have same parity (differ by 8)
bool use_low = (k0 & 1u);
const device uint8_t *row0 = packed_w + (k0 >> 1) * N + n;
const device uint8_t *row1 = packed_w + (k1 >> 1) * N + n;
if (use_low) {
for (short j = 0; j < kElemCols; j++) {
dst[j] = int8_t(row0[j] & 0xF) - 8;
dst[kElemCols + j] = int8_t(row1[j] & 0xF) - 8;
}
} else {
for (short j = 0; j < kElemCols; j++) {
dst[j] = int8_t(row0[j] >> 4) - 8;
dst[kElemCols + j] = int8_t(row1[j] >> 4) - 8;
}
}
}
// ── Fragment store with fused dequant ───────────────────────────
inline void nax_frag_store_dequant(const thread int32_t *src, device half *dst,
int ld, short2 sc, short off_m, short off_n,
uint M, uint N, uint m_base, uint n_base,
const device float *scale_a,
const device float *scale_w) {
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
uint mi = m_base + sc.y + off_m + i * kElemRowsJump;
uint ni = n_base + sc.x + off_n + j;
if (mi < M && ni < N) {
float val = float(src[i * kElemCols + j]) * scale_a[mi] * scale_w[ni];
dst[(sc.y + off_m + i * kElemRowsJump) * ld + (sc.x + off_n + j)] =
half(val);
}
}
}
}
// ── W4A8 GEMM V3 ───────────────────────────────────────────────
template <int BM, int BN, int BK, int SK, int WM, int WN>
void w4a8_gemm_impl(const device int8_t *A, const device uint8_t *packed_w,
device half *C, uint M, uint N, uint K,
const device float *scale_a, const device float *scale_w,
uint swizzle_log, uint tiles_m, uint tiles_n, uint2 tgid,
uint sgid, uint lid) {
constexpr int SM = BM / WM;
constexpr int SN = BN / WN;
constexpr short TM = SM / 16;
constexpr short TN = SN / 16;
constexpr short TK = SK / 16;
uint tid_y = (tgid.y << swizzle_log) + (tgid.x & ((1u << swizzle_log) - 1u));
uint tid_x = tgid.x >> swizzle_log;
if (tid_x >= tiles_n || tid_y >= tiles_m) {
return;
}
short2 sc = nax_get_coord(ushort(lid));
uint sg_row = sgid / WN;
uint sg_col = sgid % WN;
uint m_base = tid_y * BM + sg_row * SM;
uint n_base = tid_x * BN + sg_col * SN;
const device int8_t *sg_A = A + m_base * K;
constexpr auto matmul_desc = mpp::tensor_ops::matmul2d_descriptor(
16, 32, 16, false, false, true,
mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate);
mpp::tensor_ops::matmul2d<matmul_desc, metal::execution_simdgroup> gemm_op;
auto ct_a =
gemm_op.get_left_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_b =
gemm_op.get_right_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_c =
gemm_op.get_destination_cooperative_tensor<decltype(ct_a), decltype(ct_b),
int32_t>();
int32_t c_frags[TM * TN][kElemsPerFrag];
for (int f = 0; f < TM * TN; f++) {
for (int e = 0; e < kElemsPerFrag; e++) {
c_frags[f][e] = 0;
}
}
int gemm_k_iters = int(K) / BK;
for (int kk0 = 0; kk0 < gemm_k_iters; kk0++) {
uint k_offset = uint(kk0) * BK;
for (int kk1 = 0; kk1 < BK; kk1 += SK) {
int8_t a_frags[TM][TK][kElemsPerFrag];
int8_t b_frags[TK][TN][kElemsPerFrag];
volatile int compiler_barrier;
for (short mm = 0; mm < TM; mm++) {
for (short kk = 0; kk < TK; kk++) {
frag_load_a(a_frags[mm][kk], sg_A + k_offset + kk1, int(K), sc,
short(mm * 16), short(kk * 16));
}
}
for (short kk = 0; kk < TK; kk++) {
for (short nn = 0; nn < TN; nn++) {
frag_load_b_w4(b_frags[kk][nn], packed_w, N,
k_offset + uint(kk1) + uint(kk * 16),
n_base + uint(nn * 16), sc);
}
}
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn += 2) {
for (short kk = 0; kk < TK; kk++) {
for (short i = 0; i < kElemsPerFrag; i++) {
ct_a[i] = a_frags[mm][kk][i];
}
for (short i = 0; i < kElemsPerFrag; i++) {
ct_b[i] = b_frags[kk][nn][i];
ct_b[kElemsPerFrag + i] = b_frags[kk][nn + 1][i];
}
short c0 = mm * TN + nn, c1 = c0 + 1;
for (short i = 0; i < kElemsPerFrag; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kElemsPerFrag + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kElemsPerFrag; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kElemsPerFrag + i];
}
}
}
}
(void)compiler_barrier;
}
}
// Remainder
int rem_k = int(K) - gemm_k_iters * BK;
for (int kk1 = 0; kk1 < rem_k; kk1 += 16) {
int8_t a_frag[TM][kElemsPerFrag];
int8_t b_frag[TN][kElemsPerFrag];
short psk = short(max(0, rem_k - kk1));
uint k_abs = uint(gemm_k_iters * BK + kk1);
for (short mm = 0; mm < TM; mm++) {
const device int8_t *ptr = sg_A + k_abs + (sc.y + mm * 16) * K + sc.x;
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
short ki = short(sc.x + j);
a_frag[mm][i * kElemCols + j] =
(ki < psk) ? ptr[(i * kElemRowsJump) * K + j] : int8_t(0);
}
}
}
for (short nn = 0; nn < TN; nn++) {
uint n = n_base + uint(nn * 16);
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
short ki = short(sc.y + i * kElemRowsJump);
if (ki < psk) {
uint k = k_abs + uint(ki);
uint ni = n + uint(sc.x) + uint(j);
uint byte_row = k >> 1;
uint8_t packed = packed_w[byte_row * N + ni];
uint8_t nibble = (k & 1) == 0 ? (packed >> 4) : (packed & 0xF);
b_frag[nn][i * kElemCols + j] = int8_t(nibble) - 8;
} else {
b_frag[nn][i * kElemCols + j] = 0;
}
}
}
}
for (short mm = 0; mm < TM; mm++) {
for (short i = 0; i < kElemsPerFrag; i++) {
ct_a[i] = a_frag[mm][i];
}
for (short i = 0; i < kElemsPerFrag; i++) {
ct_b[i] = b_frag[0][i];
ct_b[kElemsPerFrag + i] = b_frag[1][i];
}
short c0 = mm * TN, c1 = c0 + 1;
for (short i = 0; i < kElemsPerFrag; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kElemsPerFrag + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kElemsPerFrag; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kElemsPerFrag + i];
}
}
}
device half *D = C + m_base * N + n_base;
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn++) {
nax_frag_store_dequant(c_frags[mm * TN + nn], D, int(N), sc,
short(mm * 16), short(nn * 16), M, N, m_base,
n_base, scale_a, scale_w);
}
}
}
kernel void w4a8_matmul_fused_dequant(
const device int8_t *A [[buffer(0)]],
const device uint8_t *packed_w [[buffer(1)]], device half *C [[buffer(2)]],
constant uint &M [[buffer(3)]], constant uint &N [[buffer(4)]],
constant uint &K [[buffer(5)]], const device float *scale_a [[buffer(6)]],
const device float *scale_w [[buffer(7)]],
constant uint &swizzle_log [[buffer(8)]],
constant uint &tiles_m [[buffer(9)]], constant uint &tiles_n [[buffer(10)]],
uint2 tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]]) {
w4a8_gemm_impl<128, 128, 512, 32, 4, 4>(A, packed_w, C, M, N, K, scale_a,
scale_w, swizzle_log, tiles_m,
tiles_n, tgid, sgid, lid);
}
kernel void w4a8_matmul_fused_dequant_small(
const device int8_t *A [[buffer(0)]],
const device uint8_t *packed_w [[buffer(1)]], device half *C [[buffer(2)]],
constant uint &M [[buffer(3)]], constant uint &N [[buffer(4)]],
constant uint &K [[buffer(5)]], const device float *scale_a [[buffer(6)]],
const device float *scale_w [[buffer(7)]],
constant uint &swizzle_log [[buffer(8)]],
constant uint &tiles_m [[buffer(9)]], constant uint &tiles_n [[buffer(10)]],
uint2 tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]]) {
w4a8_gemm_impl<32, 128, 512, 32, 1, 4>(A, packed_w, C, M, N, K, scale_a,
scale_w, swizzle_log, tiles_m, tiles_n,
tgid, sgid, lid);
}
+71
View File
@@ -0,0 +1,71 @@
// w8a8_int8_mv.metal — Per-channel symmetric INT8 MV (optimized)
// Matches MLX qmv_fast structure: 2 SG, VPT=8, block=256
// y[n] = scale_w[n] * dot(W_int8[n,:], x[:]) + bias[n]
#include <metal_stdlib>
using namespace metal;
constant constexpr int SIMD_SIZE = 32;
constant constexpr int NUM_SIMDGROUPS = 2;
constant constexpr int RESULTS_PER_SG = 4;
constant constexpr int VPT = 8;
constant constexpr int BLOCK_K = VPT * SIMD_SIZE; // 256
kernel void w8a8_int8_mv(
const device half *x [[buffer(0)]],
const device int8_t *W [[buffer(1)]],
device half *y [[buffer(2)]],
const device float *scale_w [[buffer(3)]],
constant uint &N [[buffer(4)]],
constant uint &K [[buffer(5)]],
const device half *bias [[buffer(6)]],
uint tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]])
{
const uint rows_per_tg = NUM_SIMDGROUPS * RESULTS_PER_SG; // 8
const uint out_row = tgid * rows_per_tg + sgid * RESULTS_PER_SG;
float result[RESULTS_PER_SG] = {0.0f, 0.0f, 0.0f, 0.0f};
// Pointers: each thread handles VPT consecutive elements per block
const device half *xp = x + lid * VPT;
for (uint k = 0; k < K; k += BLOCK_K) {
// Load x tile into registers
float xv[VPT];
for (int i = 0; i < VPT; i++) {
xv[i] = float(xp[i]);
}
// Dot product with weight rows
for (int r = 0; r < RESULTS_PER_SG; r++) {
uint n_idx = out_row + r;
if (n_idx < N) {
const device int8_t *wp = W + n_idx * K + k + lid * VPT;
float dot = 0.0f;
for (int i = 0; i < VPT; i++) {
dot += float(wp[i]) * xv[i];
}
result[r] += dot;
}
}
xp += BLOCK_K;
}
// Reduce across simdgroup
for (int r = 0; r < RESULTS_PER_SG; r++) {
result[r] = simd_sum(result[r]);
}
// Write back: y = scale * dot + bias
if (lid == 0) {
for (int r = 0; r < RESULTS_PER_SG; r++) {
uint n_idx = out_row + r;
if (n_idx < N) {
float val = scale_w[n_idx] * result[r] + float(bias[n_idx]);
y[n_idx] = half(val);
}
}
}
}
+483
View File
@@ -0,0 +1,483 @@
// ============================================================
// W8A8 INT8×INT8→INT32 TensorOps GEMM
// Target: Apple M5 (G17G), Metal 4
//
// Weight layout: B is [N, K] (row-major), transpose_b=true
// Computes: C[M,N] = A[M,K] × B[N,K]^T
//
// Variants:
// - fused dequant: INT8×INT8→FP16, with per-token/per-channel scales
// - raw INT32: INT8×INT8→INT32, no scale (pure integer GEMM)
// Multi-config: large (BM=128) and small (BM=32) tiles
// Swizzle dispatch for L2 cache locality
//
// matmul2d(16,32,16) via MPP cooperative_tensor
// ============================================================
#include <MetalPerformancePrimitives/MetalPerformancePrimitives.h>
#include <metal_stdlib>
using namespace metal;
// ── NAXFrag layout constants ────────────────────────────────────
constant constexpr short kElemsPerFrag = 8;
constant constexpr short kElemCols = 4;
constant constexpr short kElemRowsJump = 8;
// ── NAXFrag coordinate mapping ──────────────────────────────────
inline short2 nax_get_coord(ushort lid) {
short qid = short(lid >> 2);
short fm = ((qid & 4) | ((short(lid) >> 1) & 3));
short fn = ((qid & 2) | (short(lid) & 1)) * 4;
return short2{fn, fm};
}
// ── Fragment load: device → register ────────────────────────────
template <typename T>
inline void nax_frag_load(thread T *dst, const device T *src, int ld, short2 sc,
short off_m = 0, short off_n = 0) {
src += (sc.y + off_m) * ld + (sc.x + off_n);
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
dst[i * kElemCols + j] = src[(i * kElemRowsJump) * ld + j];
}
}
}
// ── Fragment store: raw INT32 (no dequant) ───────────────────
inline void nax_frag_store_int32(const thread int32_t *src, device int32_t *dst,
int ld, short2 sc, short off_m, short off_n,
uint M, uint N, uint m_base, uint n_base) {
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
uint mi = m_base + sc.y + off_m + i * kElemRowsJump;
uint ni = n_base + sc.x + off_n + j;
if (mi < M && ni < N) {
dst[(sc.y + off_m + i * kElemRowsJump) * ld + (sc.x + off_n + j)] =
src[i * kElemCols + j];
}
}
}
}
// ── Fragment store with bounds check and dequant ────────────────
inline void nax_frag_store_dequant(const thread int32_t *src, device half *dst,
int ld, short2 sc, short off_m, short off_n,
uint M, uint N, uint m_base, uint n_base,
const device float *scale_a,
const device float *scale_w,
const device half *bias) {
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
uint mi = m_base + sc.y + off_m + i * kElemRowsJump;
uint ni = n_base + sc.x + off_n + j;
if (mi < M && ni < N) {
float val = float(src[i * kElemCols + j]) * scale_a[mi] * scale_w[ni] +
float(bias[ni]);
dst[(sc.y + off_m + i * kElemRowsJump) * ld + (sc.x + off_n + j)] =
half(val);
}
}
}
}
// ── Generic GEMM kernel (B is [N, K], transpose_b) ─────────────
// Computes: C[M,N] = A[M,K] × B[N,K]^T
// A is [M, K] row-major, B is [N, K] row-major
// B fragments loaded as [N_tile, K_tile] and hardware transposes via
// transpose_b=true
template <int BM, int BN, int BK, int SK, int WM, int WN>
void w8a8_gemm_impl(const device int8_t *A, const device int8_t *B,
device half *C, uint M, uint N, uint K,
const device float *scale_a, const device float *scale_w,
const device half *bias, uint swizzle_log, uint tiles_m,
uint tiles_n, uint2 tgid, uint sgid, uint lid) {
constexpr int SM = BM / WM; // 32
constexpr int SN = BN / WN; // 32
constexpr short TM = SM / 16; // 2
constexpr short TN = SN / 16; // 2
constexpr short TK = SK / 16; // 2
// Swizzle decode
uint tid_y = (tgid.y << swizzle_log) + (tgid.x & ((1u << swizzle_log) - 1u));
uint tid_x = tgid.x >> swizzle_log;
if (tid_x >= tiles_n || tid_y >= tiles_m) {
return;
}
short2 sc = nax_get_coord(ushort(lid));
uint sg_row = sgid / WN;
uint sg_col = sgid % WN;
uint m_base = tid_y * BM + sg_row * SM;
uint n_base = tid_x * BN + sg_col * SN;
// A: [M, K] row-major — same as before
const device int8_t *sg_A = A + m_base * K;
// B: [N, K] row-major — pointer to start of n_base-th row
const device int8_t *sg_B = B + n_base * K;
// transpose_b=true: right operand is [N_frag, K_frag], hardware transposes
constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor(
16, 32, 16, false, true, true,
mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate);
mpp::tensor_ops::matmul2d<desc, metal::execution_simdgroup> gemm_op;
auto ct_a =
gemm_op.get_left_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_b =
gemm_op.get_right_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_c =
gemm_op.get_destination_cooperative_tensor<decltype(ct_a), decltype(ct_b),
int32_t>();
int32_t c_frags[TM * TN][kElemsPerFrag];
for (int f = 0; f < TM * TN; f++) {
for (int i = 0; i < kElemsPerFrag; i++) {
c_frags[f][i] = 0;
}
}
// ── Main K loop ─────────────────────────────────────────────
int gemm_k_iters = int(K) / BK;
for (int kk0 = 0; kk0 < gemm_k_iters; kk0++) {
threadgroup_barrier(mem_flags::mem_none);
for (int kk1 = 0; kk1 < BK; kk1 += SK) {
int8_t a_frags[TM][TK][kElemsPerFrag];
int8_t b_frags[TN][TK][kElemsPerFrag]; // [N_tile, K_tile] for transpose_b
volatile int compiler_barrier;
// Load A fragments: [M_tile, K_tile], ld=K
for (short mm = 0; mm < TM; mm++) {
for (short kk = 0; kk < TK; kk++) {
nax_frag_load(a_frags[mm][kk], sg_A + kk1, int(K), sc, short(mm * 16),
short(kk * 16));
}
}
// Load B fragments: [N_tile, K_tile], ld=K (B is [N, K] row-major)
for (short nn = 0; nn < TN; nn++) {
for (short kk = 0; kk < TK; kk++) {
nax_frag_load(b_frags[nn][kk], sg_B + kk1, int(K), sc, short(nn * 16),
short(kk * 16));
}
}
// Compute: ct_a=[M_frag, K_frag], ct_b=[N_frag, K_frag] (transpose_b)
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn += 2) {
for (short kk = 0; kk < TK; kk++) {
for (short i = 0; i < kElemsPerFrag; i++) {
ct_a[i] = a_frags[mm][kk][i];
}
for (short i = 0; i < kElemsPerFrag; i++) {
ct_b[i] = b_frags[nn][kk][i];
ct_b[kElemsPerFrag + i] = b_frags[nn + 1][kk][i];
}
short c0 = mm * TN + nn, c1 = c0 + 1;
for (short i = 0; i < kElemsPerFrag; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kElemsPerFrag + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kElemsPerFrag; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kElemsPerFrag + i];
}
}
}
}
(void)compiler_barrier;
}
sg_A += BK;
sg_B += BK; // B is [N, K]: K advances by BK along columns
}
// ── Remainder K ─────────────────────────────────────────────
int rem_k = int(K) - gemm_k_iters * BK;
for (int kk1 = 0; kk1 < rem_k; kk1 += 16) {
int8_t a_frag[TM][kElemsPerFrag];
int8_t b_frag[TN][kElemsPerFrag];
short psk = short(max(0, rem_k - kk1));
// Load A remainder: same as before
for (short mm = 0; mm < TM; mm++) {
const device int8_t *ptr = sg_A + kk1 + (sc.y + mm * 16) * K + sc.x;
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
short ki = short(sc.x + j);
a_frag[mm][i * kElemCols + j] =
(ki < psk) ? ptr[(i * kElemRowsJump) * K + j] : int8_t(0);
}
}
}
// Load B remainder: B is [N, K], reading [N_tile, K_rem]
for (short nn = 0; nn < TN; nn++) {
const device int8_t *ptr = sg_B + kk1 + (sc.y + nn * 16) * K + sc.x;
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
short ki = short(sc.x + j);
b_frag[nn][i * kElemCols + j] =
(ki < psk) ? ptr[(i * kElemRowsJump) * K + j] : int8_t(0);
}
}
}
for (short mm = 0; mm < TM; mm++) {
for (short i = 0; i < kElemsPerFrag; i++) {
ct_a[i] = a_frag[mm][i];
}
for (short i = 0; i < kElemsPerFrag; i++) {
ct_b[i] = b_frag[0][i];
ct_b[kElemsPerFrag + i] = b_frag[1][i];
}
short c0 = mm * TN, c1 = c0 + 1;
for (short i = 0; i < kElemsPerFrag; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kElemsPerFrag + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kElemsPerFrag; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kElemsPerFrag + i];
}
}
}
// ── Store with fused dequant ────────────────────────────────
device half *D = C + m_base * N + n_base;
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn++) {
nax_frag_store_dequant(c_frags[mm * TN + nn], D, int(N), sc,
short(mm * 16), short(nn * 16), M, N, m_base,
n_base, scale_a, scale_w, bias);
}
}
}
// ============================================================
// Kernel entry points — fused dequant
// ============================================================
kernel void w8a8_matmul_fused_dequant(
const device int8_t *A [[buffer(0)]], const device int8_t *B [[buffer(1)]],
device half *C [[buffer(2)]], constant uint &M [[buffer(3)]],
constant uint &N [[buffer(4)]], constant uint &K [[buffer(5)]],
const device float *scale_a [[buffer(6)]],
const device float *scale_w [[buffer(7)]],
constant uint &swizzle_log [[buffer(8)]],
constant uint &tiles_m [[buffer(9)]], constant uint &tiles_n [[buffer(10)]],
const device half *bias [[buffer(11)]],
uint2 tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]]) {
w8a8_gemm_impl<128, 128, 512, 32, 4, 4>(A, B, C, M, N, K, scale_a, scale_w,
bias, swizzle_log, tiles_m, tiles_n,
tgid, sgid, lid);
}
kernel void w8a8_matmul_fused_dequant_small(
const device int8_t *A [[buffer(0)]], const device int8_t *B [[buffer(1)]],
device half *C [[buffer(2)]], constant uint &M [[buffer(3)]],
constant uint &N [[buffer(4)]], constant uint &K [[buffer(5)]],
const device float *scale_a [[buffer(6)]],
const device float *scale_w [[buffer(7)]],
constant uint &swizzle_log [[buffer(8)]],
constant uint &tiles_m [[buffer(9)]], constant uint &tiles_n [[buffer(10)]],
const device half *bias [[buffer(11)]],
uint2 tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]]) {
w8a8_gemm_impl<32, 128, 512, 32, 1, 4>(A, B, C, M, N, K, scale_a, scale_w,
bias, swizzle_log, tiles_m, tiles_n,
tgid, sgid, lid);
}
// ============================================================
// Raw INT32 GEMM impl (B is [N, K], transpose_b=true)
// ============================================================
template <int BM, int BN, int BK, int SK, int WM, int WN>
void w8a8_gemm_int32_impl(const device int8_t *A, const device int8_t *B,
device int32_t *C, uint M, uint N, uint K,
uint swizzle_log, uint tiles_m, uint tiles_n,
uint2 tgid, uint sgid, uint lid) {
constexpr int SM = BM / WM;
constexpr int SN = BN / WN;
constexpr short TM = SM / 16;
constexpr short TN = SN / 16;
constexpr short TK = SK / 16;
uint tid_y = (tgid.y << swizzle_log) + (tgid.x & ((1u << swizzle_log) - 1u));
uint tid_x = tgid.x >> swizzle_log;
if (tid_x >= tiles_n || tid_y >= tiles_m) {
return;
}
short2 sc = nax_get_coord(ushort(lid));
uint sg_row = sgid / WN;
uint sg_col = sgid % WN;
uint m_base = tid_y * BM + sg_row * SM;
uint n_base = tid_x * BN + sg_col * SN;
const device int8_t *sg_A = A + m_base * K;
const device int8_t *sg_B = B + n_base * K;
constexpr auto desc = mpp::tensor_ops::matmul2d_descriptor(
16, 32, 16, false, true, true,
mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate);
mpp::tensor_ops::matmul2d<desc, metal::execution_simdgroup> gemm_op;
auto ct_a =
gemm_op.get_left_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_b =
gemm_op.get_right_input_cooperative_tensor<int8_t, int8_t, int32_t>();
auto ct_c =
gemm_op.get_destination_cooperative_tensor<decltype(ct_a), decltype(ct_b),
int32_t>();
int32_t c_frags[TM * TN][kElemsPerFrag];
for (int f = 0; f < TM * TN; f++) {
for (int i = 0; i < kElemsPerFrag; i++) {
c_frags[f][i] = 0;
}
}
int gemm_k_iters = int(K) / BK;
for (int kk0 = 0; kk0 < gemm_k_iters; kk0++) {
threadgroup_barrier(mem_flags::mem_none);
for (int kk1 = 0; kk1 < BK; kk1 += SK) {
int8_t a_frags[TM][TK][kElemsPerFrag];
int8_t b_frags[TN][TK][kElemsPerFrag];
volatile int compiler_barrier;
for (short mm = 0; mm < TM; mm++) {
for (short kk = 0; kk < TK; kk++) {
nax_frag_load(a_frags[mm][kk], sg_A + kk1, int(K), sc, short(mm * 16),
short(kk * 16));
}
}
for (short nn = 0; nn < TN; nn++) {
for (short kk = 0; kk < TK; kk++) {
nax_frag_load(b_frags[nn][kk], sg_B + kk1, int(K), sc, short(nn * 16),
short(kk * 16));
}
}
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn += 2) {
for (short kk = 0; kk < TK; kk++) {
for (short i = 0; i < kElemsPerFrag; i++) {
ct_a[i] = a_frags[mm][kk][i];
}
for (short i = 0; i < kElemsPerFrag; i++) {
ct_b[i] = b_frags[nn][kk][i];
ct_b[kElemsPerFrag + i] = b_frags[nn + 1][kk][i];
}
short c0 = mm * TN + nn, c1 = c0 + 1;
for (short i = 0; i < kElemsPerFrag; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kElemsPerFrag + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kElemsPerFrag; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kElemsPerFrag + i];
}
}
}
}
(void)compiler_barrier;
}
sg_A += BK;
sg_B += BK;
}
// Remainder K
int rem_k = int(K) - gemm_k_iters * BK;
for (int kk1 = 0; kk1 < rem_k; kk1 += 16) {
int8_t a_frag[TM][kElemsPerFrag];
int8_t b_frag[TN][kElemsPerFrag];
short psk = short(max(0, rem_k - kk1));
for (short mm = 0; mm < TM; mm++) {
const device int8_t *ptr = sg_A + kk1 + (sc.y + mm * 16) * K + sc.x;
for (short i = 0; i < 2; i++)
for (short j = 0; j < kElemCols; j++) {
short ki = short(sc.x + j);
a_frag[mm][i * kElemCols + j] =
(ki < psk) ? ptr[(i * kElemRowsJump) * K + j] : int8_t(0);
}
}
for (short nn = 0; nn < TN; nn++) {
const device int8_t *ptr = sg_B + kk1 + (sc.y + nn * 16) * K + sc.x;
for (short i = 0; i < 2; i++) {
for (short j = 0; j < kElemCols; j++) {
short ki = short(sc.x + j);
b_frag[nn][i * kElemCols + j] =
(ki < psk) ? ptr[(i * kElemRowsJump) * K + j] : int8_t(0);
}
}
}
for (short mm = 0; mm < TM; mm++) {
for (short i = 0; i < kElemsPerFrag; i++) {
ct_a[i] = a_frag[mm][i];
}
for (short i = 0; i < kElemsPerFrag; i++) {
ct_b[i] = b_frag[0][i];
ct_b[kElemsPerFrag + i] = b_frag[1][i];
}
short c0 = mm * TN, c1 = c0 + 1;
for (short i = 0; i < kElemsPerFrag; i++) {
ct_c[i] = c_frags[c0][i];
ct_c[kElemsPerFrag + i] = c_frags[c1][i];
}
gemm_op.run(ct_a, ct_b, ct_c);
for (short i = 0; i < kElemsPerFrag; i++) {
c_frags[c0][i] = ct_c[i];
c_frags[c1][i] = ct_c[kElemsPerFrag + i];
}
}
}
// Store raw INT32
device int32_t *D = C + m_base * N + n_base;
for (short mm = 0; mm < TM; mm++) {
for (short nn = 0; nn < TN; nn++) {
nax_frag_store_int32(c_frags[mm * TN + nn], D, int(N), sc, short(mm * 16),
short(nn * 16), M, N, m_base, n_base);
}
}
}
// ============================================================
// Kernel entry points — raw INT32 output
// ============================================================
kernel void int8_matmul_int32(
const device int8_t *A [[buffer(0)]], const device int8_t *B [[buffer(1)]],
device int32_t *C [[buffer(2)]], constant uint &M [[buffer(3)]],
constant uint &N [[buffer(4)]], constant uint &K [[buffer(5)]],
constant uint &swizzle_log [[buffer(6)]],
constant uint &tiles_m [[buffer(7)]], constant uint &tiles_n [[buffer(8)]],
uint2 tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]]) {
w8a8_gemm_int32_impl<128, 128, 512, 32, 4, 4>(
A, B, C, M, N, K, swizzle_log, tiles_m, tiles_n, tgid, sgid, lid);
}
kernel void int8_matmul_int32_small(
const device int8_t *A [[buffer(0)]], const device int8_t *B [[buffer(1)]],
device int32_t *C [[buffer(2)]], constant uint &M [[buffer(3)]],
constant uint &N [[buffer(4)]], constant uint &K [[buffer(5)]],
constant uint &swizzle_log [[buffer(6)]],
constant uint &tiles_m [[buffer(7)]], constant uint &tiles_n [[buffer(8)]],
uint2 tgid [[threadgroup_position_in_grid]],
uint sgid [[simdgroup_index_in_threadgroup]],
uint lid [[thread_index_in_simdgroup]]) {
w8a8_gemm_int32_impl<32, 128, 512, 32, 1, 4>(
A, B, C, M, N, K, swizzle_log, tiles_m, tiles_n, tgid, sgid, lid);
}
+87
View File
@@ -0,0 +1,87 @@
// ============================================================
// Per-token quantization: FP16 → INT8 + float32 scale
// Target: Apple M5, Metal 4
//
// Each threadgroup handles one row (one token).
// Threads cooperate to find absmax via simdgroup reduce,
// then quantize in parallel.
//
// Host dispatch:
// threadgroup = (min(256, ceil(K/32)*32), 1, 1)
// grid = (M, 1, 1)
// ============================================================
#include <metal_stdlib>
using namespace metal;
kernel void
quantize_per_token(const device half *X [[buffer(0)]], // [M, K] FP16 input
device int8_t *A [[buffer(1)]], // [M, K] INT8 output
device float *scale [[buffer(2)]], // [M] float32 scale
constant uint &M [[buffer(3)]],
constant uint &K [[buffer(4)]],
uint gid [[threadgroup_position_in_grid]], // row index
uint lid [[thread_index_in_threadgroup]],
uint tg_size [[threads_per_threadgroup]]) {
if (gid >= M) {
return;
}
const device half *row_in = X + gid * K;
device int8_t *row_out = A + gid * K;
// Step 1: Find local absmax
float local_max = 0.0f;
for (uint i = lid; i < K; i += tg_size) {
float v = abs(float(row_in[i]));
local_max = max(local_max, v);
}
// Step 2: Simdgroup reduce max
float sg_max = simd_max(local_max);
// Step 3: Threadgroup reduce across simdgroups via shared memory
threadgroup float sg_maxes[8]; // up to 8 simdgroups (256/32)
threadgroup float shared_scale;
threadgroup float shared_inv_scale;
uint sg_id = lid / 32;
uint sg_lid = lid % 32;
if (sg_lid == 0) {
sg_maxes[sg_id] = sg_max;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
// Final reduce (first simdgroup only)
if (sg_id == 0) {
float row_max = 0.0f;
uint num_sgs = (tg_size + 31) / 32;
if (sg_lid < num_sgs) {
row_max = sg_maxes[sg_lid];
}
row_max = simd_max(row_max);
// Compute and broadcast scale
float s = row_max / 127.0f;
if (s == 0.0f) {
s = 1.0f;
}
if (sg_lid == 0) {
shared_scale = s;
shared_inv_scale = 1.0f / s;
// Store scale to output
scale[gid] = s;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
// Step 4: All threads read broadcasted scale
float inv_s = shared_inv_scale;
// Step 5: Quantize
for (uint i = lid; i < K; i += tg_size) {
float v = float(row_in[i]) * inv_s;
v = clamp(round(v), -128.0f, 127.0f);
row_out[i] = int8_t(v);
}
}