Skip to content

Commit fbf7823

Browse files
committed
build: bind native extension to isolated Python
1 parent 4703f89 commit fbf7823

4 files changed

Lines changed: 91 additions & 10 deletions

File tree

cppmega_mlx/training/native_optim/CMakeLists.txt

Lines changed: 14 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -12,16 +12,22 @@ find_package(
1212
Python 3.11
1313
COMPONENTS Interpreter Development.Module
1414
REQUIRED)
15-
execute_process(
16-
COMMAND "${Python_EXECUTABLE}" -m nanobind --cmake_dir
17-
OUTPUT_STRIP_TRAILING_WHITESPACE
18-
OUTPUT_VARIABLE nanobind_ROOT)
15+
if(NOT DEFINED nanobind_DIR)
16+
execute_process(
17+
COMMAND "${Python_EXECUTABLE}" -m nanobind --cmake_dir
18+
OUTPUT_STRIP_TRAILING_WHITESPACE
19+
OUTPUT_VARIABLE nanobind_ROOT
20+
COMMAND_ERROR_IS_FATAL ANY)
21+
endif()
1922
find_package(nanobind CONFIG REQUIRED)
2023

21-
execute_process(
22-
COMMAND "${Python_EXECUTABLE}" -m mlx --cmake-dir
23-
OUTPUT_STRIP_TRAILING_WHITESPACE
24-
OUTPUT_VARIABLE MLX_ROOT)
24+
if(NOT DEFINED MLX_ROOT)
25+
execute_process(
26+
COMMAND "${Python_EXECUTABLE}" -m mlx --cmake-dir
27+
OUTPUT_STRIP_TRAILING_WHITESPACE
28+
OUTPUT_VARIABLE MLX_ROOT
29+
COMMAND_ERROR_IS_FATAL ANY)
30+
endif()
2531
find_package(MLX CONFIG REQUIRED)
2632

2733
add_library(cppmega_native_optim)

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
[build-system]
2-
requires = ["setuptools>=68", "mlx>=0.31"]
2+
requires = ["setuptools>=68", "mlx>=0.31", "nanobind>=2.4"]
33
build-backend = "setuptools.build_meta"
44

55
[project]

setup.py

Lines changed: 27 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,34 @@
1+
import os
2+
import sys
3+
14
from setuptools import setup
25

6+
import nanobind
37
from mlx import extension
48

59

10+
class CppmegaCMakeBuild(extension.CMakeBuild):
11+
"""Build native extensions with the interpreter that invoked setuptools."""
12+
13+
def build_extension(self, ext: extension.CMakeExtension) -> None:
14+
previous_args = os.environ.get("CMAKE_ARGS")
15+
build_contract_args = (
16+
f"-DPython_EXECUTABLE={sys.executable}",
17+
f"-Dnanobind_DIR={nanobind.cmake_dir()}",
18+
f"-DMLX_ROOT={extension._MLX_PATH}",
19+
)
20+
os.environ["CMAKE_ARGS"] = " ".join(
21+
part for part in (previous_args, *build_contract_args) if part
22+
)
23+
try:
24+
super().build_extension(ext)
25+
finally:
26+
if previous_args is None:
27+
os.environ.pop("CMAKE_ARGS", None)
28+
else:
29+
os.environ["CMAKE_ARGS"] = previous_args
30+
31+
632
if __name__ == "__main__":
733
setup(
834
ext_modules=[
@@ -11,7 +37,7 @@
1137
sourcedir="cppmega_mlx/training/native_optim",
1238
)
1339
],
14-
cmdclass={"build_ext": extension.CMakeBuild},
40+
cmdclass={"build_ext": CppmegaCMakeBuild},
1541
package_data={
1642
"cppmega_mlx.training.native_optim": [
1743
"*.so",

tests/test_setup_native_build.py

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,49 @@
1+
from __future__ import annotations
2+
3+
import os
4+
import sys
5+
import tomllib
6+
from pathlib import Path
7+
8+
from setuptools import Distribution
9+
10+
11+
def test_native_build_dependencies_are_declared() -> None:
12+
pyproject = tomllib.loads(
13+
(Path(__file__).resolve().parents[1] / "pyproject.toml").read_text()
14+
)
15+
requirements = pyproject["build-system"]["requires"]
16+
17+
assert any(requirement.startswith("mlx>=") for requirement in requirements)
18+
assert any(requirement.startswith("nanobind>=") for requirement in requirements)
19+
20+
21+
def test_native_build_pins_cmake_to_invoking_python(monkeypatch) -> None:
22+
repo_root = Path(__file__).resolve().parents[1]
23+
monkeypatch.syspath_prepend(str(repo_root))
24+
25+
import setup as cppmega_setup
26+
27+
captured: dict[str, str | None] = {}
28+
29+
def capture_build_extension(_command, _extension) -> None:
30+
captured["cmake_args"] = os.environ.get("CMAKE_ARGS")
31+
32+
monkeypatch.setattr(
33+
cppmega_setup.extension.CMakeBuild,
34+
"build_extension",
35+
capture_build_extension,
36+
)
37+
monkeypatch.setenv("CMAKE_ARGS", "-DEXISTING_OPTION=ON")
38+
39+
command = cppmega_setup.CppmegaCMakeBuild(Distribution())
40+
command.build_extension(object())
41+
42+
cmake_args = captured["cmake_args"]
43+
assert cmake_args is not None
44+
assert cmake_args.startswith(
45+
f"-DEXISTING_OPTION=ON -DPython_EXECUTABLE={sys.executable}"
46+
)
47+
assert "-Dnanobind_DIR=" in cmake_args
48+
assert "-DMLX_ROOT=" in cmake_args
49+
assert os.environ["CMAKE_ARGS"] == "-DEXISTING_OPTION=ON"

0 commit comments

Comments
 (0)