diff --git a/studio/backend/core/_torchao_stub.py b/studio/backend/core/_torchao_stub.py
index 6336954bd5..2250ea4125 100644
--- a/studio/backend/core/_torchao_stub.py
+++ b/studio/backend/core/_torchao_stub.py
@@ -100,36 +100,43 @@ class _StubSubpackageFinder(importlib.abc.MetaPathFinder):
)
+def is_win32_rocm() -> bool:
+ """True on Windows ROCm, where torch.distributed (and thus torchao) is unavailable.
+
+ Gate on the runtime torch, not env vars (HIP_PATH persists after a CUDA revert). AMD SDK
+ wheels lack torch.version.hip but tag "rocm" in __version__, so accept either. Shared by the
+ import stub and the export gate so they can't drift.
+ """
+ if sys.platform != "win32":
+ return False
+ try:
+ import torch
+ return bool(
+ getattr(getattr(torch, "version", None), "hip", None)
+ or "rocm" in getattr(torch, "__version__", "").lower()
+ )
+ except Exception:
+ return False
+
+
def install_torchao_windows_rocm_stub() -> None:
"""Pre-stub torchao on Windows ROCm so transformers/peft imports don't crash.
No-op elsewhere (incl. Windows CUDA, where torchao is real). Must run before
importing transformers / unsloth_zoo. Safe to call once per worker.
"""
- # Gate on the active torch runtime, not env-var presence -- HIP_PATH/ROCM_PATH
- # persist after reverting to a CUDA wheel. Some ROCm wheels lack
- # torch.version.hip but still encode "rocm" in __version__, so accept either.
- _is_win32_rocm = False
- if sys.platform == "win32":
- try:
- import torch as _torch_probe
- _is_win32_rocm = bool(
- getattr(getattr(_torch_probe, "version", None), "hip", None)
- or "rocm" in getattr(_torch_probe, "__version__", "").lower()
- )
- del _torch_probe
- except Exception:
- pass
- if _is_win32_rocm:
- # Register the finder only on Windows ROCm.
+ if not is_win32_rocm():
+ return
+ # Register the finder only on Windows ROCm, and only once (no duplicates on re-call).
+ if not any(isinstance(_f, _StubSubpackageFinder) for _f in sys.meta_path):
sys.meta_path.append(_StubSubpackageFinder())
- # Seed torchao top-level + key submodules; the finder handles the rest.
- for _tao_name in (
- "torchao",
- "torchao.quantization",
- "torchao.dtypes",
- "torchao.float8",
- "torchao.utils",
- ):
- if _tao_name not in sys.modules:
- sys.modules[_tao_name] = _make_mod_stub(_tao_name)
+ # Seed torchao top-level + key submodules; the finder handles the rest.
+ for _tao_name in (
+ "torchao",
+ "torchao.quantization",
+ "torchao.dtypes",
+ "torchao.float8",
+ "torchao.utils",
+ ):
+ if _tao_name not in sys.modules:
+ sys.modules[_tao_name] = _make_mod_stub(_tao_name)
diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py
index c8be50b08b..46c9aac787 100644
--- a/studio/backend/core/export/export.py
+++ b/studio/backend/core/export/export.py
@@ -102,14 +102,52 @@ def _compressed_export_supported():
def _torchao_export_supported():
- """True if the installed unsloth build has the portable torchao FP8/INT8 export path."""
+ """True if the installed unsloth build has the portable torchao FP8/INT8 export path.
+
+ Forced False on Windows ROCm, where torchao is import-stubbed (no torch.distributed) and its
+ config classes return None. Unchanged on Windows CUDA / Linux / macOS (torchao is real)."""
try:
+ from core._torchao_stub import is_win32_rocm
+
+ if is_win32_rocm():
+ return False
import unsloth.save as _us
+
return hasattr(_us, "_normalize_torchao_method")
except Exception:
return False
+def _torchao_runtime_unavailable():
+ """True where portable torchao export cannot run (Windows ROCm): torchao is import-stubbed
+ (its config classes return None) or torch.distributed is absent. False everywhere else."""
+ import sys
+ try:
+ from core._torchao_stub import is_win32_rocm, _STUB_SENTINEL
+
+ if is_win32_rocm():
+ return True
+ _tao = sys.modules.get("torchao")
+ return _tao is not None and getattr(_tao, "_unsloth_stub", None) is _STUB_SENTINEL
+ except Exception:
+ return False
+
+
+def _is_torchao_alias(alias):
+ """True if `alias` is any torchao export form (torchao_fp8, portable_int8, hyphen/space
+ variants) per unsloth's normalizer, with a torchao_ prefix fallback. Catches a torchao request
+ before the Windows-ROCm gate misclassifies it as compressed-tensors."""
+ if not alias:
+ return False
+ try:
+ import unsloth.save as _us
+ if _us._normalize_torchao_method(alias) is not None:
+ return True
+ except Exception:
+ pass
+ return str(alias).lower().startswith("torchao")
+
+
def _has_nvidia_gpu():
"""True only on a real NVIDIA CUDA box (not ROCm/XPU/CPU/MLX); compressed-tensors needs it."""
try:
@@ -495,6 +533,19 @@ class ExportBackend:
"NVFP4 (compressed-tensors)": "nvfp4",
}
compressed_alias = compressed_method or _LABEL_TO_ALIAS.get(format_type)
+
+ # Portable torchao is unavailable on Windows ROCm (stubbed, no torch.distributed). Reject
+ # any torchao alias early with a clear message instead of the cryptic NoneType crash or a
+ # misleading NVIDIA error. Other formats (16-bit/GGUF/compressed-tensors) are unaffected.
+ if _is_torchao_alias(compressed_alias) and _torchao_runtime_unavailable():
+ return (
+ False,
+ "Portable torchao FP8/INT8 export is not supported on Windows ROCm: "
+ "torch.distributed and torchao are unavailable on this build. Use 16-bit "
+ "merged or GGUF quantization instead.",
+ None,
+ )
+
compressed_suffix: Optional[str] = None
# Classify the alias: torchao-portable vs compressed-tensors.
torchao_info = None
diff --git a/studio/backend/tests/test_torchao_select.py b/studio/backend/tests/test_torchao_select.py
index e4775a10a6..fa57b08874 100644
--- a/studio/backend/tests/test_torchao_select.py
+++ b/studio/backend/tests/test_torchao_select.py
@@ -1,16 +1,19 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
-"""Tests for _select_torchao_spec in install_python_stack.py.
+"""Tests for torchao version selection and the Windows-ROCm export gate.
-torchao's C++ extensions are built against one exact torch release, so the
-installer must pick the torchao version matching the torch installed in the
-venv (otherwise the cpp kernels are skipped). This pins that mapping.
+First half: the installer must pin the torchao version matching the installed torch (its cpp
+kernels are built per torch release). Second half: torch.distributed is unsupported on Windows
+ROCm, so torchao is import-stubbed and the portable FP8/INT8 export must be gated off there
+(shared is_win32_rocm() helper) with a clear defensive error.
"""
from __future__ import annotations
+import ast
import sys
+import types
from pathlib import Path
from unittest.mock import MagicMock
@@ -19,6 +22,9 @@ import pytest
# install_python_stack.py lives at repo_root/studio/install_python_stack.py
_INSTALL_SCRIPT = Path(__file__).resolve().parents[2] / "install_python_stack.py"
+# backend root (studio/backend), for reading/exec-ing backend sources.
+_BACKEND = Path(__file__).resolve().parents[1]
+
def _load_module(monkeypatch):
"""(Re-)import install_python_stack and return it (mirrors test_pytorch_mirror)."""
@@ -134,3 +140,225 @@ def test_skips_torchao_on_windows_rocm(
assert not any(spec.startswith("torchao") for spec in installed_specs)
assert "dependency overrides (skipped, Windows ROCm)" in progress_labels
+
+
+# -- Windows-ROCm torchao export gate -----------------------------------------------------------
+# torchao is import-stubbed on Windows ROCm (no torch.distributed) and its config classes return
+# None, which made TorchAoConfig(quant_type=None) crash. These prove the shared is_win32_rocm()
+# gate hides the torchao formats and the defensive path raises a clear error instead.
+
+import core._torchao_stub as _stub
+
+
+def _func_src(rel, name):
+ src = (_BACKEND / rel).read_text(encoding = "utf-8")
+ node = next(
+ n for n in ast.walk(ast.parse(src)) if isinstance(n, ast.FunctionDef) and n.name == name
+ )
+ return ast.get_source_segment(src, node)
+
+
+def _exec_func(rel, name):
+ """Exec one backend function in isolation, avoiding export.py's heavy import chain."""
+ ns: dict = {}
+ exec(_func_src(rel, name), ns)
+ return ns[name]
+
+
+@pytest.mark.parametrize(
+ ("platform", "hip", "version", "expected"),
+ [
+ ("win32", "6.4.0", "2.10.0+rocm6.4", True), # ROCm via torch.version.hip
+ ("win32", None, "2.10.0+rocm6.4", True), # ROCm via __version__ tag only
+ ("win32", None, "2.10.0+cu128", False), # Windows CUDA -> real torchao
+ ("linux", "6.4.0", "2.10.0+rocm6.4", False), # Linux ROCm -> real torchao
+ ("darwin", None, "2.10.0", False), # macOS
+ ],
+)
+def test_is_win32_rocm(monkeypatch, platform, hip, version, expected):
+ fake_torch = types.SimpleNamespace(version = types.SimpleNamespace(hip = hip), __version__ = version)
+ monkeypatch.setattr(sys, "platform", platform)
+ monkeypatch.setitem(sys.modules, "torch", fake_torch)
+ assert _stub.is_win32_rocm() is expected
+
+
+def test_gate_and_stub_share_helper():
+ # The stub installer and the export gate must both route through is_win32_rocm() so they can't
+ # drift (the gate off while the stub is still active, or the reverse).
+ stub_src = (_BACKEND / "core" / "_torchao_stub.py").read_text(encoding = "utf-8")
+ assert "def is_win32_rocm(" in stub_src
+ assert "is_win32_rocm()" in _func_src(
+ "core/_torchao_stub.py", "install_torchao_windows_rocm_stub"
+ )
+ assert "is_win32_rocm()" in _func_src("core/export/export.py", "_torchao_export_supported")
+
+
+def test_installer_noop_off_windows_rocm(monkeypatch):
+ # is_win32_rocm() False -> installer must not register the finder or seed torchao stubs.
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: False)
+ before = list(sys.meta_path)
+ _stub.install_torchao_windows_rocm_stub()
+ assert list(sys.meta_path) == before
+
+
+# (a) gate off on Windows ROCm; (b) unchanged elsewhere
+
+
+def test_torchao_gate_false_on_windows_rocm(monkeypatch):
+ # (a) On Windows ROCm the portable torchao formats are not offered, without importing unsloth.
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: True)
+ assert _exec_func("core/export/export.py", "_torchao_export_supported")() is False
+
+
+_TORCHAO_ALIASES = {"torchao_fp8", "torchao_int8", "portable_fp8", "portable_int8"}
+
+
+def _fake_normalize_torchao(save_method):
+ # Mirrors unsloth.save._normalize_torchao_method (lower/strip, - and space -> _).
+ if not isinstance(save_method, str):
+ return None
+ key = save_method.lower().strip().replace("-", "_").replace(" ", "_")
+ return ("fp8", "torchao-fp8") if key in _TORCHAO_ALIASES else None
+
+
+def _install_fake_unsloth_save(monkeypatch, *, has_method):
+ unsloth = types.ModuleType("unsloth")
+ save = types.ModuleType("unsloth.save")
+ if has_method:
+ save._normalize_torchao_method = _fake_normalize_torchao
+ unsloth.save = save
+ monkeypatch.setitem(sys.modules, "unsloth", unsloth)
+ monkeypatch.setitem(sys.modules, "unsloth.save", save)
+
+
+def test_torchao_gate_supported_off_windows_rocm(monkeypatch):
+ # (b) Off Windows ROCm the gate is unchanged: True when the unsloth build has the method.
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: False)
+ _install_fake_unsloth_save(monkeypatch, has_method = True)
+ assert _exec_func("core/export/export.py", "_torchao_export_supported")() is True
+
+
+def test_torchao_gate_false_when_build_lacks_method(monkeypatch):
+ # (b) Off Windows ROCm, an older unsloth without the method is still unsupported.
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: False)
+ _install_fake_unsloth_save(monkeypatch, has_method = False)
+ assert _exec_func("core/export/export.py", "_torchao_export_supported")() is False
+
+
+# (c) defensive early error when torchao is stubbed / unavailable
+
+
+def _load_export_module_no_torch(monkeypatch):
+ """Import core.export.export with torch/unsloth blocked (mirrors test_export_capability), so
+ the defensive path runs on CPU with no GPU and no torchao."""
+ import builtins
+ import importlib
+
+ real_import = builtins.__import__
+
+ def blocking_import(name, *args, **kwargs):
+ # Block real torch/unsloth, but honor injected fakes already in sys.modules.
+ top = name.split(".")[0]
+ if top in {"torch", "unsloth"} and top not in sys.modules:
+ raise ImportError(f"blocked: {name}")
+ return real_import(name, *args, **kwargs)
+
+ for m in [k for k in list(sys.modules) if k.split(".")[0] in {"torch", "unsloth"}]:
+ monkeypatch.delitem(sys.modules, m, raising = False)
+ monkeypatch.delitem(sys.modules, "core.export.export", raising = False)
+ monkeypatch.setattr(builtins, "__import__", blocking_import)
+ return importlib.import_module("core.export.export")
+
+
+def _bare_backend(mod):
+ be = mod.ExportBackend.__new__(mod.ExportBackend)
+ be.current_model = object()
+ be.current_tokenizer = object()
+ be._audio_type = None
+ be.is_peft = True
+ return be
+
+
+def test_torchao_defensive_error_on_windows_rocm(monkeypatch):
+ # (c) A forced torchao request reaches the merged path -> clear error, not the NoneType crash.
+ mod = _load_export_module_no_torch(monkeypatch)
+ monkeypatch.setattr(mod, "_export_runtime_available", lambda: True)
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: True)
+
+ ok, message, out = _bare_backend(mod).export_merged_model(
+ "/tmp/x", compressed_method = "torchao_fp8"
+ )
+ assert ok is False and out is None
+ assert "Windows ROCm" in message and "torchao" in message.lower()
+
+
+def test_torchao_defensive_error_alias_form_on_windows_rocm(monkeypatch):
+ # An equivalent alias unsloth accepts (portable_fp8) must hit the same rejection, not fall
+ # through to the misleading NVIDIA compressed-tensors error.
+ mod = _load_export_module_no_torch(monkeypatch)
+ monkeypatch.setattr(mod, "_export_runtime_available", lambda: True)
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: True)
+ _install_fake_unsloth_save(monkeypatch, has_method = True)
+
+ ok, message, out = _bare_backend(mod).export_merged_model(
+ "/tmp/x", compressed_method = "portable_fp8"
+ )
+ assert ok is False and out is None
+ assert "Windows ROCm" in message and "torchao" in message.lower()
+
+
+def test_is_torchao_alias_recognizes_all_forms(monkeypatch):
+ _install_fake_unsloth_save(monkeypatch, has_method = True)
+ fn = _exec_func("core/export/export.py", "_is_torchao_alias")
+ for alias in ("torchao_fp8", "portable_int8", "portable-fp8", "Portable FP8"):
+ assert fn(alias) is True
+ for alias in ("fp8", "nvfp4", "w8a8", "", None):
+ assert fn(alias) is False
+
+
+def test_torchao_defensive_error_wired_early():
+ # Guard is in export_merged_model before the merge/quant work; alias is normalized (not just the
+ # torchao_ prefix) so every torchao form is caught.
+ m = _func_src("core/export/export.py", "export_merged_model")
+ assert "_is_torchao_alias(compressed_alias)" in m
+ assert "_torchao_runtime_unavailable()" in m
+ alias_fn = _func_src("core/export/export.py", "_is_torchao_alias")
+ assert "_normalize_torchao_method(alias)" in alias_fn
+ assert 'startswith("torchao")' in alias_fn
+
+
+# (issue 2/4) backend win32_rocm flag + single finder registration
+
+
+def test_export_capability_exposes_win32_rocm(monkeypatch):
+ import utils.hardware.hardware as hw
+
+ monkeypatch.setattr(hw, "get_device", lambda: hw.DeviceType.CUDA)
+ monkeypatch.setattr(sys, "platform", "win32")
+ monkeypatch.setattr(hw, "IS_ROCM", True)
+ assert hw.export_capability()["win32_rocm"] is True
+ monkeypatch.setattr(hw, "IS_ROCM", False)
+ assert hw.export_capability()["win32_rocm"] is False
+ monkeypatch.setattr(sys, "platform", "linux")
+ monkeypatch.setattr(hw, "IS_ROCM", True)
+ assert hw.export_capability()["win32_rocm"] is False
+
+
+def test_installer_registers_finder_once(monkeypatch):
+ # Repeated install must not stack duplicate finders. Restore global state after.
+ monkeypatch.setattr(_stub, "is_win32_rocm", lambda: True)
+ meta_before = list(sys.meta_path)
+ tao_before = {k for k in sys.modules if k == "torchao" or k.startswith("torchao.")}
+ try:
+ _stub.install_torchao_windows_rocm_stub()
+ _stub.install_torchao_windows_rocm_stub()
+ finders = [f for f in sys.meta_path if isinstance(f, _stub._StubSubpackageFinder)]
+ assert len(finders) == 1
+ finally:
+ sys.meta_path[:] = meta_before
+ for k in [
+ k
+ for k in sys.modules
+ if (k == "torchao" or k.startswith("torchao.")) and k not in tao_before
+ ]:
+ del sys.modules[k]
diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py
index 8d6c919ebd..1516ab7f96 100644
--- a/studio/backend/utils/hardware/hardware.py
+++ b/studio/backend/utils/hardware/hardware.py
@@ -270,13 +270,19 @@ def export_capability() -> dict:
import and has no CPU path), so it is supported iff ``get_device() in {CUDA, XPU, MLX}``. The
reason distinguishes a --no-torch install from a bare-CPU host. Safe to call without torch.
- Returns {export_supported, export_unsupported_reason, export_unsupported_message}.
+ Returns {export_supported, export_unsupported_reason, export_unsupported_message, win32_rocm}.
+ ``win32_rocm`` is the UI's single source of truth for the torchao gate (mirrors
+ is_win32_rocm()): torchao is unavailable on Windows ROCm.
"""
- if get_device() in (DeviceType.CUDA, DeviceType.XPU, DeviceType.MLX):
+ device = get_device()
+ # get_device() ran detect_hardware(), so IS_ROCM (hip OR "rocm" tag) is authoritative here.
+ win32_rocm = sys.platform == "win32" and IS_ROCM
+ if device in (DeviceType.CUDA, DeviceType.XPU, DeviceType.MLX):
return {
"export_supported": True,
"export_unsupported_reason": None,
"export_unsupported_message": None,
+ "win32_rocm": win32_rocm,
}
# No accelerator: name the blocker. Apple Silicon first -- its path is MLX, so "install PyTorch"
# would be wrong advice on a Mac even when torch is also absent.
@@ -303,6 +309,7 @@ def export_capability() -> dict:
"export_supported": False,
"export_unsupported_reason": reason,
"export_unsupported_message": message,
+ "win32_rocm": win32_rocm,
}
diff --git a/studio/frontend/src/features/export/export-page.tsx b/studio/frontend/src/features/export/export-page.tsx
index 07606a26ed..61e3efa9fb 100644
--- a/studio/frontend/src/features/export/export-page.tsx
+++ b/studio/frontend/src/features/export/export-page.tsx
@@ -223,8 +223,12 @@ export function ExportPage() {
const [ggufTarget, setGgufTarget] = useState<"model" | "lora">("model");
const hardware = useHardwareInfo();
+ const deviceType = usePlatformStore((s) => s.deviceType);
// GGUF LoRA conversion is rejected on the macOS / MLX path, so gate it out on a Mac host.
- const isMacHost = usePlatformStore((s) => s.deviceType) === "mac";
+ const isMacHost = deviceType === "mac";
+ // Backend truth for the torchao gate (single source). Not re-derived from `rocm`: AMD SDK
+ // wheels leave torch.version.hip unset, so `rocm` alone would miss Windows ROCm.
+ const isWindowsRocm = hardware.win32Rocm;
// Real CUDA (not ROCm); gates the NVIDIA-only compressed-tensors formats.
const hasNvidia = hardware.cuda != null && hardware.rocm == null;
// Only gray out on an authoritative unsupported response; while unloaded the backend route guard
@@ -239,14 +243,13 @@ export function ExportPage() {
MERGED_FORMATS.filter((f) => {
// compressed-tensors (llm-compressor) is the NVIDIA path; shown only on an NVIDIA GPU.
if (f.backend === "compressed") return hasNvidia;
- // Portable torchao is the fallback for hosts without the NVIDIA compressed path, i.e. a
- // CPU / non-NVIDIA box. Hidden on NVIDIA (use compressed-tensors) and on macOS/MLX (the
- // backend rejects quantized export there).
- if (f.backend === "torchao") return !hasNvidia && !isMacHost;
+ // Portable torchao: shown on non-NVIDIA hosts. Hidden on NVIDIA (use compressed-tensors),
+ // macOS/MLX (rejected), and Windows ROCm (torchao unavailable: no torch.distributed).
+ if (f.backend === "torchao") return !hasNvidia && !isMacHost && !isWindowsRocm;
// Plain 16-bit is available everywhere.
return true;
}),
- [hasNvidia, isMacHost],
+ [hasNvidia, isMacHost, isWindowsRocm],
);
const toggleFormat = useCallback((value: string) => {
setSelectedFormats((prev) =>
@@ -255,7 +258,17 @@ export function ExportPage() {
: [...prev, value],
);
}, []);
- // availableFormats already drops NVIDIA-only formats on other hardware, so no pruning needed.
+ // Drop a selected format the gate removed (e.g. torchao once win32Rocm resolves). Gate on
+ // hardware.loaded: before the authoritative response hasNvidia is false, so pruning would
+ // permanently drop a running NVIDIA FP8/NVFP4 pick that the later response can't restore.
+ useEffect(() => {
+ if (!hardware.loaded) return;
+ const allowed = new Set(availableFormats.map((f) => f.value));
+ setSelectedFormats((prev) => {
+ const next = prev.filter((v) => allowed.has(v));
+ return next.length === prev.length ? prev : next;
+ });
+ }, [availableFormats, hardware.loaded]);
// IQ quants are imatrix-only: force imatrix on when one is selected, else llama.cpp rejects it.
const requiresImatrix = quantLevels.some(
(q) => QUANT_OPTIONS.find((o) => o.value === q)?.imatrix,
@@ -1434,13 +1447,20 @@ export function ExportPage() {
)}
- {!hasNvidia && (
+ {!hasNvidia && !isWindowsRocm && (
No NVIDIA GPU detected: compressed-tensors formats are
hidden. 16-bit and portable FP8/INT8 (torchao) still
work here and load in vLLM.
)}
+
+ {isWindowsRocm && (
+
+ Windows ROCm: quantized FP8/INT8 (torchao) export is
+ unavailable (no torch.distributed). Use 16-bit or GGUF.
+
+ )}
)}
diff --git a/studio/frontend/src/hooks/use-hardware-info.ts b/studio/frontend/src/hooks/use-hardware-info.ts
index 4d63d4d6af..ff34eef089 100644
--- a/studio/frontend/src/hooks/use-hardware-info.ts
+++ b/studio/frontend/src/hooks/use-hardware-info.ts
@@ -31,6 +31,9 @@ export interface HardwareInfo {
exportSupported: boolean | null;
exportUnsupportedReason: string | null;
exportUnsupportedMessage: string | null;
+ // Backend truth for the torchao gate (mirrors is_win32_rocm(): torch.version.hip OR a "rocm"
+ // build tag). Single source; the UI must not re-derive Windows ROCm from `rocm` alone.
+ win32Rocm: boolean;
loaded: boolean;
}
@@ -48,6 +51,7 @@ const DEFAULT: HardwareInfo = {
exportSupported: null,
exportUnsupportedReason: null,
exportUnsupportedMessage: null,
+ win32Rocm: false,
loaded: false,
};
@@ -101,6 +105,7 @@ async function fetchOnce(): Promise {
exportSupported: data?.export_supported ?? null,
exportUnsupportedReason: data?.export_unsupported_reason ?? null,
exportUnsupportedMessage: data?.export_unsupported_message ?? null,
+ win32Rocm: data?.win32_rocm ?? false,
loaded: true,
};
if (generation === cacheGeneration) {
diff --git a/tests/test_import_fixes_drift.py b/tests/test_import_fixes_drift.py
index 0bee68f940..074e3c90fe 100644
--- a/tests/test_import_fixes_drift.py
+++ b/tests/test_import_fixes_drift.py
@@ -758,3 +758,248 @@ def test_bitsandbytes_rocm_detection_helpers_recognizable():
"decline to patch it and Windows ROCm import-time noise / "
"wrong ROCM_GPU_ARCH may return."
)
+
+
+# ===========================================================================
+# torchao Windows-ROCm import shim -- fix_torchao_windows_rocm_import
+# ===========================================================================
+# The shim FRAGMENT-registers the `_c10d_functional` op schemas so real torchao
+# imports on a distributed-less Windows ROCm wheel. These verify the schema table
+# tracks the installed torch, the FRAGMENT (not DEF) collision semantics, the
+# strict no-op / capability gating, and that it is wired into startup. Windows ROCm
+# itself cannot be reproduced here, so the transactional acceptance-import + rollback
+# is what guarantees no regression on the real device.
+
+
+def _torch_minor_tuple():
+ import torch
+ base = torch.__version__.split("+", 1)[0].split(".")
+ return (int(base[0]), int(base[1]))
+
+
+def _live_c10d_functional_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("_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
+ shim fail-closes there (safe, no coverage) -> skip; a present row must be exact."""
+ from unsloth.import_fixes import _C10D_FUNCTIONAL_SCHEMAS, _schema_op_name
+
+ import torch
+
+ native = _live_c10d_functional_ops()
+ if not native:
+ pytest.skip("no native _c10d_functional ops (distributed-less torch build).")
+
+ minor = _torch_minor_tuple()
+ schemas = _C10D_FUNCTIONAL_SCHEMAS.get(minor)
+ if schemas is None:
+ pytest.skip(
+ f"no shim schema 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 _c10d_functional table for torch {minor} lists "
+ f"{table_ops} but the installed torch registers {native}. Update "
+ f"_C10D_FUNCTIONAL_SCHEMAS."
+ )
+
+ parse = getattr(torch._C, "parse_schema", None)
+ if not callable(parse):
+ return
+ real = {}
+ for op in native:
+ packet = getattr(torch.ops._c10d_functional, op)
+ overload = packet.overloads()[0]
+ real[op] = str(getattr(packet, overload)._schema)
+ for s in schemas:
+ parsed = parse(f"_c10d_functional::{s}") # must not raise
+ name = _schema_op_name(s)
+ assert str(parsed) == real[name], (
+ f"DRIFT DETECTED: torchao shim schema for _c10d_functional::{name}\n"
+ f" shim: {parsed}\n torch: {real[name]}"
+ )
+
+
+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()."""
+ from unsloth.import_fixes import fix_torchao_windows_rocm_import
+
+ import torch
+
+ assert sys.platform != "win32"
+ before_ext = "torch._C._distributed_c10d" in sys.modules
+ before_avail = torch.distributed.is_available()
+ before_ops = set(_live_c10d_functional_ops())
+
+ fix_torchao_windows_rocm_import()
+
+ assert ("torch._C._distributed_c10d" in sys.modules) == before_ext
+ assert torch.distributed.is_available() == before_avail
+ assert set(_live_c10d_functional_ops()) == before_ops
+
+
+def test_torchao_rocm_shim_native_present_builds_no_library(monkeypatch):
+ """Even with the platform gates spoofed to look like Windows ROCm, a box that already
+ has real distributed (native _c10d_functional ops / is_available) must trip a guard
+ before any torch.library.Library is constructed."""
+ from unsloth.import_fixes import (
+ fix_torchao_windows_rocm_import,
+ _native_c10d_functional_present,
+ )
+
+ import torch
+
+ monkeypatch.setattr(sys, "platform", "win32")
+ if not getattr(getattr(torch, "version", None), "hip", None):
+ monkeypatch.setattr(torch.version, "hip", "6.4.0", raising = False)
+
+ assert _native_c10d_functional_present(torch) is True
+
+ calls = {"n": 0}
+ real_library = torch.library.Library
+
+ def _tripwire(*args, **kwargs):
+ calls["n"] += 1
+ return real_library(*args, **kwargs)
+
+ monkeypatch.setattr(torch.library, "Library", _tripwire)
+ fix_torchao_windows_rocm_import()
+ assert (
+ calls["n"] == 0
+ ), "torchao shim constructed a torch.library.Library despite a real distributed build."
+
+
+_TORCHAO_ROCM_FRAGMENT_PROBE = """
+import torch
+ns = "_unsloth_torchao_shim_probe_ns"
+# A second DEF on a namespace raises; FRAGMENT must not -- that is why the shim uses
+# FRAGMENT (no fatal collision with a native C++ TORCH_LIBRARY). NB: the first DEF must be
+# held by a strong ref, else CPython GCs it (its __del__ calls _destroy) and releases the
+# namespace before the second call -- the same reason the shim keeps a strong ref to its
+# FRAGMENT Library so its registered schemas are not dropped.
+_hold = torch.library.Library(ns, "DEF")
+raised = False
+try:
+ torch.library.Library(ns, "DEF")
+except Exception:
+ raised = True
+assert raised, "second DEF unexpectedly did not raise"
+frag = torch.library.Library(ns, "FRAGMENT") # must not raise
+frag.define("myop(Tensor x) -> Tensor")
+assert hasattr(torch.ops, ns) and hasattr(getattr(torch.ops, ns), "myop")
+frag._destroy()
+print("FRAGMENT_OK")
+"""
+
+
+def test_torchao_rocm_shim_fragment_semantics_subprocess():
+ """FRAGMENT-define + resolve + _destroy work and FRAGMENT (unlike a second DEF) never
+ collides -- run in a subprocess since dispatcher registration is process-global."""
+ import subprocess
+
+ result = subprocess.run(
+ [sys.executable, "-c", _TORCHAO_ROCM_FRAGMENT_PROBE],
+ capture_output = True,
+ text = True,
+ timeout = 300,
+ )
+ assert (
+ "FRAGMENT_OK" in result.stdout
+ ), f"FRAGMENT probe failed:\nSTDOUT:{result.stdout}\nSTDERR:{result.stderr}"
+
+
+def test_torchao_rocm_shim_source_has_guards_fragment_and_rollback():
+ """The shim source must keep its win32 + HIP + is_available() gates, use FRAGMENT (not a
+ DEF on _c10d_functional), and roll back via Library._destroy."""
+ import inspect
+
+ from unsloth import import_fixes
+
+ 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"
+
+
+def test_torchao_rocm_shim_wired_into_gpu_init():
+ """The shim must be called at startup (before `import unsloth_zoo`), not merely
+ importable (mirrors test_accelerate_patch_wired_into_gpu_init)."""
+ source = Path(__file__).resolve().parent.parent / "unsloth" / "_gpu_init.py"
+ text = source.read_text()
+ assert "fix_torchao_windows_rocm_import()" in text, (
+ "DRIFT DETECTED: fix_torchao_windows_rocm_import is defined but never called in "
+ "_gpu_init.py, so real imports never install it."
+ )
diff --git a/unsloth/_gpu_init.py b/unsloth/_gpu_init.py
index 984057e9f7..ec14af82e9 100644
--- a/unsloth/_gpu_init.py
+++ b/unsloth/_gpu_init.py
@@ -32,6 +32,7 @@ from .import_fixes import (
disable_broken_vllm,
configure_amdgpu_asic_id_table_path,
fix_bitsandbytes_rocm_arch_detection,
+ fix_torchao_windows_rocm_import,
torchvision_compatibility_check,
fix_diffusers_warnings,
fix_huggingface_hub,
@@ -70,6 +71,10 @@ except Exception:
configure_amdgpu_asic_id_table_path()
# Must precede `import unsloth_zoo` below, which imports bnb on ROCm.
fix_bitsandbytes_rocm_arch_detection()
+# Must also precede `import unsloth_zoo` below (it triggers the transformers/torchao
+# import chain): makes real torchao importable on legacy Windows ROCm wheels so
+# unsloth_zoo's torchao stub self-disables. Strict no-op elsewhere.
+fix_torchao_windows_rocm_import()
disable_broken_causal_conv1d()
disable_broken_vllm()
fix_message_factory_issue()
@@ -80,6 +85,7 @@ fix_diffusers_warnings()
fix_huggingface_hub()
del configure_amdgpu_asic_id_table_path
del fix_bitsandbytes_rocm_arch_detection
+del fix_torchao_windows_rocm_import
del disable_broken_causal_conv1d
del disable_broken_vllm
del fix_message_factory_issue
diff --git a/unsloth/import_fixes.py b/unsloth/import_fixes.py
index 09de248c7b..e36a074135 100644
--- a/unsloth/import_fixes.py
+++ b/unsloth/import_fixes.py
@@ -3129,3 +3129,366 @@ def patch_accelerate_recursively_apply():
setattr(mod, "find_device", _patched_find_device)
except Exception:
pass
+
+
+# ---------------------------------------------------------------------------
+# torchao Windows-ROCm import shim
+# ---------------------------------------------------------------------------
+# Legacy Windows ROCm wheels ship without the torch.distributed C-extension
+# (torch._C._distributed_c10d absent, torch.ops._c10d_functional.* unregistered), yet torchao
+# imports the distributed chain unconditionally at module load (float8/distributed_utils.py:
+# `import torch.distributed._functional_collectives` + `from torch.distributed._tensor import
+# DTensor`), so `import torchao` raises `No module named 'torch._C._distributed_c10d'` even on
+# paths that never use distributed. unsloth_zoo/Studio work around that by stubbing torchao off,
+# disabling FP8/INT8 export; this shim instead makes REAL torchao importable by faking the absent
+# C-ext module and FRAGMENT-registering the missing _c10d_functional/_dtensor schemas (schema-only
+# -- torch attaches its own Meta kernels, an actual collective still fails loudly, and weight-only
+# export invokes none). Fully transactional (any failure rolls back to torchao-unimportable, so
+# unsloth_zoo's stub still catches it: no regression) and capability-gated -- a strict no-op off
+# Windows-ROCm and once real torch.distributed is present (ROCm/TheRock#5694, torch >= 2.9).
+# FRAGMENT (not DEF) never collides with a native TORCH_LIBRARY. Opt out
+# UNSLOTH_DISABLE_TORCHAO_ROCM_SHIM=1; retire once torchao guards its float8 imports (pytorch/ao#1066).
+# ---------------------------------------------------------------------------
+
+_TORCHAO_ROCM_SHIM_SENTINEL = "__unsloth_torchao_rocm_shim__"
+_C10D_EXT_MODULE = "torch._C._distributed_c10d"
+# (fake_module, [Library, ...]) after a successful install; retained for the process so
+# GC does not drop the FRAGMENT-defined schemas.
+_TORCHAO_ROCM_SHIM_STATE = None
+
+# Per torch (major, minor): the exact `_c10d_functional` op schemas (namespace prefix added by
+# the Library). torch DEFs these only in C++, so they are absent on a distributed-less ROCm wheel
+# and torch's own _functional_collectives IMPL registrations fail at import. Verified: 2.9 vs the
+# installed dispatcher, 2.10/2.11 vs the v2.10.0/v2.11.0 Functional.cpp source. Fail closed on any
+# other minor (shim no-ops -> torchao stays stubbed, no regression).
+_C10D_FUNCTIONAL_SCHEMAS = {
+ (2, 9): (
+ "all_reduce(Tensor input, str reduce_op, str group_name) -> Tensor",
+ "all_reduce_(Tensor(a!) input, str reduce_op, str group_name) -> Tensor(a!)",
+ "all_reduce_coalesced(Tensor[] inputs, str reduce_op, str group_name) -> Tensor[]",
+ "all_reduce_coalesced_(Tensor[](a!) inputs, str reduce_op, str group_name) -> Tensor[](a!)",
+ "wait_tensor(Tensor tensor) -> Tensor",
+ "all_gather_into_tensor(Tensor input, int group_size, str group_name) -> Tensor",
+ "all_gather_into_tensor_out(Tensor input, int group_size, str group_name, *, Tensor(a!) out) -> Tensor(a!)",
+ "all_gather_into_tensor_coalesced(Tensor[] inputs, int group_size, str group_name) -> Tensor[]",
+ "reduce_scatter_tensor(Tensor input, str reduce_op, int group_size, str group_name) -> Tensor",
+ "reduce_scatter_tensor_coalesced(Tensor[] inputs, str reduce_op, int group_size, str group_name) -> Tensor[]",
+ "all_to_all_single(Tensor input, SymInt[] output_split_sizes, SymInt[] input_split_sizes, str group_name) -> Tensor",
+ "broadcast(Tensor input, int src, str group_name) -> Tensor",
+ "broadcast_(Tensor(a!) input, int src, str group_name) -> Tensor(a!)",
+ ),
+}
+# 2.10 and 2.11 add reduce_scatter_tensor_out; the other 13 schemas are unchanged (str
+# group_name), verified against the v2.10.0 / v2.11.0 Functional.cpp source.
+_C10D_FUNCTIONAL_SCHEMAS[(2, 10)] = _C10D_FUNCTIONAL_SCHEMAS[(2, 9)] + (
+ "reduce_scatter_tensor_out(Tensor input, str reduce_op, int group_size, str group_name, *, Tensor(a!) out) -> Tensor(a!)",
+)
+_C10D_FUNCTIONAL_SCHEMAS[(2, 11)] = _C10D_FUNCTIONAL_SCHEMAS[(2, 10)]
+
+# `_dtensor` is likewise C++-only, so absent on a distributed-less wheel. torchao's
+# `from torch.distributed._tensor import DTensor` loads tensor._collective_utils, whose
+# module-level `register_fake("_dtensor::shard_dim_alltoall")` raises unless the op is defined --
+# so the shim must define it in the same transaction or `import torchao` still rolls back. Schema
+# verified vs the live 2.9 dispatcher and v2.11.0 Functional.cpp (stable 2.9-2.11); else fail closed.
+_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`."""
+ return schema.split("(", 1)[0].strip()
+
+
+def _torchao_shim_torch_minor(torch):
+ base = torch.__version__.split("+", 1)[0].split(".")
+ 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::` 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(_TORCHAO_ROCM_NATIVE_PREFIXES) for n in get_ops())
+ except Exception:
+ return True
+
+
+_TORCHAO_ROCM_LEGACY_ERR = (
+ "Unsloth: torch.distributed is unavailable on this legacy Windows ROCm build; this is "
+ "an import-only compatibility shim for torchao and cannot perform real distributed / "
+ "collective work. Upgrade to a Windows ROCm PyTorch wheel built with GLOO "
+ "(ROCm/TheRock#5694)."
+)
+
+# Optional transport backends: must stay ABSENT so torch's guarded
+# `from torch._C._distributed_c10d import ProcessGroupX` probes conclude the backend is
+# unavailable (they sit behind try/except or capability checks).
+_TORCHAO_ROCM_OPTIONAL_BACKENDS = frozenset(
+ {
+ "ProcessGroupNCCL",
+ "ProcessGroupGloo",
+ "ProcessGroupMPI",
+ "ProcessGroupUCC",
+ "ProcessGroupXCCL",
+ "_ProcessGroupWrapper",
+ "_c10d_init",
+ }
+)
+
+
+class _TorchaoRocmSentinelMeta(type):
+ """Metaclass for import-only sentinel types: isinstance-safe (never matches a real
+ object), chainable via attribute access, and loud on construction (an accidental
+ runtime use raises rather than silently returning a wrong result)."""
+
+ def __instancecheck__(cls, instance):
+ return False
+
+ def __subclasscheck__(cls, subclass):
+ return False
+
+ def __getattr__(cls, name):
+ if name.startswith("__") and name.endswith("__"):
+ raise AttributeError(name)
+ child = _TorchaoRocmSentinelMeta(f"{cls.__name__}.{name}", (), {})
+ setattr(cls, name, child)
+ return child
+
+ def __call__(cls, *args, **kwargs):
+ raise RuntimeError(_TORCHAO_ROCM_LEGACY_ERR)
+
+
+def _torchao_rocm_sentinel(name):
+ return _TorchaoRocmSentinelMeta(str(name), (), {})
+
+
+def _make_torchao_rocm_fake_c10d():
+ """Build a fake `torch._C._distributed_c10d` so torchao's unconditional
+ `from torch._C._distributed_c10d import (...)` (via torch.distributed.distributed_c10d)
+ resolves. Explicit semantic fakes for the names read while torch's distributed Python
+ modules initialize; a bounded-dynamic `__getattr__` gives isinstance-safe loud sentinel
+ types for unknown import-only names but raises AttributeError for dunders and optional
+ backends. No `_c10d_init` and is_available() stays False: device_mesh.py self-stubs when
+ unavailable and the crash site imports the C-ext unconditionally, so the fake alone
+ unblocks the import."""
+ import types
+
+ mod = types.ModuleType(_C10D_EXT_MODULE)
+ mod.__package__ = "torch._C"
+ setattr(mod, _TORCHAO_ROCM_SHIM_SENTINEL, True)
+
+ # Option data-holders: constructible no-ops (torch builds them at import), kept distinct
+ # from the loud sentinels so import-time construction can't raise.
+ for name in (
+ "_DistributedBackendOptions",
+ "AllgatherOptions",
+ "AllreduceCoalescedOptions",
+ "AllreduceOptions",
+ "AllToAllOptions",
+ "BarrierOptions",
+ "BroadcastOptions",
+ "GatherOptions",
+ "ReduceOptions",
+ "ReduceScatterOptions",
+ "ScatterOptions",
+ ):
+ setattr(mod, name, type(name, (), {"__init__": lambda self, *a, **k: None}))
+ for name in ("PrefixStore", "Store", "HashStore", "Work"):
+ setattr(mod, name, _torchao_rocm_sentinel(name))
+ for name in (
+ "_register_process_group",
+ "_resolve_process_group",
+ "_unregister_all_process_groups",
+ "_unregister_process_group",
+ ):
+
+ def _loud(*a, **k):
+ raise RuntimeError(_TORCHAO_ROCM_LEGACY_ERR)
+
+ setattr(mod, name, _loud)
+
+ ReduceOp = _torchao_rocm_sentinel("ReduceOp")
+ _reduce_members = (
+ "SUM",
+ "AVG",
+ "PRODUCT",
+ "MIN",
+ "MAX",
+ "BAND",
+ "BOR",
+ "BXOR",
+ "PREMUL_SUM",
+ "UNUSED",
+ )
+ for m in _reduce_members:
+ setattr(ReduceOp, m, _torchao_rocm_sentinel(f"ReduceOp.{m}"))
+ RedOpType = _torchao_rocm_sentinel("ReduceOp.RedOpType")
+ RedOpType.__members__ = {m: getattr(ReduceOp, m) for m in _reduce_members}
+ ReduceOp.RedOpType = RedOpType
+ mod.ReduceOp = ReduceOp
+
+ ProcessGroup = _torchao_rocm_sentinel("ProcessGroup")
+ BackendType = _torchao_rocm_sentinel("ProcessGroup.BackendType")
+ for m in ("UNDEFINED", "GLOO", "NCCL", "UCC", "MPI", "XCCL", "CUSTOM"):
+ setattr(BackendType, m, _torchao_rocm_sentinel(f"ProcessGroup.BackendType.{m}"))
+ ProcessGroup.BackendType = BackendType
+ mod.ProcessGroup = ProcessGroup
+
+ DebugLevel = _torchao_rocm_sentinel("DebugLevel")
+ for m in ("OFF", "INFO", "DETAIL"):
+ setattr(DebugLevel, m, _torchao_rocm_sentinel(f"DebugLevel.{m}"))
+ mod.DebugLevel = DebugLevel
+ mod.get_debug_level = lambda *a, **k: DebugLevel.OFF
+
+ _dynamic_cache = {}
+
+ def _module_getattr(name):
+ if name.startswith("__") and name.endswith("__"):
+ raise AttributeError(name)
+ if name in _TORCHAO_ROCM_OPTIONAL_BACKENDS:
+ raise AttributeError(name)
+ if name in _dynamic_cache:
+ return _dynamic_cache[name]
+ child = _torchao_rocm_sentinel(name)
+ _dynamic_cache[name] = child
+ return child
+
+ mod.__getattr__ = _module_getattr
+ return mod
+
+
+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` 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
+ transactional: any failure rolls back and leaves torchao unimportable, so unsloth_zoo's
+ stub still catches it (no regression). Idempotent. Opt out:
+ UNSLOTH_DISABLE_TORCHAO_ROCM_SHIM=1."""
+ global _TORCHAO_ROCM_SHIM_STATE
+
+ if os.environ.get("UNSLOTH_DISABLE_TORCHAO_ROCM_SHIM", "0") == "1":
+ return
+ if sys.platform != "win32":
+ return
+
+ with _get_torchao_rocm_shim_lock():
+ if _TORCHAO_ROCM_SHIM_STATE is not None:
+ return # already installed this process
+ try:
+ import torch
+
+ # ROCm build: mirror the Studio is_win32_rocm() detector (HIP field OR a "rocm"
+ # __version__ tag -- AMD SDK wheels lack torch.version.hip) so the shim re-enables
+ # exactly the wheels the export gate disables. Capability guards below rule out false positives.
+ if not (
+ getattr(getattr(torch, "version", None), "hip", None)
+ or "rocm" in getattr(torch, "__version__", "").lower()
+ ):
+ return
+ 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
+ if any(n == "torchao" or n.startswith("torchao.") for n in list(sys.modules)):
+ return # real- or stub-imported already; cannot safely convert
+ dist = getattr(torch, "distributed", None)
+ if dist is None or dist.is_available() is not False:
+ return # real distributed present (fixed wheel) -> nothing to do
+ if hasattr(torch._C, "_c10d_init"):
+ return
+ if hasattr(torch._C, "_distributed_c10d") or _C10D_EXT_MODULE in sys.modules:
+ return
+ if _native_c10d_functional_present(torch):
+ return # dispatcher already has the ops -> never register over them
+ except Exception:
+ return
+
+ # ---- transaction: snapshot -> register -> acceptance import -> commit/rollback ---
+ modules_before = set(sys.modules)
+ had_c10d_attr = hasattr(torch._C, "_distributed_c10d")
+ fake = None
+ libs = []
+ try:
+ # Re-check the dispatcher immediately before touching it (TOCTOU guard).
+ if _native_c10d_functional_present(torch):
+ return
+ # 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)
+ # Acceptance test: the real torchao must import end to end.
+ importlib.import_module("torchao")
+ except BaseException:
+ # Atomic rollback: destroy schemas, drop the fake, purge only the torchao /
+ # distributed submodules this transaction newly created.
+ for lib in libs:
+ try:
+ lib._destroy()
+ except Exception:
+ pass
+ if fake is not None:
+ if not had_c10d_attr and getattr(torch._C, "_distributed_c10d", None) is fake:
+ try:
+ delattr(torch._C, "_distributed_c10d")
+ except Exception:
+ pass
+ if sys.modules.get(_C10D_EXT_MODULE) is fake:
+ del sys.modules[_C10D_EXT_MODULE]
+ for name in [
+ n
+ for n in set(sys.modules) - modules_before
+ if n == "torchao"
+ or n.startswith("torchao.")
+ or n == "torch.distributed"
+ or n.startswith("torch.distributed.")
+ ]:
+ sys.modules.pop(name, None)
+ return
+
+ # Commit: keep strong refs so the FRAGMENT schemas outlive GC.
+ _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/_dtensor schemas)."
+ )
+
+
+_torchao_rocm_shim_lock = None
+
+
+def _get_torchao_rocm_shim_lock():
+ global _torchao_rocm_shim_lock
+ if _torchao_rocm_shim_lock is None:
+ import threading
+ _torchao_rocm_shim_lock = threading.RLock()
+ return _torchao_rocm_shim_lock