diff --git a/install.sh b/install.sh index b8f49abf51..85268b14a4 100755 --- a/install.sh +++ b/install.sh @@ -1000,10 +1000,10 @@ get_torch_index_url() { ver="$(rpm -q --qf '%{VERSION}\n' rocm-core 2>/dev/null)" && \ [ -n "$ver" ] && \ printf '%s\n' "$ver" | awk -F'[.-]' '{print "rocm"$1"."$2; exit}'; }) 2>/dev/null - # Validate _rocm_tag: must match "rocmX.Y" with leading digits + # Validate _rocm_tag: must match "rocmX.Y" with major >= 1 case "$_rocm_tag" in - rocm[0-9]*.[0-9]*) : ;; # valid - *) _rocm_tag="" ;; # reject malformed (empty version, garbled output) + rocm[1-9]*.[0-9]*) : ;; # valid (major >= 1) + *) _rocm_tag="" ;; # reject malformed (empty, garbled, or major=0) esac if [ -n "$_rocm_tag" ]; then # ROCm 7.2 only has torch 2.11.0 which exceeds current bounds (<2.11.0). @@ -1011,10 +1011,11 @@ get_torch_index_url() { # TODO: uncomment the next line when torch upper bound is bumped to >=2.11.0 # echo "$_base/$_rocm_tag"; return case "$_rocm_tag" in - rocm7.2*|rocm7.3*|rocm7.4*|rocm7.5*|rocm8*|rocm9*) - echo "$_base/rocm7.1" ;; - *) + rocm6.*|rocm7.0*|rocm7.1*) echo "$_base/$_rocm_tag" ;; + *) + # ROCm 7.2+ (including future 10.x+): cap to rocm7.1 + echo "$_base/rocm7.1" ;; esac return fi diff --git a/studio/install_python_stack.py b/studio/install_python_stack.py index 1f17a2e79c..16b9877da9 100644 --- a/studio/install_python_stack.py +++ b/studio/install_python_stack.py @@ -47,13 +47,14 @@ _PYTORCH_WHL_BASE = "https://download.pytorch.org/whl" def _detect_rocm_version() -> tuple[int, int] | None: """Return (major, minor) of the installed ROCm stack, or None.""" # Check /opt/rocm/.info/version or ROCM_PATH equivalent - rocm_root = os.environ.get("ROCM_PATH", "/opt/rocm") + rocm_root = os.environ.get("ROCM_PATH") or "/opt/rocm" for path in ( os.path.join(rocm_root, ".info", "version"), os.path.join(rocm_root, "lib", "rocm_version"), ): try: - parts = open(path).read().strip().split("-")[0].split(".") + with open(path) as fh: + parts = fh.read().strip().split("-")[0].split(".") return int(parts[0]), int(parts[1]) except Exception: pass @@ -86,7 +87,7 @@ def _ensure_rocm_torch() -> None: links against HIP (ROCm) or CUDA (NVIDIA). Skips on Windows/macOS. Uses pip_install() to respect uv, constraints, and --python targeting. """ - rocm_root = os.environ.get("ROCM_PATH", "/opt/rocm") + rocm_root = os.environ.get("ROCM_PATH") or "/opt/rocm" if not os.path.isdir(rocm_root) and not shutil.which("hipcc"): return # no ROCm toolchain @@ -96,22 +97,29 @@ def _ensure_rocm_torch() -> None: return # Skip if torch is already GPU-enabled (HIP or CUDA) - probe = subprocess.run( - [ - sys.executable, - "-c", - "import torch; print(torch.version.hip or torch.version.cuda or '')", - ], - stdout = subprocess.PIPE, - stderr = subprocess.DEVNULL, - timeout = 30, - ) - if probe.returncode == 0 and probe.stdout.decode().strip(): + try: + probe = subprocess.run( + [ + sys.executable, + "-c", + "import torch; print(torch.version.hip or torch.version.cuda or '')", + ], + stdout = subprocess.PIPE, + stderr = subprocess.DEVNULL, + timeout = 30, + ) + except (OSError, subprocess.TimeoutExpired): + probe = None + if probe is not None and probe.returncode == 0 and probe.stdout.decode().strip(): return # torch already GPU-enabled # Select best matching wheel tag (newest ROCm version <= installed) tag = next( - (t for (maj, mn), t in _ROCM_TORCH_INDEX.items() if ver >= (maj, mn)), + ( + t + for (maj, mn), t in sorted(_ROCM_TORCH_INDEX.items(), reverse = True) + if ver >= (maj, mn) + ), None, ) if tag is None: @@ -124,7 +132,7 @@ def _ensure_rocm_torch() -> None: f"ROCm torch ({tag})", "--force-reinstall", "--no-cache-dir", - "torch", + "torch>=2.4,<2.11.0", "torchvision", "torchaudio", "--index-url", diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py index 62ceece457..5c45fb2ca9 100644 --- a/tests/studio/install/test_rocm_support.py +++ b/tests/studio/install/test_rocm_support.py @@ -609,19 +609,16 @@ class TestEnsureRocmTorch: @patch.object(stack_mod, "pip_install") @patch.object(stack_mod, "_detect_rocm_version", return_value = (7, 1)) - def test_probe_timeout_handled(self, mock_ver, mock_pip): - """Probe subprocess timeout should be handled gracefully.""" + def test_probe_timeout_triggers_reinstall(self, mock_ver, mock_pip): + """Probe subprocess timeout should not crash; should proceed to reinstall.""" with patch("os.path.isdir", return_value = True): with patch( "subprocess.run", side_effect = subprocess.TimeoutExpired("python", 30) ): - # Should not crash -- timeout on probe means torch not importable - # The function will get an exception from subprocess.run and - # proceed to reinstall - try: - _ensure_rocm_torch() - except subprocess.TimeoutExpired: - pass # Acceptable -- the fix is about the timeout being set + _ensure_rocm_torch() + # If probe times out, the function should treat torch as unusable and reinstall + assert mock_pip.call_count == 2 + assert "rocm7.1" in str(mock_pip.call_args_list[0]) # ============================================================================= @@ -831,14 +828,15 @@ class TestInstallShStructure: """ROCm 7.2+ should fall back to rocm7.1 index.""" sh_path = PACKAGE_ROOT / "install.sh" source = sh_path.read_text() - assert "rocm7.2" in source # case pattern - assert 'echo "$_base/rocm7.1"' in source # fallback + assert 'echo "$_base/rocm7.1"' in source # fallback for unknown versions + # Allowlisted versions should pass through directly + assert "rocm6.*|rocm7.0*|rocm7.1*)" in source def test_rocm_tag_validation_guard_exists(self): """install.sh should validate _rocm_tag with a case guard.""" sh_path = PACKAGE_ROOT / "install.sh" source = sh_path.read_text() - assert "rocm[0-9]*.[0-9]*)" in source + assert "rocm[1-9]*.[0-9]*)" in source assert '_rocm_tag=""' in source # rejection path def test_dpkg_epoch_handling(self):