[CUDA] Implement SegmentedMM (#3238)

This commit is contained in:
Long Yixing
2026-03-11 13:31:43 -07:00
committed by GitHub
parent 1c2d7041ab
commit a9573f92f6
7 changed files with 399 additions and 4 deletions
+209
View File
@@ -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()
+14
View File
@@ -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
@@ -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<int>(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<int64_t>(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<ProblemSize>(gemm_args);
int64_t* a_lds = gpu_ptr<int64_t>(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<void**>(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<uint32_t>(segments),
num_segments,
M,
N,
static_cast<int>(lda),
static_cast<int>(ldb),
static_cast<int>(out.itemsize()),
a_transposed,
b_transposed,
gpu_ptr<int8_t>(a),
gpu_ptr<int8_t>(b),
gpu_ptr<int8_t>(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
+50
View File
@@ -370,4 +370,54 @@ void GatherMM::eval_gpu(const std::vector<array>& inputs, array& out) {
throw std::runtime_error("NYI");
}
void SegmentedMM::eval_gpu(const std::vector<array>& 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
-1
View File
@@ -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)
-2
View File
@@ -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",
+1 -1
View File
@@ -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))