mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-12 21:38:59 -04:00
Add MLX tinygrad interop benchmark workspace
This commit is contained in:
1 parent
9ea9bf0aca
commit
4ab2b7efda
4 files changed
+355
No files matched your search
@@ -0,0 +1,134 @@
|
||||
# Tensor Conversion Benchmark Notes
|
||||
|
||||
## Current Goal
|
||||
|
||||
Benchmark raw tinygrad `<->` MLX tensor transformation latency on Apple Silicon
|
||||
for tensors that are already:
|
||||
|
||||
- synchronized
|
||||
- allocated
|
||||
- materialized / realized
|
||||
|
||||
The timed region should measure only the transformation itself.
|
||||
|
||||
## Repo Layout
|
||||
|
||||
- Root notes file:
|
||||
- `CONVERSION_BENCH_NOTES.md`
|
||||
- Interop code:
|
||||
- `mlx_tinygrad_interop/`
|
||||
- Historical benchmark kept as-is:
|
||||
- `tmp/bench_pingpong.py`
|
||||
|
||||
## Current Fast-Path Design
|
||||
|
||||
The first direct benchmark path is intentionally narrow.
|
||||
|
||||
- `MLX -> tinygrad`
|
||||
- export MLX Metal storage metadata
|
||||
- import into tinygrad by aliasing the existing `MTLBuffer*`
|
||||
- `tinygrad -> MLX`
|
||||
- export tinygrad Metal storage metadata
|
||||
- import into MLX by wrapping the underlying unified-memory pointer with a
|
||||
no-copy MLX array constructor path
|
||||
|
||||
This is asymmetric internally, but both directions aim to avoid copying tensor
|
||||
bytes.
|
||||
|
||||
## Implemented Helper Surface
|
||||
|
||||
- MLX
|
||||
- `mx.metal._unsafe_export_storage(array)`
|
||||
- `mx.metal._unsafe_array_from_ptr(raw_ptr, shape, dtype, owner=None)`
|
||||
- tinygrad
|
||||
- `Tensor._unsafe_metal_storage()`
|
||||
- `Tensor._unsafe_from_metal_buffer(mtl_buffer_ptr, shape, dtype=..., byte_offset=0, owner=None)`
|
||||
|
||||
These helpers are intentionally private and unsafe.
|
||||
|
||||
## Temporary Eligibility Rules
|
||||
|
||||
The current direct path should only accept tensors that are:
|
||||
|
||||
- backed by Metal storage
|
||||
- already realized / available
|
||||
- single-device
|
||||
- dense row-major contiguous
|
||||
- concrete-shaped
|
||||
- dtype-compatible without conversion
|
||||
|
||||
Non-contiguous views, broadcasts, dtype casts, and multi-device tensors should
|
||||
fall back to slower paths.
|
||||
|
||||
## Workflow
|
||||
|
||||
Use the `exo` devshell and `uv` workflow.
|
||||
|
||||
1. Change code locally.
|
||||
2. Push the updated `mlx` and `tinygrad` fork branches.
|
||||
3. On the remote Mac, pull the updated repos.
|
||||
4. Enter the devshell with `nix develop`.
|
||||
5. Refresh dependency resolution with `uv lock && uv sync`.
|
||||
6. Run tests and benchmarks with `uv run ...`.
|
||||
|
||||
Do not rely on ad-hoc per-host build environments when the flake / devshell can
|
||||
carry the needed toolchain.
|
||||
|
||||
## Known Nuances / Footguns
|
||||
|
||||
- Unified memory does not mean both frameworks consume shared storage in the
|
||||
same way. Metal kernels still bind `MTLBuffer` objects.
|
||||
- Synchronization can dominate measured latency if it leaks into the timed path.
|
||||
- Python overhead matters at the `1-10 us` scale, so helper calls and wrapper
|
||||
construction can dominate tiny tensors even when no tensor bytes are copied.
|
||||
- tinygrad tensors are graph objects, but once realized they do have concrete
|
||||
underlying storage.
|
||||
- External mutation and aliasing can bypass autograd expectations in both
|
||||
frameworks.
|
||||
- The current fast path is asymmetric:
|
||||
- MLX exports `MTLBuffer*` for the `MLX -> tinygrad` direction.
|
||||
- tinygrad exports raw unified-memory pointer for the `tinygrad -> MLX`
|
||||
direction.
|
||||
- The first tinygrad import helper supports byte offsets.
|
||||
- The first MLX import helper is raw-pointer based rather than foreign
|
||||
`MTLBuffer*` based.
|
||||
- `mx.metal._unsafe_export_storage(...)` currently expects an MLX array that is
|
||||
already in the C++ `available` state. In practice, `mx.array(np_array)` met
|
||||
that precondition for local smoke testing, while `mx.arange(...)` did not.
|
||||
|
||||
## Current Findings
|
||||
|
||||
The unsafe bridge was smoke-tested successfully on a remote Mac:
|
||||
|
||||
- `MLX -> tinygrad` direct alias path returned correct values.
|
||||
- `tinygrad -> MLX` direct alias path returned correct values.
|
||||
|
||||
Preliminary latency measurements for `float32` and `7168` bytes were:
|
||||
|
||||
- `direct_alias`
|
||||
- `mlx_to_tinygrad`: about `33.5 us`
|
||||
- `tinygrad_to_mlx`: about `28.0 us`
|
||||
- `numpy_fallback`
|
||||
- `mlx_to_tinygrad`: about `269 us`
|
||||
- `tinygrad_to_mlx`: about `12.9 us`
|
||||
|
||||
These numbers were gathered before switching fully to the repo-standard
|
||||
`nix develop` + `uv run` flow, so they should be treated as provisional rather
|
||||
than canonical benchmark results.
|
||||
|
||||
## Near-Term Plan
|
||||
|
||||
1. Push the current helper changes in `mlx` and `tinygrad`.
|
||||
2. Re-run the benchmark from `mlx_tinygrad_interop/` using the repo-standard
|
||||
remote flow.
|
||||
3. Separate true alias-path cost from Python wrapper overhead more aggressively
|
||||
if the current numbers remain too high.
|
||||
|
||||
## Open Questions
|
||||
|
||||
- Whether MLX should eventually import foreign `MTLBuffer*` handles directly,
|
||||
rather than only raw pointers, for a more symmetric bridge.
|
||||
- Whether the first fast path should support contiguous slices with byte
|
||||
offsets, or only base-contiguous tensors.
|
||||
- Whether a native-copy middle path should be benchmarked immediately, or only
|
||||
the direct path and Python / NumPy fallback.
|
||||
@@ -0,0 +1,37 @@
|
||||
# MLX Tinygrad Interop
|
||||
|
||||
Private code for benchmarking MLX `<->` tinygrad tensor conversions.
|
||||
|
||||
## Workflow
|
||||
|
||||
Use the repo devshell and top-level dependency graph. Do not install ad-hoc
|
||||
build dependencies or patch around them with one-off environment setups.
|
||||
|
||||
1. Change code locally.
|
||||
2. Push the `mlx` and `tinygrad` fork changes.
|
||||
3. On the remote Mac, pull the updated repos.
|
||||
4. Enter the devshell with `nix develop`.
|
||||
5. Refresh dependencies with `uv lock && uv sync`.
|
||||
6. Run tests or benchmarks with `uv run ...`.
|
||||
|
||||
## Benchmark
|
||||
|
||||
The current benchmark script isolates the raw transformation only.
|
||||
|
||||
- Inputs are assumed to already be synchronized.
|
||||
- Inputs are assumed to already be allocated.
|
||||
- Inputs are assumed to already be materialized / realized.
|
||||
- Setup stays outside the timed loop.
|
||||
|
||||
Example:
|
||||
|
||||
```bash
|
||||
uv run python mlx_tinygrad_interop/bench_raw_conversion.py --dtype float32 --sizes 256,512,1024,2048,4096,7168
|
||||
```
|
||||
|
||||
## Current Scope
|
||||
|
||||
- Private / unsafe helpers only.
|
||||
- Metal / unified-memory path only.
|
||||
- Dense contiguous tensors only.
|
||||
- Same-dtype conversions only.
|
||||
@@ -0,0 +1 @@
|
||||
"""Private MLX <-> tinygrad interop experiments and benchmarks."""
|
||||
@@ -0,0 +1,183 @@
|
||||
import argparse
|
||||
import gc
|
||||
import platform
|
||||
import statistics
|
||||
import sys
|
||||
import time
|
||||
from typing import Any, Callable, cast
|
||||
|
||||
import mlx.core as mx
|
||||
import numpy as np
|
||||
from tinygrad import Device, Tensor, dtypes
|
||||
from tinygrad.device import Buffer
|
||||
|
||||
blackhole: Any = None
|
||||
|
||||
DTYPES: dict[str, tuple[Any, Any, np.dtype[Any]]] = {
|
||||
"float16": (mx.float16, dtypes.float16, np.dtype(np.float16)),
|
||||
"float32": (mx.float32, dtypes.float32, np.dtype(np.float32)),
|
||||
"int32": (mx.int32, dtypes.int32, np.dtype(np.int32)),
|
||||
"uint8": (mx.uint8, dtypes.uint8, np.dtype(np.uint8)),
|
||||
}
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Benchmark raw tinygrad <-> MLX tensor conversion overhead.")
|
||||
parser.add_argument("--dtype", choices=sorted(DTYPES), default="float32")
|
||||
parser.add_argument("--sizes", default="256,512,1024,2048,4096,7168,8192,16384,32768,65536,262144,1048576",
|
||||
help="Comma-separated tensor sizes in bytes.")
|
||||
parser.add_argument("--warmup", type=int, default=128)
|
||||
parser.add_argument("--samples", type=int, default=12)
|
||||
parser.add_argument("--min-batch-us", type=float, default=2000.0,
|
||||
help="Minimum target batch duration per sample in microseconds.")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def bytes_view(mv: memoryview) -> memoryview:
|
||||
return mv if mv.format == "B" and mv.ndim == 1 else mv.cast("B")
|
||||
|
||||
|
||||
def tinygrad_zero_copy_memoryview(t: Tensor) -> memoryview:
|
||||
assert t.device == "METAL", f"expected METAL tensor, got {t.device}"
|
||||
buf = cast(Buffer, t.uop.buffer).ensure_allocated()
|
||||
assert t.dtype.base.fmt is not None, f"no buffer format for dtype {t.dtype.base}"
|
||||
return buf.as_memoryview(force_zero_copy=True).cast(t.dtype.base.fmt, t.shape)
|
||||
|
||||
|
||||
def tinygrad_from_mlx_direct(x: Any, tg_dtype: Any) -> Tensor:
|
||||
storage = mx.metal._unsafe_export_storage(x)
|
||||
return Tensor._unsafe_from_metal_buffer(
|
||||
int(storage["mtl_buffer_ptr"]),
|
||||
tuple(storage["shape"]),
|
||||
dtype=tg_dtype,
|
||||
byte_offset=int(storage["offset_bytes"]),
|
||||
owner=x,
|
||||
)
|
||||
|
||||
|
||||
def mlx_from_tinygrad_direct(t: Tensor, mx_dtype: Any) -> Any:
|
||||
storage = t._unsafe_metal_storage()
|
||||
return mx.metal._unsafe_array_from_ptr(
|
||||
int(storage["raw_ptr"]),
|
||||
tuple(storage["shape"]),
|
||||
mx_dtype,
|
||||
owner=t,
|
||||
)
|
||||
|
||||
|
||||
def tinygrad_from_mlx_copy(x: Any, tg_dtype: Any) -> Tensor:
|
||||
out = Tensor.empty(*tuple(int(dim) for dim in x.shape), device="METAL", dtype=tg_dtype)
|
||||
cast(Buffer, out.uop.buffer).ensure_allocated().copyin(bytes_view(memoryview(x)))
|
||||
return out
|
||||
|
||||
|
||||
def mlx_from_tinygrad_copy(t: Tensor) -> Any:
|
||||
return mx.array(tinygrad_zero_copy_memoryview(t))
|
||||
|
||||
|
||||
def tinygrad_from_mlx_numpy(x: Any) -> Tensor:
|
||||
out = Tensor(np.array(x, copy=True), device="METAL")
|
||||
out.realize()
|
||||
Device["METAL"].synchronize()
|
||||
return out
|
||||
|
||||
|
||||
def mlx_from_tinygrad_numpy(t: Tensor) -> Any:
|
||||
return mx.array(t.numpy())
|
||||
|
||||
|
||||
def bench_callable(fn: Callable[[], Any], warmup: int, samples: int, min_batch_us: float) -> dict[str, float]:
|
||||
global blackhole
|
||||
for _ in range(warmup):
|
||||
blackhole = fn()
|
||||
|
||||
min_batch_ns = int(min_batch_us * 1000.0)
|
||||
iters = 1
|
||||
while True:
|
||||
start = time.perf_counter_ns()
|
||||
for _ in range(iters):
|
||||
blackhole = fn()
|
||||
elapsed = time.perf_counter_ns() - start
|
||||
if elapsed >= min_batch_ns or iters >= (1 << 20):
|
||||
break
|
||||
iters *= 2
|
||||
|
||||
vals_us: list[float] = []
|
||||
for _ in range(samples):
|
||||
start = time.perf_counter_ns()
|
||||
for _ in range(iters):
|
||||
blackhole = fn()
|
||||
elapsed = time.perf_counter_ns() - start
|
||||
vals_us.append(elapsed / iters / 1000.0)
|
||||
|
||||
return {
|
||||
"iters": float(iters),
|
||||
"min_us": min(vals_us),
|
||||
"median_us": statistics.median(vals_us),
|
||||
"mean_us": statistics.mean(vals_us),
|
||||
"stdev_us": statistics.stdev(vals_us) if len(vals_us) > 1 else 0.0,
|
||||
}
|
||||
|
||||
|
||||
def print_header(dtype_name: str) -> None:
|
||||
print(f"# python={platform.python_version()} platform={platform.platform()}")
|
||||
print(f"# dtype={dtype_name} mlx_metal_available={mx.metal.is_available()} tinygrad_device=METAL")
|
||||
print("# sizes are source tensor sizes in bytes")
|
||||
print("size_bytes,method,direction,min_us,median_us,mean_us,stdev_us,iters")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
mx_dtype, tg_dtype, np_dtype = DTYPES[args.dtype]
|
||||
sizes = [int(x.strip()) for x in args.sizes.split(",") if x.strip()]
|
||||
|
||||
required = [
|
||||
("mx.metal._unsafe_export_storage", getattr(mx.metal, "_unsafe_export_storage", None)),
|
||||
("mx.metal._unsafe_array_from_ptr", getattr(mx.metal, "_unsafe_array_from_ptr", None)),
|
||||
("Tensor._unsafe_from_metal_buffer", getattr(Tensor, "_unsafe_from_metal_buffer", None)),
|
||||
("Tensor._unsafe_metal_storage", getattr(Tensor, "_unsafe_metal_storage", None)),
|
||||
]
|
||||
missing = [name for name, value in required if value is None]
|
||||
if missing:
|
||||
raise RuntimeError(f"Missing required helper(s): {', '.join(missing)}")
|
||||
|
||||
gc.disable()
|
||||
try:
|
||||
print_header(args.dtype)
|
||||
for size_bytes in sizes:
|
||||
if size_bytes <= 0:
|
||||
continue
|
||||
if size_bytes % np_dtype.itemsize != 0:
|
||||
print(f"# skipping size {size_bytes}: not divisible by dtype itemsize {np_dtype.itemsize}", file=sys.stderr)
|
||||
continue
|
||||
|
||||
numel = size_bytes // np_dtype.itemsize
|
||||
|
||||
# Build one realized source tensor on each side. The timed region measures
|
||||
# only the transformation, not creation / synchronization.
|
||||
src_mx = mx.array(np.arange(numel, dtype=np_dtype), dtype=mx_dtype)
|
||||
|
||||
src_tg = Tensor(np.arange(numel, dtype=np_dtype), device="METAL").realize()
|
||||
Device["METAL"].synchronize()
|
||||
|
||||
benches: list[tuple[str, str, Callable[[], Any]]] = [
|
||||
("direct_alias", "mlx_to_tinygrad", lambda s=src_mx: tinygrad_from_mlx_direct(s, tg_dtype)),
|
||||
("memoryview_copy", "mlx_to_tinygrad", lambda s=src_mx: tinygrad_from_mlx_copy(s, tg_dtype)),
|
||||
("numpy_fallback", "mlx_to_tinygrad", lambda s=src_mx: tinygrad_from_mlx_numpy(s)),
|
||||
("direct_alias", "tinygrad_to_mlx", lambda s=src_tg: mlx_from_tinygrad_direct(s, mx_dtype)),
|
||||
("memoryview_copy", "tinygrad_to_mlx", lambda s=src_tg: mlx_from_tinygrad_copy(s)),
|
||||
("numpy_fallback", "tinygrad_to_mlx", lambda s=src_tg: mlx_from_tinygrad_numpy(s)),
|
||||
]
|
||||
|
||||
for method, direction, fn in benches:
|
||||
stats = bench_callable(fn, warmup=args.warmup, samples=args.samples, min_batch_us=args.min_batch_us)
|
||||
print(
|
||||
f"{size_bytes},{method},{direction},"
|
||||
f"{stats['min_us']:.3f},{stats['median_us']:.3f},{stats['mean_us']:.3f},{stats['stdev_us']:.3f},{int(stats['iters'])}"
|
||||
)
|
||||
finally:
|
||||
gc.enable()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in new issue
Block a user