[CUDA] Implement SegmentedMM (#3238)
This commit is contained in:
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user