chore: import upstream snapshot with attribution
This commit is contained in:
@@ -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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user