Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -154,7 +154,7 @@ This starts the exo dashboard and API at http://localhost:52415/
```bash
# Install Node.js and npm
sudo apt update
sudo apt install nodejs npm
sudo apt install nodejs npm python3-dev build-essential

# Install uv
curl -LsSf https://astral.sh/uv/install.sh | sh
Expand Down Expand Up @@ -198,7 +198,7 @@ uv run exo

This starts the exo dashboard and API at http://localhost:52415/

**Important note for Linux users:** Currently, exo runs on CPU on Linux. GPU support for Linux platforms is under development. If you'd like to see support for your specific Linux hardware, please [search for existing feature requests](https://github.com/exo-explore/exo/issues) or create a new one.
**Important note for Linux users:** Currently, exo runs on the MLX CPU backend on Linux. Large-model inference can be slow, and CPU core utilization depends on the model operations and backend. GPU support for Linux platforms is under development. If you'd like to see support for your specific Linux hardware, please [search for existing feature requests](https://github.com/exo-explore/exo/issues) or create a new one.

**Configuration Options:**

Expand Down
6 changes: 3 additions & 3 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ mlx = [
"torchvision==0.25.0; sys_platform == 'darwin'",
"torchvision==0.25.0; sys_platform == 'linux'",
]
mlx-cpu = ["exo[mlx]", "mlx-cpu==0.31.2; sys_platform == 'linux'"]
mlx-cpu = ["exo[mlx]", "mlx-cpu==0.32.0; sys_platform == 'linux'"]
mlx-cuda12 = [
"exo[mlx]",
"mlx-cuda-12==0.32.0; sys_platform == 'linux'",
Expand All @@ -82,8 +82,8 @@ members = ["rust/exo_rs", "bench", "tools"]
exo-rs = { workspace = true }
mlx = [
{ git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git", branch = "address-rdma-gpu-locks", marker = "sys_platform == 'darwin'" },
{ url = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv/releases/download/mlx_cuda/mlx-0.32.0-cp313-cp313-manylinux_2_35_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64'" },
{ url = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv/releases/download/mlx_cuda/mlx-0.32.0-cp313-cp313-manylinux_2_35_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine != 'aarch64'" },
{ url = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv/releases/download/mlx_cuda/mlx-0.32.0-cp313-cp313-manylinux_2_35_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64' and (extra == 'mlx-cuda12' or extra == 'mlx-cuda13')" },
{ url = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv/releases/download/mlx_cuda/mlx-0.32.0-cp313-cp313-manylinux_2_35_x86_64.whl", marker = "sys_platform == 'linux' and platform_machine != 'aarch64' and (extra == 'mlx-cuda12' or extra == 'mlx-cuda13')" },
]
mlx-cuda-12 = [
{ url = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv/releases/download/mlx_cuda/mlx_cuda_12-0.32.0-py3-none-manylinux_2_35_aarch64.whl", marker = "sys_platform == 'linux' and platform_machine == 'aarch64'" },
Expand Down
39 changes: 39 additions & 0 deletions tests/test_linux_mlx_cpu_dependency_contract.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
import tomllib
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]


def test_linux_mlx_cpu_extra_uses_a_matching_registry_pair() -> None:
pyproject = tomllib.loads((ROOT / "pyproject.toml").read_text())

cpu_requirements = pyproject["project"]["optional-dependencies"]["mlx-cpu"]
mlx_requirement = next(
requirement
for requirement in pyproject["project"]["optional-dependencies"]["mlx"]
if requirement.startswith("mlx==")
)
cpu_requirement = next(
requirement
for requirement in cpu_requirements
if requirement.startswith("mlx-cpu==")
)
mlx_version = mlx_requirement.removeprefix("mlx==")
cpu_version = cpu_requirement.removeprefix("mlx-cpu==").split(";", 1)[0]

assert cpu_version == mlx_version
linux_sources = [
source
for source in pyproject["tool"]["uv"]["sources"]["mlx"]
if "sys_platform == 'linux'" in source.get("marker", "")
]
assert linux_sources
assert all(
"extra == 'mlx-cuda12'" in source["marker"]
and "extra == 'mlx-cuda13'" in source["marker"]
for source in linux_sources
)


if __name__ == "__main__":
test_linux_mlx_cpu_extra_uses_a_matching_registry_pair()