From 1e31181ae42e1de10bb2ab7326f64c442f3292cb Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 27 Jun 2024 13:00:51 +0300 Subject: [PATCH] handle sampling interrupt better --- nodes.py | 20 ++++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) diff --git a/nodes.py b/nodes.py index 32682c0..8decede 100644 --- a/nodes.py +++ b/nodes.py @@ -6,6 +6,7 @@ import gc import comfy.model_management as mm from comfy.utils import ProgressBar, load_torch_file + import folder_paths script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -485,16 +486,23 @@ class LuminaT2ISampler: model_kwargs["scale_factor"] = 1.0 model_kwargs["scale_watershed"] = 1.0 - #inference - model.to(device) - - samples = ode.sample(z, model.forward_with_cfg, **model_kwargs)[-1] - - if not keep_model_loaded: + def offload_model(): print("Offloading Lumina model...") model.to(offload_device) mm.soft_empty_cache() gc.collect() + + #inference + model.to(device) + try: + samples = ode.sample(z, model.forward_with_cfg, **model_kwargs)[-1] + except: + if not keep_model_loaded: + offload_model() + raise mm.InterruptProcessingException() + + if not keep_model_loaded: + offload_model() samples = samples[:len(samples) // 2] samples = samples / vae_scaling_factor