diff --git a/pyproject.toml b/pyproject.toml index 33ad345..39389d6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,3 +33,25 @@ requires = [ ] build-backend = "setuptools.build_meta" +[tool.ruff] +target-version = "py310" +line-length = 88 +exclude = [ + ".git", + ".venv", + "__pycache__", + "build", + "dist", +] + +[tool.ruff.lint] +select = ["E", "F", "I", "B"] +ignore = [] +fixable = ["ALL"] +unfixable = [] + +[tool.ruff.format] +quote-style = "double" +indent-style = "space" +line-ending = "auto" +skip-magic-trailing-comma = false diff --git a/setup.py b/setup.py index cd2fd86..17f698f 100644 --- a/setup.py +++ b/setup.py @@ -1,6 +1,7 @@ -from setuptools import setup, Extension import importlib.util +from setuptools import Extension, setup + core_ext = Extension( "rixa.PMIx_core", sources=[ @@ -21,7 +22,7 @@ cmdclass = {} ext_modules = [core_ext] if importlib.util.find_spec("torch"): - from torch.utils.cpp_extension import CppExtension, BuildExtension, library_paths + from torch.utils.cpp_extension import BuildExtension, CppExtension, library_paths torch_lib_path = library_paths() diff --git a/src/rixa/__init__.py b/src/rixa/__init__.py index 9a0c00c..1048594 100644 --- a/src/rixa/__init__.py +++ b/src/rixa/__init__.py @@ -1,6 +1,8 @@ import importlib.util -from .PMIx_core import PMIxStore import os +import warnings + +from .PMIx_core import PMIxStore is_strict = False # We don't want to be strict maybe? # Ninja has build lock, first process will acquire the lock @@ -19,10 +21,11 @@ if importlib.util.find_spec("torch") is not None: if not (pytorch.is_compiled_source() or pytorch.is_compiled_lazy()): if is_distributed_env() and is_strict: raise RuntimeError( - "Pytorch support and distributed enviroment is set up but noting is yet compiled!\n" + "Pytorch support and pmix is set up but noting is yet compiled!\n" 'In order to compile the extension do python -c "import rixa" ' ) else: + warnings.warning("Distributed env detected but extension is not compiled!") pytorch.get_pmix_store() __all__ += ["pytorch"] diff --git a/src/rixa/nvshmem.py b/src/rixa/nvshmem.py index 4beb18e..b3880e3 100644 --- a/src/rixa/nvshmem.py +++ b/src/rixa/nvshmem.py @@ -1,6 +1,7 @@ -from cuda.core import Device -import nvshmem.core as nvshmem import numpy as np +import nvshmem.core as nvshmem +from cuda.core import Device + from rixa.PMIx_core import PMIxStore diff --git a/src/rixa/pytorch.py b/src/rixa/pytorch.py index 74a8df0..d12dc3b 100644 --- a/src/rixa/pytorch.py +++ b/src/rixa/pytorch.py @@ -1,11 +1,12 @@ -import torch.distributed as dist -import os +import ctypes import importlib.util +import os +import warnings from pathlib import Path -import ctypes + import torch +import torch.distributed as dist import torch.utils.cpp_extension as cpp_ext -import warnings _ext = None @@ -46,7 +47,7 @@ def _get_ext(): try: load_pmix_wheel() except RuntimeError: - warnings.warn("Using global pmix to compile torch extension") + warnings.warn("Using global pmix to compile torch extension", stacklevel=2) _ext = cpp_ext.load( name="_rixa_torch", diff --git a/tests/test_check_positive.py b/tests/test_check_positive.py index 5e7c8b4..b3dc02c 100644 --- a/tests/test_check_positive.py +++ b/tests/test_check_positive.py @@ -1,9 +1,11 @@ -import sys import os import subprocess -from rixa.pytorch import get_pmix_store +import sys + import pytest +from rixa.pytorch import get_pmix_store + def is_worker(): return os.environ.get("RIXA_WORKER_MODE") == "1" diff --git a/tests/test_nvshmem_init.py b/tests/test_nvshmem_init.py index d08110f..b76e050 100644 --- a/tests/test_nvshmem_init.py +++ b/tests/test_nvshmem_init.py @@ -1,6 +1,7 @@ import os -import sys import subprocess +import sys + import pytest WORKER_ENV_VAR = "RIXA_NVSHMEM_WORKER" @@ -12,11 +13,11 @@ def is_worker(): if is_worker(): try: - import rixa.nvshmem import nvshmem.core as nvshmem - from cuda.core import Device + import rixa.nvshmem + store = rixa.PMIxStore(30) dev = Device(0) rixa.nvshmem.init(dev, store, 30) diff --git a/tests/test_pmix.py b/tests/test_pmix.py index 7a034a8..637e5bc 100644 --- a/tests/test_pmix.py +++ b/tests/test_pmix.py @@ -1,9 +1,11 @@ -from rixa.PMIx_core import PMIxStore -import sys import os import subprocess +import sys + import pytest +from rixa.PMIx_core import PMIxStore + def is_worker(): return os.environ.get("RIXA_WORKER_MODE") == "1" diff --git a/tests/test_pytorch_init.py b/tests/test_pytorch_init.py index 0495ee3..9897aaf 100644 --- a/tests/test_pytorch_init.py +++ b/tests/test_pytorch_init.py @@ -1,8 +1,10 @@ import os -import sys import subprocess +import sys + import pytest import torch + import rixa WORKER_ENV_VAR = "RIXA_PYTORCH_WORKER" diff --git a/tests/test_pytorch_init_gpu.py b/tests/test_pytorch_init_gpu.py index 73974de..afb8211 100644 --- a/tests/test_pytorch_init_gpu.py +++ b/tests/test_pytorch_init_gpu.py @@ -1,8 +1,10 @@ import os -import sys import subprocess +import sys + import pytest import torch + import rixa WORKER_ENV_VAR = "RIXA_PYTORCH_WORKER" diff --git a/tests/test_set_get_wait.py b/tests/test_set_get_wait.py index e149715..24219d6 100644 --- a/tests/test_set_get_wait.py +++ b/tests/test_set_get_wait.py @@ -1,9 +1,11 @@ -import sys import os import subprocess -from rixa.pytorch import get_pmix_store +import sys + import pytest +from rixa.pytorch import get_pmix_store + def is_worker(): return os.environ.get("RIXA_WORKER_MODE") == "1"