From d1bb51eee0fd6da2ca546a63a6902c40f1676c66 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 25 Jan 2024 12:08:45 -0500 Subject: [PATCH] enable lama offloading --- modules/masking.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/modules/masking.py b/modules/masking.py index ad19f5106..c98fb21b6 100644 --- a/modules/masking.py +++ b/modules/masking.py @@ -8,6 +8,7 @@ import cv2 from PIL import Image, ImageFilter, ImageOps from transformers import SamModel, SamImageProcessor, MaskGenerationPipeline from modules import shared, errors, devices, ui_components, ui_symbols, paths +from modules.memstats import memory_stats def get_crop_region(mask, pad=0): @@ -238,7 +239,6 @@ def run_rembg(input_image: Image, input_mask: np.ndarray): except Exception as e: shared.log.error(f'Segment Rembg load failed: {e}') return input_mask - if "U2NET_HOME" not in os.environ: os.environ["U2NET_HOME"] = os.path.join(paths.models_path, "Rembg") args = { @@ -417,9 +417,14 @@ def run_lama(input_image: gr.Image, input_mask: gr.Image = None): return None input_mask = run_mask(input_image, input_mask, return_type='Grayscale') if lama_model is None: - from modules.lama import SimpleLama - lama_model = SimpleLama() + import modules.lama + shared.log.debug(f'LaMa loading: model={modules.lama.LAMA_MODEL_URL}') + lama_model = modules.lama.SimpleLama() + shared.log.debug(f'LaMa loaded: {memory_stats()}') + lama_model.model.to(devices.device) result = lama_model(input_image, input_mask) + if shared.opts.control_move_processor: + lama_model.model.to('cpu') return result