history prototype

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2024-10-07 19:09:44 -04:00
parent 072a837132
commit a3a177277a
6 changed files with 101 additions and 11 deletions
+59
View File
@@ -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
+2 -2
View File
@@ -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
+5 -7
View File
@@ -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):
+3 -1
View File
@@ -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":
+3 -1
View File
@@ -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:
+29
View File
@@ -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 = '<h1>buttons</h1>'
for ts in shared.history.list:
self.html += '<p>' + str(ts) + '</p>'
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)