Warn users when release_memory_occupation is called without memory saver enabled (#4566)

This commit is contained in:
fzyzcjy
2025-03-26 00:18:14 -07:00
committed by GitHub
parent 34e07a65f1
commit 26f07294f1
10 changed files with 50 additions and 12 deletions
@@ -1,3 +1,4 @@
import logging
from abc import ABC
from contextlib import contextmanager
@@ -8,6 +9,8 @@ try:
except ImportError:
pass
logger = logging.getLogger(__name__)
class TorchMemorySaverAdapter(ABC):
@staticmethod
@@ -16,6 +19,13 @@ class TorchMemorySaverAdapter(ABC):
_TorchMemorySaverAdapterReal() if enable else _TorchMemorySaverAdapterNoop()
)
def check_validity(self, caller_name):
if not self.enabled:
logger.warning(
f"`{caller_name}` will not save memory because torch_memory_saver is not enabled. "
f"Potential causes: `enable_memory_saver` is false, or torch_memory_saver has installation issues."
)
def configure_subprocess(self):
raise NotImplementedError
@@ -28,6 +38,10 @@ class TorchMemorySaverAdapter(ABC):
def resume(self):
raise NotImplementedError
@property
def enabled(self):
raise NotImplementedError
class _TorchMemorySaverAdapterReal(TorchMemorySaverAdapter):
def configure_subprocess(self):
@@ -42,6 +56,10 @@ class _TorchMemorySaverAdapterReal(TorchMemorySaverAdapter):
def resume(self):
return _primary_memory_saver.resume()
@property
def enabled(self):
return _primary_memory_saver.enabled
class _TorchMemorySaverAdapterNoop(TorchMemorySaverAdapter):
@contextmanager
@@ -57,3 +75,7 @@ class _TorchMemorySaverAdapterNoop(TorchMemorySaverAdapter):
def resume(self):
pass
@property
def enabled(self):
return False