diff --git a/scripts/mlxprobe.py b/scripts/mlxprobe.py new file mode 100644 index 0000000..2400b94 --- /dev/null +++ b/scripts/mlxprobe.py @@ -0,0 +1,90 @@ +"""MLX fleet probe — one JSON line of capability + micro-bench per machine. + +Benches (each ~seconds, sized for 16GB boxes): + matmul fp16/bf16 4096^2 -> TFLOPS (bf16-on-M1 question) + elementwise add 256MB -> GB/s (memory bandwidth proxy) + sdpa head_dim 56 vs 64 -> ratio (fused fast-path check, corridorkey lesson) +""" +import json +import platform +import subprocess +import time + +R = {"host": platform.node().split(".")[0]} +try: + R["chip"] = subprocess.check_output( + ["sysctl", "-n", "machdep.cpu.brand_string"], text=True).strip() + R["ram_gb"] = round(int(subprocess.check_output( + ["sysctl", "-n", "hw.memsize"], text=True)) / 2**30) + R["macos"] = subprocess.check_output( + ["sw_vers", "-productVersion"], text=True).strip() +except Exception: + pass + +try: + import mlx.core as mx + R["mlx"] = getattr(mx, "__version__", "?") +except Exception as e: + R["mlx"] = None + R["error"] = f"mlx import failed: {e}" + print(json.dumps(R)) + raise SystemExit(0) + +try: + di = mx.metal.device_info() if hasattr(mx, "metal") else {} + R["gpu"] = {k: di[k] for k in ("architecture", "max_recommended_working_set_size", + "memory_size") if k in di} +except Exception: + pass + + +def timeit(fn, warmup=3, iters=10): + for _ in range(warmup): + mx.eval(fn()) + t0 = time.perf_counter() + for _ in range(iters): + mx.eval(fn()) + return (time.perf_counter() - t0) / iters + + +N = 4096 +flops = 2 * N * N * N +for dt, name in ((mx.float16, "fp16"), (mx.bfloat16, "bf16")): + try: + a = mx.random.normal((N, N)).astype(dt) + b = mx.random.normal((N, N)).astype(dt) + mx.eval(a, b) + t = timeit(lambda: a @ b) + R[f"matmul_{name}_tflops"] = round(flops / t / 1e12, 2) + except Exception as e: + R[f"matmul_{name}_tflops"] = f"fail: {e}" + +try: + M = 64 * 1024 * 1024 # 64M floats = 256MB per array + x = mx.random.normal((M,)) + y = mx.random.normal((M,)) + mx.eval(x, y) + t = timeit(lambda: x + y, warmup=2, iters=8) + R["bandwidth_gbs"] = round(3 * M * 4 / t / 1e9) # 2 reads + 1 write, fp32 +except Exception as e: + R["bandwidth_gbs"] = f"fail: {e}" + +try: + B, H, L = 1, 8, 2048 + sdpa = {} + for D in (56, 64): + q = mx.random.normal((B, H, L, D)).astype(mx.float16) + k = mx.random.normal((B, H, L, D)).astype(mx.float16) + v = mx.random.normal((B, H, L, D)).astype(mx.float16) + mx.eval(q, k, v) + t = timeit(lambda: mx.fast.scaled_dot_product_attention( + q, k, v, scale=D ** -0.5), warmup=3, iters=20) + sdpa[f"d{D}_us"] = round(t * 1e6) + sdpa["d56_vs_d64"] = round(sdpa["d56_us"] / max(sdpa["d64_us"], 1), 2) + R["sdpa"] = sdpa +except Exception as e: + R["sdpa"] = f"fail: {e}" + +R["peak_mem_gb"] = round(mx.get_peak_memory() / 2**30, 2) \ + if hasattr(mx, "get_peak_memory") else None +print(json.dumps(R))