Warn users when release_memory_occupation is called without memory saver enabled (#4566)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user