diff --git a/tests/test_import_fixes_drift.py b/tests/test_import_fixes_drift.py index 7faabf851b..074e3c90fe 100644 --- a/tests/test_import_fixes_drift.py +++ b/tests/test_import_fixes_drift.py @@ -786,6 +786,15 @@ def _live_c10d_functional_ops(): return sorted({n.split("::", 1)[1] for n in get_ops() if n.startswith("_c10d_functional::")}) +def _live_dtensor_ops(): + import torch + + get_ops = getattr(torch._C, "_dispatch_get_all_op_names", None) + if not callable(get_ops): + pytest.skip("dispatcher op enumeration unavailable") + return sorted({n.split("::", 1)[1] for n in get_ops() if n.startswith("_dtensor::")}) + + def test_torchao_rocm_shim_schema_table_matches_installed_torch(): """The `_c10d_functional` schema table must exactly match the ops the installed torch registers (op set + canonical schema strings). A minor with no row means the @@ -830,6 +839,50 @@ def test_torchao_rocm_shim_schema_table_matches_installed_torch(): ) +def test_torchao_rocm_shim_dtensor_schema_matches_installed_torch(): + """The `_dtensor` schema table must exactly match the ops the installed torch registers. + torchao's `from torch.distributed._tensor import DTensor` runs + `register_fake("_dtensor::shard_dim_alltoall")` at import, which raises unless the op is + defined, so the shim must define this namespace too (not only _c10d_functional).""" + from unsloth.import_fixes import _DTENSOR_SCHEMAS, _schema_op_name + + import torch + + native = _live_dtensor_ops() + if not native: + pytest.skip("no native _dtensor ops (distributed-less torch build).") + + minor = _torch_minor_tuple() + schemas = _DTENSOR_SCHEMAS.get(minor) + if schemas is None: + pytest.skip( + f"no shim _dtensor row for torch {minor}; fix_torchao_windows_rocm_import " + f"fail-closes here (safe). Add a reviewed tuple to enable it (ops: {native})." + ) + + table_ops = sorted(_schema_op_name(s) for s in schemas) + assert table_ops == native, ( + f"DRIFT DETECTED: torchao shim _dtensor table for torch {minor} lists {table_ops} " + f"but the installed torch registers {native}. Update _DTENSOR_SCHEMAS." + ) + + parse = getattr(torch._C, "parse_schema", None) + if not callable(parse): + return + real = {} + for op in native: + packet = getattr(torch.ops._dtensor, op) + overload = packet.overloads()[0] + real[op] = str(getattr(packet, overload)._schema) + for s in schemas: + parsed = parse(f"_dtensor::{s}") # must not raise + name = _schema_op_name(s) + assert str(parsed) == real[name], ( + f"DRIFT DETECTED: torchao shim schema for _dtensor::{name}\n" + f" shim: {parsed}\n torch: {real[name]}" + ) + + def test_torchao_rocm_shim_strict_noop_on_non_windows(): """On a non-Windows / distributed-present box the shim must not touch sys.modules, torch.ops, or torch.distributed.is_available().""" @@ -929,9 +982,15 @@ def test_torchao_rocm_shim_source_has_guards_fragment_and_rollback(): src = inspect.getsource(import_fixes.fix_torchao_windows_rocm_import) assert "win32" in src, "missing win32 guard" assert "hip" in src, "missing torch.version.hip guard" + assert '"rocm" in' in src, ( + "shim ROCm detection must also accept a 'rocm'-tagged __version__ wheel (parity with " + "the Studio is_win32_rocm() helper), not gate on torch.version.hip alone" + ) assert "is_available()" in src, "missing is_available() guard" assert '"FRAGMENT"' in src, "shim must register with FRAGMENT, not DEF" assert '"_c10d_functional", "DEF"' not in src, "shim must never DEF _c10d_functional" + assert '"_dtensor", "DEF"' not in src, "shim must never DEF _dtensor" + assert '"_dtensor"' in src, "shim must also register the _dtensor namespace" assert "_destroy" in src, "missing rollback via Library._destroy" diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py index d96c2065a7..39618f90be 100644 --- a/unsloth/import_fixes.py +++ b/unsloth/import_fixes.py @@ -3197,6 +3197,22 @@ _C10D_FUNCTIONAL_SCHEMAS[(2, 10)] = _C10D_FUNCTIONAL_SCHEMAS[(2, 9)] + ( ) _C10D_FUNCTIONAL_SCHEMAS[(2, 11)] = _C10D_FUNCTIONAL_SCHEMAS[(2, 10)] +# The `_dtensor` namespace is likewise DEF'd only in C++ (the same Functional.cpp), so it too +# is absent on a distributed-less ROCm wheel. torchao's `from torch.distributed._tensor import +# DTensor` loads torch.distributed.tensor._collective_utils, which at import does +# `@torch.library.register_fake("_dtensor::shard_dim_alltoall")`; register_fake raises +# "operator _dtensor::shard_dim_alltoall does not exist" unless the op is already defined, so +# the shim must define it in the same transaction or `import torchao` still fails and rolls +# back. Schema verified against the live 2.9 dispatcher and the v2.11.0 Functional.cpp source +# (stable across 2.9-2.11). Fail closed on any other minor. +_DTENSOR_SCHEMAS = { + (2, 9): ( + "shard_dim_alltoall(Tensor input, int gather_dim, int shard_dim, str group_name) -> Tensor", + ), +} +_DTENSOR_SCHEMAS[(2, 10)] = _DTENSOR_SCHEMAS[(2, 9)] +_DTENSOR_SCHEMAS[(2, 11)] = _DTENSOR_SCHEMAS[(2, 9)] + def _schema_op_name(schema): """`all_reduce(Tensor ...) -> Tensor` -> `all_reduce`.""" @@ -3208,15 +3224,19 @@ def _torchao_shim_torch_minor(torch): return (int(base[0]), int(base[1])) +_TORCHAO_ROCM_NATIVE_PREFIXES = ("_c10d_functional::", "_dtensor::") + + def _native_c10d_functional_present(torch): - """True if the dispatcher already has any `_c10d_functional::` op (real torch - distributed present). Fail closed: an unexpected error counts as present, so the shim - never registers over a real namespace.""" + """True if the dispatcher already has any `_c10d_functional::` or `_dtensor::` op (real + torch distributed present). Both namespaces are DEF'd in the same C++ Functional.cpp, so + either being present means the shim must not register over it. Fail closed: an unexpected + error counts as present, so the shim never registers over a real namespace.""" get_ops = getattr(torch._C, "_dispatch_get_all_op_names", None) if not callable(get_ops): return True try: - return any(n.startswith("_c10d_functional::") for n in get_ops()) + return any(n.startswith(_TORCHAO_ROCM_NATIVE_PREFIXES) for n in get_ops()) except Exception: return True @@ -3368,8 +3388,8 @@ def _make_torchao_rocm_fake_c10d(): def fix_torchao_windows_rocm_import(): """On a legacy Windows ROCm wheel (no torch.distributed C-extension), make real torchao importable by faking `torch._C._distributed_c10d` and FRAGMENT-registering the - `_c10d_functional` op schemas, so torchao's module-top distributed imports resolve and - portable FP8/INT8 export works instead of torchao being stubbed off. + `_c10d_functional` and `_dtensor` op schemas, so torchao's module-top distributed imports + resolve and portable FP8/INT8 export works instead of torchao being stubbed off. Strict no-op unless every capability guard holds (Windows + HIP torch + distributed genuinely absent + known torch minor + torchao installed and not yet imported). Fully @@ -3389,11 +3409,20 @@ def fix_torchao_windows_rocm_import(): try: import torch - # ROCm build only (authoritative runtime HIP field, not the loose version tag). - if not getattr(getattr(torch, "version", None), "hip", None): + # ROCm build: mirror the Studio is_win32_rocm() detector -- HIP field OR a "rocm" + # __version__ tag (AMD SDK wheels lack torch.version.hip but tag "rocm"). Matching + # it keeps the shim and the export gate from drifting so the same wheels the gate + # disables torchao on are the ones the shim re-enables. The capability guards below + # (distributed absent, no _c10d_functional/_dtensor ops) rule out any false positive. + if not ( + getattr(getattr(torch, "version", None), "hip", None) + or "rocm" in getattr(torch, "__version__", "").lower() + ): return - schemas = _C10D_FUNCTIONAL_SCHEMAS.get(_torchao_shim_torch_minor(torch)) - if schemas is None: + minor = _torchao_shim_torch_minor(torch) + schemas = _C10D_FUNCTIONAL_SCHEMAS.get(minor) + dtensor_schemas = _DTENSOR_SCHEMAS.get(minor) + if schemas is None or dtensor_schemas is None: return # unknown torch minor -> fail closed if importlib.util.find_spec("torchao") is None: return @@ -3415,14 +3444,22 @@ def fix_torchao_windows_rocm_import(): modules_before = set(sys.modules) had_c10d_attr = hasattr(torch._C, "_distributed_c10d") fake = None - lib = None + libs = [] try: # Re-check the dispatcher immediately before touching it (TOCTOU guard). if _native_c10d_functional_present(torch): return - lib = torch.library.Library("_c10d_functional", "FRAGMENT") # FRAGMENT, never DEF - for schema in schemas: - lib.define(schema) + # FRAGMENT (never DEF): defines the schemas torch's distributed Python modules + # register impls / fakes against at import. `_c10d_functional` for + # _functional_collectives, `_dtensor` for tensor._collective_utils. + for namespace, ns_schemas in ( + ("_c10d_functional", schemas), + ("_dtensor", dtensor_schemas), + ): + lib = torch.library.Library(namespace, "FRAGMENT") + libs.append(lib) + for schema in ns_schemas: + lib.define(schema) fake = _make_torchao_rocm_fake_c10d() sys.modules[_C10D_EXT_MODULE] = fake setattr(torch._C, "_distributed_c10d", fake) @@ -3431,7 +3468,7 @@ def fix_torchao_windows_rocm_import(): except BaseException: # Atomic rollback: destroy schemas, drop the fake, purge only the torchao / # distributed submodules this transaction newly created. - if lib is not None: + for lib in libs: try: lib._destroy() except Exception: @@ -3456,10 +3493,10 @@ def fix_torchao_windows_rocm_import(): return # Commit: keep strong refs so the FRAGMENT schemas outlive GC. - _TORCHAO_ROCM_SHIM_STATE = {"fake_module": fake, "libraries": [lib]} + _TORCHAO_ROCM_SHIM_STATE = {"fake_module": fake, "libraries": libs} _log_rocm_detection( "Unsloth: Installed the torchao Windows-ROCm import shim " - "(fake torch._C._distributed_c10d + _c10d_functional schemas)." + "(fake torch._C._distributed_c10d + _c10d_functional/_dtensor schemas)." )