Initial tiled sampling support
No tiled autoprompting yet
This commit is contained in:
+30
-17
@@ -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
|
||||
|
||||
@@ -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}]')
|
||||
|
||||
@@ -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: ~
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user