Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -79,3 +79,4 @@ uv.lock
.cache/
# vim
*.swp
wip
72 changes: 72 additions & 0 deletions benchmarks/python/swiglu_qmm_bench.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
# Copyright © 2026 Apple Inc.

import mlx.core as mx
import mlx.nn as nn
from time_utils import time_fn

SEQ_LENS = [512, 713, 1024, 2048, 4123, 8192, 8192 * 2]
MODES = ["affine", "nvfp4", "mxfp8"]

# https://huggingface.co/zai-org/GLM-5.3-Flash-BF16/blob/main/config.json
# https://huggingface.co/zai-org/GLM-5.3/blob/main/config.json
# https://huggingface.co/Qwen/Qwen3.5-35B-A3B/blob/main/config.json
# https://huggingface.co/Qwen/Qwen3.5-122B-A10B/blob/main/config.json
# https://huggingface.co/Qwen/Qwen3.5-397B-A17B/blob/main/config.json
CONFIGS = {
"glm-5.3-flash": (4096, 2048, 288, 8),
"glm-5.3": (6144, 2048, 256, 8),
"qwen3.5-35b-a3b": (2048, 512, 256, 8),
"qwen3.5-122b-a10b": (3072, 1024, 256, 8),
"qwen3.5-397b-a17b": (4096, 1024, 512, 10),
}


def gather_sort(x, indices):
N, M = indices.shape
indices = indices.flatten()
order = mx.argsort(indices)
inv_order = mx.argsort(order)
return x.flatten(0, -3)[order // M], indices[order], inv_order


def scatter_unsort(x, inv_order, shape=None):
x = x[inv_order]
if shape is not None:
x = mx.unflatten(x, 0, shape)
return x


def time_gather_qmm(name, D, M, E, I, mode):
w1 = mx.random.normal((E, M, D), dtype=mx.bfloat16, scale=D**-0.5)
w2 = mx.random.normal((E, M, D), dtype=mx.bfloat16, scale=D**-0.5)
w3 = mx.random.normal((E, D, M), dtype=mx.bfloat16, scale=M**-0.5)
w1, w2, w3 = (mx.quantize(w, mode=mode) for w in (w1, w2, w3))
mx.eval(w1, w2, w3)

def gather_qmm(x, w1, w2, w3, indices, sort):
idx = indices
inv_order = None
if sort:
x, idx, inv_order = gather_sort(x, indices)
kwargs = dict(transpose=True, mode=mode, rhs_indices=idx, sorted_indices=sort)
gate = mx.gather_qmm(x, *w1, **kwargs)
up = mx.gather_qmm(x, *w2, **kwargs)
x = mx.gather_qmm(nn.silu(gate) * up, *w3, **kwargs)
if sort:
x = scatter_unsort(x, inv_order, indices.shape)
return x

for N in SEQ_LENS:
x = mx.random.normal((N, 1, 1, D), dtype=mx.bfloat16)
scores = mx.random.uniform(shape=(N, E))
indices = mx.argpartition(scores, E - I, axis=-1)[:, -I:].astype(mx.uint32)
mx.eval(x, indices)

label = f"{name} {mode} N={N}"
time_fn(gather_qmm, x, w1, w2, w3, indices, True, msg=f"{label} swiglu")


if __name__ == "__main__":
for mode in MODES:
for name, config in CONFIGS.items():
time_gather_qmm(name, *config, mode)
164 changes: 59 additions & 105 deletions mlx/backend/metal/kernels/fp_quantized.h
Original file line number Diff line number Diff line change
Expand Up @@ -2050,11 +2050,12 @@ template <
const device uint32_t* w,
const device uint8_t* scales,
const device float* global_scale,
const device uint32_t* indices,
const device int32_t* offsets,
device T* y,
const constant int& M,
const constant int& N,
const constant int& K,
const constant int& num_groups,
uint3 tid [[threadgroup_position_in_grid]],
uint simd_group_id [[simdgroup_index_in_threadgroup]],
uint simd_lane_id [[thread_index_in_simdgroup]]) {
Expand Down Expand Up @@ -2099,13 +2100,18 @@ template <
const int K_it = K / BK;
const size_t stride_w = transpose ? N * K_w : K * N_w;
const size_t stride_s = transpose ? N * K_g : K * N_g;
const int y_row = tid.y * BM;
int y_row;
int group;
short tgp_bm;
if (!schedule_row_tile<BM>(
offsets, num_groups, M, tid.y, simd_lane_id, y_row, group, tgp_bm)) {
return;
}
const int y_col = tid.x * BN;
const size_t y_row_long = size_t(y_row);
const size_t y_col_long = size_t(y_col);

// Prepare threadgroup bounds
const short tgp_bm = align_M ? BM : short(min(BM, M - y_row));
const short tgp_bn = align_N ? BN : short(min(BN, N - y_col));

// Calculate the final tiles in the case that K is not aligned
Expand All @@ -2121,113 +2127,61 @@ template <
wl += transpose ? y_col_long * K_w : y_col * bytes_per_pack / pack_factor;
scales += transpose ? y_col_long * K_g : y_col / group_size;

// Do as many matmuls as necessary
uint32_t index;
short offset;
uint32_t index_next = indices[y_row];
short offset_next = 0;
int n = 0;
while (n < tgp_bm) {
n++;
offset = offset_next;
index = index_next;
offset_next = tgp_bm;
for (; n < tgp_bm; n++) {
if (indices[y_row + n] != index) {
offset_next = n;
index_next = indices[y_row + n];
break;
}
}
threadgroup_barrier(mem_flags::mem_none);

// Prepare threadgroup mma operation
thread mma_t mma_op(simd_group_id, simd_lane_id);

// Prepare threadgroup loading operations
thread loader_x_t loader_x(x, K, Xs, simd_group_id, simd_lane_id);
thread loader_w_t loader_w(
wl + index * stride_w,
scales + index * stride_s,
transpose ? K : N,
Ws,
simd_group_id,
simd_lane_id,
global_scale + index);

// Matrices are all aligned check nothing
if (align_M && align_N) {
gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
// Prepare threadgroup mma operation
thread mma_t mma_op(simd_group_id, simd_lane_id);

// Store results to device memory
if (offset_next - offset == BM) {
mma_op.store_result(y, N);
} else {
mma_op.store_result_slice(
y, N, short2(0, offset), short2(BN, offset_next));
}
} else {
// Tile aligned so check outside of the hot loop
if ((align_M || tgp_bm == BM) && (align_N || tgp_bn == BN)) {
gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(
Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}

// Store results to device memory
if (offset_next - offset == BM) {
mma_op.store_result(y, N);
} else {
mma_op.store_result_slice(
y, N, short2(0, offset), short2(BN, offset_next));
}
}
// Prepare threadgroup loading operations
thread loader_x_t loader_x(x, K, Xs, simd_group_id, simd_lane_id);
thread loader_w_t loader_w(
wl + group * stride_w,
scales + group * stride_s,
transpose ? K : N,
Ws,
simd_group_id,
simd_lane_id,
global_scale + group);

// Tile aligned so check outside of the hot loop
if (tgp_bm == BM && (align_N || tgp_bn == BN)) {
gemm_loop_aligned(Xs, Ws, mma_op, loader_x, loader_w, K_it);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result(y, N);
}

// Tile partially aligned check rows
else if (align_N || tgp_bn == BN) {
gemm_loop_unaligned<false, true, transpose>(
Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(
Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result_slice(
y, N, short2(0, offset), short2(BN, offset_next));
}
// Tile partially aligned check rows
else if (align_N || tgp_bn == BN) {
gemm_loop_unaligned<false, true, transpose>(
Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result_safe(y, N, short2(BN, tgp_bm));
}

// Tile partially aligned check cols
else if (align_M || tgp_bm == BM) {
gemm_loop_unaligned<true, false, transpose>(
Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(
Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result_slice(
y, N, short2(0, offset), short2(tgp_bn, offset_next));
}
// Tile partially aligned check cols
else if (tgp_bm == BM) {
gemm_loop_unaligned<true, false, transpose>(
Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result_safe(y, N, short2(tgp_bn, BM));
}

// Nothing aligned so check both rows and cols
else {
gemm_loop_unaligned<false, false, transpose>(
Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(
Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result_slice(
y, N, short2(0, offset), short2(tgp_bn, offset_next));
}
// Nothing aligned so check both rows and cols
else {
gemm_loop_unaligned<false, false, transpose>(
Xs, Ws, mma_op, loader_x, loader_w, K_it, tgp_bm, tgp_bn, BK);
if (!align_K) {
threadgroup_barrier(mem_flags::mem_threadgroup);
gemm_loop_finalize(Xs, Ws, mma_op, loader_x, loader_w, tile_x, tile_w);
}
mma_op.store_result_safe(y, N, short2(tgp_bn, tgp_bm));
}
}

Expand Down
Loading
Loading