fix(device): harden dynamic runtime GPU failure recovery with irreversible fallback to CPU
This commit is contained in:
parent
45870eb0ec
commit
146b9c00db
6 changed files with 100 additions and 33 deletions
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue