from pathlib import Path import setuptools from torch.utils.cpp_extension import BuildExtension, CUDAExtension def _nvidia_include_dirs() -> list[str]: """Include dirs from pip-installed nvidia packages (e.g. cusparse headers).""" try: import nvidia # noqa: PLC0415 return [str(p) for pkg in Path(nvidia.__path__[0]).iterdir() if (p := pkg / "include").is_dir()] except ImportError: return [] if __name__ == "__main__": cxx_flags = ["-O3", "-Wall", "-Wextra", "-Werror", "-Wno-unused-parameter", "-Wno-attributes"] nvcc_flags = ["-O3"] extra_compile_args = { "cxx": cxx_flags, "nvcc": nvcc_flags, } ROOT = Path(__file__).resolve().parent setuptools.setup( ext_modules=[ CUDAExtension( name="all2all_cpp", include_dirs=[str(ROOT / "csrc/all2all"), str(ROOT / "csrc/include"), *_nvidia_include_dirs()], sources=[ "csrc/all2all/all2all.cpp", "csrc/all2all/cuda/all2all_heads.cu", "csrc/all2all/cuda/allgather.cu", ], extra_compile_args=extra_compile_args, ) ], cmdclass={"build_ext": BuildExtension}, )