main #2

Merged
Dan4ick merged 5 commits from Zhirik1337/Brainrot_Muxa:main into main 2026-09-22 19:31:28 +00:00
6 changed files with 100 additions and 33 deletions
Showing only changes of commit 146b9c00db - Show all commits

View file

@ -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.

View file

@ -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)

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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()