diff --git a/README.md b/README.md index 0f62ddf..f16e206 100644 --- a/README.md +++ b/README.md @@ -408,7 +408,7 @@ See `src/torchada/_mappings/` for 400+ mapping rules grouped by API domain. ``` # pyproject.toml or requirements.txt -torchada>=0.1.92 +torchada>=0.1.93 ``` ### Step 2: Conditional Import diff --git a/README_CN.md b/README_CN.md index 069dd57..8194523 100644 --- a/README_CN.md +++ b/README_CN.md @@ -387,7 +387,7 @@ if torchada.is_gpu_device(device): # 在 CUDA 和 MUSA 上都能工作 ``` # pyproject.toml 或 requirements.txt -torchada>=0.1.92 +torchada>=0.1.93 ``` ### 步骤 2:条件导入 diff --git a/benchmarks/benchmark_history.json b/benchmarks/benchmark_history.json index b4eddb1..40c6b84 100644 --- a/benchmarks/benchmark_history.json +++ b/benchmarks/benchmark_history.json @@ -3,7 +3,7 @@ "description": "Historical benchmark results for torchada performance tracking", "results": [ { - "version": "0.1.92", + "version": "0.1.93", "date": "2026-01-29", "platform": "MUSA", "pytorch_version": "2.7.1", diff --git a/examples/extension_setup.py b/examples/extension_setup.py index 9359b5a..442da81 100755 --- a/examples/extension_setup.py +++ b/examples/extension_setup.py @@ -69,7 +69,7 @@ def get_extensions(): if extensions: setup( name="my_cuda_extension", - version="0.1.92", + version="0.1.93", ext_modules=extensions, cmdclass={"build_ext": BuildExtension.with_options(use_ninja=True)}, python_requires=">=3.8", diff --git a/pyproject.toml b/pyproject.toml index f237e88..96f059f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "torchada" -version = "0.1.92" +version = "0.1.93" description = "Adapter package for torch_musa to act exactly like PyTorch CUDA" readme = "README.md" license = {text = "MIT"} diff --git a/src/torchada/__init__.py b/src/torchada/__init__.py index cb13f5f..20faca7 100644 --- a/src/torchada/__init__.py +++ b/src/torchada/__init__.py @@ -24,7 +24,7 @@ from torch.utils.cpp_extension import CUDAExtension, BuildExtension, CUDA_HOME """ -__version__ = "0.1.92" +__version__ = "0.1.93" from . import cuda, utils diff --git a/src/torchada/csrc/ops.h b/src/torchada/csrc/ops.h index 6cd98af..e14b84f 100644 --- a/src/torchada/csrc/ops.h +++ b/src/torchada/csrc/ops.h @@ -40,7 +40,7 @@ namespace torchada { // Version information -constexpr const char* VERSION = "0.1.92"; +constexpr const char* VERSION = "0.1.93"; // Check if operator override is enabled via environment variable inline bool is_override_enabled(const char* op_name) { diff --git a/src/torchada/utils/cpp_extension.py b/src/torchada/utils/cpp_extension.py index ec7cbdc..03ef202 100644 --- a/src/torchada/utils/cpp_extension.py +++ b/src/torchada/utils/cpp_extension.py @@ -838,20 +838,46 @@ def _port_cuda_source(source_code: str, mapping_rules: Optional[Dict[str, str]] ) -def include_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None) -> List[str]: +def _split_device_argument( + device_type: Optional[Any], cuda: Optional[bool] +) -> Tuple[Optional[str], Optional[bool]]: + """Return ``(device_type, cuda)``, moving a PyTorch < 2.6 positional ``cuda`` bool.""" + if isinstance(device_type, bool): + return None, device_type + return device_type, cuda + + +def _is_torch_path(path: str, subdir: str) -> bool: + """Whether ``path`` lies in PyTorch's own ``include`` or ``lib`` directory.""" + import torch + + root = os.path.realpath(os.path.join(os.path.dirname(torch.__file__), subdir)) + path = os.path.realpath(path) + return path == root or path.startswith(root + os.sep) + + +def include_paths( + device_type: Optional[str] = None, + torch_include_dirs: bool = True, + *, + cuda: Optional[bool] = None, +) -> List[str]: """ Get include paths for compiling extensions. Supports both PyTorch < 2.6 (cuda=True) and PyTorch 2.6+ (device_type="cuda") - signatures for compatibility. + signatures for compatibility, called positionally or by keyword. Args: - cuda: (PyTorch < 2.6) Whether to include CUDA/MUSA paths. Deprecated in 2.6+. device_type: (PyTorch 2.6+) Device type string, e.g. "cuda", "cpu", "musa". + A bool here is the PyTorch < 2.6 positional ``cuda`` argument. + torch_include_dirs: (PyTorch 2.10+) Whether to include PyTorch's own headers. + cuda: (PyTorch < 2.6) Whether to include CUDA/MUSA paths. Deprecated in 2.6+. Returns: List of include paths """ + device_type, cuda = _split_device_argument(device_type, cuda) # Handle both old (cuda=bool) and new (device_type=str) signatures if device_type is not None: # PyTorch 2.6+ style: device_type="cuda" or "cpu" @@ -888,6 +914,8 @@ def include_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None # torch_musa shipping the real header wins. if include_device: paths.append(stable_compat_include_dir()) + if not torch_include_dirs: + paths = [path for path in paths if not _is_torch_path(path, "include")] return paths else: @@ -899,10 +927,13 @@ def include_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None sig = inspect.signature(torch_include_paths) if "device_type" in sig.parameters: # PyTorch 2.6+ - if device_type is not None: - return torch_include_paths(device_type=device_type) - else: - return torch_include_paths(device_type="cuda" if include_device else "cpu") + if device_type is None: + device_type = "cuda" if include_device else "cpu" + if "torch_include_dirs" in sig.parameters: + return torch_include_paths( + device_type=device_type, torch_include_dirs=torch_include_dirs + ) + return torch_include_paths(device_type=device_type) else: # PyTorch < 2.6 return torch_include_paths(cuda=include_device) @@ -933,20 +964,30 @@ def stable_compat_box_header() -> str: return os.path.join(stable_compat_include_dir(), "torchada_stable_box.h") -def library_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None) -> List[str]: +def library_paths( + device_type: Optional[str] = None, + torch_include_dirs: bool = True, + cross_target_platform: Optional[str] = None, + *, + cuda: Optional[bool] = None, +) -> List[str]: """ Get library paths for compiling extensions. Supports both PyTorch < 2.6 (cuda=True) and PyTorch 2.6+ (device_type="cuda") - signatures for compatibility. + signatures for compatibility, called positionally or by keyword. Args: - cuda: (PyTorch < 2.6) Whether to include CUDA/MUSA library paths. Deprecated in 2.6+. device_type: (PyTorch 2.6+) Device type string, e.g. "cuda", "cpu", "musa". + A bool here is the PyTorch < 2.6 positional ``cuda`` argument. + torch_include_dirs: (PyTorch 2.10+) Whether to include PyTorch's own libraries. + cross_target_platform: (PyTorch 2.10+) Passed through to PyTorch off MUSA. + cuda: (PyTorch < 2.6) Whether to include CUDA/MUSA library paths. Deprecated in 2.6+. Returns: List of library paths """ + device_type, cuda = _split_device_argument(device_type, cuda) # Handle both old (cuda=bool) and new (device_type=str) signatures if device_type is not None: # PyTorch 2.6+ style: device_type="cuda" or "cpu" @@ -969,7 +1010,10 @@ def library_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None if hasattr(musa_ext, "library_paths"): # musa_ext uses musa=bool parameter, not cuda= or device_type= - return musa_ext.library_paths(musa=include_device) + paths = list(musa_ext.library_paths(musa=include_device)) + if not torch_include_dirs: + paths = [path for path in paths if not _is_torch_path(path, "lib")] + return paths except ImportError: pass @@ -990,10 +1034,14 @@ def library_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None sig = inspect.signature(torch_library_paths) if "device_type" in sig.parameters: # PyTorch 2.6+ - if device_type is not None: - return torch_library_paths(device_type=device_type) - else: - return torch_library_paths(device_type="cuda" if include_device else "cpu") + if device_type is None: + device_type = "cuda" if include_device else "cpu" + kwargs: Dict[str, Any] = {} + if "torch_include_dirs" in sig.parameters: + kwargs["torch_include_dirs"] = torch_include_dirs + if "cross_target_platform" in sig.parameters: + kwargs["cross_target_platform"] = cross_target_platform + return torch_library_paths(device_type=device_type, **kwargs) else: # PyTorch < 2.6 return torch_library_paths(cuda=include_device) diff --git a/tests/test_cuda_patching.py b/tests/test_cuda_patching.py index e47cdf0..7491c42 100644 --- a/tests/test_cuda_patching.py +++ b/tests/test_cuda_patching.py @@ -2212,6 +2212,66 @@ def test_library_paths_default(self): assert isinstance(paths, list) assert len(paths) > 0 + def test_paths_accept_the_positional_torch_signature(self): + """torch 2.10+ Inductor calls include_paths(device_type, torch_include_dirs) by position.""" + import torch.utils.cpp_extension as cpp_extension + + import torchada + + if not torchada.is_musa_platform(): + pytest.skip("Only applicable on MUSA platform") + + assert cpp_extension.include_paths("cpu", True) == cpp_extension.include_paths( + device_type="cpu" + ) + assert cpp_extension.include_paths("cuda") == cpp_extension.include_paths( + device_type="cuda" + ) + assert cpp_extension.library_paths( + "cuda", torch_include_dirs=True, cross_target_platform=None + ) == cpp_extension.library_paths(device_type="cuda") + # A positional bool is the PyTorch < 2.6 ``cuda`` argument. + assert cpp_extension.include_paths(True) == cpp_extension.include_paths(cuda=True) + assert cpp_extension.library_paths(False) == cpp_extension.library_paths(cuda=False) + + def test_paths_without_torch_dirs(self): + """torch_include_dirs=False leaves out PyTorch's own include and lib directories.""" + import torch + import torch.utils.cpp_extension as cpp_extension + + import torchada + + if not torchada.is_musa_platform(): + pytest.skip("Only applicable on MUSA platform") + + torch_root = os.path.realpath(os.path.dirname(torch.__file__)) + for function in (cpp_extension.include_paths, cpp_extension.library_paths): + paths = function("cuda", torch_include_dirs=False) + assert paths + assert not any(os.path.realpath(path).startswith(torch_root + os.sep) for path in paths) + assert set(paths) <= set(function("cuda")) + + def test_inductor_cpu_compile_with_a_cold_cache(self, tmp_path): + """A cold Inductor cache builds CPU kernels through the patched path helpers.""" + import subprocess + import sys + + import torchada + + if not torchada.is_musa_platform(): + pytest.skip("Only applicable on MUSA platform") + + code = ( + "import torchada, torch\n" + "out = torch.compile(lambda x: (x * 2 + 1).sum(-1), fullgraph=True)(torch.ones(8, 16))\n" + "assert torch.equal(out, torch.full((8,), 48.0))\n" + ) + env = dict(os.environ, TORCHINDUCTOR_CACHE_DIR=str(tmp_path / "inductor")) + result = subprocess.run( + [sys.executable, "-c", code], env=env, capture_output=True, text=True, timeout=600 + ) + assert result.returncode == 0, result.stderr[-2000:] + def test_include_paths_patched_in_torch_module(self): """Test that include_paths is properly patched in torch.utils.cpp_extension.""" import torch.utils.cpp_extension diff --git a/tests/test_platform.py b/tests/test_platform.py index e07956a..40193dd 100644 --- a/tests/test_platform.py +++ b/tests/test_platform.py @@ -59,7 +59,7 @@ def test_get_version(self): version = torchada.get_version() assert version == torchada.__version__ - assert version == "0.1.92" + assert version == "0.1.93" assert isinstance(version, str) def test_project_version_matches_runtime_version(self):