Initial tiled sampling support

No tiled autoprompting yet
This commit is contained in:
kijai
2024-03-02 23:20:11 +02:00
parent baa5b286ad
commit 6713939b20
4 changed files with 316 additions and 28 deletions
+30 -17
View File
@@ -113,6 +113,7 @@ class SUPIRModel(DiffusionEngine):
self.sampler_config.params.guider_config.params.scale_min = cfg_scale
self.sampler_config.params.guider_config.params.scale = cfg_scale_start
else:
self.sampler_config.params.guider_config.params.scale_min = cfg_scale
self.sampler_config.params.guider_config.params.scale = cfg_scale
self.sampler_config.params.restore_cfg = restoration_scale
self.sampler_config.params.s_churn = s_churn
@@ -126,26 +127,10 @@ class SUPIRModel(DiffusionEngine):
_z = self.encode_first_stage_with_denoise(x, use_sample=False)
x_stage1 = self.decode_first_stage(_z)
# x_stage1 = interpolate(x_stage1, scale_factor=scale_factor, mode='bilinear', antialias=True)
# _z = self.encode_first_stage_with_denoise(x_stage1)
z_stage1 = self.encode_first_stage(x_stage1)
batch = {}
batch['txt'] = [''.join([_p, p_p]) for _p in p]
batch['original_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(x.device)
batch['crop_coords_top_left'] = torch.tensor([0, 0]).repeat(N, 1).to(x.device)
batch['target_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(x.device)
batch['aesthetic_score'] = torch.tensor([9.0]).repeat(N, 1).to(x.device)
batch['control'] = _z
batch_uc = copy.deepcopy(batch)
batch_uc['txt'] = [n_p for _ in p]
#with torch.cuda.amp.autocast(dtype=self.ae_dtype):
autocast_condition = (self.model.dtype == torch.float16 or self.model.dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.model.dtype) if autocast_condition else nullcontext():
c, uc = self.conditioner.get_unconditional_conditioning(batch, batch_uc)
c, uc = self.prepare_condition(_z, p, p_p, n_p, N)
denoiser = lambda input, sigma, c, control_scale: self.denoiser(
self.model, input, sigma, c, control_scale, **kwargs
@@ -176,6 +161,34 @@ class SUPIRModel(DiffusionEngine):
self.first_stage_model.decoder, decoder_tile_size, is_decoder=True, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
def prepare_condition(self, _z, p, p_p, n_p, N):
batch = {}
batch['original_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(_z.device)
batch['crop_coords_top_left'] = torch.tensor([0, 0]).repeat(N, 1).to(_z.device)
batch['target_size_as_tuple'] = torch.tensor([1024, 1024]).repeat(N, 1).to(_z.device)
batch['aesthetic_score'] = torch.tensor([9.0]).repeat(N, 1).to(_z.device)
batch['control'] = _z
batch_uc = copy.deepcopy(batch)
batch_uc['txt'] = [n_p for _ in p]
autocast_condition = (self.model.dtype == torch.float16 or self.model.dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device)
if not isinstance(p[0], list):
batch['txt'] = [''.join([_p, p_p]) for _p in p]
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.model.dtype) if autocast_condition else nullcontext():
c, uc = self.conditioner.get_unconditional_conditioning(batch, batch_uc)
else:
assert len(p) == 1, 'Support bs=1 only for local prompt conditioning.'
p_tiles = p[0]
c = []
for i, p_tile in enumerate(p_tiles):
batch['txt'] = [''.join([p_tile, p_p])]
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.model.dtype) if autocast_condition else nullcontext():
if i == 0:
_c, uc = self.conditioner.get_unconditional_conditioning(batch, batch_uc)
else:
_c, _ = self.conditioner.get_unconditional_conditioning(batch, None)
c.append(_c)
return c, uc
# if __name__ == '__main__':
# from SUPIR.util import create_model, load_state_dict
+19 -2
View File
@@ -22,6 +22,7 @@ class SUPIR_Upscale:
self.current_diffusion_dtype = None
self.current_encoder_dtype = None
self.tiled_vae_state = None
self.tiled_sampling_state = None
upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"]
@@ -77,6 +78,9 @@ class SUPIR_Upscale:
"default": 'auto'
}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 128, "step": 1}),
"use_tiled_sampling": ("BOOLEAN", {"default": False}),
"sampler_tile_size": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 32}),
"sampler_tile_stride": ("INT", {"default": 512, "min": 32, "max": 2048, "step": 32}),
}
}
@@ -89,7 +93,7 @@ class SUPIR_Upscale:
def process(self, steps, image, color_fix_type, seed, scale_by, cfg_scale, resize_method, s_churn, s_noise,
encoder_tile_size_pixels, decoder_tile_size_latent,
control_scale, cfg_scale_start, control_scale_start, restoration_scale, keep_model_loaded,
a_prompt, n_prompt, sdxl_model, supir_model, use_tiled_vae, captions="", diffusion_dtype="auto",
a_prompt, n_prompt, sdxl_model, supir_model, use_tiled_vae, use_tiled_sampling=False, sampler_tile_size=128, sampler_tile_stride=64, captions="", diffusion_dtype="auto",
encoder_dtype="auto", batch_size=1):
device = mm.get_torch_device()
@@ -99,6 +103,7 @@ class SUPIR_Upscale:
SDXL_MODEL_PATH = folder_paths.get_full_path("checkpoints", sdxl_model)
config_path = os.path.join(script_directory, "options/SUPIR_v0.yaml")
config_path_tiled = os.path.join(script_directory, "options/SUPIR_v0_tiled.yaml")
if diffusion_dtype == 'auto':
try:
@@ -135,16 +140,28 @@ class SUPIR_Upscale:
vae_dtype = encoder_dtype
print(f"Encoder using using {vae_dtype}")
if not hasattr(self, "model") or self.model is None or self.current_sdxl_model != sdxl_model or self.current_diffusion_dtype != diffusion_dtype or self.current_encoder_dtype != encoder_dtype or self.tiled_vae_state != use_tiled_vae or self.current_supir_model != supir_model:
if not hasattr(self, "model") or self.model is None or self.current_sdxl_model != sdxl_model or self.current_diffusion_dtype != diffusion_dtype or self.current_encoder_dtype != encoder_dtype or self.tiled_vae_state != use_tiled_vae or self.current_supir_model != supir_model or self.tiled_sampling_state != use_tiled_sampling:
self.model = None
mm.soft_empty_cache()
self.current_diffusion_dtype = diffusion_dtype
self.current_encoder_dtype = encoder_dtype
self.current_sdxl_model = sdxl_model
self.current_supir_model = supir_model
if use_tiled_sampling:
self.tiled_sampling_state = True
config = OmegaConf.load(config_path_tiled)
config.model.params.sampler_config.params.tile_size = sampler_tile_size // 8
config.model.params.sampler_config.params.tile_stride = sampler_tile_stride // 8
print("Using tiled sampling")
else:
self.tiled_sampling_state = False
config = OmegaConf.load(config_path)
print("Using non-tiled sampling")
config.model.params.ae_dtype = vae_dtype
config.model.params.diffusion_dtype = model_dtype
self.model = instantiate_from_config(config.model).cpu()
try:
print(f'Attempting to load SUPIR model: [{SUPIR_MODEL_PATH}]')
+158
View File
@@ -0,0 +1,158 @@
model:
target: .SUPIR.models.SUPIR_model.SUPIRModel
params:
ae_dtype: bf16
diffusion_dtype: fp16
scale_factor: 0.13025
disable_first_stage_autocast: True
network_wrapper: .sgm.modules.diffusionmodules.wrappers.ControlWrapper
denoiser_config:
target: .sgm.modules.diffusionmodules.denoiser.DiscreteDenoiserWithControl
params:
num_idx: 1000
weighting_config:
target: .sgm.modules.diffusionmodules.denoiser_weighting.EpsWeighting
scaling_config:
target: .sgm.modules.diffusionmodules.denoiser_scaling.EpsScaling
discretization_config:
target: .sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization
control_stage_config:
target: .SUPIR.modules.SUPIR_v0.GLVControl
params:
adm_in_channels: 2816
num_classes: sequential
use_checkpoint: True
in_channels: 4
out_channels: 4
model_channels: 320
attention_resolutions: [4, 2]
num_res_blocks: 2
channel_mult: [1, 2, 4]
num_head_channels: 64
use_spatial_transformer: True
use_linear_in_transformer: True
transformer_depth: [1, 2, 10] # note: the first is unused (due to attn_res starting at 2) 32, 16, 8 --> 64, 32, 16
# transformer_depth: [1, 1, 4]
context_dim: 2048
spatial_transformer_attn_type: softmax-xformers
legacy: False
input_upscale: 1
network_config:
target: .SUPIR.modules.SUPIR_v0.LightGLVUNet
params:
mode: XL-base
project_type: ZeroSFT
project_channel_scale: 2
adm_in_channels: 2816
num_classes: sequential
use_checkpoint: True
in_channels: 4
out_channels: 4
model_channels: 320
attention_resolutions: [4, 2]
num_res_blocks: 2
channel_mult: [1, 2, 4]
num_head_channels: 64
use_spatial_transformer: True
use_linear_in_transformer: True
transformer_depth: [1, 2, 10] # note: the first is unused (due to attn_res starting at 2) 32, 16, 8 --> 64, 32, 16
context_dim: 2048
spatial_transformer_attn_type: softmax-xformers
legacy: False
conditioner_config:
target: .sgm.modules.GeneralConditionerWithControl
params:
emb_models:
# crossattn cond
- is_trainable: False
input_key: txt
target: .sgm.modules.encoders.modules.FrozenCLIPEmbedder
params:
layer: hidden
layer_idx: 11
# crossattn and vector cond
- is_trainable: False
input_key: txt
target: .sgm.modules.encoders.modules.FrozenOpenCLIPEmbedder2
params:
arch: ViT-bigG-14
version: laion2b_s39b_b160k
freeze: True
layer: penultimate
always_return_pooled: True
legacy: False
# vector cond
- is_trainable: False
input_key: original_size_as_tuple
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
params:
outdim: 256 # multiplied by two
# vector cond
- is_trainable: False
input_key: crop_coords_top_left
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
params:
outdim: 256 # multiplied by two
# vector cond
- is_trainable: False
input_key: target_size_as_tuple
target: .sgm.modules.encoders.modules.ConcatTimestepEmbedderND
params:
outdim: 256 # multiplied by two
first_stage_config:
target: .sgm.models.autoencoder.AutoencoderKLInferenceWrapper
params:
ckpt_path: ~
embed_dim: 4
monitor: val/rec_loss
ddconfig:
attn_type: vanilla-xformers
double_z: true
z_channels: 4
resolution: 256
in_channels: 3
out_ch: 3
ch: 128
ch_mult: [ 1, 2, 4, 4 ]
num_res_blocks: 2
attn_resolutions: [ ]
dropout: 0.0
lossconfig:
target: torch.nn.Identity
sampler_config:
target: .sgm.modules.diffusionmodules.sampling.TiledRestoreEDMSampler
params:
num_steps: 100
restore_cfg: 4.0
s_churn: 0
s_noise: 1.003
tile_size: 128
tile_stride: 64
discretization_config:
target: .sgm.modules.diffusionmodules.discretizer.LegacyDDPMDiscretization
guider_config:
target: .sgm.modules.diffusionmodules.guiders.LinearCFG
params:
scale: 7.5
scale_min: 4.0
p_p:
'Cinematic, High Contrast, highly detailed, taken using a Canon EOS R camera,
hyper detailed photo - realistic maximum detail, 32k, Color Grading, ultra HD, extreme meticulous detailing,
skin pore detailing, hyper sharpness, perfect without deformations.'
n_p:
'painting, oil painting, illustration, drawing, art, sketch, oil painting, cartoon, CG Style, 3D render,
unreal engine, blurring, dirty, messy, worst quality, low quality, frames, watermark, signature,
jpeg artifacts, deformed, lowres, over-smooth'
SDXL_CKPT: /opt/data/private/AIGC_pretrain/SDXL_cache/sd_xl_base_1.0_0.9vae.safetensors
SUPIR_CKPT_F: /opt/data/private/AIGC_pretrain/SUPIR_cache/SUPIR-v0F.ckpt
SUPIR_CKPT_Q: /opt/data/private/AIGC_pretrain/SUPIR_cache/SUPIR-v0Q.ckpt
SUPIR_CKPT: ~
+108 -8
View File
@@ -366,6 +366,15 @@ class DPMPP2MSampler(BaseDiffusionSampler):
return x
def to_d_center(denoised, x_center, x):
b = denoised.shape[0]
v_center = (denoised - x_center).view(b, -1)
v_denoise = (x - denoised).view(b, -1)
d_center = v_center - v_denoise * (v_center * v_denoise).sum(dim=1).view(b, 1) / \
(v_denoise * v_denoise).sum(dim=1).view(b, 1)
d_center = d_center / d_center.view(x.shape[0], -1).norm(dim=1).view(-1, 1)
return d_center.view(denoised.shape)
import comfy.utils
class RestoreEDMSampler(SingleStepDiffusionSampler):
@@ -441,11 +450,102 @@ class RestoreEDMSampler(SingleStepDiffusionSampler):
pbar_comfy.update(1)
return x
def to_d_center(denoised, x_center, x):
b = denoised.shape[0]
v_center = (denoised - x_center).view(b, -1)
v_denoise = (x - denoised).view(b, -1)
d_center = v_center - v_denoise * (v_center * v_denoise).sum(dim=1).view(b, 1) / \
(v_denoise * v_denoise).sum(dim=1).view(b, 1)
d_center = d_center / d_center.view(x.shape[0], -1).norm(dim=1).view(-1, 1)
return d_center.view(denoised.shape)
class TiledRestoreEDMSampler(RestoreEDMSampler):
def __init__(self, tile_size=128, tile_stride=64, *args, **kwargs):
super().__init__(*args, **kwargs)
self.tile_size = tile_size
self.tile_stride = tile_stride
self.tile_weights = gaussian_weights(self.tile_size, self.tile_size, 1)
def __call__(self, denoiser, x, cond, uc=None, num_steps=None, x_center=None, control_scale=1.0,
use_linear_control_scale=False, control_scale_start=0.0):
use_local_prompt = isinstance(cond, list)
b, _, h, w = x.shape
latent_tiles_iterator = _sliding_windows(h, w, self.tile_size, self.tile_stride)
tile_weights = self.tile_weights.repeat(b, 1, 1, 1)
if not use_local_prompt:
LQ_latent = cond['control']
else:
assert len(cond) == len(latent_tiles_iterator), "Number of local prompts should be equal to number of tiles"
LQ_latent = cond[0]['control']
clean_LQ_latent = x_center
x, s_in, sigmas, num_sigmas, cond, uc = self.prepare_sampling_loop(
x, cond, uc, num_steps
)
pbar_comfy = comfy.utils.ProgressBar(num_sigmas)
for _idx, i in enumerate(self.get_sigma_gen(num_sigmas)):
gamma = (
min(self.s_churn / (num_sigmas - 1), 2**0.5 - 1)
if self.s_tmin <= sigmas[i] <= self.s_tmax
else 0.0
)
x_next = torch.zeros_like(x)
count = torch.zeros_like(x)
eps_noise = torch.randn_like(x)
for j, (hi, hi_end, wi, wi_end) in enumerate(latent_tiles_iterator):
x_tile = x[:, :, hi:hi_end, wi:wi_end]
_eps_noise = eps_noise[:, :, hi:hi_end, wi:wi_end]
x_center_tile = clean_LQ_latent[:, :, hi:hi_end, wi:wi_end]
if use_local_prompt:
_cond = cond[j]
else:
_cond = cond
_cond['control'] = LQ_latent[:, :, hi:hi_end, wi:wi_end]
uc['control'] = LQ_latent[:, :, hi:hi_end, wi:wi_end]
_x = self.sampler_step(
s_in * sigmas[i],
s_in * sigmas[i + 1],
denoiser,
x_tile,
_cond,
uc,
gamma,
x_center_tile,
eps_noise=_eps_noise,
control_scale=control_scale,
use_linear_control_scale=use_linear_control_scale,
control_scale_start=control_scale_start,
)
x_next[:, :, hi:hi_end, wi:wi_end] += _x * tile_weights
count[:, :, hi:hi_end, wi:wi_end] += tile_weights
x_next /= count
x = x_next
pbar_comfy.update(1)
return x
def gaussian_weights(tile_width, tile_height, nbatches):
"""Generates a gaussian mask of weights for tile contributions"""
from numpy import pi, exp, sqrt
import numpy as np
latent_width = tile_width
latent_height = tile_height
var = 0.01
midpoint = (latent_width - 1) / 2 # -1 because index goes from 0 to latent_width - 1
x_probs = [exp(-(x - midpoint) * (x - midpoint) / (latent_width * latent_width) / (2 * var)) / sqrt(2 * pi * var)
for x in range(latent_width)]
midpoint = latent_height / 2
y_probs = [exp(-(y - midpoint) * (y - midpoint) / (latent_height * latent_height) / (2 * var)) / sqrt(2 * pi * var)
for y in range(latent_height)]
weights = np.outer(y_probs, x_probs)
return torch.tile(torch.tensor(weights, device='cuda'), (nbatches, 4, 1, 1))
def _sliding_windows(h: int, w: int, tile_size: int, tile_stride: int):
hi_list = list(range(0, h - tile_size + 1, tile_stride))
if (h - tile_size) % tile_stride != 0:
hi_list.append(h - tile_size)
wi_list = list(range(0, w - tile_size + 1, tile_stride))
if (w - tile_size) % tile_stride != 0:
wi_list.append(w - tile_size)
coords = []
for hi in hi_list:
for wi in wi_list:
coords.append((hi, hi + tile_size, wi, wi + tile_size))
return coords