Skip to content

Commit 6c6ce44

Browse files
committed
bench: add linalg inverse device benchmark
1 parent 568e514 commit 6c6ce44

1 file changed

Lines changed: 116 additions & 0 deletions

File tree

benchmarks/python/inv_bench.py

Lines changed: 116 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,116 @@
1+
# Copyright © 2026 Apple Inc.
2+
3+
import argparse
4+
import platform
5+
import time
6+
7+
import mlx.core as mx
8+
import numpy as np
9+
10+
11+
def make_input(order, batch_size, seed):
12+
rng = np.random.default_rng(seed)
13+
shape = (order, order) if batch_size == 1 else (batch_size, order, order)
14+
matrix = rng.standard_normal(shape, dtype=np.float32)
15+
return matrix + order * np.eye(order, dtype=np.float32)
16+
17+
18+
def array_on_device(matrix, device):
19+
previous_device = mx.default_device()
20+
mx.set_default_device(device)
21+
try:
22+
result = mx.array(matrix)
23+
mx.eval(result)
24+
mx.synchronize(device)
25+
finally:
26+
mx.set_default_device(previous_device)
27+
return result
28+
29+
30+
def benchmark(a, device, warmup, iters):
31+
for _ in range(warmup):
32+
output = mx.linalg.inv(a, stream=device)
33+
mx.eval(output)
34+
mx.synchronize(device)
35+
36+
samples = []
37+
for _ in range(iters):
38+
start = time.perf_counter_ns()
39+
output = mx.linalg.inv(a, stream=device)
40+
mx.eval(output)
41+
mx.synchronize(device)
42+
samples.append((time.perf_counter_ns() - start) / 1e6)
43+
return float(np.median(samples))
44+
45+
46+
def print_table(headers, rows):
47+
widths = [len(header) for header in headers]
48+
for row in rows:
49+
for index, cell in enumerate(row):
50+
widths[index] = max(widths[index], len(cell))
51+
52+
def format_row(row):
53+
return (
54+
"| "
55+
+ " | ".join(f"{cell:<{widths[index]}}" for index, cell in enumerate(row))
56+
+ " |"
57+
)
58+
59+
print(format_row(headers))
60+
print("|-" + "-|-".join("-" * width for width in widths) + "-|")
61+
for row in rows:
62+
print(format_row(row))
63+
64+
65+
def main():
66+
parser = argparse.ArgumentParser(
67+
description="Compare mx.linalg.inv on the CPU and Metal backends."
68+
)
69+
parser.add_argument("--sizes", default="1,3,16,64,256,1024")
70+
parser.add_argument("--batch-size", type=int, default=1)
71+
parser.add_argument("--warmup", type=int, default=5)
72+
parser.add_argument("--iters", type=int, default=20)
73+
parser.add_argument("--seed", type=int, default=0)
74+
args = parser.parse_args()
75+
76+
if not mx.metal.is_available():
77+
raise RuntimeError("This benchmark requires a Metal device.")
78+
if args.batch_size < 1:
79+
raise ValueError("--batch-size must be positive.")
80+
81+
orders = [int(value) for value in args.sizes.split(",")]
82+
if any(order < 1 for order in orders):
83+
raise ValueError("--sizes must contain positive matrix orders.")
84+
85+
print(
86+
f"machine={platform.machine()} system={platform.platform()} "
87+
f"dtype=float32 batch_size={args.batch_size} warmup={args.warmup} "
88+
f"iters={args.iters}"
89+
)
90+
91+
rows = []
92+
for index, order in enumerate(orders):
93+
matrix = make_input(order, args.batch_size, args.seed + index)
94+
cpu_time = benchmark(
95+
array_on_device(matrix, mx.cpu), mx.cpu, args.warmup, args.iters
96+
)
97+
metal_time = benchmark(
98+
array_on_device(matrix, mx.gpu), mx.gpu, args.warmup, args.iters
99+
)
100+
shape = f"{order}x{order}"
101+
if args.batch_size > 1:
102+
shape = f"{args.batch_size}x{shape}"
103+
rows.append(
104+
[
105+
shape,
106+
f"{cpu_time:.3f}",
107+
f"{metal_time:.3f}",
108+
f"{cpu_time / metal_time:.2f}x",
109+
]
110+
)
111+
112+
print_table(["Shape", "CPU ms", "Metal ms", "CPU / Metal"], rows)
113+
114+
115+
if __name__ == "__main__":
116+
main()

0 commit comments

Comments
 (0)