Improve warmup for memory bandwidth profiling

- Add 3 full warmup iterations before benchmarking
- Increase benchmark runs to 4 and take best result
- Fixes slow first run issue on M3 Ultra

Co-Authored-By: Claude Opus 4.5 <[email protected]>
This commit is contained in:
Alex Cheema
2026-01-16 15:30:32 +00:00
co-authored by Claude Opus 4.5
parent ae3086167f
commit d1f80c9e86
2 changed files with 116 additions and 63 deletions
+66 -2
View File
@@ -4,6 +4,7 @@ import platform
from typing import Any, Callable, Coroutine
import anyio
from anyio import to_thread
from loguru import logger
from exo.shared.types.memory import Memory
@@ -22,11 +23,63 @@ from .macmon import (
)
from .system_info import (
get_friendly_name,
get_memory_bandwidth,
get_model_and_chip,
get_network_interfaces,
profile_memory_bandwidth,
)
# Module-level cache for memory bandwidth (doesn't change at runtime)
_cached_bandwidth: int | None = None
_bandwidth_profiled: bool = False
_bandwidth_profiling_task: asyncio.Task[int | None] | None = None
async def profile_bandwidth_once() -> int | None:
"""Profile bandwidth once in a background thread and cache the result.
This function is non-blocking - it runs the profiling in a thread pool.
Subsequent calls return the cached result immediately.
"""
global _cached_bandwidth, _bandwidth_profiled, _bandwidth_profiling_task
# Already profiled, return cached value
if _bandwidth_profiled:
return _cached_bandwidth
# Profiling already in progress, wait for it
if _bandwidth_profiling_task is not None:
return await _bandwidth_profiling_task
# Start profiling in background thread
async def _do_profile() -> int | None:
global _cached_bandwidth, _bandwidth_profiled
try:
logger.info("Starting memory bandwidth profiling in background thread...")
bandwidth = await to_thread.run_sync(profile_memory_bandwidth, cancellable=True)
_cached_bandwidth = bandwidth
_bandwidth_profiled = True
if bandwidth:
logger.info(f"Memory bandwidth profiled: {bandwidth / 1e9:.1f} GB/s")
else:
logger.warning("Memory bandwidth profiling returned None")
return bandwidth
except Exception as e:
logger.opt(exception=e).error("Memory bandwidth profiling failed")
_bandwidth_profiled = True # Mark as done to avoid retrying
return None
_bandwidth_profiling_task = asyncio.create_task(_do_profile())
return await _bandwidth_profiling_task
def get_memory_bandwidth_cached() -> int | None:
"""Return cached bandwidth or None if not yet profiled.
This is a non-blocking synchronous function that returns immediately.
Call profile_bandwidth_once() first to trigger profiling.
"""
return _cached_bandwidth if _bandwidth_profiled else None
async def get_metrics_async() -> Metrics | None:
"""Return detailed Metrics on macOS or a minimal fallback elsewhere."""
@@ -72,6 +125,8 @@ async def start_polling_node_metrics(
callback: Callable[[NodePerformanceProfile], Coroutine[Any, Any, None]],
):
poll_interval_s = 1.0
bandwidth_profile_started = False
while True:
try:
metrics = await get_metrics_async()
@@ -86,6 +141,15 @@ async def start_polling_node_metrics(
# do the memory profile last to get a fresh reading to not conflict with the other memory profiling loop
memory_profile = get_memory_profile()
# Start bandwidth profiling in background on first poll (non-blocking)
if not bandwidth_profile_started:
bandwidth_profile_started = True
# Fire and forget - don't await, let it run in background
asyncio.create_task(profile_bandwidth_once())
# Use cached bandwidth (None until profiling completes)
memory_bandwidth = get_memory_bandwidth_cached()
await callback(
NodePerformanceProfile(
model_id=model_id,
@@ -93,7 +157,7 @@ async def start_polling_node_metrics(
friendly_name=friendly_name,
network_interfaces=network_interfaces,
memory=memory_profile,
memory_bandwidth=get_memory_bandwidth(chip_id),
memory_bandwidth=memory_bandwidth,
system=SystemPerformanceProfile(
gpu_usage=metrics.gpu_usage[1],
temp=metrics.temp.gpu_temp_avg,
+50 -61
View File
@@ -1,5 +1,6 @@
import socket
import sys
import time
from subprocess import CalledProcessError
import psutil
@@ -83,78 +84,66 @@ async def get_model_and_chip() -> tuple[str, str]:
return (model, chip)
def _profile_memory_bandwidth_numpy() -> int | None:
"""Profile memory bandwidth using 1GB array benchmark."""
try:
import numpy as np
import time
size = 1024 * 1024 * 1024 // 8
num_runs = 3
best_bandwidth = 0.0
for _ in range(num_runs):
src = np.random.random(size)
start = time.perf_counter()
dst = src.copy()
end = time.perf_counter()
_ = dst[0]
bandwidth = (size * 8 * 2) / (end - start)
best_bandwidth = max(best_bandwidth, bandwidth)
del src, dst
return int(best_bandwidth)
except Exception:
return None
def _profile_memory_bandwidth_simple() -> int | None:
"""Fallback memory bandwidth benchmark using 200MB array."""
try:
import numpy as np
import time
size = 200 * 1024 * 1024 // 8
best_bandwidth = 0.0
num_runs = 5
for _ in range(num_runs):
src = np.random.random(size)
start = time.perf_counter()
dst = src.copy()
end = time.perf_counter()
_ = dst[0]
bandwidth = (size * 8 * 2) / (end - start)
best_bandwidth = max(best_bandwidth, bandwidth)
del src, dst
return int(best_bandwidth)
except Exception:
return None
def profile_memory_bandwidth() -> int | None:
"""
Profile device memory bandwidth using numpy benchmarks.
Profile device memory bandwidth using MLX GPU operations.
Returns measured bandwidth which may be lower than theoretical peak.
Relative ratios between devices remain accurate for placement decisions.
Uses a large array copy on the GPU to measure unified memory bandwidth.
Returns measured bandwidth in bytes/second, or None if MLX is unavailable.
"""
bandwidth = _profile_memory_bandwidth_numpy()
if bandwidth and bandwidth > 0:
return bandwidth
try:
import mlx.core as mx
bandwidth = _profile_memory_bandwidth_simple()
if bandwidth and bandwidth > 0:
return bandwidth
if not mx.metal.is_available():
return None
return None
# Use 512MB buffer - large enough to bypass cache
size_bytes = 512 * 1024 * 1024
num_elements = size_bytes // 4 # float32 = 4 bytes
# Warm-up: run the full benchmark operation multiple times to stabilize GPU
for _ in range(3):
src = mx.random.uniform(shape=(num_elements,), dtype=mx.float32)
mx.eval(src)
dst = src + 0.0
mx.eval(dst)
mx.synchronize()
del src, dst
# Benchmark: measure time to copy array (skip first run as it may be slow)
best_bandwidth = 0.0
num_runs = 4 # First run may still be slow, take best of 4
for _ in range(num_runs):
# Create source array
src = mx.random.uniform(shape=(num_elements,), dtype=mx.float32)
mx.eval(src)
mx.synchronize()
# Time the copy operation (src + 0.0 forces read of src, write of dst)
start = time.perf_counter()
dst = src + 0.0
mx.eval(dst)
mx.synchronize()
end = time.perf_counter()
# Bandwidth = bytes transferred / time
# Operation reads size_bytes and writes size_bytes
bytes_transferred = size_bytes * 2
bandwidth = bytes_transferred / (end - start)
best_bandwidth = max(best_bandwidth, bandwidth)
del src, dst
return int(best_bandwidth)
except Exception:
return None
def get_memory_bandwidth(_chip_id: str) -> int | None:
"""
Returns measured memory bandwidth in bytes/second.
Uses runtime profiling via numpy benchmarks. Works on any platform.
Uses MLX GPU operations for accurate unified memory bandwidth measurement.
"""
return profile_memory_bandwidth()