Support DeepSeek V4 Flash 0731

This commit is contained in:
Georg Bauer
2026-08-29 20:28:50 +02:00
parent ad855b321e
commit f1c177b754
23 changed files with 10510 additions and 2324 deletions

View File

@@ -197,33 +197,6 @@ kernel void kernel_mul_mv_q8_0_f32(
kernel_mul_mv_q8_0_f32_impl<N_R0_Q8_0, constant ds4_metal_args_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
[[host_name("kernel_mul_mv_q8_0_f32_r4")]]
kernel void kernel_mul_mv_q8_0_f32_r4(
constant ds4_metal_args_mul_mv & args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_q8_0_f32_impl<4, constant ds4_metal_args_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
// Output projection alias used by the optimized host dispatch.
[[host_name("kernel_mul_mv_q8_0_f32_nr4")]]
kernel void kernel_mul_mv_q8_0_f32_nr4(
constant ds4_metal_args_mul_mv & args,
device const char * src0,
device const char * src1,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_mul_mv_q8_0_f32_impl<4, constant ds4_metal_args_mul_mv &>(
args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
}
// Decode Q-A/KV pair. Both projections consume the same activation row but
// have independent weight ranges and output extents. Keep the standalone Q8_0
@@ -497,25 +470,190 @@ kernel void kernel_dsv4_shared_gate_up_swiglu_q8_0(
clamp_value, shmem, tgpig, tiisg, sgitg);
}
[[host_name("kernel_dsv4_shared_gate_up_swiglu_q8_0_r4")]]
kernel void kernel_dsv4_shared_gate_up_swiglu_q8_0_r4(
// Decode-only fusion of the router logits matvec (F16, embd -> n_expert)
// with the shared-expert gate/up SwiGLU (Q8_0, embd -> shared). Both read
// the same normalized FFN input back to back; one dispatch removes one
// launch per decode layer. Router threadgroups replicate
// kernel_mul_mv_f16_f32_4 (nsg=8, nr0=2); shared threadgroups host two
// virtual 4-simdgroup cohorts replicating
// kernel_dsv4_shared_gate_up_swiglu_q8_0 (nsg=4, nr0=2), including its
// per-row simd/shmem reduction trees. Bit-exact by construction.
kernel void kernel_dsv4_router_shared_gate_up_q8_0(
constant ds4_metal_args_mul_mv & args,
constant ds4_metal_args_mul_mv & sargs,
device const char * src0_router,
device const char * src0_gate,
device const char * src0_up,
device const char * src1,
device char * dst_router,
device char * dst_gate,
device char * dst_up,
device char * dst_mid,
constant float &clamp_value,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_dsv4_shared_gate_up_swiglu_q8_0_impl<4, true>(
args, src0_gate, src0_up, src1, dst_gate, dst_up, dst_mid,
clamp_value, shmem, tgpig, tiisg, sgitg);
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NW = N_SIMDWIDTH;
const uint router_tgs = ((uint)args.ne01 + 1u) / 2u;
if (tgpig.x < router_tgs) {
// Exact replica of kernel_mul_mv_f16_f32_4 with NSG=8, NR0=2.
constexpr short NSG = 8;
constexpr short NR0 = 2;
constexpr short NB = 32;
constexpr short NF = 16;
constexpr short NF4 = NF/4;
const int nb = args.ne00/NB;
const int r0 = tgpig.x*NR0;
device const float4 * y4 = (device const float4 *) src1;
device const half4 * ax4[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
ax4[row] = (device const half4 *)
(src0_router + (uint64_t)(r0 + row)*args.nb01);
}
float sumf[NR0] = { 0.f };
const short ix = tiisg/(NW/NF);
const short il = tiisg%(NW/NF);
const int ib0 = sgitg*NF + ix;
device const float4 * yb4 = y4 + (ib0*NB + il*NF)/4;
for (int ib = ib0; ib < nb; ib += NSG*NF) {
float4 yl4[NF4];
FOR_UNROLL (short i = 0; i < NF4; ++i) {
yl4[i] = yb4[i];
}
FOR_UNROLL (short row = 0; row < NR0; row++) {
device const half4 * xb4 = ax4[row] + (ib*NB + il*NF)/4;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NF4; ++i) {
sumq += dot(float4(xb4[i]), yl4[i]);
}
sumf[row] += sumq;
}
yb4 += NSG*NF*NW/4;
}
device float * dst_f32 = (device float *) dst_router;
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01,
tiisg, sgitg, shmem);
return;
}
// Shared-expert part: two virtual nsg=4 cohorts per threadgroup, each an
// exact replica of kernel_dsv4_shared_gate_up_swiglu_q8_0 (NR0=2).
constexpr short NSG = 4;
constexpr short NR0 = 2;
constexpr short NQ = 8;
const uint cohort = sgitg >> 2;
const ushort vsg = sgitg & 3u;
const uint vt = (tgpig.x - router_tgs) * 2u + cohort;
const int nb = sargs.ne00 / QK8_0;
const int r0 = vt * NR0;
device const float *y = (device const float *) src1;
device const block_q8_0 *ag[NR0];
device const block_q8_0 *au[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const uint64_t offset0 = (uint64_t)(r0 + row) * sargs.nb01;
ag[row] = (device const block_q8_0 *)(src0_gate + offset0);
au[row] = (device const block_q8_0 *)(src0_up + offset0);
}
float sumg[NR0] = { 0.f };
float sumu[NR0] = { 0.f };
const short ix = tiisg / (NW / NQ);
const short il = tiisg % (NW / NQ);
const int ib0 = vsg * NQ + ix;
float yl[NQ];
device const float *yb = y + ib0 * QK8_0 + il * NQ;
for (int ib = ib0; ib < nb; ib += NSG * NQ) {
FOR_UNROLL (short i = 0; i < NQ; ++i) {
yl[i] = yb[i];
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
device const int8_t *qg = ag[row][ib].qs + il * NQ;
device const int8_t *qu = au[row][ib].qs + il * NQ;
float sg = 0.f;
float su = 0.f;
FOR_UNROLL (short i = 0; i < NQ; ++i) {
sg += qg[i] * yl[i];
su += qu[i] * yl[i];
}
sumg[row] += sg * ag[row][ib].d;
sumu[row] += su * au[row][ib].d;
}
yb += NSG * NQ * QK8_0;
}
threadgroup float *shmem_f32 = (threadgroup float *)shmem + cohort * (2*NR0*NW);
threadgroup float *sh_gate[NR0];
threadgroup float *sh_up[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
sh_gate[row] = shmem_f32 + NW * row;
sh_up[row] = shmem_f32 + NW * (NR0 + row);
if (vsg == 0) {
sh_gate[row][tiisg] = 0.0f;
sh_up[row][tiisg] = 0.0f;
}
sumg[row] = simd_sum(sumg[row]);
sumu[row] = simd_sum(sumu[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
sh_gate[row][vsg] = sumg[row];
sh_up[row][vsg] = sumu[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
device float *gate_f32 = (device float *)dst_gate;
device float *up_f32 = (device float *)dst_up;
device float *mid_f32 = (device float *)dst_mid;
FOR_UNROLL (short row = 0; row < NR0 && r0 + row < sargs.ne01; ++row) {
const float gate = simd_sum(sh_gate[row][tiisg]);
const float up = simd_sum(sh_up[row][tiisg]);
if (tiisg == 0 && vsg == 0) {
const uint out_row = r0 + row;
gate_f32[out_row] = gate;
up_f32[out_row] = up;
float g = gate;
float u = up;
if (clamp_value > 1.0e-6f) {
g = min(g, clamp_value);
u = clamp(u, -clamp_value, clamp_value);
}
const float silu = g / (1.0f + exp(-g));
mid_f32[out_row] = silu * u;
}
}
}
[[host_name("kernel_dsv4_shared_mid_swiglu_q8_0")]]
kernel void kernel_dsv4_shared_mid_swiglu_q8_0(
constant ds4_metal_args_mul_mv & args,
@@ -535,24 +673,6 @@ kernel void kernel_dsv4_shared_mid_swiglu_q8_0(
clamp_value, shmem, tgpig, tiisg, sgitg);
}
[[host_name("kernel_dsv4_shared_mid_swiglu_q8_0_r4")]]
kernel void kernel_dsv4_shared_mid_swiglu_q8_0_r4(
constant ds4_metal_args_mul_mv & args,
device const char * src0_gate,
device const char * src0_up,
device const char * src1,
device char * dst_gate,
device char * dst_up,
device char * dst_mid,
constant float &clamp_value,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiisg[[thread_index_in_simdgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
kernel_dsv4_shared_gate_up_swiglu_q8_0_impl<4, false>(
args, src0_gate, src0_up, src1, dst_gate, dst_up, dst_mid,
clamp_value, shmem, tgpig, tiisg, sgitg);
}
template<typename T0, typename T1, short NR0, typename args_t>
void kernel_mul_mv_t_t_impl(
@@ -974,6 +1094,311 @@ kernel void kernel_mul_mv_f16_f32_pair_compressor_store_4(
state_score[dst] = projected_score[col] + ape_v;
}
// Decode compressor + indexer-compressor projection in one dispatch. Both
// pairs read the same normalized activation with the same F16 matvec shape,
// so one launch covers all four matrices: threadgroups below the first
// range boundary run the exact paired matvec + state store of
// kernel_mul_mv_f16_f32_pair_compressor_store_4 for the attention
// compressor, the rest for the indexer compressor. Per-row reduction trees
// and the per-threadgroup state stores are unchanged, keeping the fused
// result bit-identical to the two separate dispatches while removing one
// dispatch per decode layer.
kernel void kernel_mul_mv_f16_f32_quad_compressor_store_4(
constant ds4_metal_args_mul_mv & args,
constant ds4_metal_args_compressor_pair_store & store0,
constant ds4_metal_args_compressor_pair_store & store1,
device const char * src0_a0,
device const char * src0_b0,
device const char * src0_a1,
device const char * src0_b1,
device const char * src1,
device char * dst_a0,
device char * dst_b0,
device char * dst_a1,
device char * dst_b1,
device const char * ape0,
device const char * ape1,
device float * state0_kv,
device float * state0_score,
device float * state1_kv,
device float * state1_score,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NR0 = 2;
const uint tgs0 = ((uint)store0.width + NR0 - 1u) / NR0;
const bool second = tgpig.x >= tgs0;
uint3 local_tgpig = tgpig;
if (second) local_tgpig.x = tgpig.x - tgs0;
ds4_metal_args_mul_mv largs = args;
largs.nr0 = NR0;
largs.ne01 = second ? (int32_t)store1.width : (int32_t)store0.width;
if (!second) {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, src0_a0, src0_b0, src1, dst_a0, dst_b0,
shmem, local_tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, src0_a1, src0_b1, src1, dst_a1, dst_b1,
shmem, local_tgpig, tiisg, sgitg);
}
threadgroup_barrier(mem_flags::mem_device);
// State append: identical to the paired store kernel, scoped to the
// range this threadgroup just projected (its own outputs only).
constant ds4_metal_args_compressor_pair_store & store = second ? store1 : store0;
if (tiitg >= NR0 || store.width == 0u || store.ratio == 0u) {
return;
}
const uint col = local_tgpig.x * (uint)NR0 + tiitg;
if (col >= store.width) return;
const uint pos_mod = store.pos % store.ratio;
const uint dst_row = store.ratio == 4u ? store.ratio + pos_mod : pos_mod;
const uint dst = dst_row * store.width + col;
const uint ape_i = pos_mod * store.width + col;
device volatile const float * projected_kv = second
? (device volatile const float *)dst_a1
: (device volatile const float *)dst_a0;
device volatile const float * projected_score = second
? (device volatile const float *)dst_b1
: (device volatile const float *)dst_b0;
device const char * ape = second ? ape1 : ape0;
device float * state_kv = second ? state1_kv : state0_kv;
device float * state_score = second ? state1_score : state0_score;
float ape_v;
if (store.ape_type == 1u) {
ape_v = (float)(((device const half *)ape)[ape_i]);
} else {
ape_v = ((device const float *)ape)[ape_i];
}
state_kv[dst] = projected_kv[col];
state_score[dst] = projected_score[col] + ape_v;
}
/* Decode-only fusion: one dispatch covers the q_a/kv Q8 pair projection and
* the four F16 compressor projections (attention + indexer) with their
* state-store epilogue. Both stages read the same normalized attention
* input and write disjoint outputs. The q_a/kv range hosts two virtual
* NSG=4 cohorts per threadgroup, each an exact replica of
* kernel_mul_mv_q8_0_f32_pair (same per-lane K walk and reduction tree, cf.
* kernel_dsv4_router_shared_gate_up_q8_0); the compressor ranges run
* kernel_mul_mv_f16_f32_pair_4_impl<2> and the paired store epilogue
* verbatim, so every output bit matches the two separate dispatches. */
kernel void kernel_dsv4_qkv_pair_quad_compressor_store_q8_0(
constant ds4_metal_args_mul_mv & args0,
constant ds4_metal_args_mul_mv & args1,
constant ds4_metal_args_mul_mv & cargs,
constant ds4_metal_args_compressor_pair_store & store0,
constant ds4_metal_args_compressor_pair_store & store1,
constant uint & pair_vtgs,
device const char * qw0,
device const char * qw1,
device const char * cw0a,
device const char * cw0b,
device const char * cw1a,
device const char * cw1b,
device const char * src1,
device char * dst0,
device char * dst1,
device char * cdst_a0,
device char * cdst_b0,
device char * cdst_a1,
device char * cdst_b1,
device const char * ape0,
device const char * ape1,
device float * state0_kv,
device float * state0_score,
device float * state1_kv,
device float * state1_score,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort tiisg [[thread_index_in_simdgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
constexpr short NW = N_SIMDWIDTH;
const uint pair_ctgs = (pair_vtgs + 1u) / 2u;
if (tgpig.x < pair_ctgs) {
/* Q8 pair range: cohort c of threadgroup t runs virtual pair
* threadgroup 2t+c with the original NSG=4 mapping. */
constexpr short NSG = 4;
constexpr short NQ = 8;
constexpr short NR0 = 2;
const uint cohort = sgitg >> 2;
const ushort vsg = sgitg & 3u;
const uint vt = tgpig.x * 2u + cohort;
const bool valid = vt < pair_vtgs;
const int r0 = vt * NR0;
const bool active_a = valid && r0 < args0.ne01;
const bool active_b = valid && r0 < args1.ne01;
const int nb = args0.ne00 / QK8_0;
device const float *y = (device const float *)src1;
device const block_q8_0 *ax_a[NR0];
device const block_q8_0 *ax_b[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const int out_row = r0 + row;
ax_a[row] = active_a && out_row < args0.ne01
? (device const block_q8_0 *)(qw0 + (uint64_t)out_row * args0.nb01)
: (device const block_q8_0 *)qw0;
ax_b[row] = active_b && out_row < args1.ne01
? (device const block_q8_0 *)(qw1 + (uint64_t)out_row * args1.nb01)
: (device const block_q8_0 *)qw1;
}
float suma[NR0] = { 0.f };
float sumb[NR0] = { 0.f };
const short ix = tiisg / (NW / NQ);
const short il = tiisg % (NW / NQ);
const int ib0 = vsg * NQ + ix;
float yl[NQ];
device const float *yb = y + ib0 * QK8_0 + il * NQ;
if (valid) {
for (int ib = ib0; ib < nb; ib += NSG * NQ) {
FOR_UNROLL (short i = 0; i < NQ; ++i) {
yl[i] = yb[i];
}
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const int out_row = r0 + row;
if (active_a && out_row < args0.ne01) {
device const int8_t *qs = ax_a[row][ib].qs + il * NQ;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NQ; ++i) {
sumq += qs[i] * yl[i];
}
suma[row] += sumq * ax_a[row][ib].d;
}
if (active_b && out_row < args1.ne01) {
device const int8_t *qs = ax_b[row][ib].qs + il * NQ;
float sumq = 0.f;
FOR_UNROLL (short i = 0; i < NQ; ++i) {
sumq += qs[i] * yl[i];
}
sumb[row] += sumq * ax_b[row][ib].d;
}
}
yb += NSG * NQ * QK8_0;
}
}
threadgroup float *shared =
(threadgroup float *)shmem + cohort * (2 * NR0 * NW);
threadgroup float *sha[NR0];
threadgroup float *shb[NR0];
FOR_UNROLL (short row = 0; row < NR0; ++row) {
sha[row] = shared + NW * row;
shb[row] = shared + NW * (NR0 + row);
if (vsg == 0) {
sha[row][tiisg] = 0.0f;
if (active_b) shb[row][tiisg] = 0.0f;
}
suma[row] = simd_sum(suma[row]);
if (active_b) sumb[row] = simd_sum(sumb[row]);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short row = 0; row < NR0; ++row) {
if (tiisg == 0) {
sha[row][vsg] = suma[row];
if (active_b) shb[row][vsg] = sumb[row];
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
device float *out_a = (device float *)dst0;
device float *out_b = (device float *)dst1;
FOR_UNROLL (short row = 0; row < NR0; ++row) {
const float total_a = simd_sum(sha[row][tiisg]);
if (tiisg == 0 && vsg == 0) {
const int out_row = r0 + row;
if (active_a && out_row < args0.ne01) out_a[out_row] = total_a;
}
if (active_b) {
const float total_b = simd_sum(shb[row][tiisg]);
if (tiisg == 0 && vsg == 0) {
const int out_row = r0 + row;
if (out_row < args1.ne01) out_b[out_row] = total_b;
}
}
}
return;
}
/* Compressor quad range: verbatim body of
* kernel_mul_mv_f16_f32_quad_compressor_store_4 on the shifted grid. */
constexpr short NR0 = 2;
const uint lx = tgpig.x - pair_ctgs;
const uint tgs0 = ((uint)store0.width + NR0 - 1u) / NR0;
const bool second = lx >= tgs0;
uint3 local_tgpig = tgpig;
local_tgpig.x = second ? lx - tgs0 : lx;
ds4_metal_args_mul_mv largs = cargs;
largs.nr0 = NR0;
largs.ne01 = second ? (int32_t)store1.width : (int32_t)store0.width;
if (!second) {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, cw0a, cw0b, src1, cdst_a0, cdst_b0,
shmem, local_tgpig, tiisg, sgitg);
} else {
kernel_mul_mv_f16_f32_pair_4_impl<NR0>(
largs, cw1a, cw1b, src1, cdst_a1, cdst_b1,
shmem, local_tgpig, tiisg, sgitg);
}
threadgroup_barrier(mem_flags::mem_device);
// State append: identical to the paired store kernel, scoped to the
// range this threadgroup just projected (its own outputs only).
constant ds4_metal_args_compressor_pair_store & store = second ? store1 : store0;
if (tiitg >= NR0 || store.width == 0u || store.ratio == 0u) {
return;
}
const uint col = local_tgpig.x * (uint)NR0 + tiitg;
if (col >= store.width) return;
const uint pos_mod = store.pos % store.ratio;
const uint dst_row = store.ratio == 4u ? store.ratio + pos_mod : pos_mod;
const uint dst = dst_row * store.width + col;
const uint ape_i = pos_mod * store.width + col;
device volatile const float * projected_kv = second
? (device volatile const float *)cdst_a1
: (device volatile const float *)cdst_a0;
device volatile const float * projected_score = second
? (device volatile const float *)cdst_b1
: (device volatile const float *)cdst_b0;
device const char * ape = second ? ape1 : ape0;
device float * state_kv = second ? state1_kv : state0_kv;
device float * state_score = second ? state1_score : state0_score;
float ape_v;
if (store.ape_type == 1u) {
ape_v = (float)(((device const half *)ape)[ape_i]);
} else {
ape_v = ((device const float *)ape)[ape_i];
}
state_kv[dst] = projected_kv[col];
state_score[dst] = projected_score[col] + ape_v;
}
template<typename T0, typename T1, typename args_t>
void kernel_mul_mv_t_t_short_impl(
args_t args,
@@ -1476,125 +1901,6 @@ constant bool FC_mul_mm_bc_inp [[function_constant(FC_MUL_MM + 0)]];
constant bool FC_mul_mm_bc_out [[function_constant(FC_MUL_MM + 1)]];
#ifdef DS4_METAL_HAS_TENSOR
template<
short NR0, short NR1,
typename SA, typename SA_4x4, typename block_q, short nl,
void (*dequantize_func)(device const block_q *, short, thread SA_4x4 &),
typename T0, typename T0_4x4, typename T1>
kernel void kernel_mul_mm_mpp(
constant ds4_metal_args_mul_mm & args,
device const char * srcA,
device const char * srcB,
device char * dst,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig [[threadgroup_position_in_grid]],
ushort tiitg [[thread_index_in_threadgroup]],
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
(void) sgitg;
constexpr int NK = 32;
constexpr int NL = NK/16;
constexpr int NUM_THREADS = 128;
const int K = args.ne00;
const int M = args.ne0;
const int N = args.ne1;
const int im = tgpig.z;
const int i12 = im%args.ne12;
const int i13 = im/args.ne12;
const int r0 = tgpig.y*NR0;
const int r1 = tgpig.x*NR1;
const uint64_t offset0 = (i12/args.r2)*args.nb02 + (i13/args.r3)*args.nb03;
threadgroup SA *sa = (threadgroup SA *)shmem;
threadgroup SA *sb = sa + NR0*NK;
auto tA = tensor(sa, dextents<int32_t, 2>(NK, NR0));
auto tB = tensor(sb, dextents<int32_t, 2>(NK, NR1));
device const T1 *ptrB = (device const T1 *)(srcB + args.nb12*i12 + args.nb13*i13);
const int strideB = args.nb11/sizeof(T1);
matmul2d<
matmul2d_descriptor(NR1, NR0, NK, false, true, false,
matmul2d_descriptor::mode::multiply_accumulate),
execution_simdgroups<4>> mm;
auto cT = mm.template get_destination_cooperative_tensor<decltype(tB), decltype(tA), float>();
#pragma unroll
for (uint16_t i = 0; i < cT.get_capacity(); ++i) {
if (cT.is_valid_element(i)) {
cT[i] = 0.0f;
}
}
for (int loop_k = 0; loop_k < K; loop_k += NK) {
for (int work = tiitg; work < NR0*NL; work += NUM_THREADS) {
const int row = work/NL;
const int k_chunk = work%NL;
const int k_pos = loop_k + k_chunk*16;
const short k_base = k_chunk*16;
if (!FC_mul_mm_bc_out || r0 + row < M) {
if (is_same<T0_4x4, block_q>::value && FC_mul_mm_bc_inp) {
device const T0 *row_ptr = (device const T0 *)(srcA + args.nb01*(r0 + row) + offset0);
FOR_UNROLL (short i = 0; i < 16; i++) {
sa[row*NK + k_base + i] = (k_pos + i < K) ? (SA)row_ptr[k_pos + i] : (SA)0;
}
} else {
const int block_idx = k_pos/(16*nl);
const short il = (k_pos/16)%nl;
device const block_q *row_ptr = (device const block_q *)(srcA + args.nb01*(r0 + row) + offset0);
SA_4x4 temp_a;
dequantize_func(row_ptr + block_idx, il, temp_a);
FOR_UNROLL (short i = 0; i < 16; i++) {
sa[row*NK + k_base + i] = (k_pos + i < K) ? temp_a[i/4][i%4] : (SA)0;
}
}
} else {
FOR_UNROLL (short i = 0; i < 16; i++) {
sa[row*NK + k_base + i] = (SA)0;
}
}
}
for (int work = tiitg; work < NK*NR1; work += NUM_THREADS) {
const int col = work/NK;
const int k = work%NK;
if ((!FC_mul_mm_bc_out && !FC_mul_mm_bc_inp) ||
(r1 + col < N && loop_k + k < K)) {
sb[col*NK + k] = (SA)ptrB[(uint64_t)(r1 + col)*strideB + loop_k + k];
} else {
sb[col*NK + k] = (SA)0;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
auto mA = tA.slice(0, 0);
auto mB = tB.slice(0, 0);
mm.run(mB, mA, cT);
threadgroup_barrier(mem_flags::mem_threadgroup);
}
device float *dst_batch = (device float *)dst + im*N*M;
if (!FC_mul_mm_bc_out) {
device float *dst_tile = dst_batch + r0 + (uint64_t)r1*M;
auto tD = tensor(dst_tile, dextents<int32_t, 2>(NR0, NR1), array<int, 2>({1, M}));
cT.store(tD);
} else {
auto tD = tensor(dst_batch, dextents<int32_t, 2>(M, N), array<int, 2>({1, M}));
auto mD = tD.slice(r0, r1);
cT.store(mD);
}
}
typedef decltype(kernel_mul_mm_mpp<64, 32, half, half4x4, float4x4, 1, dequantize_f32, float, float4x4, float>) mul_mm_mpp_t;
template [[host_name("kernel_mul_mm_f16_f32_mpp")]] kernel mul_mm_mpp_t kernel_mul_mm_mpp<64, 32, half, half4x4, half4x4, 1, dequantize_f16, half, half4x4, float>;
// Retained Metal4/TensorOps dense prefill kernel. The legacy MPP prototype
// staged both operands in threadgroup memory; this version stages only the
// model weight tile and lets MPP read the dense RHS activation matrix directly
@@ -2144,242 +2450,6 @@ kernel void kernel_mul_mm_f16_f32_scaled(
}
}
kernel void kernel_mul_mm_f16_f32_pair(
constant ds4_metal_args_mul_mm & args,
device const char * src0_a,
device const char * src0_b,
device const char * src1,
device char * dst_a,
device char * dst_b,
threadgroup char * shmem [[threadgroup(0)]],
uint3 tgpig[[threadgroup_position_in_grid]],
ushort tiitg[[thread_index_in_threadgroup]],
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
threadgroup half * sa_a = (threadgroup half *)(shmem);
threadgroup half * sa_b = (threadgroup half *)(shmem + 4096);
threadgroup half * sb = (threadgroup half *)(shmem + 8192);
constexpr int NR0 = 64;
constexpr int NR1 = 32;
constexpr int NK = 32;
constexpr int NL0 = NK/16;
constexpr int NL1 = NK/8;
const int im = tgpig.z;
const int r0 = tgpig.y*NR0;
const int r1 = tgpig.x*NR1;
const short nr0 = (args.ne0 - r0 < NR0) ? (args.ne0 - r0) : NR0;
const short nr1 = (args.ne1 - r1 < NR1) ? (args.ne1 - r1) : NR1;
const short lr0 = ((short)tiitg/NL0) < nr0 ? ((short)tiitg/NL0) : nr0 - 1;
const short lr1 = ((short)tiitg/NL1) < nr1 ? ((short)tiitg/NL1) : nr1 - 1;
const short il0 = (tiitg % NL0);
short il = il0;
const int i12 = im%args.ne12;
const int i13 = im/args.ne12;
const uint64_t offset0 = (i12/args.r2)*args.nb02 + (i13/args.r3)*args.nb03;
const short offset1 = il0;
device const half4x4 * xa = (device const half4x4 *)(src0_a + args.nb01*(r0 + lr0) + offset0) + offset1;
device const half4x4 * xb = (device const half4x4 *)(src0_b + args.nb01*(r0 + lr0) + offset0) + offset1;
const short iy = 8*(tiitg % NL1);
device const float * y = (device const float *)(src1
+ args.nb13*i13
+ args.nb12*i12
+ args.nb11*(r1 + lr1)
+ args.nb10*iy);
simdgroup_half8x8 ma[4];
simdgroup_half8x8 mb[2];
simdgroup_float8x8 mc_a[8];
simdgroup_float8x8 mc_b[8];
for (short i = 0; i < 8; i++) {
mc_a[i] = make_filled_simdgroup_matrix<float, 8>(0.f);
mc_b[i] = make_filled_simdgroup_matrix<float, 8>(0.f);
}
for (int loop_k = 0; loop_k < args.ne00; loop_k += NK) {
half4x4 temp_a;
half4x4 temp_b;
dequantize_f16(xa, il, temp_a);
dequantize_f16(xb, il, temp_b);
threadgroup_barrier(mem_flags::mem_threadgroup);
FOR_UNROLL (short i = 0; i < 16; i++) {
const short sx = 2*il0 + i/8;
const short sy = (tiitg/NL0)/8;
const short lx = (tiitg/NL0)%8;
const short ly = i%8;
const short ib = 8*sx + sy;
*(sa_a + 64*ib + 8*ly + lx) = temp_a[i/4][i%4];
*(sa_b + 64*ib + 8*ly + lx) = temp_b[i/4][i%4];
}
if (FC_mul_mm_bc_inp) {
for (short i = 0; i < 8; ++i) {
const short sx = (tiitg%NL1);
const short sy = (tiitg/NL1)/8;
const short lx = i;
const short ly = (tiitg/NL1)%8;
const short ib = 4*sx + sy;
*(sb + 64*ib + 8*ly + lx) = loop_k + iy + i < args.ne00 ? (half) *((device float *) y + i) : 0;
}
} else {
const short sx = (tiitg%NL1);
const short sy = (tiitg/NL1)/8;
const short ly = (tiitg/NL1)%8;
const short ib = 4*sx + sy;
*(threadgroup half2x4 *)(sb + 64*ib + 8*ly) = (half2x4)(*((device float2x4 *) y));
}
il = (il + 2 < 1) ? il + 2 : il % 2;
xa = (il < 2) ? xa + 2 : xa;
xb = (il < 2) ? xb + 2 : xb;
y += NK;
threadgroup_barrier(mem_flags::mem_threadgroup);
threadgroup const half * lsma_a = (sa_a + 4*64*(sgitg%2));
threadgroup const half * lsma_b = (sa_b + 4*64*(sgitg%2));
threadgroup const half * lsmb = (sb + 2*64*(sgitg/2));
FOR_UNROLL (short ik = 0; ik < NK/8; ik++) {
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 2; i++) {
simdgroup_load(mb[i], lsmb + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 4; i++) {
simdgroup_load(ma[i], lsma_a + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 8; i++) {
simdgroup_multiply_accumulate(mc_a[i], mb[i/4], ma[i%4], mc_a[i]);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 4; i++) {
simdgroup_load(ma[i], lsma_b + 64*i, 8, 0, false);
}
simdgroup_barrier(mem_flags::mem_none);
FOR_UNROLL (short i = 0; i < 8; i++) {
simdgroup_multiply_accumulate(mc_b[i], mb[i/4], ma[i%4], mc_b[i]);
}
lsma_a += 8*64;
lsma_b += 8*64;
lsmb += 4*64;
}
}
if (!FC_mul_mm_bc_out || (r0 + NR0 <= args.ne0 && r1 + NR1 <= args.ne1)) {
device float * C_a = (device float *) dst_a +
(r0 + 32*(sgitg & 1)) +
(r1 + 16*(sgitg >> 1)) * args.ne0 + im*args.ne1*args.ne0;
device float * C_b = (device float *) dst_b +
(r0 + 32*(sgitg & 1)) +
(r1 + 16*(sgitg >> 1)) * args.ne0 + im*args.ne1*args.ne0;
for (short i = 0; i < 8; i++) {
simdgroup_store(mc_a[i], C_a + 8*(i%4) + 8*args.ne0*(i/4), args.ne0, 0, false);
simdgroup_store(mc_b[i], C_b + 8*(i%4) + 8*args.ne0*(i/4), args.ne0, 0, false);
}
} else {
threadgroup_barrier(mem_flags::mem_threadgroup);
threadgroup float * temp_str = (threadgroup float *) shmem;
for (short i = 0; i < 8; i++) {
simdgroup_store(mc_a[i],
temp_str + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0 + 8*(i%4) + 8*NR0*(i/4),
NR0,
0,
false);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sgitg == 0) {
for (int j = tiitg; j < nr1; j += NR1) {
device float * D = (device float *) dst_a + r0 + (r1 + j)*args.ne0 + im*args.ne1*args.ne0;
device float4 * D4 = (device float4 *) D;
threadgroup float * C = temp_str + (j*NR0);
threadgroup float4 * C4 = (threadgroup float4 *) C;
int i = 0;
for (; i < nr0/4; i++) {
*(D4 + i) = *(C4 + i);
}
i *= 4;
for (; i < nr0; i++) {
*(D + i) = *(C + i);
}
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
for (short i = 0; i < 8; i++) {
simdgroup_store(mc_b[i],
temp_str + 32*(sgitg&1) + (16*(sgitg >> 1))*NR0 + 8*(i%4) + 8*NR0*(i/4),
NR0,
0,
false);
}
threadgroup_barrier(mem_flags::mem_threadgroup);
if (sgitg == 0) {
for (int j = tiitg; j < nr1; j += NR1) {
device float * D = (device float *) dst_b + r0 + (r1 + j)*args.ne0 + im*args.ne1*args.ne0;
device float4 * D4 = (device float4 *) D;
threadgroup float * C = temp_str + (j*NR0);
threadgroup float4 * C4 = (threadgroup float4 *) C;
int i = 0;
for (; i < nr0/4; i++) {
*(D4 + i) = *(C4 + i);
}
i *= 4;
for (; i < nr0; i++) {
*(D + i) = *(C + i);
}
}
}
}
}
typedef decltype(kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, float, float2x4>) mul_mm_t;
// Host-visible prefill matmul variants for F16 and Q8_0 weights.