reset DEVICE type on fastapi lifespan exit
This commit is contained in:
parent
e7c3e7b48d
commit
7db31723b9
2 changed files with 22 additions and 4 deletions
|
|
@ -16,6 +16,7 @@ from datetime import datetime
|
|||
from routes import training_router, models_router, inference_router, datasets_router, auth_router
|
||||
from auth import storage
|
||||
from utils.hardware import detect_hardware
|
||||
import utils.hardware.hardware as _hw_module
|
||||
|
||||
UNSLOTH_CACHE_DIR = Path(__file__).parent / "unsloth_compiled_cache"
|
||||
|
||||
|
|
@ -36,7 +37,8 @@ async def lifespan(app: FastAPI):
|
|||
print("This token can only be used once.")
|
||||
print("=" * 60 + "\n")
|
||||
yield
|
||||
# Cleanup: remove Unsloth compiled cache on shutdown
|
||||
# Cleanup
|
||||
_hw_module.DEVICE = None
|
||||
shutil.rmtree(UNSLOTH_CACHE_DIR, ignore_errors=True)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -29,12 +29,14 @@ needs_torch = pytest.mark.skipif(not HAS_TORCH, reason="PyTorch not installed")
|
|||
|
||||
from utils.hardware import (
|
||||
get_device,
|
||||
detect_hardware,
|
||||
is_apple_silicon,
|
||||
clear_gpu_cache,
|
||||
get_gpu_memory_info,
|
||||
log_gpu_memory,
|
||||
DeviceType,
|
||||
)
|
||||
import utils.hardware.hardware as _hw_module
|
||||
from utils.utils import format_error_message
|
||||
|
||||
|
||||
|
|
@ -50,11 +52,25 @@ def _actual_device() -> str:
|
|||
return "cpu"
|
||||
|
||||
|
||||
def _reset_and_detect():
|
||||
"""Reset the cached DEVICE global and re-run detection."""
|
||||
_hw_module.DEVICE = None
|
||||
return detect_hardware()
|
||||
|
||||
|
||||
# ========== get_device() ==========
|
||||
|
||||
class TestGetDevice:
|
||||
"""Tests for get_device() — should agree with the real hardware."""
|
||||
|
||||
def setup_method(self):
|
||||
"""Save DEVICE before each test."""
|
||||
self._saved_device = _hw_module.DEVICE
|
||||
|
||||
def teardown_method(self):
|
||||
"""Restore DEVICE after each test so mocked tests don't poison later ones."""
|
||||
_hw_module.DEVICE = self._saved_device
|
||||
|
||||
def test_returns_valid_device_type(self):
|
||||
result = get_device()
|
||||
assert result in (DeviceType.CUDA, DeviceType.MPS, DeviceType.CPU)
|
||||
|
|
@ -67,7 +83,7 @@ class TestGetDevice:
|
|||
@needs_torch
|
||||
def test_returns_cuda_when_cuda_available(self):
|
||||
with patch("torch.cuda.is_available", return_value=True):
|
||||
assert get_device() == DeviceType.CUDA
|
||||
assert _reset_and_detect() == DeviceType.CUDA
|
||||
|
||||
@needs_torch
|
||||
def test_returns_mps_when_only_mps_available(self):
|
||||
|
|
@ -75,13 +91,13 @@ class TestGetDevice:
|
|||
mock_mps.is_available.return_value = True
|
||||
with patch("torch.cuda.is_available", return_value=False), \
|
||||
patch.object(torch.backends, "mps", mock_mps, create=True):
|
||||
assert get_device() == DeviceType.MPS
|
||||
assert _reset_and_detect() == DeviceType.MPS
|
||||
|
||||
@needs_torch
|
||||
def test_returns_cpu_when_nothing_available(self):
|
||||
with patch("torch.cuda.is_available", return_value=False), \
|
||||
patch("builtins.hasattr", side_effect=lambda obj, name: False if name == "mps" else hasattr(obj, name)):
|
||||
assert get_device() == DeviceType.CPU
|
||||
assert _reset_and_detect() == DeviceType.CPU
|
||||
|
||||
|
||||
# ========== is_apple_silicon() ==========
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue