diff --git a/studio/backend/utils/hardware/hardware.py b/studio/backend/utils/hardware/hardware.py index df105f7a59..9bce231f45 100644 --- a/studio/backend/utils/hardware/hardware.py +++ b/studio/backend/utils/hardware/hardware.py @@ -648,7 +648,10 @@ def _get_parent_visible_gpu_spec() -> Dict[str, Any]: # Use explicit None checks (not `or`) so empty string "" is honoured # as "no visible GPUs" rather than falling through to CUDA_VISIBLE_DEVICES. cuda_visible = None - if IS_ROCM: + _is_rocm_spec = IS_ROCM or ( + "HIP_VISIBLE_DEVICES" in os.environ or "ROCR_VISIBLE_DEVICES" in os.environ + ) + if _is_rocm_spec: hip_vis = os.environ.get("HIP_VISIBLE_DEVICES") rocr_vis = os.environ.get("ROCR_VISIBLE_DEVICES") if hip_vis is not None: @@ -1511,14 +1514,38 @@ def apply_gpu_ids(gpu_ids) -> None: # parent process already set a ROCm visibility variable -- that # way a downstream ROCm process inherits the narrowed mask even # before Studio's hardware detection has classified the host. + # As a final fallback, probe torch.version.hip directly so spawned + # training workers on AMD hosts where the user never set HIP_VISIBLE_DEVICES + # still get the correct ROCm visibility mask (mirrors the llama_cpp.py + # approach for llama-server subprocess GPU pinning). _inherits_rocm_visibility = ( "HIP_VISIBLE_DEVICES" in os.environ or "ROCR_VISIBLE_DEVICES" in os.environ ) - if IS_ROCM or _inherits_rocm_visibility: + _is_rocm = IS_ROCM or _inherits_rocm_visibility + if not _is_rocm: + # Use ``is not None`` here to match the detect_hardware() check at + # module top -- torch ships HIP version as a non-empty string on + # ROCm builds and None on CUDA builds, so the two forms agree on + # every shipping torch wheel; the ``is not None`` form is the one + # the rest of the codebase reads for "this torch was built with + # HIP". Keep the broad ``except`` as a safety net (we never want + # apply_gpu_ids to crash a worker over a probe failure) but log at + # debug level so the skip is observable when needed. + try: + import torch as _torch + + _is_rocm = getattr(_torch.version, "hip", None) is not None + except Exception as e: + logger.debug( + "apply_gpu_ids: torch.version.hip probe skipped (%s: %s)", + type(e).__name__, + e, + ) + if _is_rocm: os.environ["HIP_VISIBLE_DEVICES"] = value os.environ["ROCR_VISIBLE_DEVICES"] = value _visible_gpu_count = None - if IS_ROCM or _inherits_rocm_visibility: + if _is_rocm: logger.info("Applied gpu_ids: CUDA_VISIBLE_DEVICES='%s' (rocm)", value) else: logger.info("Applied gpu_ids: CUDA_VISIBLE_DEVICES='%s'", value) diff --git a/tests/studio/install/test_rocm_support.py b/tests/studio/install/test_rocm_support.py index 48831fd57b..99bc9c11bc 100644 --- a/tests/studio/install/test_rocm_support.py +++ b/tests/studio/install/test_rocm_support.py @@ -1250,6 +1250,48 @@ class TestHardwareAmdBranching: assert "amd.get_physical_gpu_count" in func_body +# ============================================================================= +# TEST: hardware.py -- apply_gpu_ids ROCm fallback (issue #5180) +# ============================================================================= + + +class TestApplyGpuIdsRocmFallback: + """Verify apply_gpu_ids sets HIP_VISIBLE_DEVICES on ROCm hosts even when + IS_ROCM is still False (worker subprocess before detect_hardware runs).""" + + def test_apply_gpu_ids_falls_back_to_torch_version_hip(self): + """apply_gpu_ids should probe torch.version.hip when IS_ROCM is False and no ROCm env vars are set.""" + hw_path = ( + PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py" + ) + source = hw_path.read_text() + func_start = source.find("def apply_gpu_ids") + func_body = source[func_start : source.find("\ndef ", func_start + 1)] + assert 'getattr(_torch.version, "hip", None)' in func_body + + def test_apply_gpu_ids_sets_hip_and_rocr_visible_devices(self): + """apply_gpu_ids should set both HIP_VISIBLE_DEVICES and ROCR_VISIBLE_DEVICES on ROCm.""" + hw_path = ( + PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py" + ) + source = hw_path.read_text() + func_start = source.find("def apply_gpu_ids") + func_body = source[func_start : source.find("\ndef ", func_start + 1)] + assert 'os.environ["HIP_VISIBLE_DEVICES"] = value' in func_body + assert 'os.environ["ROCR_VISIBLE_DEVICES"] = value' in func_body + + def test_apply_gpu_ids_rocm_fallback_is_guarded_by_try_except(self): + """torch import in apply_gpu_ids must be wrapped in try/except so a missing torch never crashes.""" + hw_path = ( + PACKAGE_ROOT / "studio" / "backend" / "utils" / "hardware" / "hardware.py" + ) + source = hw_path.read_text() + func_start = source.find("def apply_gpu_ids") + func_body = source[func_start : source.find("\ndef ", func_start + 1)] + assert "import torch as _torch" in func_body + assert "except Exception" in func_body + + # ============================================================================= # TEST: install_python_stack.py -- Windows AMD warning # =============================================================================