mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
implement styles extra field
This commit is contained in:
+2
-1
@@ -5,7 +5,6 @@
|
||||
**TBD**: Candidates before release:
|
||||
- Note: Free-U requires unreleased diffusers
|
||||
- Update Lora handler for *backend:diffusers*
|
||||
- Implement styles extra field
|
||||
- Merge parallel batch processing
|
||||
|
||||
This is a big one, with some major changes and new functionality...
|
||||
@@ -48,6 +47,8 @@ or even free speedups and quality improvements (regardless of which workflows yo
|
||||
if style if an exact match, it will be used
|
||||
otherwise it will rotate between styles that match the start of the name
|
||||
that way you can use different styles as wildcards when processing batches
|
||||
- styles can have **extra** fields, not just prompt and negative prompt
|
||||
for example: *"Extra: sampler: Euler a, width: 480, height: 640, steps: 30, cfg scale: 10, clip skip: 2"*
|
||||
- **VAE**
|
||||
- VAEs are now also listed as part of extra networks
|
||||
- **LoRA**
|
||||
|
||||
@@ -35,7 +35,7 @@ import modules.sd_models
|
||||
import modules.sd_vae
|
||||
import modules.sd_vae_approx
|
||||
import modules.generation_parameters_copypaste
|
||||
from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet
|
||||
from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet, hypertile_set
|
||||
|
||||
|
||||
opt_C = 4
|
||||
@@ -639,6 +639,8 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
|
||||
if k == 'sd_vae':
|
||||
modules.sd_vae.reload_vae_weights()
|
||||
|
||||
shared.prompt_styles.apply_styles_to_extra(p)
|
||||
|
||||
if not shared.opts.cuda_compile:
|
||||
modules.sd_models.apply_token_merging(p.sd_model, p.get_token_merging_ratio())
|
||||
modules.sd_hijack_freeu.apply_freeu(p.sd_model, shared.backend == shared.Backend.ORIGINAL)
|
||||
@@ -1032,6 +1034,7 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
|
||||
self.enable_hr = False
|
||||
|
||||
self.ops.append('txt2img')
|
||||
hypertile_set(self)
|
||||
self.sampler = modules.sd_samplers.create_sampler(self.sampler_name, self.sd_model)
|
||||
x = create_random_tensors([4, self.height // 8, self.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=self.subseed_strength, seed_resize_from_h=self.seed_resize_from_h, seed_resize_from_w=self.seed_resize_from_w, p=self)
|
||||
samples = self.sampler.sample(self, x, conditioning, unconditional_conditioning, image_conditioning=self.txt2img_image_conditioning(x))
|
||||
@@ -1243,6 +1246,7 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
|
||||
self.image_conditioning = self.img2img_image_conditioning(image, self.init_latent, image_mask)
|
||||
|
||||
def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts):
|
||||
hypertile_set(self)
|
||||
x = create_random_tensors([4, self.height // 8, self.width // 8], seeds=seeds, subseeds=subseeds, subseed_strength=self.subseed_strength, seed_resize_from_h=self.seed_resize_from_h, seed_resize_from_w=self.seed_resize_from_w, p=self)
|
||||
x *= self.initial_noise_multiplier
|
||||
samples = self.sampler.sample_img2img(self, self.init_latent, x, conditioning, unconditional_conditioning, image_conditioning=self.image_conditioning)
|
||||
|
||||
@@ -10,6 +10,7 @@ import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from modules.shared import log
|
||||
|
||||
|
||||
# global variables to keep track of changing image size in multiple passes
|
||||
@@ -17,6 +18,8 @@ height = None
|
||||
width = None
|
||||
max_h = 0
|
||||
max_w = 0
|
||||
error_reported = False
|
||||
reset_needed = False
|
||||
|
||||
|
||||
def possible_tile_sizes(dimension: int, tile_size: int, min_tile_size: int, tile_options: int) -> list[int]:
|
||||
@@ -57,7 +60,7 @@ def split_attention(layer: nn.Module, tile_size: int=256, min_tile_size: int=256
|
||||
def self_attn_forward(forward: Callable) -> Callable:
|
||||
@wraps(forward)
|
||||
def wrapper(*args, **kwargs):
|
||||
global height, width, max_h, max_w # pylint: disable=global-statement
|
||||
global height, width, max_h, max_w, reset_needed # pylint: disable=global-statement
|
||||
nh, nw = make_ns()
|
||||
x = args[0]
|
||||
if x.ndim == 4: # VAE
|
||||
@@ -72,26 +75,40 @@ def split_attention(layer: nn.Module, tile_size: int=256, min_tile_size: int=256
|
||||
h, w = round(math.sqrt(ar * hw)), round(math.sqrt(hw / ar))
|
||||
# dynamic height/width based on fact that first two forward calls contain actual height/width
|
||||
# and reset if latest hw is larger since we're never downscaling in 2nd pass
|
||||
if h > max_h:
|
||||
height = 8 * h
|
||||
max_h = max(max_h, h)
|
||||
if reset_needed:
|
||||
reset_nhs()
|
||||
if w > max_w:
|
||||
width = 8 * w
|
||||
max_w = max(max_w, w)
|
||||
reset_nws()
|
||||
max_h = height
|
||||
max_w = width
|
||||
reset_needed = False
|
||||
else:
|
||||
if h > max_h:
|
||||
height = 8 * h
|
||||
max_h = max(max_h, h)
|
||||
reset_nhs()
|
||||
if w > max_w:
|
||||
width = 8 * w
|
||||
max_w = max(max_w, w)
|
||||
reset_nws()
|
||||
down_ratio = height // 8 // h
|
||||
curr_depth = round(math.log(down_ratio, 2))
|
||||
# scale-up the tile-size the deeper we go
|
||||
nh = max(1, nh // down_ratio)
|
||||
nw = max(1, nw // down_ratio)
|
||||
do_split = curr_depth <= depth and h % nh == 0 and w % nw == 0 and nh * nw > 1
|
||||
if do_split:
|
||||
x = rearrange(x, "b (nh h nw w) c -> (b nh nw) (h w) c", h=h // nh, w=w // nw, nh=nh, nw=nw)
|
||||
out = forward(x, *args[1:], **kwargs)
|
||||
if do_split:
|
||||
out = rearrange(out, "(b nh nw) hw c -> b nh nw hw c", nh=nh, nw=nw)
|
||||
out = rearrange(out, "b nh nw (h w) c -> b (nh h nw w) c", h=h // nh, w=w // nw)
|
||||
try:
|
||||
if do_split:
|
||||
x = rearrange(x, "b (nh h nw w) c -> (b nh nw) (h w) c", h=h // nh, w=w // nw, nh=nh, nw=nw)
|
||||
out = forward(x, *args[1:], **kwargs)
|
||||
if do_split:
|
||||
out = rearrange(out, "(b nh nw) hw c -> b nh nw hw c", nh=nh, nw=nw)
|
||||
out = rearrange(out, "b nh nw (h w) c -> b (nh h nw w) c", h=h // nh, w=w // nw)
|
||||
except Exception as e:
|
||||
global error_reported # pylint: disable=global-statement
|
||||
if not error_reported:
|
||||
error_reported = True
|
||||
log.error(f'Hypertile error: width={width} height={height} {e}')
|
||||
out = forward(x, *args[1:], **kwargs)
|
||||
return out
|
||||
return wrapper
|
||||
try: # hijack forward method and restore
|
||||
@@ -110,7 +127,8 @@ def split_attention(layer: nn.Module, tile_size: int=256, min_tile_size: int=256
|
||||
|
||||
|
||||
def context_hypertile_vae(p):
|
||||
global height, width, max_h, max_w # pylint: disable=global-statement
|
||||
global height, width, max_h, max_w, error_reported # pylint: disable=global-statement
|
||||
error_reported = False
|
||||
height=p.height
|
||||
width=p.width
|
||||
max_h = 0
|
||||
@@ -128,7 +146,8 @@ def context_hypertile_vae(p):
|
||||
|
||||
|
||||
def context_hypertile_unet(p):
|
||||
global height, width, max_h, max_w # pylint: disable=global-statement
|
||||
global height, width, max_h, max_w, error_reported # pylint: disable=global-statement
|
||||
error_reported = False
|
||||
height=p.height
|
||||
width=p.width
|
||||
max_h = 0
|
||||
@@ -143,3 +162,11 @@ def context_hypertile_unet(p):
|
||||
else:
|
||||
shared.log.info(f'Applying hypertile: unet={shared.opts.hypertile_unet_tile}')
|
||||
return split_attention(unet, tile_size=shared.opts.hypertile_unet_tile, min_tile_size=128, swap_size=1)
|
||||
|
||||
|
||||
def hypertile_set(p):
|
||||
global height, width, error_reported, reset_needed # pylint: disable=global-statement
|
||||
error_reported = False
|
||||
height=p.height
|
||||
width=p.width
|
||||
reset_needed = True
|
||||
|
||||
@@ -38,6 +38,31 @@ def apply_styles_to_prompt(prompt, styles):
|
||||
return prompt
|
||||
|
||||
|
||||
def apply_styles_to_extra(p, style: Style):
|
||||
if style is None:
|
||||
return
|
||||
name_map = {
|
||||
'sampler': 'sampler_name',
|
||||
}
|
||||
from modules.generation_parameters_copypaste import parse_generation_parameters
|
||||
extra = parse_generation_parameters(style.extra)
|
||||
extra.pop('Prompt', None)
|
||||
extra.pop('Negative prompt', None)
|
||||
fields = []
|
||||
for k, v in extra.items():
|
||||
k = k.lower()
|
||||
k = k.replace(' ', '_')
|
||||
if k in name_map: # rename some fields
|
||||
k = name_map[k]
|
||||
if hasattr(p, k):
|
||||
orig = getattr(p, k)
|
||||
if type(orig) != type(v) and orig is not None:
|
||||
v = type(orig)(v)
|
||||
setattr(p, k, v)
|
||||
fields.append(f'{k}={v}')
|
||||
log.debug(f'Applied style: {style.name} extra={fields}')
|
||||
|
||||
|
||||
class StyleDatabase:
|
||||
def __init__(self, opts):
|
||||
self.no_style = Style("None")
|
||||
@@ -119,6 +144,11 @@ class StyleDatabase:
|
||||
def apply_negative_styles_to_prompt(self, prompt, styles):
|
||||
return apply_styles_to_prompt(prompt, [self.find_style(x).negative_prompt for x in styles])
|
||||
|
||||
def apply_styles_to_extra(self, p):
|
||||
for style in p.styles:
|
||||
s = self.find_style(style)
|
||||
apply_styles_to_extra(p, s)
|
||||
|
||||
def save_styles(self, path, verbose=False):
|
||||
for name in list(self.styles):
|
||||
style = {
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import html
|
||||
import json
|
||||
from modules import shared, script_callbacks, extra_networks, ui_extra_networks, styles
|
||||
from modules import shared, extra_networks, ui_extra_networks, styles
|
||||
|
||||
|
||||
class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage):
|
||||
@@ -57,7 +57,7 @@ class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage):
|
||||
"description": '',
|
||||
"prompt": params.get('Prompt', ''),
|
||||
"negative": params.get('Negative prompt', ''),
|
||||
"extra": '', # TODO add extras to styles
|
||||
"extra": '',
|
||||
"local_preview": f"{name}.{shared.opts.samples_format}",
|
||||
}
|
||||
return item
|
||||
@@ -118,6 +118,7 @@ class ExtraNetworkStyles(extra_networks.ExtraNetwork):
|
||||
p.styles.append(style.name)
|
||||
p.prompts = [styles.merge_prompts(style.prompt, prompt) for prompt in p.prompts]
|
||||
p.negative_prompts = [styles.merge_prompts(style.negative_prompt, prompt) for prompt in p.negative_prompts]
|
||||
styles.apply_styles_to_extra(p, style)
|
||||
|
||||
|
||||
def deactivate(self, p):
|
||||
|
||||
Reference in New Issue
Block a user