Skip to content

Commit 9346f0b

Browse files
author
fengchao.wei
committed
fix: protect dependency headers during extension builds
1 parent 86d842c commit 9346f0b

7 files changed

Lines changed: 248 additions & 18 deletions

File tree

‎README.md‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -201,9 +201,12 @@ torch_musa 2.9. The patched `torch.utils.cpp_extension.include_paths()` exposes
201201
the compatibility include directory on MUSA. Custom stable-ABI builds should
202202
add `stable_compat_include_dir()` explicitly, and kernels that use `TORCH_BOX`
203203
must force-include the path returned by `stable_compat_box_header()`. Both
204-
helpers are available from `torchada.utils.cpp_extension`. The torch_musa 2.9
205-
header backport runs lazily and best-effort at extension-build time; read-only
206-
headers are left unchanged. A plain `import torchada` does not modify PyTorch or
204+
helpers are available from `torchada.utils.cpp_extension`. Mark custom stable
205+
ABI extensions explicitly with `CUDAExtension(..., torchada_stable_abi=True)`.
206+
The torch_musa 2.9 header backport then runs lazily and best-effort only for
207+
stable-ABI extensions; read-only headers are left unchanged. Existing vLLM and
208+
SGLang extension names/source layouts are recognized for backward compatibility.
209+
Ordinary extension builds and a plain `import torchada` do not modify PyTorch or
207210
torch_musa headers.
208211

209212
### Custom Ops

‎README_CN.md‎

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -192,9 +192,11 @@ torchada 还为近期 vLLM 和 SGLang 在 torch_musa 2.9 上使用的 libtorch
192192
`torch.utils.cpp_extension.include_paths()` 会返回该兼容 include 目录。自定义
193193
稳定 ABI 构建应显式加入 `stable_compat_include_dir()`;使用 `TORCH_BOX` 的内核
194194
还必须通过编译器强制 include `stable_compat_box_header()` 返回的头文件。这两个
195-
辅助函数都位于 `torchada.utils.cpp_extension`。torch_musa 2.9 头文件回补只会在
196-
扩展构建时延迟、尽力执行;只读头文件保持不变。单纯 `import torchada` 不会修改
197-
PyTorch 或 torch_musa 头文件。
195+
辅助函数都位于 `torchada.utils.cpp_extension`。自定义 stable ABI 扩展应通过
196+
`CUDAExtension(..., torchada_stable_abi=True)` 显式标记。torch_musa 2.9 头文件
197+
回补只会针对 stable ABI 扩展延迟、尽力执行;为保持兼容,现有 vLLM/SGLang 的
198+
扩展名称和源码目录结构也会被自动识别。普通扩展构建以及单纯
199+
`import torchada` 都不会修改 PyTorch 或 torch_musa 头文件。
198200

199201
### 自定义算子
200202

‎pyproject.toml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
44

55
[project]
66
name = "torchada"
7-
version = "0.1.76"
7+
version = "0.1.77"
88
description = "Adapter package for torch_musa to act exactly like PyTorch CUDA"
99
readme = "README.md"
1010
license = {text = "MIT"}

‎src/torchada/__init__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@
2424
from torch.utils.cpp_extension import CUDAExtension, BuildExtension, CUDA_HOME
2525
"""
2626

27-
__version__ = "0.1.76"
27+
__version__ = "0.1.77"
2828

2929
from . import cuda, utils
3030

‎src/torchada/utils/cpp_extension.py‎

Lines changed: 90 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
"""
2727

2828
import functools
29+
import importlib
2930
import inspect
3031
import logging
3132
import os
@@ -243,6 +244,34 @@ def _coalesce_port_roots(paths):
243244
return result
244245

245246

247+
def _runtime_dependency_roots() -> List[str]:
248+
"""Return package roots that may be used as extension include paths.
249+
250+
``torch_musa`` is often installed in editable mode during development. In
251+
that case ``torch_musa.__file__`` resolves into the checkout, even though
252+
the directory is a dependency of the extension being built. Such a path
253+
must remain available to the compiler, but it must never become an
254+
in-place CUDA-to-MUSA porting root.
255+
"""
256+
roots = []
257+
for module_name in ("torch", "torch_musa"):
258+
try:
259+
module = importlib.import_module(module_name)
260+
module_file = getattr(module, "__file__", None)
261+
if module_file:
262+
roots.append(os.path.realpath(os.path.dirname(module_file)))
263+
except (ImportError, AttributeError, OSError):
264+
continue
265+
return roots
266+
267+
268+
def _path_overlaps_any(path: str, roots: List[str]) -> bool:
269+
"""Return whether ``path`` contains or is contained by a protected root."""
270+
return any(
271+
_path_is_within(path, root) or _path_is_within(root, path) for root in roots
272+
)
273+
274+
246275
def _validate_portable_symlinks(source_dir: str) -> None:
247276
"""Reject portable symlink files before upstream ``realpath`` can escape.
248277
@@ -926,6 +955,38 @@ def library_paths(cuda: Optional[bool] = None, device_type: Optional[str] = None
926955
return torch_library_paths(cuda=include_device)
927956

928957

958+
def _extension_requires_stable_header_backport(extension) -> bool:
959+
"""Return whether an extension needs the torch_musa stable-ABI backport.
960+
961+
Callers can set ``torchada_stable_abi=True`` on ``CUDAExtension`` for an
962+
explicit decision. The name/source fallback preserves compatibility with
963+
existing vLLM and SGLang extensions that predate that keyword.
964+
"""
965+
explicit = getattr(extension, "_torchada_stable_abi", None)
966+
if explicit is not None:
967+
return explicit
968+
969+
markers = ("libtorch_stable", "stable_libtorch", "stable_abi")
970+
candidates = [getattr(extension, "name", "")]
971+
candidates.extend(getattr(extension, "sources", None) or [])
972+
if any(any(marker in str(value).lower() for marker in markers) for value in candidates):
973+
return True
974+
975+
compile_args = getattr(extension, "extra_compile_args", None) or []
976+
if isinstance(compile_args, dict):
977+
compile_args = [arg for args in compile_args.values() for arg in args]
978+
elif isinstance(compile_args, str):
979+
compile_args = [compile_args]
980+
if any("torchada_stable_box.h" in str(arg) for arg in compile_args):
981+
return True
982+
983+
return any(
984+
macro[0] == "TORCHADA_STABLE_ACCESSORS"
985+
for macro in (getattr(extension, "define_macros", None) or [])
986+
if macro
987+
)
988+
989+
929990
class CUDAExtension:
930991
"""
931992
A wrapper that creates either a torch CUDAExtension or MUSA MUSAExtension.
@@ -944,12 +1005,21 @@ def __new__(cls, name: str, sources: List[str], *args, **kwargs):
9441005
*args: Additional positional arguments
9451006
**kwargs: Additional keyword arguments
9461007
"""
947-
platform = detect_platform()
1008+
stable_abi = kwargs.pop("torchada_stable_abi", None)
1009+
if stable_abi is not None and not isinstance(stable_abi, bool):
1010+
raise TypeError("torchada_stable_abi must be a bool")
9481011

1012+
platform = detect_platform()
9491013
if platform == Platform.MUSA:
950-
return _create_musa_extension(name, sources, *args, **kwargs)
1014+
extension = _create_musa_extension(name, sources, *args, **kwargs)
9511015
else:
952-
return _create_cuda_extension(name, sources, *args, **kwargs)
1016+
extension = _create_cuda_extension(name, sources, *args, **kwargs)
1017+
1018+
if stable_abi is not None:
1019+
extension._torchada_stable_abi = stable_abi
1020+
if platform == Platform.MUSA and _extension_requires_stable_header_backport(extension):
1021+
_ensure_stable_headers_patched()
1022+
return extension
9531023

9541024

9551025
class CppExtension:
@@ -1036,9 +1106,6 @@ def _create_musa_extension(name: str, sources: List[str], *args, **kwargs):
10361106
"""
10371107
# Ensure patches are applied
10381108
_apply_musa_patches()
1039-
# Building a MUSA extension: apply the (on-disk) libtorch-stable header
1040-
# backport now so csrc/libtorch_stable/*.cu can compile.
1041-
_ensure_stable_headers_patched()
10421109

10431110
# Translate CUDA compiler, library, and feature-macro names to MUSA.
10441111
kwargs = _translate_compile_args(kwargs)
@@ -1131,9 +1198,14 @@ def get_mapping_rule(self):
11311198
return _MAPPING_RULE.copy()
11321199

11331200
def build_extensions(self):
1134-
# Building now: apply the (on-disk) libtorch-stable header
1135-
# backport before the compiler reads the torch headers.
1136-
_ensure_stable_headers_patched()
1201+
# Only stable-ABI extensions need the legacy torch_musa
1202+
# header backport. Ordinary extensions (for example,
1203+
# torchvision) must not mutate dependency headers.
1204+
if any(
1205+
_extension_requires_stable_header_backport(extension)
1206+
for extension in self.extensions
1207+
):
1208+
_ensure_stable_headers_patched()
11371209
# Register .cu, .cuh as valid source extensions
11381210
self.compiler.src_extensions += [".cu", ".cuh"]
11391211
super().build_extensions()
@@ -1180,12 +1252,20 @@ def _dir_has_portable_sources(path):
11801252

11811253
@staticmethod
11821254
def _is_system_include_dir(path):
1183-
return (
1255+
is_system_path = (
11841256
path.startswith("/usr/")
11851257
or path.startswith("/opt/")
11861258
or "site-packages" in path
11871259
or "dist-packages" in path
11881260
)
1261+
if is_system_path:
1262+
return True
1263+
1264+
# Include roots belonging to the installed PyTorch or
1265+
# torch_musa dependency are compiler inputs, not project
1266+
# sources. This also covers editable installs whose real
1267+
# path is a development checkout such as /home/torch_musa.
1268+
return _path_overlaps_any(path, _runtime_dependency_roots())
11891269

11901270
def run(self):
11911271
"""Port each project-local include root's CUDA sources to MUSA

‎tests/test_cpp_extension.py‎

Lines changed: 102 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -145,6 +145,108 @@ def test_porting_translates_torch_cuda_header(self):
145145
assert _replace_porting_line(project_include, rules) == project_include
146146

147147

148+
class TestStableHeaderBackportSelection:
149+
"""Stable header mutation is limited to extensions that opt in or match."""
150+
151+
@staticmethod
152+
def _requires(extension):
153+
from torchada.utils.cpp_extension import (
154+
_extension_requires_stable_header_backport,
155+
)
156+
157+
return _extension_requires_stable_header_backport(extension)
158+
159+
def test_ordinary_extension_does_not_require_backport(self):
160+
extension = SimpleNamespace(
161+
name="torchvision_image",
162+
sources=["torchvision/csrc/io/image/decode_jpegs_cuda.cpp"],
163+
extra_compile_args={"cxx": ["-O3"], "mcc": ["-O3"]},
164+
define_macros=[("NVJPEG_FOUND", 1)],
165+
)
166+
167+
assert self._requires(extension) is False
168+
169+
def test_explicit_stable_abi_marker_takes_precedence(self):
170+
extension = SimpleNamespace(
171+
name="ordinary_name",
172+
sources=["kernel.cu"],
173+
_torchada_stable_abi=True,
174+
)
175+
assert self._requires(extension) is True
176+
177+
extension._torchada_stable_abi = False
178+
extension.name = "_C_stable_libtorch"
179+
assert self._requires(extension) is False
180+
181+
def test_existing_stable_extension_names_and_sources_are_detected(self):
182+
by_name = SimpleNamespace(name="_C_stable_libtorch", sources=[])
183+
by_source = SimpleNamespace(
184+
name="vllm_extension",
185+
sources=["vllm/csrc/libtorch_stable/gemm_kernels.cu"],
186+
)
187+
188+
assert self._requires(by_name) is True
189+
assert self._requires(by_source) is True
190+
191+
def test_stable_box_force_include_is_detected(self):
192+
extension = SimpleNamespace(
193+
name="custom_ops",
194+
sources=["ops.cu"],
195+
extra_compile_args={
196+
"cxx": ["-include", "/tmp/torchada_stable_box.h"],
197+
},
198+
)
199+
200+
assert self._requires(extension) is True
201+
202+
def test_cuda_extension_only_patches_headers_for_stable_abi(self, monkeypatch):
203+
from torchada.utils import cpp_extension
204+
from torchada.utils.cpp_extension import Platform
205+
206+
patched = []
207+
monkeypatch.setattr(cpp_extension, "detect_platform", lambda: Platform.MUSA)
208+
monkeypatch.setattr(
209+
cpp_extension,
210+
"_create_musa_extension",
211+
lambda name, sources, *args, **kwargs: SimpleNamespace(
212+
name=name,
213+
sources=sources,
214+
),
215+
)
216+
monkeypatch.setattr(
217+
cpp_extension,
218+
"_ensure_stable_headers_patched",
219+
lambda: patched.append(True),
220+
)
221+
222+
cpp_extension.CUDAExtension("torchvision_image", ["decode_jpegs_cuda.cpp"])
223+
assert patched == []
224+
225+
extension = cpp_extension.CUDAExtension(
226+
"custom_stable_ops",
227+
["ops.cu"],
228+
torchada_stable_abi=True,
229+
)
230+
assert extension._torchada_stable_abi is True
231+
assert patched == [True]
232+
233+
234+
class TestDependencyIncludeProtection:
235+
"""Editable dependency paths are never treated as project porting roots."""
236+
237+
def test_dependency_package_and_checkout_paths_overlap(self):
238+
from torchada.utils.cpp_extension import _path_overlaps_any
239+
240+
dependency_package = "/home/torch_musa/torch_musa"
241+
242+
assert _path_overlaps_any("/home/torch_musa", [dependency_package])
243+
assert _path_overlaps_any(
244+
"/home/torch_musa/torch_musa/share/generated_cuda_compatible",
245+
[dependency_package],
246+
)
247+
assert not _path_overlaps_any("/home/torchvision", [dependency_package])
248+
249+
148250
class TestMusaPatches:
149251
"""Test patches applied to torch_musa for extension building."""
150252

‎tests/test_inplace_porting.py‎

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -244,3 +244,46 @@ def get_mapping_rule(self):
244244

245245
assert source.read_text(encoding="utf-8") == f"{PORTED_TOKEN}\n"
246246
assert command._ported_dirs == {os.path.realpath(source_dir)}
247+
248+
249+
def test_editable_dependency_include_root_is_not_ported(tmp_path, monkeypatch):
250+
"""An editable torch_musa checkout stays an include path, not a port root."""
251+
from torchada.utils import cpp_extension
252+
from torchada.utils.cpp_extension import BuildExtension
253+
254+
project_dir = tmp_path / "vision"
255+
dependency_dir = tmp_path / "torch_musa"
256+
project_dir.mkdir()
257+
dependency_package = dependency_dir / "torch_musa"
258+
dependency_package.mkdir(parents=True)
259+
260+
source = project_dir / "kernel.cu"
261+
dependency_header = dependency_package / "dependency.h"
262+
source.write_text(f"{TOKEN}\n", encoding="utf-8")
263+
dependency_header.write_text(f"{TOKEN}\n", encoding="utf-8")
264+
265+
monkeypatch.setattr(
266+
cpp_extension,
267+
"_runtime_dependency_roots",
268+
lambda: [str(dependency_package)],
269+
)
270+
271+
class CustomBuildExtension(BuildExtension):
272+
def get_mapping_rule(self):
273+
return MAPPING
274+
275+
command = CustomBuildExtension(Distribution())
276+
command.extensions = [
277+
Extension(
278+
"test_editable_dependency",
279+
sources=[str(source)],
280+
include_dirs=[str(dependency_dir)],
281+
)
282+
]
283+
monkeypatch.setattr(BuildExtension.__mro__[1], "run", lambda self: None)
284+
285+
command.run()
286+
287+
assert source.read_text(encoding="utf-8") == f"{PORTED_TOKEN}\n"
288+
assert dependency_header.read_text(encoding="utf-8") == f"{TOKEN}\n"
289+
assert command._ported_dirs == {os.path.realpath(project_dir)}

0 commit comments

Comments
 (0)