diff --git a/studio/backend/core/export/export.py b/studio/backend/core/export/export.py index a3d0f12856..4b0ecb0840 100644 --- a/studio/backend/core/export/export.py +++ b/studio/backend/core/export/export.py @@ -51,6 +51,7 @@ def _multi_gpu_device_map_kwargs() -> dict: return {} try: from utils.hardware import get_device_map, get_parent_visible_gpu_ids + visible = get_parent_visible_gpu_ids() if len(visible) > 1: device_map = get_device_map(visible) diff --git a/tests/test_compressed_export_gpu_release.py b/tests/test_compressed_export_gpu_release.py index d6fb6ee0bf..d41f7e4f3e 100644 --- a/tests/test_compressed_export_gpu_release.py +++ b/tests/test_compressed_export_gpu_release.py @@ -38,9 +38,7 @@ class _FakeLogger: def _load_helpers(fake_torch, fake_logger): tree = ast.parse(_SAVE_PY.read_text(encoding = "utf-8")) keep = [ - node - for node in tree.body - if isinstance(node, ast.FunctionDef) and node.name in _WANTED + node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name in _WANTED ] assert len(keep) == len(_WANTED), "release helpers missing from save.py" namespace = {"torch": fake_torch, "logger": fake_logger} @@ -58,7 +56,12 @@ def _fake_torch(cuda_available = True): class _FakeModel: - def __init__(self, device_map = None, devices = ("cuda:0",), quantized = False): + def __init__( + self, + device_map = None, + devices = ("cuda:0",), + quantized = False, + ): if device_map is not None: self.hf_device_map = device_map self._devices = [types.SimpleNamespace(device = d) for d in devices]