diff --git a/studio/backend/main.py b/studio/backend/main.py index 488ce936ff..5d1670c182 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -16,7 +16,6 @@ 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": - # Retained at module scope -- os.add_dll_directory returns a handle that # removes the search-path entry when garbage collected. _ROCM_DLL_HANDLES: list = [] @@ -66,6 +65,7 @@ if sys.platform == "win32": if os.environ.get("HIP_PATH") or os.environ.get("ROCM_PATH"): try: import torch as _torch_probe + _is_rocm_host = bool(getattr(_torch_probe.version, "hip", None)) del _torch_probe except Exception: diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py index 035a9ea3c6..bf6ab2a9bc 100644 --- a/tests/studio/install/test_rocm_support.py +++ b/tests/studio/install/test_rocm_support.py @@ -1799,8 +1799,10 @@ class TestRocmTorchInstalledEnvVar: @patch.object(stack_mod, "pip_install") def test_env_var_skips_main_pip_install(self, mock_pip, mock_bnb): """UNSLOTH_ROCM_TORCH_INSTALLED=1 should not trigger torch pip_install.""" - with patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), \ - patch.object(stack_mod.subprocess, "run", side_effect=self._ok_torch_probe): + with ( + patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), + patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe), + ): stack_mod._ensure_rocm_torch() mock_pip.assert_not_called() @@ -1808,8 +1810,10 @@ class TestRocmTorchInstalledEnvVar: @patch.object(stack_mod, "pip_install") def test_env_var_calls_bnb_install(self, mock_pip, mock_bnb): """UNSLOTH_ROCM_TORCH_INSTALLED=1 should still call _install_bnb_windows_rocm.""" - with patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), \ - patch.object(stack_mod.subprocess, "run", side_effect=self._ok_torch_probe): + with ( + patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), + patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe), + ): stack_mod._ensure_rocm_torch() mock_bnb.assert_called_once() @@ -1818,8 +1822,10 @@ class TestRocmTorchInstalledEnvVar: def test_env_var_sets_rocm_windows_flag(self, mock_pip, mock_bnb): """UNSLOTH_ROCM_TORCH_INSTALLED=1 should set _rocm_windows_torch_installed.""" stack_mod._rocm_windows_torch_installed = False - with patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), \ - patch.object(stack_mod.subprocess, "run", side_effect=self._ok_torch_probe): + with ( + patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), + patch.object(stack_mod.subprocess, "run", side_effect = self._ok_torch_probe), + ): stack_mod._ensure_rocm_torch() assert stack_mod._rocm_windows_torch_installed is True @@ -1828,14 +1834,18 @@ class TestRocmTorchInstalledEnvVar: def test_env_var_falls_through_when_torch_missing(self, mock_pip, mock_bnb): """If the venv was wiped between runs, the stale env-var must not suppress reinstall.""" stack_mod._rocm_windows_torch_installed = False + def _bad_probe(*a, **kw): rv = MagicMock() rv.returncode = 1 return rv - with patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), \ - patch.object(stack_mod.subprocess, "run", side_effect=_bad_probe), \ - patch.object(stack_mod, "IS_WINDOWS", False), \ - patch.object(stack_mod, "IS_MACOS", True): + + with ( + patch.dict(os.environ, {"UNSLOTH_ROCM_TORCH_INSTALLED": "1"}), + patch.object(stack_mod.subprocess, "run", side_effect = _bad_probe), + patch.object(stack_mod, "IS_WINDOWS", False), + patch.object(stack_mod, "IS_MACOS", True), + ): stack_mod._ensure_rocm_torch() # macOS branch is the next exit -- but the point is the early-return did NOT fire. mock_bnb.assert_not_called()