Add GLM 5.3 Flash support

This commit is contained in:
Georg Bauer
2026-09-01 19:06:20 +02:00
parent 9a33c61ea6
commit 46d6a976a5
31 changed files with 7806 additions and 447 deletions

View File

@@ -423,6 +423,7 @@ struct ds4_metal_glm_routed_moe_args {
uint32_t n_tokens;
uint32_t mid_token_stride;
uint32_t down_type;
float swiglu_clamp;
/* Expert ownership under tensor parallelism: tp_world 0/1 = full
* compute; otherwise each rank owns a contiguous expert range. */
int32_t tp_rank;
@@ -436,6 +437,14 @@ struct ds4_metal_glm_routed_moe_args {
uint64_t down_row_bytes;
};
static inline float ds4_glm_swiglu(float gate, float up, float limit) {
if (limit > 1.0e-6f) {
gate = min(gate, limit);
up = clamp(up, -limit, limit);
}
return (gate / (1.0f + exp(-gate))) * up;
}
static inline bool ds4_tp_owns_expert(int expert, int n_total,
int tp_rank, int tp_world) {
@@ -712,18 +721,20 @@ kernel void kernel_glm_q4_K_pair_swiglu_f32(
const uint64_t mid_off = (uint64_t)token * args.mid_token_stride +
(uint64_t)slot * args.mid_dim + row;
const int expert = selected[selected_off];
if (expert < 0 || (uint)expert >= args.n_total_expert) {
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) {
if (tid == 0u) mid[mid_off] = 0.0f;
return;
}
const int local_expert = expert - args.tp_expert_base;
device const block_q4_K *gate_row =
(device const block_q4_K *)(gate +
(uint64_t)(uint)expert * args.gate_expert_bytes +
(uint64_t)(uint)local_expert * args.gate_expert_bytes +
(uint64_t)row * args.gate_row_bytes);
device const block_q4_K *up_row =
(device const block_q4_K *)(up +
(uint64_t)(uint)expert * args.up_expert_bytes +
(uint64_t)(uint)local_expert * args.up_expert_bytes +
(uint64_t)row * args.up_row_bytes);
float acc_gate = 0.0f;
@@ -747,10 +758,8 @@ kernel void kernel_glm_q4_K_pair_swiglu_f32(
}
if (tid == 0u) {
const float g = scratch[0];
const float u = scratch[ntg];
const float sw = g / (1.0f + exp(-g));
mid[mid_off] = sw * u * weights[selected_off];
mid[mid_off] = ds4_glm_swiglu(scratch[0], scratch[ntg],
args.swiglu_clamp) * weights[selected_off];
}
}
@@ -887,8 +896,8 @@ static inline void glm_q2_K_pair_swiglu_simd_f32_impl(
const float g = simd_sum(sumg[row]);
const float u = simd_sum(sumu[row]);
if (tiisg == 0u) {
const float sw = g / (1.0f + exp(-g));
mid[mid_base + row0 + (uint)row] = sw * u * weights[selected_off];
mid[mid_base + row0 + (uint)row] =
ds4_glm_swiglu(g, u, args.swiglu_clamp) * weights[selected_off];
}
}
@@ -944,7 +953,9 @@ kernel void kernel_glm_q2_K_addr_pair_swiglu2_f32(
(uint64_t)slot * args.mid_dim;
if (row0 >= args.mid_dim) return;
if (expert < 0 || (uint)expert >= args.n_total_expert) {
if (expert < 0 || (uint)expert >= args.n_total_expert ||
!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) {
if (tiisg == 0u) {
for (short row = 0;
row < N_R0_GLM_Q2_PAIR2_K && row0 + (uint)row < args.mid_dim;
@@ -1005,7 +1016,9 @@ kernel void kernel_glm_q2_K_addr_pair_swiglu2_f32_masked(
(uint64_t)slot * args.mid_dim;
if (row0 >= args.mid_dim) return;
if (expert < 0 || (uint)expert >= args.n_total_expert) {
if (expert < 0 || (uint)expert >= args.n_total_expert ||
!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) {
if (tiisg == 0u) {
for (short row = 0;
row < N_R0_GLM_Q2_PAIR2_K && row0 + (uint)row < args.mid_dim;
@@ -1191,8 +1204,8 @@ static inline void glm_q4_K_pair_swiglu_simd_f32_impl(
const float g = simd_sum(sumg[row]);
const float u = simd_sum(sumu[row]);
if (tiisg == 0u) {
const float sw = g / (1.0f + exp(-g));
mid[mid_base + row0 + (uint)row] = sw * u * weights[selected_off];
mid[mid_base + row0 + (uint)row] =
ds4_glm_swiglu(g, u, args.swiglu_clamp) * weights[selected_off];
}
}
@@ -1216,9 +1229,12 @@ kernel void kernel_glm_q4_K_pair_swiglu2_f32(
if (slot >= args.n_expert_used || token >= args.n_tokens) return;
const uint64_t selected_off = (uint64_t)token * args.n_expert_used + slot;
const int expert = selected[selected_off];
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) return;
glm_q4_K_pair_swiglu_simd_f32_impl<N_R0_GLM_Q4_PAIR2_K>(
args, gate, up, x, weights, mid, scratch,
tgpig, slot, token, selected_off, expert, tiisg, sgitg);
tgpig, slot, token, selected_off,
expert - args.tp_expert_base, tiisg, sgitg);
}
kernel void kernel_glm_q4_K_addr_pair_swiglu_f32(
@@ -1245,7 +1261,9 @@ kernel void kernel_glm_q4_K_addr_pair_swiglu_f32(
(uint64_t)slot * args.mid_dim;
if (row0 >= args.mid_dim) return;
if (expert < 0 || (uint)expert >= args.n_total_expert) {
if (expert < 0 || (uint)expert >= args.n_total_expert ||
!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) {
if (tiisg == 0u) {
for (short row = 0;
row < N_R0_Q4_K && row0 + (uint)row < args.mid_dim;
@@ -1306,7 +1324,9 @@ kernel void kernel_glm_q4_K_addr_pair_swiglu_f32_masked(
(uint64_t)slot * args.mid_dim;
if (row0 >= args.mid_dim) return;
if (expert < 0 || (uint)expert >= args.n_total_expert) {
if (expert < 0 || (uint)expert >= args.n_total_expert ||
!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) {
if (tiisg == 0u) {
for (short row = 0;
row < N_R0_Q4_K && row0 + (uint)row < args.mid_dim;
@@ -1358,9 +1378,12 @@ kernel void kernel_glm_q4_K_pair_swiglu4_f32(
if (slot >= args.n_expert_used || token >= args.n_tokens) return;
const uint64_t selected_off = (uint64_t)token * args.n_expert_used + slot;
const int expert = selected[selected_off];
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) return;
glm_q4_K_pair_swiglu_simd_f32_impl<N_R0_GLM_Q4_PAIR_K>(
args, gate, up, x, weights, mid, scratch,
tgpig, slot, token, selected_off, expert, tiisg, sgitg);
tgpig, slot, token, selected_off,
expert - args.tp_expert_base, tiisg, sgitg);
}
kernel void kernel_glm_q4_K_pair_swiglu2_mapped_f32(
@@ -1378,6 +1401,8 @@ kernel void kernel_glm_q4_K_pair_swiglu2_mapped_f32(
ushort sgitg [[simdgroup_index_in_threadgroup]]) {
const uint expert = tgpig.z;
if (expert >= args.n_total_expert) return;
if (!ds4_tp_owns_expert((int)expert, args.n_total_expert,
args.tp_rank, args.tp_world)) return;
const uint count = htpe[expert];
const uint map_base = tgpig.y * 32u;
for (uint i = 0; i < 32u; i++) {
@@ -1391,7 +1416,8 @@ kernel void kernel_glm_q4_K_pair_swiglu2_mapped_f32(
const uint64_t selected_off = (uint64_t)token * args.n_expert_used + slot;
glm_q4_K_pair_swiglu_simd_f32_impl<N_R0_Q4_K>(
args, gate, up, x, weights, mid, scratch,
tgpig, slot, token, selected_off, (int)expert, tiisg, sgitg);
tgpig, slot, token, selected_off,
(int)expert - args.tp_expert_base, tiisg, sgitg);
}
}
@@ -1411,6 +1437,8 @@ kernel void kernel_glm_q4_K_pair_swiglu2_mapped_row_f32(
const uint expert = tgpig.z;
const uint map_row = tgpig.y;
if (expert >= args.n_total_expert || map_row >= htpe[expert]) return;
if (!ds4_tp_owns_expert((int)expert, args.n_total_expert,
args.tp_rank, args.tp_world)) return;
const int id = hids[(uint64_t)expert * args.n_tokens + map_row];
if (id < 0) return;
const uint token = (uint)id / args.n_expert_used;
@@ -1419,7 +1447,8 @@ kernel void kernel_glm_q4_K_pair_swiglu2_mapped_row_f32(
const uint64_t selected_off = (uint64_t)token * args.n_expert_used + slot;
glm_q4_K_pair_swiglu_simd_f32_impl<N_R0_Q4_K>(
args, gate, up, x, weights, mid, scratch,
tgpig, slot, token, selected_off, (int)expert, tiisg, sgitg);
tgpig, slot, token, selected_off,
(int)expert - args.tp_expert_base, tiisg, sgitg);
}
static inline void glm_q5_K_pair_swiglu_f32_impl(
@@ -1591,8 +1620,8 @@ static inline void glm_q5_K_pair_swiglu_f32_impl(
const float g = simd_sum(sumg[row]);
const float u = simd_sum(sumu[row]);
if (tiisg == 0u) {
const float sw = g / (1.0f + exp(-g));
mid[mid_base + row0 + (uint)row] = sw * u * weights[selected_off];
mid[mid_base + row0 + (uint)row] =
ds4_glm_swiglu(g, u, args.swiglu_clamp) * weights[selected_off];
}
}
@@ -1913,6 +1942,8 @@ kernel void kernel_glm_q2_K_addr_down_f32(
for (uint slot = 0; slot < args.n_expert_used; slot++) {
const int expert = selected[selected_base + slot];
if (expert < 0 || (uint)expert >= args.n_total_expert) continue;
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) continue;
const uint64_t down_addr = down_addrs[(uint)expert];
if (down_addr == 0) continue;
device const block_q2_K *x =
@@ -1996,10 +2027,12 @@ kernel void kernel_glm_q4_K_down_f32(
const uint64_t mid_base = (uint64_t)token * args.mid_token_stride;
for (uint slot = 0; slot < args.n_expert_used; slot++) {
const int expert = selected[selected_base + slot];
if (expert < 0 || (uint)expert >= args.n_total_expert) continue;
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) continue;
device const block_q4_K *down_row =
(device const block_q4_K *)(down +
(uint64_t)(uint)expert * args.down_expert_bytes +
(uint64_t)(uint)(expert - args.tp_expert_base) *
args.down_expert_bytes +
(uint64_t)row * args.down_row_bytes);
device const float *slot_mid = mid + mid_base + (uint64_t)slot * args.mid_dim;
for (uint k = tid; k < args.mid_dim; k += ntg) {
@@ -2038,6 +2071,8 @@ kernel void kernel_glm_q4_K_addr_down_f32(
for (uint slot = 0; slot < args.n_expert_used; slot++) {
const int expert = selected[selected_base + slot];
if (expert < 0 || (uint)expert >= args.n_total_expert) continue;
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) continue;
const uint64_t down_addr = down_addrs[(uint)expert];
if (down_addr == 0) continue;
device const block_q4_K *down_row =
@@ -2093,11 +2128,13 @@ kernel void kernel_glm_q4_K_down_simd_f32(
const uint64_t mid_base = (uint64_t)token * args.mid_token_stride;
for (uint slot = 0; slot < args.n_expert_used; slot++) {
const int expert = selected[selected_base + slot];
if (expert < 0 || (uint)expert >= args.n_total_expert) continue;
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) continue;
device const block_q4_K *x =
(device const block_q4_K *)(down +
(uint64_t)(uint)expert * args.down_expert_bytes +
(uint64_t)(uint)(expert - args.tp_expert_base) *
args.down_expert_bytes +
(uint64_t)row0 * args.down_row_bytes);
device const float *y = mid + mid_base + (uint64_t)slot * args.mid_dim;
device const float *y4 = y + ix * QK_K + 64 * iq + 8 * ir;
@@ -2197,6 +2234,8 @@ kernel void kernel_glm_q4_K_addr_down_simd_f32(
for (uint slot = 0; slot < args.n_expert_used; slot++) {
const int expert = selected[selected_base + slot];
if (expert < 0 || (uint)expert >= args.n_total_expert) continue;
if (!ds4_tp_owns_expert(expert, args.n_total_expert,
args.tp_rank, args.tp_world)) continue;
const uint64_t down_addr = down_addrs[(uint)expert];
if (down_addr == 0) continue;