Skip to content
Open
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
84 changes: 49 additions & 35 deletions mlx/backend/metal/kernels/quantized_nax.h
Original file line number Diff line number Diff line change
Expand Up @@ -1094,7 +1094,6 @@ METAL_FUNC void qmm_n_nax_tgp_impl(
uint simd_gid [[simdgroup_index_in_threadgroup]],
uint simd_lid [[thread_index_in_simdgroup]]) {
(void)lid;
(void)M;

static_assert(BK >= SIMD_SIZE, "BK should be larger than SIMD_SIZE");
static_assert(BK % SIMD_SIZE == 0, "BK should be divisible by SIMD_SIZE");
Expand All @@ -1115,23 +1114,20 @@ METAL_FUNC void qmm_n_nax_tgp_impl(
bits>;

// Set the block
const int K_w = K * bytes_per_pack / pack_factor;
const int K_g = K / group_size;
const int y_row = tid.y * BM;
const int y_col = tid.x * BN;

auto wl = (const device uint8_t*)w;

// Here w is [K, N]: packed and group-quantized along N, with row stride N.
x += y_row * static_cast<int64_t>(K);
wl += y_col * K_w;
scales += y_col * K_g;
biases += y_col * K_g;
wl += y_col * bytes_per_pack / pack_factor;
scales += y_col / group_size;
biases += y_col / group_size;
y += y_row * static_cast<int64_t>(N) + y_col;

// Make the x loader and mma operation
// const short num_els = min(BM, M - y_row);
// const short num_outs = min(BN, N - y_col);
loader_w_t loader_w(wl, scales, biases, K, Ws, simd_gid, simd_lid);
// Make the weight loader
loader_w_t loader_w(wl, scales, biases, N, Ws, simd_gid, simd_lid);

constexpr short SM = BM / WM;
constexpr short SN = BN / WN;
Expand All @@ -1144,6 +1140,13 @@ METAL_FUNC void qmm_n_nax_tgp_impl(
const short tm = SM * (simd_gid / WN);
const short tn = SN * (simd_gid % WN);

// Rows of this simdgroup's SM-row slice that actually exist. This can go
// <= 0 when a whole simdgroup sits past the end of the matrix, which
// load_safe/store_safe handle, as they already do for the transposed
// kernel. N needs no equivalent: the dispatch guard keeps N % BN == 0.
const short sgp_sm = min(int(SM), M - (y_row + tm));
const bool is_unaligned_sm = (sgp_sm != SM);

const short ldb_tgp = BN_padded;

constexpr bool transpose_a = false;
Expand All @@ -1156,39 +1159,50 @@ METAL_FUNC void qmm_n_nax_tgp_impl(

x += tm * K;

for (int k = 0; k < K; k += BK) {
threadgroup_barrier(mem_flags::mem_threadgroup);
loader_w.load_unsafe();
threadgroup_barrier(mem_flags::mem_threadgroup);
dispatch_bool(!is_unaligned_sm, [&](auto kAlignedM) {
for (int k = 0; k < K; k += BK) {
threadgroup_barrier(mem_flags::mem_threadgroup);
loader_w.load_unsafe();
threadgroup_barrier(mem_flags::mem_threadgroup);

STEEL_PRAGMA_NO_UNROLL
for (int kk1 = 0; kk1 < BK; kk1 += SK) {
NAXTile<T, TM, TK> Atile;
NAXTile<T, TK, TN> Btile;
STEEL_PRAGMA_NO_UNROLL
for (int kk1 = 0; kk1 < BK; kk1 += SK) {
NAXTile<T, TM, TK> Atile;
NAXTile<T, TK, TN> Btile;

volatile int compiler_barrier;
volatile int compiler_barrier;

Atile.load(x + kk1, K);
Btile.template load<T, BN_padded, 1>(Ws + tn + kk1 * ldb_tgp);
if constexpr (kAlignedM.value) {
Atile.load(x + kk1, K);
} else {
Atile.load_safe(x + kk1, K, short2(SK, sgp_sm));
}

tile_matmad_nax(
Dtile,
Atile,
metal::bool_constant<transpose_a>{},
Btile,
metal::bool_constant<transpose_b>{});
Btile.template load<T, BN_padded, 1>(Ws + tn + kk1 * ldb_tgp);

(void)compiler_barrier;
}
tile_matmad_nax(
Dtile,
Atile,
metal::bool_constant<transpose_a>{},
Btile,
metal::bool_constant<transpose_b>{});

x += BK;
loader_w.next();
}
(void)compiler_barrier;
}

// Store results to device memory
threadgroup_barrier(mem_flags::mem_threadgroup);
x += BK;
loader_w.next();
}

// Store results to device memory
threadgroup_barrier(mem_flags::mem_threadgroup);

Dtile.store(y + tm * N + tn, N);
if constexpr (kAlignedM.value) {
Dtile.store(y + tm * N + tn, N);
} else {
Dtile.store_safe(y + tm * N + tn, N, short2(SN, sgp_sm));
}
});
}

template <
Expand Down
8 changes: 6 additions & 2 deletions mlx/backend/metal/quantized.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1011,8 +1011,12 @@ void qmm(
metal::Device& d,
const Stream& s,
const std::string& mode) {
if (metal::is_nax_available() && transpose && (K % 64 == 0) &&
(env::enable_tf32() || x.dtype() != float32)) {
// The non-transposed kernel loads a BK x BN tile of w = [K, N] without
// clamping in N, so it additionally requires N % 64 == 0. Affine
// quantization with group_size >= 64 satisfies that structurally, since
// the groups run along N.
if (metal::is_nax_available() && (transpose || (N % 64 == 0)) &&
(K % 64 == 0) && (env::enable_tf32() || x.dtype() != float32)) {
return qmm_nax(
/* const array& x = */ x,
/* const array& w = */ w,
Expand Down
47 changes: 47 additions & 0 deletions python/tests/test_quantized.py
Original file line number Diff line number Diff line change
Expand Up @@ -353,6 +353,53 @@ def test_qmm_large_dims(self):
tol = 1e-3 if dtype == mx.float32 else 1.5e-3
self.assertLess((y_q - y_hat).abs().max(), tol)

def test_qmm_non_transposed(self):
# The non-transposed matmul (w is [K, N]) is reachable mainly from the
# vjp of a quantized linear layer, so it gets much less coverage than
# the transposed one. Sweep it over transformer-sized K/N and over M
# values that leave a partial M-tile.
key = mx.random.key(0)
k1, k2 = mx.random.split(key)
dtype = mx.float16 if (mx.default_device() == mx.gpu) else mx.float32
tol = 1e-3 if dtype == mx.float32 else 1.5e-3

def check(M, K, N, group_size, bits, batch=()):
x = mx.random.normal(shape=(*batch, M, K), key=k1) / K**0.5
w = mx.random.normal(shape=(K, N), key=k2) / K**0.5
x = x.astype(dtype)
w = w.astype(dtype)
w_q, scales, biases = mx.quantize(w, group_size, bits)
w_hat = mx.dequantize(w_q, scales, biases, group_size, bits)
y_q = mx.quantized_matmul(x, w_q, scales, biases, False, group_size, bits)
y_hat = x @ w_hat
self.assertEqual(y_q.shape, y_hat.shape)
self.assertLess((y_q - y_hat).abs().max(), tol)

# M sweep. 33..63 is the interesting range: a whole simdgroup of the
# threadgroup's M-tile falls past the end of the matrix.
for M in [1, 2, 31, 32, 33, 63, 64, 65, 96, 97, 100, 127, 128, 129]:
for group_size, bits in [(64, 4), (128, 4), (64, 8)]:
with self.subTest(M=M, group_size=group_size, bits=bits):
check(M, 512, 1024, group_size, bits)

# Transformer-sized K/N, aligned and unaligned M.
for K, N in [(2048, 2048), (512, 2048), (2048, 512), (11008, 2048)]:
for M in [100, 256]:
with self.subTest(shape=(M, K, N)):
check(M, K, N, 64, 4)

# Batched x, unaligned M.
for batch in [(2,), (2, 3)]:
for M in [33, 250]:
with self.subTest(batch=batch, M=M):
check(M, 512, 1024, 64, 4, batch=batch)

# M > 2**15 with a partial M-tile, so the per-simdgroup row count is a
# distance that does not fit in an int16. Same failure mode as the one
# test_qmm_large_dims covers for the transposed kernel.
with self.subTest(shape=(33000, 128, 64)):
check(33000, 128, 64, 64, 4)

def test_qmm_vjp(self):
key = mx.random.key(0)
k1, k2 = mx.random.split(key)
Expand Down
Loading