From 146b9c00db6672257f6d5cfa9480eb723a5fad0a Mon Sep 17 00:00:00 2001 From: Zhirik1337 Date: Tue, 22 Sep 2026 22:21:51 +0300 Subject: [PATCH] fix(device): harden dynamic runtime GPU failure recovery with irreversible fallback to CPU --- flyguard/device.py | 26 ++++++++++++++++++++++++-- flyguard/lamina.py | 5 +++-- flyguard/mbon_readout.py | 37 ++++++++++++++++++++++--------------- flyguard/mushroom_body.py | 32 +++++++++++++++++++------------- flyguard/pipeline.py | 7 ++++++- tests/test_pipeline.py | 26 ++++++++++++++++++++++++++ 6 files changed, 100 insertions(+), 33 deletions(-) diff --git a/flyguard/device.py b/flyguard/device.py index 15b7740..b99dea3 100644 --- a/flyguard/device.py +++ b/flyguard/device.py @@ -37,11 +37,33 @@ def is_cuda_available() -> bool: if not is_torch_available(): _CUDA_AVAILABLE = False else: - import torch - _CUDA_AVAILABLE = bool(torch.cuda.is_available()) + try: + import torch + _CUDA_AVAILABLE = bool(torch.cuda.is_available() and torch.cuda.device_count() > 0) + except Exception as e: + logger.warning("Проверка CUDA завершилась ошибкой: %s. Fallback на CPU.", e) + _CUDA_AVAILABLE = False return _CUDA_AVAILABLE +def notify_cuda_error(exc: Exception | None = None) -> None: + """Зафиксировать сбой CUDA в рантайме и принудительно перевести систему в режим CPU fallback. + + Вызывается, если во время работы на GPU произошёл OOM, таймаут или сбой драйвера. + Последующие вызовы конвейера будут прозрачно исполняться на CPU. + """ + global _CUDA_AVAILABLE + _CUDA_AVAILABLE = False + logger.warning("Зафиксирован сбой GPU в рантайме (%s). Выполнен необратимый Fallback на CPU.", exc) + + +def reset_device_cache() -> None: + """Сбросить кэш состояния устройств (для юнит-тестов).""" + global _TORCH_AVAILABLE, _CUDA_AVAILABLE + _TORCH_AVAILABLE = None + _CUDA_AVAILABLE = None + + def get_device(preferred: str = "auto") -> str: """Выбрать вычислительное устройство с автоматическим fallback на CPU. diff --git a/flyguard/lamina.py b/flyguard/lamina.py index fe460d5..2250715 100644 --- a/flyguard/lamina.py +++ b/flyguard/lamina.py @@ -196,8 +196,9 @@ def process(r: np.ndarray, valid: np.ndarray, *, r_max: float = 300.0, if target_dev.startswith("cuda"): try: return _process_gpu(r, valid, r_max=r_max, device=target_dev) - except Exception: - # При любых непредвиденных сбоях GPU — прозрачный откат на CPU + except Exception as e: + from .device import notify_cuda_error + notify_cuda_error(e) return _process_cpu(r, valid, r_max=r_max) return _process_cpu(r, valid, r_max=r_max) diff --git a/flyguard/mbon_readout.py b/flyguard/mbon_readout.py index 4ccf2fc..89746b0 100644 --- a/flyguard/mbon_readout.py +++ b/flyguard/mbon_readout.py @@ -126,21 +126,26 @@ class MbonReadout: chunk = max(1, int(2 ** 26 // max(self.cfg.n_kc, 1))) if device and device != "cpu": - import torch - out_i = np.empty((X.shape[0], k), np.int64) - out_v = np.empty((X.shape[0], k), np.float32) - with torch.no_grad(): - m = torch.as_tensor(self.mean, device=device) - s = torch.as_tensor(self.scale, device=device) - w = torch.as_tensor(self.W, device=device).T.contiguous() - for i in range(0, X.shape[0], chunk): - t = torch.as_tensor(X[i:i + chunk], device=device) - y = torch.relu(((t - m) / s) @ w) - v, a = torch.topk(y, k, dim=1) - v = v * (k / v.sum(1, keepdim=True).clamp_min(1e-6)) - out_i[i:i + chunk] = a.cpu().numpy() - out_v[i:i + chunk] = v.cpu().numpy() - return out_i, out_v + try: + import torch + out_i = np.empty((X.shape[0], k), np.int64) + out_v = np.empty((X.shape[0], k), np.float32) + with torch.no_grad(): + m = torch.as_tensor(self.mean, device=device) + s = torch.as_tensor(self.scale, device=device) + w = torch.as_tensor(self.W, device=device).T.contiguous() + for i in range(0, X.shape[0], chunk): + t = torch.as_tensor(X[i:i + chunk], device=device) + y = torch.relu(((t - m) / s) @ w) + v, a = torch.topk(y, k, dim=1) + v = v * (k / v.sum(1, keepdim=True).clamp_min(1e-6)) + out_i[i:i + chunk] = a.cpu().numpy() + out_v[i:i + chunk] = v.cpu().numpy() + return out_i, out_v + except Exception as e: + from .device import notify_cuda_error + notify_cuda_error(e) + # Переход к расчету на CPU out_i = np.empty((X.shape[0], k), np.int64) out_v = np.empty((X.shape[0], k), np.float32) @@ -333,6 +338,8 @@ class MbonReadout: device=target_dev, verbose=verbose) return except Exception as e: + from .device import notify_cuda_error + notify_cuda_error(e) import logging logging.getLogger("flyguard.mbon").warning("GPU learning failed (%s), fallback to CPU", e) diff --git a/flyguard/mushroom_body.py b/flyguard/mushroom_body.py index 2443bc0..6624cb7 100644 --- a/flyguard/mushroom_body.py +++ b/flyguard/mushroom_body.py @@ -158,17 +158,22 @@ class MushroomBody: target_dev = get_device("auto") if target_dev and target_dev != "cpu": - import torch - out = np.empty((X.shape[0], k), np.int64) - with torch.no_grad(): - m = torch.as_tensor(self.mean, device=target_dev) - s = torch.as_tensor(self.scale, device=target_dev) - w = torch.as_tensor(self.W, device=target_dev).T.contiguous() - for i in range(0, X.shape[0], chunk): - t = torch.as_tensor(X[i:i + chunk], device=target_dev) - y = ((t - m) / s) @ w - out[i:i + chunk] = torch.topk(y, k, dim=1).indices.cpu().numpy() - return out + try: + import torch + out = np.empty((X.shape[0], k), np.int64) + with torch.no_grad(): + m = torch.as_tensor(self.mean, device=target_dev) + s = torch.as_tensor(self.scale, device=target_dev) + w = torch.as_tensor(self.W, device=target_dev).T.contiguous() + for i in range(0, X.shape[0], chunk): + t = torch.as_tensor(X[i:i + chunk], device=target_dev) + y = ((t - m) / s) @ w + out[i:i + chunk] = torch.topk(y, k, dim=1).indices.cpu().numpy() + return out + except Exception as e: + from .device import notify_cuda_error + notify_cuda_error(e) + # Fallback на CPU ниже if X.shape[0] <= chunk: z = (X - self.mean) / self.scale @@ -236,8 +241,9 @@ class MushroomBody: self.w_mbon = w_mbon_t.cpu().numpy() self.n_seen += X.shape[0] return - except Exception: - pass + except Exception as e: + from .device import notify_cuda_error + notify_cuda_error(e) act = self.encode(X, device="cpu") cnt = np.bincount(act.ravel(), minlength=self.cfg.n_kc) diff --git a/flyguard/pipeline.py b/flyguard/pipeline.py index c794649..4fb6c35 100644 --- a/flyguard/pipeline.py +++ b/flyguard/pipeline.py @@ -319,7 +319,12 @@ class FlyGuard: else STRAIGHT) with t("lamina"): - lam = lamina.process(tf.r, tf.valid, device=self.device) + dev = self.device + if dev.startswith("cuda"): + from .device import is_cuda_available + if not is_cuda_available(): + self.device = dev = "cpu" + lam = lamina.process(tf.r, tf.valid, device=dev) with t("ego"): ego = (self.ego_est.update(tf, pc.stamp) if self.p.enable_motion diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index 2c3b389..5edbd73 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -740,3 +740,29 @@ def test_lamina_device_routing(): assert out_auto.on.shape == (16, 32) assert out_cpu.on[8, 16] > 0.0 assert np.allclose(out_cpu.on, out_auto.on, atol=1e-5) + + +def test_device_runtime_failure_and_fallback(): + """Проверка динамического перехода на CPU при сбое/отвале GPU в рантайме.""" + from flyguard.device import is_cuda_available, get_device, notify_cuda_error, reset_device_cache + from flyguard import lamina + + # Симуляция критического сбоя GPU + notify_cuda_error(RuntimeError("Simulated CUDA device disconnect / OOM")) + + assert not is_cuda_available() + assert get_device("cuda") == "cpu" + assert get_device("auto") == "cpu" + + r = np.full((16, 32), 20.0, dtype=np.float32) + r[8, 16] = 4.0 + valid = np.ones((16, 32), dtype=bool) + + # Даже при явном указании device="cuda", Lamina должна успешно отработать на CPU + out = lamina.process(r, valid, device="cuda") + assert out.on.shape == (16, 32) + assert out.on[8, 16] > 0.0 + + # Восстановление кэша + reset_device_cache() +