Skip to content
Merged
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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion README_CN.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:条件导入
Expand Down
2 changes: 1 addition & 1 deletion benchmarks/benchmark_history.json
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion examples/extension_setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"}
Expand Down
2 changes: 1 addition & 1 deletion src/torchada/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion src/torchada/csrc/ops.h
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
78 changes: 63 additions & 15 deletions src/torchada/utils/cpp_extension.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -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"
Expand All @@ -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

Expand All @@ -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)
Expand Down
60 changes: 60 additions & 0 deletions tests/test_cuda_patching.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion tests/test_platform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down