2626"""
2727
2828import functools
29+ import importlib
2930import inspect
3031import logging
3132import 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+
246275def _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+
929990class 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
9551025class 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
0 commit comments