diff --git a/modules/history.py b/modules/history.py new file mode 100644 index 000000000..f98a0026d --- /dev/null +++ b/modules/history.py @@ -0,0 +1,59 @@ +import sys +import datetime +from collections import deque +from modules import shared, devices + + +class Item(): + def __init__(self, latent, preview=None, meta=None): + self.ts = datetime.datetime.now() + self.latent = latent.detach().clone().to(devices.cpu) + self.preview = preview + self.meta = meta + + +class History(): + def __init__(self): + self.latents = deque(maxlen=1000) + shared.log.debug(f'History init: max={shared.opts.latent_history}') + + @property + def count(self): + return len(self.latents) + + @property + def size(self): + s = 0 + for item in self.latents: + s += sys.getsizeof(item.latent.storage()) + return s + + @property + def list(self): + return [item.ts for item in self.latents] + + @property + def latest(self): + return self.get(0) + + def add(self, latent, preview=None, meta=None): + item = Item(latent, preview, meta) + self.latents.appendleft(item) + shared.log.debug(f'History add: shape={latent.shape} dtype={latent.dtype} count={self.count}') + if self.count >= shared.opts.latent_history: + self.latents.pop() + + def get(self, index: int = 0): + item = self.latents[index] + shared.log.debug(f'History get: index={index} time={item.ts} shape={item.latent.shape} dtype={item.latent.dtype} count={self.count}') + return item.latent.to(devices.device) + + def clear(self): + self.latents.clear() + shared.log.debug(f'History clear: count={self.count}') + + def load(self): + pass + + def save(self): + pass diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index c06c3d3de..8d19dfc4d 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -25,7 +25,7 @@ def restore_state(p: processing.StableDiffusionProcessing): if p.__class__ != last_p.__class__: shared.log.warning(f'Restore state: op={p.state} last state is different type') return p - if processing_vae.last_latent is None: + if shared.history.count == 0: shared.log.warning(f'Restore state: op={p.state} last latents missing') return p state = p.state @@ -388,7 +388,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): if 'base' not in p.skip: output = process_base(p) else: - output = SimpleNamespace(images=processing_vae.last_latent) + output = SimpleNamespace(images=shared.history.latest) if shared.state.interrupted or shared.state.skipped: shared.sd_model = orig_pipeline diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 4ad77edae..454eeed32 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -9,7 +9,6 @@ from modules import shared, devices, sd_models, sd_vae, sd_vae_taesd, errors debug = os.environ.get('SD_VAE_DEBUG', None) is not None log_debug = shared.log.trace if debug else lambda *args, **kwargs: None log_debug('Trace: VAE') -last_latent = None def create_latents(image, p, dtype=None, device=None): @@ -147,10 +146,8 @@ def taesd_vae_encode(image): def vae_decode(latents, model, output_type='np', full_quality=True, width=None, height=None, save=True): - global last_latent # pylint: disable=global-statement t0 = time.time() if latents is None or not torch.is_tensor(latents): # already decoded - last_latent = None return latents prev_job = shared.state.job shared.state.job = 'VAE' @@ -170,7 +167,7 @@ def vae_decode(latents, model, output_type='np', full_quality=True, width=None, if latents.shape[0] == 4 and latents.shape[1] != 4: # likely animatediff latent latents = latents.permute(1, 0, 2, 3) if save: - last_latent = latents.clone().detach() + shared.history.add(latents) if latents.shape[-1] <= 4: # not a latent, likely an image decoded = latents.float().cpu().numpy() @@ -216,10 +213,11 @@ def vae_encode(image, model, full_quality=True): # pylint: disable=unused-variab def reprocess(gallery): from PIL import Image from modules import images - if last_latent is None or gallery is None: + latent = shared.history.latest + if latent is None or gallery is None: return None - shared.log.info(f'Reprocessing: latent={last_latent.shape}') - reprocessed = vae_decode(last_latent, shared.sd_model, output_type='pil', full_quality=True) + shared.log.info(f'Reprocessing: latent={latent.shape}') + reprocessed = vae_decode(latent, shared.sd_model, output_type='pil', full_quality=True) outputs = [] for i0, i1 in zip(gallery, reprocessed): if isinstance(i1, np.ndarray): diff --git a/modules/shared.py b/modules/shared.py index d91d2823f..ccfca97cd 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -15,7 +15,7 @@ import fasteners import orjson import diffusers from rich.console import Console -from modules import errors, devices, shared_items, shared_state, cmd_args, theme +from modules import errors, devices, shared_items, shared_state, cmd_args, theme, history from modules.paths import models_path, script_path, data_path, sd_configs_path, sd_default_config, sd_model_file, default_sd_model_file, extensions_dir, extensions_builtin_dir # pylint: disable=W0611 from modules.dml import memory_providers, default_memory_provider, directml_do_hijack from modules.onnx_impl import initialize_onnx, execution_providers @@ -425,6 +425,7 @@ options_templates.update(options_section(('sd', "Execution & Models"), { "prompt_attention": OptionInfo("Full parser", "Prompt attention parser", gr.Radio, {"choices": ["Full parser", "Compel parser", "xhinker parser", "A1111 parser", "Fixed attention"] }), "sd_checkpoint_cache": OptionInfo(0, "Cached models", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": not native }), "sd_vae_checkpoint_cache": OptionInfo(0, "Cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": False}), + "latent_history": OptionInfo(1, "Latent history size", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1}), "sd_disable_ckpt": OptionInfo(False, "Disallow models in ckpt format", gr.Checkbox, {"visible": False}), "diffusers_version": OptionInfo("", "Diffusers version", gr.Textbox, {"visible": False}), })) @@ -1091,6 +1092,7 @@ batch_cond_uncond = opts.always_batch_cond_uncond or not (cmd_opts.lowvram or cm parallel_processing_allowed = not cmd_opts.lowvram mem_mon = modules.memmon.MemUsageMonitor("MemMon", devices.device) max_workers = 8 +history = history.History() if devices.backend == "directml": directml_do_hijack() elif devices.backend == "cuda": diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index e0d212f31..e8feb2501 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -447,11 +447,13 @@ def register_pages(): from modules.ui_extra_networks_checkpoints import ExtraNetworksPageCheckpoints from modules.ui_extra_networks_styles import ExtraNetworksPageStyles from modules.ui_extra_networks_vae import ExtraNetworksPageVAEs + from modules.ui_extra_networks_history import ExtraNetworksPageHistory debug('EN register-pages') register_page(ExtraNetworksPageCheckpoints()) register_page(ExtraNetworksPageStyles()) register_page(ExtraNetworksPageTextualInversion()) register_page(ExtraNetworksPageVAEs()) + register_page(ExtraNetworksPageHistory()) if shared.opts.hypernetwork_enabled: from modules.ui_extra_networks_hypernets import ExtraNetworksPageHypernetworks register_page(ExtraNetworksPageHypernetworks()) @@ -461,7 +463,7 @@ def get_pages(title=None): visible = shared.opts.extra_networks pages = [] if 'All' in visible or visible == []: # default en sort order - visible = ['Model', 'Lora', 'Style', 'Embedding', 'VAE', 'Hypernetwork'] + visible = ['Model', 'Lora', 'Style', 'Embedding', 'VAE', 'History', 'Hypernetwork'] titles = [page.title for page in shared.extra_networks] if title is None: diff --git a/modules/ui_extra_networks_history.py b/modules/ui_extra_networks_history.py new file mode 100644 index 000000000..7b471e394 --- /dev/null +++ b/modules/ui_extra_networks_history.py @@ -0,0 +1,29 @@ +import time +from modules import shared, ui_extra_networks + + +class ExtraNetworksPageHistory(ui_extra_networks.ExtraNetworksPage): + def __init__(self): + super().__init__('History') + # shared.log.trace('History init') + self.last_refresh = 0 + + def refresh(self): + # shared.log.trace('History refresh') + self.last_refresh = time.time() + self.html = '

buttons

' + for ts in shared.history.list: + self.html += '

' + str(ts) + '

' + + def list_items(self): + # shared.log.trace('History list') + return shared.history.list + + def create_page(self, tabname, skip = False): + # shared.log.trace(f'History page: tab={tabname} skip={skip}') + self.page_time = time.time() + if tabname == 'txt2img': + self.last_refresh = time.time() + if self.page_time <= self.last_refresh: # cached page + self.refresh() + return self.patch(self.html, tabname)