From 17fbadb0e811c7ab946ae87cf8ceadde41ca617e Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 1 Apr 2025 14:00:39 -0400 Subject: [PATCH] reset memory stats between runs Signed-off-by: Vladimir Mandic --- modules/memstats.py | 7 +++++++ modules/processing.py | 1 + modules/sd_models.py | 3 ++- 3 files changed, 10 insertions(+), 1 deletion(-) diff --git a/modules/memstats.py b/modules/memstats.py index 4fd6206f1..5880a6d97 100644 --- a/modules/memstats.py +++ b/modules/memstats.py @@ -74,6 +74,13 @@ def memory_stats(): return mem +def reset_stats(): + try: + torch.cuda.reset_memory_stats() + except Exception: + pass + + def memory_cache(): return mem diff --git a/modules/processing.py b/modules/processing.py index 67d86021b..b1a8bd548 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -81,6 +81,7 @@ class Processed: self.all_seeds = all_seeds or p.all_seeds or [self.seed] self.all_subseeds = all_subseeds or p.all_subseeds or [self.subseed] self.infotexts = infotexts or [self.info] + memstats.reset_stats() def js(self): obj = { diff --git a/modules/sd_models.py b/modules/sd_models.py index 1cdec4afb..4e85ec74c 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1056,8 +1056,9 @@ def clear_caches(): lora_common.loaded_networks.clear() lora_common.previously_loaded_networks.clear() lora_load.lora_cache.clear() - from modules import prompt_parser_diffusers + from modules import prompt_parser_diffusers, memstats prompt_parser_diffusers.cache.clear() + memstats.reset_stats() def unload_model_weights(op='model'):