Update and rewrite to use contextlib

This commit is contained in:
awsr
2026-01-23 04:56:11 -08:00
parent 0310dc8fd6
commit 3343d2e05f
4 changed files with 88 additions and 57 deletions
+1 -1
View File
@@ -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
+3 -9
View File
@@ -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}')
+42 -43
View File
@@ -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):
+42 -4
View File
@@ -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)