diff --git a/benchmarks/python/segmented_mm_bench.py b/benchmarks/python/segmented_mm_bench.py new file mode 100644 index 00000000..0823c2e0 --- /dev/null +++ b/benchmarks/python/segmented_mm_bench.py @@ -0,0 +1,209 @@ +# Copyright © 2026 Apple Inc. + +import argparse +import time + +import mlx.core as mx +import numpy as np + +MLX_DTYPES = { + "float16": mx.float16, + "bfloat16": mx.bfloat16, + "float32": mx.float32, +} + + +def parse_cases(cases): + parsed = [] + for spec in cases.split(","): + m, n, k, s = [int(x) for x in spec.split("x")] + parsed.append((m, n, k, s)) + return parsed + + +def make_segments(k, num_segments, pattern, seed): + if pattern == "equal": + cuts = np.linspace(0, k, num_segments + 1, dtype=np.int64) + else: + rng = np.random.default_rng(seed) + cuts = rng.integers(0, k + 1, size=(num_segments - 1,), dtype=np.int64) + cuts = np.sort(cuts) + cuts = np.concatenate(([0], cuts, [k])) + return np.stack([cuts[:-1], cuts[1:]], axis=1).astype(np.uint32) + + +def numpy_segmented_mm_ref(a, b, segments): + """Ground-truth reference in float64.""" + out = [] + for start, end in segments: + out.append(a[:, start:end] @ b[start:end, :]) + return np.stack(out, axis=0) + + +def mlx_segmented_mm_loop(a, b, segments): + """MLX loop-of-matmuls baseline.""" + segments_list = segments.tolist() + out = [] + for start, end in segments_list: + out.append(a[:, start:end] @ b[start:end, :]) + return mx.stack(out, axis=0) + + +def bench_mlx(a, b, segments, warmup, iters): + for _ in range(warmup): + y = mx.segmented_mm(a, b, segments) + mx.eval(y) + mx.synchronize() + + start = time.perf_counter() + for _ in range(iters): + y = mx.segmented_mm(a, b, segments) + mx.eval(y) + mx.synchronize() + end = time.perf_counter() + return (end - start) * 1e3 / iters + + +def bench_mlx_loop(a, b, segments, warmup, iters): + for _ in range(warmup): + y = mlx_segmented_mm_loop(a, b, segments) + mx.eval(y) + mx.synchronize() + + start = time.perf_counter() + for _ in range(iters): + y = mlx_segmented_mm_loop(a, b, segments) + mx.eval(y) + mx.synchronize() + end = time.perf_counter() + return (end - start) * 1e3 / iters + + +def print_table(headers, rows): + widths = [len(h) for h in headers] + for row in rows: + for i, cell in enumerate(row): + widths[i] = max(widths[i], len(cell)) + + def fmt_row(row): + return ( + "| " + + " | ".join(f"{cell:<{widths[i]}}" for i, cell in enumerate(row)) + + " |" + ) + + sep = "|-" + "-|-".join("-" * w for w in widths) + "-|" + print(fmt_row(headers)) + print(sep) + for row in rows: + print(fmt_row(row)) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument( + "--cases", + default=( + "128x128x1024x16," + "128x128x1024x32," + "256x256x2048x16," + "512x512x4096x32," + "1024x1024x4096x32," + "1024x1024x8192x64" + ), + help="Comma-separated MxNxKxS list.", + ) + parser.add_argument( + "--dtype", + default="float32", + choices=["float16", "bfloat16", "float32"], + ) + parser.add_argument("--warmup", type=int, default=10) + parser.add_argument("--iters", type=int, default=50) + parser.add_argument( + "--segments", + choices=["equal", "random"], + default="random", + help="Segment generation pattern.", + ) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--no-check", action="store_true") + args = parser.parse_args() + + mlx_dtype = MLX_DTYPES[args.dtype] + + print( + f"dtype={args.dtype} warmup={args.warmup} iters={args.iters} segments={args.segments}" + ) + + headers = [ + "Case", + "MLX ms", + "Loop ms", + "Speedup", + "MLX err", + "Loop err", + ] + rows = [] + + cases = parse_cases(args.cases) + for idx, (m, n, k, s) in enumerate(cases): + rng = np.random.default_rng(args.seed + idx) + a_np = rng.standard_normal((m, k)).astype(np.float32) + b_np = rng.standard_normal((k, n)).astype(np.float32) + seg_np = make_segments(k, s, args.segments, args.seed + idx) + + a_mx = mx.array(a_np, dtype=mlx_dtype) + b_mx = mx.array(b_np, dtype=mlx_dtype) + seg_mx = mx.array(seg_np, dtype=mx.uint32) + mx.eval(a_mx, b_mx, seg_mx) + + mlx_err_str = "" + loop_err_str = "" + if not args.no_check: + y_mlx = mx.segmented_mm(a_mx, b_mx, seg_mx) + y_loop = mlx_segmented_mm_loop(a_mx, b_mx, seg_mx) + mx.eval(y_mlx, y_loop) + + if args.dtype == "float32": + ref = numpy_segmented_mm_ref( + a_np.astype(np.float64), + b_np.astype(np.float64), + seg_np.tolist(), + ) + mlx_err = np.max(np.abs(np.array(y_mlx, dtype=np.float64) - ref)) + loop_err = np.max(np.abs(np.array(y_loop, dtype=np.float64) - ref)) + else: + a_mx_f32 = mx.array(a_np, dtype=mx.float32) + b_mx_f32 = mx.array(b_np, dtype=mx.float32) + ref = mx.segmented_mm(a_mx_f32, b_mx_f32, seg_mx) + mx.eval(ref) + mlx_err = float(mx.max(mx.abs(ref - y_mlx.astype(mx.float32))).item()) + loop_err = float(mx.max(mx.abs(ref - y_loop.astype(mx.float32))).item()) + mlx_err_str = f"{mlx_err:.2e}" + loop_err_str = f"{loop_err:.2e}" + + t_mlx = bench_mlx(a_mx, b_mx, seg_mx, args.warmup, args.iters) + t_loop = bench_mlx_loop(a_mx, b_mx, seg_mx, args.warmup, args.iters) + ratio = t_loop / t_mlx if t_mlx > 0 else float("inf") + rows.append( + [ + f"{m}x{n}x{k}x{s}", + f"{t_mlx:.3f}", + f"{t_loop:.3f}", + f"{ratio:.2f}x", + mlx_err_str, + loop_err_str, + ] + ) + + print_table(headers, rows) + if not args.no_check: + if args.dtype == "float32": + print("err: max|result - numpy_fp64_ref|") + else: + print("err: max|result - own_fp32_result|") + + +if __name__ == "__main__": + main() diff --git a/mlx/backend/cuda/gemms/grouped_gemm.h b/mlx/backend/cuda/gemms/grouped_gemm.h index 308a1ba9..844b8f44 100644 --- a/mlx/backend/cuda/gemms/grouped_gemm.h +++ b/mlx/backend/cuda/gemms/grouped_gemm.h @@ -22,4 +22,18 @@ void cutlass_grouped_gemm_unaligned( array& out, cu::CommandEncoder& encoder); +void cutlass_segmented_mm( + bool a_transposed, + int lda, + bool b_transposed, + int ldb, + int num_segments, + int M, + int N, + const array& a, + const array& b, + const array& segments, + array& out, + cu::CommandEncoder& encoder); + } // namespace mlx::core diff --git a/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu b/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu index 5c06704f..bf17f597 100644 --- a/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu +++ b/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu @@ -93,6 +93,50 @@ __global__ void prepare_grouped_mm_data( } } +__global__ void prepare_segmented_mm_data( + const uint32_t* segments, + int num_segments, + int M, + int N, + int lda, + int ldb, + int item_size, + bool a_transposed, + bool b_transposed, + int8_t* a_start, + int8_t* b_start, + int8_t* out_start, + ProblemSize* problem_sizes, + int64_t* a_lds, + int64_t* b_lds, + int64_t* out_lds, + void** a_ptrs, + void** b_ptrs, + void** out_ptrs) { + int idx = cg::this_grid().thread_rank(); + if (idx >= num_segments) + return; + + int64_t start = segments[2 * idx]; + int64_t end = segments[2 * idx + 1]; + int K_i = (end > start) ? static_cast(end - start) : 0; + + problem_sizes[idx] = {M, N, K_i}; + a_lds[idx] = lda; + b_lds[idx] = ldb; + out_lds[idx] = N; + + // Offset into K dimension depends on layout: + // A [M,K]: row-major offset = start, col-major offset = start * lda + // B [K,N]: row-major offset = start * ldb, col-major offset = start + int64_t a_offset = a_transposed ? start * lda : start; + int64_t b_offset = b_transposed ? start : start * ldb; + + a_ptrs[idx] = a_start + a_offset * item_size; + b_ptrs[idx] = b_start + b_offset * item_size; + out_ptrs[idx] = out_start + static_cast(idx) * M * N * item_size; +} + } // namespace cu namespace { @@ -357,4 +401,85 @@ void cutlass_grouped_gemm_unaligned( encoder); } +void cutlass_segmented_mm( + bool a_transposed, + int lda, + bool b_transposed, + int ldb, + int num_segments, + int M, + int N, + const array& a, + const array& b, + const array& segments, + array& out, + cu::CommandEncoder& encoder) { + // Allocate grouped GEMM args on device. + int problem_sizes_nbytes = + num_segments * cuda::ceil_div(sizeof(ProblemSize), 8) * 8; + int nbytes = problem_sizes_nbytes + + num_segments * (3 * sizeof(void*) + 3 * sizeof(int64_t)); + nbytes = cuda::ceil_div(nbytes, 256) * 256; + array gemm_args(cu::malloc_async(nbytes, encoder), {nbytes}, int8); + encoder.add_temporary(gemm_args); + + ProblemSize* problem_sizes = gpu_ptr(gemm_args); + int64_t* a_lds = gpu_ptr(gemm_args) + problem_sizes_nbytes / 8; + int64_t* b_lds = a_lds + num_segments; + int64_t* out_lds = b_lds + num_segments; + void** a_ptrs = reinterpret_cast(out_lds + num_segments); + void** b_ptrs = a_ptrs + num_segments; + void** out_ptrs = b_ptrs + num_segments; + + // Build problem descriptions from segments on the GPU. + int block_size = std::min(num_segments, 256); + int num_blocks = cuda::ceil_div(num_segments, block_size); + + encoder.set_input_array(segments); + encoder.set_output_array(gemm_args); + encoder.add_kernel_node_ex( + cu::prepare_segmented_mm_data, + dim3(num_blocks), + dim3(block_size), + {}, + 0, + gpu_ptr(segments), + num_segments, + M, + N, + static_cast(lda), + static_cast(ldb), + static_cast(out.itemsize()), + a_transposed, + b_transposed, + gpu_ptr(a), + gpu_ptr(b), + gpu_ptr(out), + problem_sizes, + a_lds, + b_lds, + out_lds, + a_ptrs, + b_ptrs, + out_ptrs); + + // Dispatch grouped GEMM. + encoder.set_input_array(a); + encoder.set_input_array(b); + encoder.set_input_array(gemm_args); + encoder.set_output_array(out); + auto* fun = get_grouped_mm_funcion(a.dtype(), N, encoder.device()); + fun(a_transposed, + b_transposed, + num_segments, + problem_sizes, + a_lds, + b_lds, + out_lds, + a_ptrs, + b_ptrs, + out_ptrs, + encoder); +} + } // namespace mlx::core diff --git a/mlx/backend/cuda/matmul.cpp b/mlx/backend/cuda/matmul.cpp index 725f1b40..83365905 100644 --- a/mlx/backend/cuda/matmul.cpp +++ b/mlx/backend/cuda/matmul.cpp @@ -370,4 +370,54 @@ void GatherMM::eval_gpu(const std::vector& inputs, array& out) { throw std::runtime_error("NYI"); } +void SegmentedMM::eval_gpu(const std::vector& inputs, array& out) { + nvtx3::scoped_range r("SegmentedMM::eval_gpu"); + auto& s = stream(); + auto& encoder = cu::get_command_encoder(s); + + assert(inputs.size() == 3); + auto& a_pre = inputs[0]; + auto& b_pre = inputs[1]; + auto& segments_pre = inputs[2]; + + // Return zeros if output is empty or either input is empty. + if (out.size() == 0 || a_pre.size() == 0 || b_pre.size() == 0) { + array zero(0, a_pre.dtype()); + encoder.add_temporary(zero); + fill_gpu(zero, out, s); + return; + } + + out.set_data(cu::malloc_async(out.nbytes(), encoder)); + + int M = a_pre.shape(-2); + int N = b_pre.shape(-1); + int num_segments = segments_pre.size() / 2; + + auto [a_transposed, lda, a] = check_transpose(encoder, s, a_pre); + auto [b_transposed, ldb, b] = check_transpose(encoder, s, b_pre); + auto segments = [&] { + if (segments_pre.flags().row_contiguous) { + return segments_pre; + } + array copy = contiguous_copy_gpu(segments_pre, s); + encoder.add_temporary(copy); + return copy; + }(); + + cutlass_segmented_mm( + a_transposed, + lda, + b_transposed, + ldb, + num_segments, + M, + N, + a, + b, + segments, + out, + encoder); +} + } // namespace mlx::core diff --git a/mlx/backend/cuda/primitives.cpp b/mlx/backend/cuda/primitives.cpp index e29b3e50..1321ec38 100644 --- a/mlx/backend/cuda/primitives.cpp +++ b/mlx/backend/cuda/primitives.cpp @@ -29,7 +29,6 @@ NO_GPU(FFT) NO_GPU(GatherQMM) NO_GPU_MULTI(LUF) NO_GPU_MULTI(QRF) -NO_GPU(SegmentedMM) NO_GPU_MULTI(SVD) NO_GPU(Inverse) NO_GPU(Cholesky) diff --git a/python/tests/cuda_skip.py b/python/tests/cuda_skip.py index 000cc84f..888be52b 100644 --- a/python/tests/cuda_skip.py +++ b/python/tests/cuda_skip.py @@ -6,8 +6,6 @@ cuda_skip = { "TestBlas.test_gather_matmul", "TestBlas.test_gather_matmul_grad", "TestBlas.test_gather_mm_sorted_vjp", - # Segmented matmul NYI - "TestBlas.test_segmented_mm", # FFTs NYI "TestFFT.test_fft", "TestFFT.test_fft_big_powers_of_two", diff --git a/python/tests/test_conv.py b/python/tests/test_conv.py index 090cf0f3..80b7d6b6 100644 --- a/python/tests/test_conv.py +++ b/python/tests/test_conv.py @@ -1151,7 +1151,7 @@ class TestConv(mlx_tests.MLXTestCase): ) self.assertEqual(grads.shape, k_shape) - def test_1d_conv_with_2d(self): + def test_conv_1d_with_2d(self): x = mx.random.uniform(shape=(2, 10, 16)) y = mx.random.normal(shape=(16, 3, 16))