From 71ac9ffe54b76ee17469266ac997f036643b3e7e Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 3 Dec 2025 11:51:30 -0500 Subject: [PATCH] fix: use max_memory_reserved for accurate VRAM peak tracking --- src/optimization/memory_manager.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/optimization/memory_manager.py b/src/optimization/memory_manager.py index 25c7383..592e37e 100644 --- a/src/optimization/memory_manager.py +++ b/src/optimization/memory_manager.py @@ -138,7 +138,7 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug debug: Optional debug instance for logging Returns: - tuple: (allocated_gb, reserved_gb, max_allocated_gb) + tuple: (allocated_gb, reserved_gb, max_reserved_gb) Returns (0, 0, 0) if no GPU available """ try: @@ -149,8 +149,8 @@ def get_vram_usage(device: Optional[torch.device] = None, debug: Optional['Debug device = torch.device(device) allocated = torch.cuda.memory_allocated(device) / (1024**3) reserved = torch.cuda.memory_reserved(device) / (1024**3) - max_allocated = torch.cuda.max_memory_allocated(device) / (1024**3) - return allocated, reserved, max_allocated + max_reserved = torch.cuda.max_memory_reserved(device) / (1024**3) + return allocated, reserved, max_reserved elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(): # MPS doesn't support per-device queries - uses global memory tracking allocated = torch.mps.current_allocated_memory() / (1024**3)