diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index 75c1bd76e8..f586081582 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -1108,14 +1108,17 @@ def run_training_process( class _StubTypeMeta(type): def __instancecheck__(cls, instance): return False + def __subclasscheck__(cls, subclass): return False + def __getattr__(cls, attr): if attr.startswith("__"): raise AttributeError(attr) child = _StubTypeMeta(attr, (), {}) setattr(cls, attr, child) return child + def __call__(cls, *args, **kwargs): return None @@ -1128,8 +1131,9 @@ def run_training_process( m.__path__ = [] m.__package__ = mod_name m._unsloth_stub = _STUB_SENTINEL - m.__spec__ = _ilm.ModuleSpec(mod_name, loader=None, is_package=True) - def _ga(attr, _m=m, _n=mod_name): + m.__spec__ = _ilm.ModuleSpec(mod_name, loader = None, is_package = True) + + def _ga(attr, _m = m, _n = mod_name): if attr.startswith("__"): raise AttributeError(attr) # Return a stub CLASS (not a module) so that isinstance(x, attr) @@ -1137,19 +1141,22 @@ def run_training_process( child = _make_stub_type(f"{_n}.{attr}") setattr(_m, attr, child) return child + m.__getattr__ = _ga return m class _StubSubpackageLoader(_ilabc.Loader): def __init__(self, mod_name): self._mod_name = mod_name + def create_module(self, spec): return _make_mod_stub(self._mod_name) + def exec_module(self, module): pass class _StubSubpackageFinder(_ilabc.MetaPathFinder): - def find_spec(self, fullname, path, target=None): + def find_spec(self, fullname, path, target = None): if "." not in fullname: return None parent = sys.modules.get(fullname.rsplit(".", 1)[0]) @@ -1157,7 +1164,9 @@ def run_training_process( return None if getattr(parent, "_unsloth_stub", None) is not _STUB_SENTINEL: return None - return _ilm.ModuleSpec(fullname, _StubSubpackageLoader(fullname), is_package=True) + return _ilm.ModuleSpec( + fullname, _StubSubpackageLoader(fullname), is_package = True + ) sys.meta_path.append(_StubSubpackageFinder()) @@ -1198,6 +1207,7 @@ def run_training_process( sys.modules["torch.distributed"] = _td_mock try: import torch as _torch + _torch.distributed = _td_mock except Exception: pass @@ -1240,7 +1250,7 @@ def run_training_process( _gm_lib = _torch_for_rocm.library.Library("aten", "IMPL") def _grouped_mm_safe_impl( - self, mat2, offs=None, bias=None, out_dtype=None + self, mat2, offs = None, bias = None, out_dtype = None ): """Safe fallback for _grouped_mm on gfx1200 (null HIP kernel).""" _t = _torch_for_rocm @@ -1272,13 +1282,13 @@ def run_training_process( ) pieces.append(_t.mm(a_tail, b_tail)) result = ( - _t.cat(pieces, dim=0) + _t.cat(pieces, dim = 0) if pieces else _t.zeros( 0, mat2.shape[-1], - device=self.device, - dtype=self.dtype, + device = self.device, + dtype = self.dtype, ) ) if bias is not None: diff --git a/studio/backend/main.py b/studio/backend/main.py index 46f5b97a86..4cce7e552a 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -16,6 +16,7 @@ os.environ["PYTHONWARNINGS"] = "ignore" # Python 3.8+ ignores PATH for extension modules; register ROCm bin dirs with # os.add_dll_directory() so amdhip64.dll etc. are found before any torch import. if sys.platform == "win32": + def _add_rocm_dll_dirs() -> None: candidates = [] # 1. HIP_PATH / ROCM_PATH -- set by the AMD HIP SDK installer diff --git a/studio/backend/utils/hardware/amd.py b/studio/backend/utils/hardware/amd.py index d13cbdb545..bcacdb57ea 100644 --- a/studio/backend/utils/hardware/amd.py +++ b/studio/backend/utils/hardware/amd.py @@ -53,14 +53,18 @@ def _run_amd_smi(*args: str, timeout: int = _AMD_SMI_DEFAULT_TIMEOUT) -> Optiona logger.warning("amd-smi query failed: %s", e) _amd_smi_consecutive_failures += 1 if _amd_smi_consecutive_failures >= _AMD_SMI_FAILURE_LIMIT: - logger.warning("amd-smi unavailable -- disabling GPU polling to avoid repeated prompts") + logger.warning( + "amd-smi unavailable -- disabling GPU polling to avoid repeated prompts" + ) _amd_smi_disabled = True return None if result.returncode != 0 or not result.stdout.strip(): logger.warning("amd-smi returned code %d", result.returncode) _amd_smi_consecutive_failures += 1 if _amd_smi_consecutive_failures >= _AMD_SMI_FAILURE_LIMIT: - logger.warning("amd-smi unavailable -- disabling GPU polling to avoid repeated prompts") + logger.warning( + "amd-smi unavailable -- disabling GPU polling to avoid repeated prompts" + ) _amd_smi_disabled = True return None _amd_smi_consecutive_failures = 0 # reset on success diff --git a/studio/install_python_stack.py b/studio/install_python_stack.py index 147a4f8fb3..95dd7d5f2c 100644 --- a/studio/install_python_stack.py +++ b/studio/install_python_stack.py @@ -73,18 +73,22 @@ _PYTORCH_WHL_BASE = ( # Format: https://repo.amd.com/rocm/whl/{arch_family}/ # Override with UNSLOTH_ROCM_WINDOWS_MIRROR for air-gapped / mirror installs. _ROCM_WINDOWS_INDEX_BASE = ( - os.environ.get("UNSLOTH_ROCM_WINDOWS_MIRROR") - or "https://repo.amd.com/rocm/whl" + os.environ.get("UNSLOTH_ROCM_WINDOWS_MIRROR") or "https://repo.amd.com/rocm/whl" ).rstrip("/") # Maps gfx arch → AMD index arch-family suffix. # Each family is a separate pip index on repo.amd.com. _GFX_TO_AMD_INDEX_ARCH: dict[str, str] = { - "gfx1201": "gfx120X-all", "gfx1200": "gfx120X-all", # RDNA 4 - "gfx1151": "gfx1151", "gfx1150": "gfx1150", # RDNA 3.5 (Strix Halo/Point) - "gfx1103": "gfx110X-all", "gfx1102": "gfx110X-all", # RDNA 3 - "gfx1101": "gfx110X-all", "gfx1100": "gfx110X-all", - "gfx90a": "gfx90a", "gfx908": "gfx908", # MI200/MI100 + "gfx1201": "gfx120X-all", + "gfx1200": "gfx120X-all", # RDNA 4 + "gfx1151": "gfx1151", + "gfx1150": "gfx1150", # RDNA 3.5 (Strix Halo/Point) + "gfx1103": "gfx110X-all", + "gfx1102": "gfx110X-all", # RDNA 3 + "gfx1101": "gfx110X-all", + "gfx1100": "gfx110X-all", + "gfx90a": "gfx90a", + "gfx908": "gfx908", # MI200/MI100 } # bitsandbytes continuous-release_main wheels with the ROCm 4-bit GEMV fix @@ -421,14 +425,19 @@ def _ensure_rocm_torch() -> None: if not _torch_already_rocm: index_url = _windows_rocm_index_url(gfx_arch) if index_url is None: - print(f" No AMD Windows torch index for GPU arch {gfx_arch} -- skipping") + print( + f" No AMD Windows torch index for GPU arch {gfx_arch} -- skipping" + ) return print(f" {gfx_arch} (Windows) -- installing torch from {index_url}") pip_install( f"ROCm torch (Windows, {gfx_arch})", "--force-reinstall", - "--index-url", index_url, - "torch", "torchvision", "torchaudio", + "--index-url", + index_url, + "torch", + "torchvision", + "torchaudio", constrain = False, ) # Always install AMD Windows bitsandbytes — the PyPI wheel ships only