From 3343d2e05fc60670c0aecb908b1c823c8ad08433 Mon Sep 17 00:00:00 2001 From: awsr <43862868+awsr@users.noreply.github.com> Date: Fri, 23 Jan 2026 04:56:11 -0800 Subject: [PATCH] Update and rewrite to use contextlib --- modules/lora/lora_apply.py | 2 +- modules/lora/networks.py | 12 ++--- modules/textual_inversion.py | 85 ++++++++++++++++++------------------ sdnext_core/errorlimiter.py | 46 +++++++++++++++++-- 4 files changed, 88 insertions(+), 57 deletions(-) diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 2d6e28abb..e2f3a8740 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -142,7 +142,7 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. if l.debug: errors.display(e, 'LoRA') raise RuntimeError('LoRA apply weight') from e - ErrorLimiter.update("network_calc_weights") + ErrorLimiter.notify("network_calc_weights") continue return batch_updown, batch_ex_bias diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 5a89220b0..65a726c52 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -1,7 +1,7 @@ from contextlib import nullcontext import time import rich.progress as rp -from sdnext_core.errorlimiter import ErrorLimiter, ErrorLimiterTrigger +from sdnext_core.errorlimiter import limit_errors from modules.lora import lora_common as l from modules.lora.lora_apply import network_apply_weights, network_apply_direct, network_backup_weights, network_calc_weights from modules import shared, devices, sd_models @@ -13,8 +13,7 @@ default_components = ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'text_ def network_activate(include=[], exclude=[]): t0 = time.time() - ErrorLimiter.start("network_calc_weights") - try: + with limit_errors("network_calc_weights"): sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) if shared.opts.diffusers_offload_mode == "sequential": sd_models.disable_offload(sd_model) @@ -70,8 +69,6 @@ def network_activate(include=[], exclude=[]): if task is not None and len(applied_layers) == 0: pbar.remove_task(task) # hide progress bar for no action - except ErrorLimiterTrigger as e: - raise RuntimeError(f"HALTING. Too many errors during {e.name}") from e l.timer.activate += time.time() - t0 if l.debug and len(l.loaded_networks) > 0: shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={round(backup_size/1024/1024/1024, 2)} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers} device={device} time={l.timer.summary}') @@ -86,8 +83,7 @@ def network_deactivate(include=[], exclude=[]): if len(l.previously_loaded_networks) == 0: return t0 = time.time() - ErrorLimiter.start("network_calc_weights") - try: + with limit_errors("network_calc_weights"): sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) if shared.opts.diffusers_offload_mode == "sequential": sd_models.disable_offload(sd_model) @@ -130,8 +126,6 @@ def network_deactivate(include=[], exclude=[]): module.network_current_names = () if task is not None: pbar.update(task, advance=1, description=f'networks={len(l.previously_loaded_networks)} modules={active_components} layers={total} unapply={len(applied_layers)}') - except ErrorLimiterTrigger as e: - raise RuntimeError(f"HALTING. Too many errors during {e.name}") from e l.timer.deactivate = time.time() - t0 if l.debug and len(l.previously_loaded_networks) > 0: shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in l.previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers} time={l.timer.summary}') diff --git a/modules/textual_inversion.py b/modules/textual_inversion.py index 9433130f0..832ce8810 100644 --- a/modules/textual_inversion.py +++ b/modules/textual_inversion.py @@ -3,7 +3,7 @@ import os import time import torch import safetensors.torch -from sdnext_core.errorlimiter import ErrorLimiter +from sdnext_core.errorlimiter import limit_errors from modules import shared, devices, errors from modules.files_cache import directory_files, directory_mtime, extension_filter @@ -259,51 +259,50 @@ class EmbeddingDatabase: File names take precidence over bundled embeddings passed as a dict. Bundled embeddings are automatically set to overwrite previous embeddings. """ - overwrite = bool(data) - if not shared.sd_loaded: - return - if not shared.opts.diffusers_enable_embed: - return - embeddings, skipped = open_embeddings(filename) or convert_bundled(data) - for skip in skipped: - self.skipped_embeddings[skip.name] = skipped - if not embeddings: - return - text_encoders, tokenizers, hiddensizes = get_text_encoders() - if not all([text_encoders, tokenizers, hiddensizes]): - return - ErrorLimiter.start("load_diffusers_embedding_1") - for embedding in embeddings: - try: - embedding.vector_sizes = [v.shape[-1] for v in embedding.vec] - if shared.opts.diffusers_convert_embed and 768 in hiddensizes and 1280 in hiddensizes and 1280 not in embedding.vector_sizes and 768 in embedding.vector_sizes: - embedding.vec.append(convert_embedding(embedding.vec[embedding.vector_sizes.index(768)], text_encoders[hiddensizes.index(768)], text_encoders[hiddensizes.index(1280)])) - embedding.vector_sizes.append(1280) - if (not all(vs in hiddensizes for vs in embedding.vector_sizes) or # Skip SD2.1 in SD1.5/SDXL/SD3 vis versa - len(embedding.vector_sizes) > len(hiddensizes) or # Skip SDXL/SD3 in SD1.5 - (len(embedding.vector_sizes) < len(hiddensizes) and len(embedding.vector_sizes) != 2)): # SD3 no T5 - embedding.tokens = [] + with limit_errors("load_diffusers_embedding") as elimit: + overwrite = bool(data) + if not shared.sd_loaded: + return + if not shared.opts.diffusers_enable_embed: + return + embeddings, skipped = open_embeddings(filename) or convert_bundled(data) + for skip in skipped: + self.skipped_embeddings[skip.name] = skipped + if not embeddings: + return + text_encoders, tokenizers, hiddensizes = get_text_encoders() + if not all([text_encoders, tokenizers, hiddensizes]): + return + for embedding in embeddings: + try: + embedding.vector_sizes = [v.shape[-1] for v in embedding.vec] + if shared.opts.diffusers_convert_embed and 768 in hiddensizes and 1280 in hiddensizes and 1280 not in embedding.vector_sizes and 768 in embedding.vector_sizes: + embedding.vec.append(convert_embedding(embedding.vec[embedding.vector_sizes.index(768)], text_encoders[hiddensizes.index(768)], text_encoders[hiddensizes.index(1280)])) + embedding.vector_sizes.append(1280) + if (not all(vs in hiddensizes for vs in embedding.vector_sizes) or # Skip SD2.1 in SD1.5/SDXL/SD3 vis versa + len(embedding.vector_sizes) > len(hiddensizes) or # Skip SDXL/SD3 in SD1.5 + (len(embedding.vector_sizes) < len(hiddensizes) and len(embedding.vector_sizes) != 2)): # SD3 no T5 + embedding.tokens = [] + self.skipped_embeddings[embedding.name] = embedding + except Exception as e: + shared.log.error(f'Load embedding invalid: name="{embedding.name}" fn="{filename}" {e}') self.skipped_embeddings[embedding.name] = embedding - except Exception as e: - shared.log.error(f'Load embedding invalid: name="{embedding.name}" fn="{filename}" {e}') - self.skipped_embeddings[embedding.name] = embedding - ErrorLimiter.update("load_diffusers_embedding_1") - if overwrite: - shared.log.info(f"Load bundled embeddings: {list(data.keys())}") + elimit() + if overwrite: + shared.log.info(f"Load bundled embeddings: {list(data.keys())}") + for embedding in embeddings: + if embedding.name not in self.skipped_embeddings: + deref_tokenizers(embedding.tokens, tokenizers) + insert_tokens(embeddings, tokenizers) for embedding in embeddings: if embedding.name not in self.skipped_embeddings: - deref_tokenizers(embedding.tokens, tokenizers) - insert_tokens(embeddings, tokenizers) - ErrorLimiter.start("load_diffusers_embedding_2") - for embedding in embeddings: - if embedding.name not in self.skipped_embeddings: - try: - insert_vectors(embedding, tokenizers, text_encoders, hiddensizes) - self.register_embedding(embedding, shared.sd_model) - except Exception as e: - shared.log.error(f'Load embedding: name="{embedding.name}" file="{embedding.filename}" {e}') - errors.display(e, f'Load embedding: name="{embedding.name}" file="{embedding.filename}"') - ErrorLimiter.update("load_diffusers_embedding_2") + try: + insert_vectors(embedding, tokenizers, text_encoders, hiddensizes) + self.register_embedding(embedding, shared.sd_model) + except Exception as e: + shared.log.error(f'Load embedding: name="{embedding.name}" file="{embedding.filename}" {e}') + errors.display(e, f'Load embedding: name="{embedding.name}" file="{embedding.filename}"') + elimit() return def load_from_dir(self, embdir): diff --git a/sdnext_core/errorlimiter.py b/sdnext_core/errorlimiter.py index ea8afd1e5..1d1078036 100644 --- a/sdnext_core/errorlimiter.py +++ b/sdnext_core/errorlimiter.py @@ -1,10 +1,13 @@ -class ErrorLimiterTrigger(Exception): +from contextlib import contextmanager + + +class ErrorLimiterTrigger(BaseException): # Use BaseException to avoid being caught by "except Exception:". def __init__(self, name: str, *args): super().__init__(*args) self.name = name -class ErrorLimiterError(Exception): +class ErrorLimiterAbort(RuntimeError): def __init__(self, msg: str): super().__init__(msg) @@ -17,10 +20,45 @@ class ErrorLimiter: cls._store[name] = limit @classmethod - def update(cls, name: str): + def notify(cls, name: str): # Can be manually triggered if execution is spread across multiple files if name in cls._store.keys(): cls._store[name] = cls._store[name] - 1 if cls._store[name] <= 0: raise ErrorLimiterTrigger(name) else: - raise ErrorLimiterError(f"ErrorLimiter for '{name}' was called before setup") + raise RuntimeError(f"ErrorLimiter for '{name}' was called before setup") + + @classmethod + def end(cls, name: str): + cls._store.pop(name) + + +@contextmanager +def limit_errors(name: str, limit: int = 5): + """Limiter for aborting execution after being triggered a specified number of times (default 5). + + >>> with limit_errors("identifier", limit=5) as elimit: + >>> while do_thing(): + >>> if (something_bad): + >>> print("Something bad happened") + >>> elimit() # In this example, raises ErrorLimiterAbort on the 5th call + >>> try: + >>> something_broken() + >>> except Exception: + >>> print("Encountered an exception") + >>> elimit() # Count is shared across all calls + + Args: + name (str): Identifier. + limit (int, optional): Abort after `limit` number of triggers. Defaults to 5. + + Raises: + ErrorLimiterAbort: Subclass of RuntimeException. + """ + try: + ErrorLimiter.start(name, limit) + yield lambda: ErrorLimiter.notify(name) + except ErrorLimiterTrigger as e: + raise ErrorLimiterAbort(f"HALTING. Too many errors during {e.name}") from None + finally: + ErrorLimiter.end(name)