mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
add explicit vae method=none and allow images in history
Signed-off-by: vladmandic <mandic00@live.com>
This commit is contained in:
@@ -415,7 +415,7 @@ class ResProgress(BaseModel):
|
||||
progress: float = Field(title="Progress", description="The progress with a range of 0 to 1")
|
||||
eta_relative: float = Field(title="ETA in secs")
|
||||
state: dict = Field(title="State", description="The current state snapshot")
|
||||
current_image: str | None = Field(default=None, title="Current image", description="The current image in base64 format. opts.show_progress_every_n_steps is required for this to work.")
|
||||
current_image: str | None = Field(default=None, title="Current image", description="The current image in base64 format")
|
||||
textinfo: str | None = Field(default=None, title="Info text", description="Info text used by WebUI.")
|
||||
|
||||
class ResHistory(BaseModel):
|
||||
|
||||
+51
-11
@@ -11,16 +11,33 @@ from modules.logger import log
|
||||
|
||||
|
||||
class Item:
|
||||
def __init__(self, latent, preview=None, info=None, ops=None):
|
||||
latent: torch.Tensor | None = None
|
||||
size: int = 0
|
||||
preview: str | None = None
|
||||
info: str | None = None
|
||||
ops: list = []
|
||||
images: list | None = None
|
||||
|
||||
def __init__(self, latent, preview=None, info=None, ops=None, images=None):
|
||||
if ops is None:
|
||||
ops = []
|
||||
self.ts = datetime.datetime.now().replace(microsecond=0)
|
||||
self.name = self.ts.strftime('%Y-%m-%d %H:%M:%S')
|
||||
self.latent = latent.detach().clone().to(devices.cpu)
|
||||
if torch.is_tensor(latent):
|
||||
self.latent = latent.detach().clone().to(devices.cpu)
|
||||
self.size = sys.getsizeof(self.latent.storage())
|
||||
self.preview = preview
|
||||
self.info = info
|
||||
self.ops = ops.copy()
|
||||
self.size = sys.getsizeof(self.latent.storage())
|
||||
self.images = images
|
||||
|
||||
def __str__(self):
|
||||
if self.latent is not None:
|
||||
return f'Item(ts="{self.name}" ops={self.ops} latent={self.latent.shape} size={self.size})'
|
||||
elif self.images is not None:
|
||||
return f'Item(ts="{self.name}" ops={self.ops} images={self.images})'
|
||||
else:
|
||||
return f'Item(ts="{self.name}" ops={self.ops} unknown content)'
|
||||
|
||||
|
||||
class History:
|
||||
@@ -42,7 +59,7 @@ class History:
|
||||
@property
|
||||
def list(self):
|
||||
log.info(f'History: items={self.count}/{shared.opts.latent_history} size={self.size}')
|
||||
return [item.name for item in self.latents]
|
||||
return [item.name for item in self.latents if item.latent is not None]
|
||||
|
||||
@property
|
||||
def selected(self):
|
||||
@@ -52,30 +69,53 @@ class History:
|
||||
else:
|
||||
current_index = 0
|
||||
item = self.latents[current_index]
|
||||
if item.latent is None:
|
||||
return None
|
||||
log.debug(f'History get: index={current_index} time={item.ts} shape={list(item.latent.shape)} dtype={item.latent.dtype} count={self.count}')
|
||||
return item.latent.to(devices.device), current_index
|
||||
|
||||
@property
|
||||
def last_item(self):
|
||||
return self.latents[0] if self.count > 0 else None
|
||||
|
||||
@property
|
||||
def last_image(self):
|
||||
if self.count == 0:
|
||||
return None
|
||||
for item in self.latents:
|
||||
if item.images is not None:
|
||||
return item.images
|
||||
return None
|
||||
|
||||
@property
|
||||
def last_latent(self):
|
||||
if self.count == 0:
|
||||
return None
|
||||
for item in self.latents:
|
||||
if item.latent is not None:
|
||||
return item.latent
|
||||
return None
|
||||
|
||||
def find(self, name):
|
||||
for i, item in enumerate(self.latents):
|
||||
if item.name == name:
|
||||
return i
|
||||
return -1
|
||||
|
||||
def add(self, latent, preview=None, info=None, ops=None):
|
||||
def add(self, latent, preview=None, info=None, ops=None, images=None):
|
||||
if ops is None:
|
||||
ops = []
|
||||
shared.state.latent_history += 1
|
||||
if shared.opts.latent_history == 0:
|
||||
return
|
||||
if torch.is_tensor(latent):
|
||||
item = Item(latent, preview, info, ops)
|
||||
self.latents.appendleft(item)
|
||||
if self.count >= shared.opts.latent_history:
|
||||
self.latents.pop()
|
||||
item = Item(latent, preview, info, ops, images)
|
||||
self.latents.appendleft(item)
|
||||
if self.count >= shared.opts.latent_history:
|
||||
self.latents.pop()
|
||||
log.debug(f'History: len={self.count} add={item}')
|
||||
|
||||
def clear(self):
|
||||
self.latents.clear()
|
||||
# log.debug(f'History clear: count={self.count}')
|
||||
|
||||
def load(self):
|
||||
pass
|
||||
|
||||
+16
-10
@@ -64,18 +64,24 @@ def get_grid_size(imgs: list, batch_size=1, rows: int | None = None, cols: int |
|
||||
|
||||
|
||||
def image_grid(imgs: list, batch_size=1, rows: int | None = None, cols: int | None = None):
|
||||
rows, cols = get_grid_size(imgs, batch_size, rows=rows, cols=cols)
|
||||
params = script_callbacks.ImageGridLoopParams(imgs, cols, rows)
|
||||
script_callbacks.image_grid_callback(params)
|
||||
imgs = [i for i in imgs if i is not None] if imgs is not None else []
|
||||
if isinstance(imgs, Image.Image):
|
||||
return imgs
|
||||
imgs = [i for i in imgs if i is not None] if imgs is not None and isinstance(imgs, list) else []
|
||||
if len(imgs) == 0:
|
||||
return None
|
||||
w, h = max(i.width for i in imgs if i is not None), max(i.height for i in imgs if i is not None)
|
||||
grid = Image.new('RGB', size=(params.cols * w, params.rows * h), color=shared.opts.grid_background)
|
||||
for i, img in enumerate(params.imgs):
|
||||
if img is not None:
|
||||
grid.paste(img, box=(i % params.cols * w, i // params.cols * h))
|
||||
return grid
|
||||
try:
|
||||
rows, cols = get_grid_size(imgs, batch_size, rows=rows, cols=cols)
|
||||
params = script_callbacks.ImageGridLoopParams(imgs, cols, rows)
|
||||
script_callbacks.image_grid_callback(params)
|
||||
w, h = max(i.width for i in imgs if i is not None), max(i.height for i in imgs if i is not None)
|
||||
grid = Image.new('RGB', size=(params.cols * w, params.rows * h), color=shared.opts.grid_background)
|
||||
for i, img in enumerate(params.imgs):
|
||||
if img is not None:
|
||||
grid.paste(img, box=(i % params.cols * w, i // params.cols * h))
|
||||
return grid
|
||||
except Exception as e:
|
||||
log.error(f'Grid: images={imgs} {e}')
|
||||
return None
|
||||
|
||||
|
||||
def split_grid(image: Image.Image, tile_w=512, tile_h=512, overlap=64):
|
||||
|
||||
@@ -315,7 +315,7 @@ def process_samples(p: StableDiffusionProcessing, samples):
|
||||
p.ops.append('detailer')
|
||||
if not p.do_not_save_samples and get_opt(p, 'save_images_before_detailer'):
|
||||
info = create_infotext(p, p.prompts, p.seeds, p.subseeds, index=i)
|
||||
images.save_image(Image.fromarray(sample), path=p.outpath_samples, basename="", seed=p.seeds[i], prompt=p.prompts[i], extension=get_opt(p, 'samples_format'), info=info, p=p, suffix="-before-detailer")
|
||||
images.save_image(image, path=p.outpath_samples, basename="", seed=p.seeds[i], prompt=p.prompts[i], extension=get_opt(p, 'samples_format'), info=info, p=p, suffix="-before-detailer")
|
||||
sample = detailer.detail(sample, p)
|
||||
if isinstance(sample, list):
|
||||
if len(sample) > 0:
|
||||
@@ -426,6 +426,7 @@ def process_samples(p: StableDiffusionProcessing, samples):
|
||||
image.info["parameters"] = info
|
||||
out_infotexts.append(info)
|
||||
out_images.append(image)
|
||||
shared.history.add(None, info=out_infotexts, ops=p.ops, images=out_images)
|
||||
return out_images, out_infotexts
|
||||
|
||||
|
||||
|
||||
+10
-4
@@ -10,17 +10,23 @@ def hijack_encode_prompt(*args, **kwargs):
|
||||
if 'max_sequence_length' in kwargs and kwargs['max_sequence_length'] is not None:
|
||||
kwargs['max_sequence_length'] = max(kwargs['max_sequence_length'], os.environ.get('MAX_SEQUENCE_LENGTH', 256))
|
||||
try:
|
||||
prompt = kwargs.get('prompt', None) or (args[0] if len(args) > 0 else None)
|
||||
args_copy = list(args)
|
||||
patch_prompt = False
|
||||
prompt = kwargs.get('prompt', None)
|
||||
if prompt is None and len(args_copy) > 0:
|
||||
prompt = args[0]
|
||||
patch_prompt = True
|
||||
res = prompt
|
||||
if prompt is not None:
|
||||
log.debug(f'Encode: prompt="{prompt}" hijack=True')
|
||||
if hasattr(shared.sd_model, 'before_prompt_encode'):
|
||||
prompt = shared.sd_model.before_prompt_encode(prompt)
|
||||
res = shared.sd_model.before_prompt_encode(prompt)
|
||||
if patch_prompt:
|
||||
args_copy[0] = res
|
||||
if hasattr(shared.sd_model, 'orig_encode_prompt'):
|
||||
res = shared.sd_model.orig_encode_prompt(*args, **kwargs)
|
||||
if hasattr(shared.sd_model, 'after_prompt_encode'):
|
||||
res = shared.sd_model.after_prompt_encode(res)
|
||||
else:
|
||||
res = prompt
|
||||
except Exception as e:
|
||||
log.error(f'Encode prompt: {e}')
|
||||
errors.display(e, 'Encode prompt')
|
||||
|
||||
@@ -875,7 +875,7 @@ def load_diffuser(checkpoint_info=None, op='model', revision=None): # pylint: di
|
||||
vae_file = None
|
||||
if model_type.startswith('Stable Diffusion') and (op == 'model' or op == 'refiner'): # preload vae for sd models
|
||||
vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename)
|
||||
vae = sd_vae.load_vae_diffusers(checkpoint_info.path, vae_file, vae_source)
|
||||
vae = sd_vae.load_vae(checkpoint_info.path, vae_file, vae_source)
|
||||
if vae is not None:
|
||||
diffusers_load_config["vae"] = vae
|
||||
timer.load.record("vae")
|
||||
|
||||
@@ -3,7 +3,7 @@ import threading
|
||||
from collections import namedtuple
|
||||
import torch
|
||||
from PIL import Image
|
||||
from modules import shared, devices, processing, images, sd_samplers, timer
|
||||
from modules import shared, processing, images, sd_samplers, timer
|
||||
from modules.logger import log
|
||||
from modules.vae import sd_vae_approx, sd_vae_taesd, sd_vae_stablecascade
|
||||
from modules.image import convert
|
||||
@@ -38,13 +38,9 @@ def setup_img2img_steps(p, steps=None):
|
||||
def single_sample_to_image(sample, approximation=None):
|
||||
with queue_lock:
|
||||
t0 = time.time()
|
||||
if approximation is None:
|
||||
approximation = approximation_indexes.get(shared.opts.show_progress_type, None)
|
||||
if approximation is None:
|
||||
warn_once('Unknown decode type')
|
||||
approximation = 0
|
||||
approximation = approximation or shared.opts.show_progress_type
|
||||
try:
|
||||
if sample.dtype == torch.bfloat16 and (approximation == 0 or approximation == 1):
|
||||
if (sample.dtype == torch.bfloat16) and (approximation in ["Simple", "Approximate"]):
|
||||
sample = sample.to(torch.float16)
|
||||
except Exception as e:
|
||||
warn_once(f'Preview: {e}')
|
||||
@@ -53,7 +49,10 @@ def single_sample_to_image(sample, approximation=None):
|
||||
return Image.new(mode="RGB", size=(512, 512))
|
||||
if len(sample.shape) == 4 and sample.shape[0]: # likely animatediff latent
|
||||
sample = sample.permute(1, 0, 2, 3)[0]
|
||||
if approximation == 2: # TAESD
|
||||
|
||||
if approximation == "None":
|
||||
return Image.new(mode="RGB", size=(512, 512)) # already handled
|
||||
elif approximation == "TAESD":
|
||||
if (len(sample.shape) == 3 or len(sample.shape) == 4) and shared.opts.live_preview_downscale and (sample.shape[-1]*sample.shape[-2] > 128*128):
|
||||
try:
|
||||
scale = (128 * 128) / (sample.shape[-1] * sample.shape[-2])
|
||||
@@ -62,19 +61,20 @@ def single_sample_to_image(sample, approximation=None):
|
||||
pass
|
||||
x_sample = sd_vae_taesd.decode(sample)
|
||||
# x_sample = (1.0 + x_sample) / 2.0 # preview requires smaller range
|
||||
elif shared.sd_model_type == 'sc' and approximation != 3:
|
||||
elif shared.sd_model_type == 'sc' and approximation != "Full":
|
||||
x_sample = sd_vae_stablecascade.decode(sample)
|
||||
elif approximation == 0: # Simple
|
||||
elif approximation == "Simple":
|
||||
x_sample = sd_vae_approx.cheap_approximation(sample) * 0.5 + 0.5
|
||||
elif approximation == 1: # Approximate
|
||||
elif approximation == "Approximate":
|
||||
x_sample = sd_vae_approx.nn_approximation(sample) * 0.5 + 0.5
|
||||
if shared.sd_model_type == "sdxl":
|
||||
x_sample = x_sample[[2, 1, 0], :, :] # BGR to RGB
|
||||
elif approximation == 3: # Full VAE
|
||||
elif approximation == "Full":
|
||||
x_sample = processing.decode_first_stage(shared.sd_model, sample.unsqueeze(0))[0]
|
||||
else:
|
||||
warn_once(f"Unknown latent decode type: {approximation}")
|
||||
warn_once(f"VAE: method={approximation} unknown")
|
||||
return Image.new(mode="RGB", size=(512, 512))
|
||||
|
||||
try:
|
||||
if isinstance(x_sample, Image.Image):
|
||||
image = x_sample
|
||||
@@ -102,33 +102,11 @@ def samples_to_image_grid(samples, approximation=None):
|
||||
return images.image_grid([single_sample_to_image(sample, approximation) for sample in samples])
|
||||
|
||||
|
||||
def images_tensor_to_samples(image, approximation=None, model=None):
|
||||
'''image[0, 1] -> latent'''
|
||||
if approximation is None:
|
||||
approximation = approximation_indexes.get(shared.opts.show_progress_type, 0)
|
||||
if approximation == 2:
|
||||
image = image.to(devices.device, devices.dtype)
|
||||
x_latent = sd_vae_taesd.encode(image)
|
||||
else:
|
||||
if model is None:
|
||||
model = shared.sd_model
|
||||
model.first_stage_model.to(devices.dtype_vae)
|
||||
image = image.to(shared.device, dtype=devices.dtype_vae)
|
||||
image = image * 2 - 1
|
||||
if len(image) > 1:
|
||||
image_latents = [model.get_first_stage_encoding(model.encode_first_stage(torch.unsqueeze(img, 0)))[0] for img in image]
|
||||
x_latent = torch.stack(image_latents)
|
||||
else:
|
||||
x_latent = model.get_first_stage_encoding(model.encode_first_stage(image))
|
||||
return x_latent
|
||||
|
||||
|
||||
def store_latent(decoded):
|
||||
shared.state.current_latent = decoded
|
||||
if shared.opts.show_progress_every_n_steps > 0 and shared.state.sampling_step % shared.opts.show_progress_every_n_steps == 0:
|
||||
if not shared.parallel_processing_allowed:
|
||||
image = sample_to_image(decoded)
|
||||
shared.state.assign_current_image(image)
|
||||
if not shared.parallel_processing_allowed:
|
||||
image = sample_to_image(decoded)
|
||||
shared.state.assign_current_image(image)
|
||||
|
||||
|
||||
def is_sampler_using_eta_noise_seed_delta(p):
|
||||
|
||||
+2
-2
@@ -145,7 +145,7 @@ def apply_vae_config(model_file, vae_file, sd_model):
|
||||
sd_model.vae.config[k] = v
|
||||
|
||||
|
||||
def load_vae_diffusers(model_file, vae_file=None, vae_source="unknown-source"):
|
||||
def load_vae(model_file, vae_file=None, vae_source="unknown-source"):
|
||||
if vae_file is None:
|
||||
return None
|
||||
if not os.path.exists(vae_file):
|
||||
@@ -222,7 +222,7 @@ def reload_vae_weights(sd_model=None, vae_file=unspecified):
|
||||
return None
|
||||
|
||||
if hasattr(sd_model, "vae") and getattr(sd_model, "sd_checkpoint_info", None) is not None:
|
||||
vae = load_vae_diffusers(sd_model.sd_checkpoint_info.filename, vae_file, vae_source)
|
||||
vae = load_vae(sd_model.sd_checkpoint_info.filename, vae_file, vae_source)
|
||||
if vae is not None:
|
||||
if not hasattr(sd_model, 'original_vae'):
|
||||
sd_model.original_vae = sd_model.vae
|
||||
|
||||
@@ -5,8 +5,8 @@ Contains optimized tensor operations and video processing functions
|
||||
Extracted from: seedvr2.py (lines 1633-1730)
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
import torch
|
||||
from typing import List, Union
|
||||
|
||||
|
||||
def optimized_video_rearrange(video_tensors: List[torch.Tensor]) -> List[torch.Tensor]:
|
||||
|
||||
+12
-8
@@ -263,19 +263,23 @@ class State:
|
||||
def set_current_image(self):
|
||||
if self.job == 'VAE' or self.job == 'Upscale': # avoid generating preview while vae is running
|
||||
return False
|
||||
from modules.shared import opts, cmd_opts
|
||||
if cmd_opts.lowvram or self.api or (opts.show_progress_every_n_steps <= 0):
|
||||
from modules.shared import cmd_opts
|
||||
if cmd_opts.lowvram or self.api or self.disable_preview:
|
||||
return False
|
||||
if (not self.disable_preview) and (abs(self.sampling_step - self.current_image_sampling_step) >= opts.show_progress_every_n_steps):
|
||||
return self.do_set_current_image()
|
||||
return False
|
||||
return self.do_set_current_image()
|
||||
|
||||
def do_set_current_image(self):
|
||||
from modules import shared, images, sd_samplers_common
|
||||
if (self.current_latent is None) or self.disable_preview or (self.preview_job == self.job_no):
|
||||
return False
|
||||
from modules import shared
|
||||
from modules.sd_samplers_common import samples_to_image_grid, sample_to_image
|
||||
self.preview_job = self.job_no
|
||||
|
||||
if (shared.opts.show_progress_type == "None") and (shared.history.last_image is not None):
|
||||
last_image = images.image_grid(shared.history.last_image)
|
||||
self.assign_current_image(last_image)
|
||||
self.preview_job = -1
|
||||
return True
|
||||
|
||||
try:
|
||||
sample = self.current_latent
|
||||
self.current_image_sampling_step = self.sampling_step
|
||||
@@ -288,7 +292,7 @@ class State:
|
||||
sample = self.current_noise_pred * (-self.current_sigma / (self.current_sigma**2 + 1) ** 0.5) + (original_sample / (self.current_sigma**2 + 1)) # pylint: disable=invalid-unary-operand-type
|
||||
except Exception:
|
||||
pass # ignore sigma errors
|
||||
image = samples_to_image_grid(sample) if shared.opts.show_progress_grid else sample_to_image(sample)
|
||||
image = sd_samplers_common.samples_to_image_grid(sample)
|
||||
self.assign_current_image(image)
|
||||
self.preview_job = -1
|
||||
return True
|
||||
|
||||
@@ -74,7 +74,7 @@ def create_settings(cmd_opts):
|
||||
"sd_model_checkpoint": OptionInfo(default_checkpoint, "Base model", DropdownEditable, lambda: {"choices": list_checkpoint_titles()}, refresh=refresh_checkpoints),
|
||||
"sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints),
|
||||
"sd_unet": OptionInfo("Default", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list),
|
||||
"latent_history": OptionInfo(16, "Latent history size", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1}),
|
||||
"latent_history": OptionInfo(20, "Latent history size", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1}),
|
||||
|
||||
"advanced_sep": OptionInfo("<h2>Advanced Options</h2>", "", gr.HTML),
|
||||
"sd_checkpoint_autoload": OptionInfo(True, "Model auto-load on start"),
|
||||
@@ -536,8 +536,8 @@ def create_settings(cmd_opts):
|
||||
|
||||
# --- Live Previews ---
|
||||
options_templates.update(options_section(('live-preview', "Live Previews"), {
|
||||
"show_progress_every_n_steps": OptionInfo(1, "Live preview display period", gr.Slider, {"minimum": 0, "maximum": 20, "step": 1}),
|
||||
"show_progress_type": OptionInfo("TAESD", "Live preview method", gr.Radio, {"choices": ["Simple", "Approximate", "TAESD", "Full VAE"]}),
|
||||
"show_progress_every_n_steps": OptionInfo(1, "Live preview display period", gr.Slider, {"minimum": 0, "maximum": 20, "step": 1, "visible": False}),
|
||||
"show_progress_type": OptionInfo("TAESD", "Live preview method", gr.Dropdown, {"choices": ["None", "Simple", "Approximate", "TAESD", "Full"]}),
|
||||
"live_preview_refresh_period": OptionInfo(500, "Progress update period", gr.Slider, {"minimum": 0, "maximum": 5000, "step": 25}),
|
||||
"taesd_variant": OptionInfo(shared_items.sd_taesd_items()[0], "TAESD variant", gr.Dropdown, {"choices": shared_items.sd_taesd_items()}),
|
||||
"taesd_layers": OptionInfo(3, "TAESD decode layers", gr.Slider, {"minimum": 1, "maximum": 3, "step": 1}),
|
||||
|
||||
@@ -150,13 +150,19 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage):
|
||||
if 'baseModel' in version:
|
||||
record['version'] = version.get("baseModel", "")
|
||||
elif '_class_name' in record['info']:
|
||||
record['version'] = record['info'].get('_class_name', '').replace('Pipeline', '').replace('Image', '')
|
||||
cls = record['info']['_class_name']
|
||||
if isinstance(cls, list):
|
||||
cls = cls[-1]
|
||||
record['version'] = cls.replace('Pipeline', '').replace('Image', '')
|
||||
else:
|
||||
record['version'] = ''
|
||||
record['version'] = version_map.get(record['version'], record['version'])
|
||||
|
||||
except Exception as e:
|
||||
log.debug(f'Networks error: type=model file="{name}" {e}')
|
||||
log.error(f'Networks error: type=model file="{name}" {e}')
|
||||
if os.environ.get('SD_EN_DEBUG', None) is not None:
|
||||
from modules import errors
|
||||
errors.display(e, 'Networks')
|
||||
return record
|
||||
|
||||
def list_items(self):
|
||||
|
||||
Reference in New Issue
Block a user