diff --git a/CHANGELOG.md b/CHANGELOG.md index e32d4e3fa..03d988c36 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2026-07-16 +## Update for 2026-07-18 - **Compute** - torch: update to `2.13.0` for CUDA, ROCm, IPEX @@ -10,9 +10,12 @@ - sdnq quantization optimizations - sdnq attention optimizations - sdnq separate dit/te settings +- **Features** + - SeedVR enhanced support - **Fixes** - upscaler auto-refresh to catch chainner upscalers that are not loaded on first attempt - - lora loader support diffusers trainer + - lora support diffusers trainer + - flux1 load t5 ## Update for 2026-07-14 diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 07890f0d7..b98735978 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 07890f0d7b1412634fb2bbde45f326b8b3760711 +Subproject commit b987359783ddd8e217f1727f7505486db6d76a73 diff --git a/modules/postprocess/seedvr_model.py b/modules/postprocess/seedvr_model.py index 39ad9c3bd..e419745cc 100644 --- a/modules/postprocess/seedvr_model.py +++ b/modules/postprocess/seedvr_model.py @@ -28,6 +28,8 @@ class UpscalerSeedVR(Upscaler): ] self.model = None self.model_loaded = None + self.tile_size = 1024 + self.tile_overlap = 0.25 self.device = devices.device def load_model(self, path: str): @@ -60,8 +62,9 @@ class UpscalerSeedVR(Upscaler): } t1 = time.time() self.model.dit.config = self.model.config.dit - self.model.vae.tile_sample_min_size = 1024 - self.model.vae.tile_latent_min_size = 128 + self.model.vae.tile_sample_min_size = self.tile_size + self.model.vae.tile_latent_min_size = self.tile_size // 8 + self.model.vae.tile_overlap_factor = self.tile_overlap self.model = do_post_load_quant(self.model, allow=True) @@ -136,26 +139,34 @@ class UpscalerSeedVR(Upscaler): devices.torch_gc() return result - def do_upscale(self, img: Image.Image, selected_file): + def do_upscale(self, img: Image.Image, selected_file, cfg_scale: float = 3.5, cfg_rescale: float = 0.0, steps: int = 1, seed: int = -1, scale: float | None = None, tile_size: int = 1024, tile_overlap: float = 0.25): self.load_model(selected_file) if self.model is None: return img from modules.seedvr.src.core import generation + self.scale = self.scale if scale is None else scale + self.tile_size = tile_size if tile_size is not None else self.tile_size + self.tile_overlap = tile_overlap if tile_overlap is not None else self.tile_overlap + self.model.vae.tile_sample_min_size = self.tile_size + self.model.vae.tile_latent_min_size = self.tile_size // 8 + self.model.vae.tile_overlap_factor = self.tile_overlap width = int(self.scale * img.width) // 8 * 8 image_tensor = np.array(img) image_tensor = torch.from_numpy(image_tensor).to(device=devices.device, dtype=devices.dtype).unsqueeze(0) / 255.0 random.seed() - seed = int(random.randrange(4294967294)) + seed = int(random.randrange(4294967294)) if seed == -1 else int(seed) t0 = time.time() with devices.inference_context(): result_tensor = generation.generation_loop( runner=self.model, images=image_tensor, - cfg_scale=opts.seedvr_cfg_scale, + cfg_scale=cfg_scale, + cfg_rescale=cfg_rescale, + steps=steps, seed=seed, res_w=width, batch_size=1, @@ -163,7 +174,8 @@ class UpscalerSeedVR(Upscaler): device=devices.device, ) t1 = time.time() - log.info(f'Upscaler: type="{self.name}" model="{selected_file}" scale={self.scale} cfg={opts.seedvr_cfg_scale} seed={seed} time={t1 - t0:.2f}') + tiles = getattr(self.model.vae, "tiles", None) + log.info(f'Upscaler: type="{self.name}" model="{selected_file}" scale={self.scale} cfg={cfg_scale} seed={seed} tiles={tiles} time={t1 - t0:.2f}') img = convert.to_pil(result_tensor.squeeze()) if opts.upscaler_unload: diff --git a/modules/seedvr/src/common/diffusion/samplers/__init__.py b/modules/seedvr/src/common/diffusion/samplers/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/modules/seedvr/src/common/diffusion/samplers/euler.py b/modules/seedvr/src/common/diffusion/samplers/euler.py index d4892a3f1..1f2da9c5a 100644 --- a/modules/seedvr/src/common/diffusion/samplers/euler.py +++ b/modules/seedvr/src/common/diffusion/samplers/euler.py @@ -18,16 +18,10 @@ Euler ODE solver. """ from typing import Callable +import itertools import torch -from einops import rearrange -from torch.nn import functional as F - -#from ....models.dit_v2 import na - -from ..types import PredictionType from ..utils import expand_dims from .base import Sampler, SamplerModelArgs -import itertools class EulerSampler(Sampler): diff --git a/modules/seedvr/src/core/generation.py b/modules/seedvr/src/core/generation.py index 06f40e872..569250c3b 100644 --- a/modules/seedvr/src/core/generation.py +++ b/modules/seedvr/src/core/generation.py @@ -106,7 +106,7 @@ def cut_videos(videos): return result -def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_size=90, temporal_overlap=0, progress_callback=None, device:str='cpu'): +def generation_loop(runner, images, cfg_scale=1.0, cfg_rescale=0.0, steps=1, seed=666, res_w=720, batch_size=90, temporal_overlap=0, progress_callback=None, device:str='cpu'): """ Main generation loop with context-aware temporal processing @@ -137,9 +137,9 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si # Configure classifier-free guidance runner.config.diffusion.cfg.scale = cfg_scale - runner.config.diffusion.cfg.rescale = 0.0 + runner.config.diffusion.cfg.rescale = cfg_rescale # Configure sampling steps - runner.config.diffusion.timesteps.sampling.steps = 1 + runner.config.diffusion.timesteps.sampling.steps = steps runner.configure_diffusion() # Set random seed diff --git a/modules/seedvr/src/core/infer.py b/modules/seedvr/src/core/infer.py index 9e95cc02d..f18ee433f 100644 --- a/modules/seedvr/src/core/infer.py +++ b/modules/seedvr/src/core/infer.py @@ -309,10 +309,7 @@ class SeedVRPipeline(): timestep=args.t.repeat(batch_size), ).vid_sample, scale=( - cfg_scale - if (args.i + 1) / len(self.sampler.timesteps) - <= self.config.diffusion.cfg.get("partial", 1) - else 1.0 + cfg_scale if (args.i + 1) / len(self.sampler.timesteps) <= self.config.diffusion.cfg.get("partial", 1) else 1.0 ), rescale=self.config.diffusion.cfg.rescale, ), diff --git a/modules/seedvr/src/data/image/transforms/side_resize.py b/modules/seedvr/src/data/image/transforms/side_resize.py index 6e07402b2..76061e248 100644 --- a/modules/seedvr/src/data/image/transforms/side_resize.py +++ b/modules/seedvr/src/data/image/transforms/side_resize.py @@ -50,5 +50,6 @@ class SideResize: size = min(width, height) else: size = self.size - - return TVF.resize(image, size, self.interpolation) + size_w = int(size) // 8 * 8 + size_h = int(size * height / width) // 8 * 8 + return TVF.resize(image, (size_h, size_w), self.interpolation) diff --git a/modules/seedvr/src/models/dit/patch.py b/modules/seedvr/src/models/dit/patch.py index 570e3ae40..c273d0d5b 100644 --- a/modules/seedvr/src/models/dit/patch.py +++ b/modules/seedvr/src/models/dit/patch.py @@ -42,6 +42,8 @@ class PatchIn(nn.Module): ) -> torch.Tensor: t, h, w = self.patch_size vid = rearrange(vid, "b c (T t) (H h) (W w) -> b T H W (t h w c)", t=t, h=h, w=w) + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) return vid @@ -63,30 +65,40 @@ class PatchOut(nn.Module): vid: torch.Tensor, ) -> torch.Tensor: t, h, w = self.patch_size + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) vid = rearrange(vid, "b T H W (t h w c) -> b c (T t) (H h) (W w)", t=t, h=h, w=w) return vid class NaPatchIn(PatchIn): - def forward( + def forward( # pylint: disable=arguments-differ self, vid: torch.Tensor, # l c vid_shape: torch.LongTensor, ) -> torch.Tensor: t, h, w = self.patch_size - if not (t == h == w == 1): + if not t == h == w == 1: vid, vid_shape = na.rearrange( vid, vid_shape, "(T t) (H h) (W w) c -> T H W (t h w c)", t=t, h=h, w=w ) + for i in range(len(vid)): + if h > 1 and vid_shape[i, 1] % h != 0: + vid[i] = torch.cat([vid[i][:, :1]] * (h - vid[i].size(1) % h) + [vid[i]], dim=1) + if w > 1 and vid_shape[i, 2] % w != 0: + vid[i] = torch.cat([vid[i][:, :, :1]] * (w - vid[i].size(2) % w) + [vid[i]], dim=2) + vid, vid_shape = na.flatten(vid) # slice vid after patching in when using sequence parallelism vid = slice_inputs(vid, dim=0) + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) return vid, vid_shape class NaPatchOut(PatchOut): - def forward( + def forward( # pylint: disable=arguments-differ self, vid: torch.FloatTensor, # l c vid_shape: torch.LongTensor, @@ -96,8 +108,10 @@ class NaPatchOut(PatchOut): torch.LongTensor, ]: t, h, w = self.patch_size + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) - # gather vid before patching out when enabling sequence parallelism + # gather vid before patchting out when enabling sequence parallelism vid = gather_outputs( vid, gather_dim=0, @@ -105,8 +119,13 @@ class NaPatchOut(PatchOut): unpad_shape=vid_shape, cache=cache.namespace("vid"), ) - if not (t == h == w == 1): + if not t == h == w == 1: vid, vid_shape = na.rearrange( vid, vid_shape, "T H W (t h w c) -> (T t) (H h) (W w) c", t=t, h=h, w=w ) + for i in range(len(vid)): + if h > 1 and vid_shape[i, 1] % h != 0: + vid[i] = vid[i][:, (h - vid_shape[i, 1] % h) :] + if w > 1 and vid_shape[i, 2] % w != 0: + vid[i] = vid[i][:, :, (w - vid_shape[i, 2] % w) :] return vid, vid_shape diff --git a/modules/seedvr/src/models/dit_v2/patch/patch_v1.py b/modules/seedvr/src/models/dit_v2/patch/patch_v1.py index 14b5a62d5..8a0f5fc88 100644 --- a/modules/seedvr/src/models/dit_v2/patch/patch_v1.py +++ b/modules/seedvr/src/models/dit_v2/patch/patch_v1.py @@ -45,6 +45,8 @@ class PatchIn(nn.Module): assert vid.size(2) % t == 1 vid = torch.cat([vid[:, :, :1]] * (t - 1) + [vid], dim=2) vid = rearrange(vid, "b c (T t) (H h) (W w) -> b T H W (t h w c)", t=t, h=h, w=w) + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) return vid @@ -83,16 +85,22 @@ class NaPatchIn(PatchIn): cache = cache.namespace("patch") vid_shape_before_patchify = cache("vid_shape_before_patchify", lambda: vid_shape) t, h, w = self.patch_size - if not (t == h == w == 1): + if not t == h == w == 1: vid = na.unflatten(vid, vid_shape) for i in range(len(vid)): if t > 1 and vid_shape_before_patchify[i, 0] % t != 0: vid[i] = torch.cat([vid[i][:1]] * (t - vid[i].size(0) % t) + [vid[i]], dim=0) + if h > 1 and vid_shape_before_patchify[i, 1] % h != 0: + vid[i] = torch.cat([vid[i][:, :1]] * (h - vid[i].size(1) % h) + [vid[i]], dim=1) + if w > 1 and vid_shape_before_patchify[i, 2] % w != 0: + vid[i] = torch.cat([vid[i][:, :, :1]] * (w - vid[i].size(2) % w) + [vid[i]], dim=2) vid[i] = rearrange(vid[i], "(T t) (H h) (W w) c -> T H W (t h w c)", t=t, h=h, w=w) vid, vid_shape = na.flatten(vid) # slice vid after patching in when using sequence parallelism vid = slice_inputs(vid, dim=0) + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) return vid, vid_shape @@ -111,17 +119,23 @@ class NaPatchOut(PatchOut): vid_shape_before_patchify = cache.get("vid_shape_before_patchify") t, h, w = self.patch_size + if vid.dtype != self.proj.weight.dtype: + vid = vid.to(self.proj.weight.dtype) vid = self.proj(vid) # gather vid before patching out when enabling sequence parallelism vid = gather_outputs( vid, gather_dim=0, padding_dim=0, unpad_shape=vid_shape, cache=cache.namespace("vid") ) - if not (t == h == w == 1): + if not t == h == w == 1: vid = na.unflatten(vid, vid_shape) for i in range(len(vid)): vid[i] = rearrange(vid[i], "T H W (t h w c) -> (T t) (H h) (W w) c", t=t, h=h, w=w) if t > 1 and vid_shape_before_patchify[i, 0] % t != 0: vid[i] = vid[i][(t - vid_shape_before_patchify[i, 0] % t) :] + if h > 1 and vid_shape_before_patchify[i, 1] % h != 0: + vid[i] = vid[i][:, (h - vid_shape_before_patchify[i, 1] % h) :] + if w > 1 and vid_shape_before_patchify[i, 2] % w != 0: + vid[i] = vid[i][:, :, (w - vid_shape_before_patchify[i, 2] % w) :] vid, vid_shape = na.flatten(vid) return vid, vid_shape diff --git a/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py b/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py index 26f85d882..f5effd8b8 100644 --- a/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py +++ b/modules/seedvr/src/models/video_vae_v3/modules/attn_video_vae.py @@ -1044,7 +1044,7 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): norm_num_groups: int = 32, sample_size: int = 32, scaling_factor: float = 0.18215, - force_upcast: float = True, + force_upcast: float = False, attention: bool = True, temporal_scale_num: int = 0, slicing_up_num: int = 0, @@ -1244,12 +1244,14 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) row_limit = self.tile_latent_min_size - blend_extent rows = [] + self.tiles = 0 for i in range(0, x.shape[3], overlap_size): row = [] for j in range(0, x.shape[4], overlap_size): tile = x[:, :, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] tile = self._encode(tile) row.append(tile) + self.tiles += 1 rows.append(row) result_rows = [] for i, row in enumerate(rows): @@ -1329,6 +1331,7 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL): self.spatial_downsample_factor = spatial_downsample_factor self.temporal_downsample_factor = temporal_downsample_factor self.freeze_encoder = freeze_encoder + self.freeze_encoder = True super().__init__(*args, **kwargs) def forward(self, x: torch.FloatTensor) -> CausalAutoencoderOutput: diff --git a/modules/seedvr/src/models/video_vae_v3/modules/video_vae.py.old b/modules/seedvr/src/models/video_vae_v3/modules/video_vae.py.old deleted file mode 100644 index 2696c32ca..000000000 --- a/modules/seedvr/src/models/video_vae_v3/modules/video_vae.py.old +++ /dev/null @@ -1,936 +0,0 @@ -# Copyright (c) 2023 HuggingFace Team -# Copyright (c) 2025 ByteDance Ltd. and/or its affiliates. -# SPDX-License-Identifier: Apache License, Version 2.0 (the "License") -# -# This file has been modified by ByteDance Ltd. and/or its affiliates. on 1st June 2025 -# -# Original file was released under Apache License, Version 2.0 (the "License"), with the full license text -# available at http://www.apache.org/licenses/LICENSE-2.0. -# -# This modified file is released under the same license. - -from contextlib import nullcontext -from typing import Optional, Tuple, Literal, Callable, Union -import torch -import torch.nn as nn -from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution -from einops import rearrange -from ....common.half_precision_fixes import safe_pad_operation -from ....common.logger import get_logger -from .causal_inflation_lib import InflatedCausalConv3d, causal_norm_wrapper, init_causal_conv3d, remove_head -from .context_parallel_lib import causal_conv_gather_outputs, causal_conv_slice_inputs -from .global_config import set_norm_limit -from .types import CausalAutoencoderOutput, CausalDecoderOutput, CausalEncoderOutput, MemoryState, _inflation_mode_t, _memory_device_t, _receptive_field_t, _selective_checkpointing_t - - -logger = get_logger(__name__) # pylint: disable=invalid-name - -# Fake func, no checkpointing is required for inference -def gradient_checkpointing(module: Union[Callable, nn.Module], *args, enabled: bool, **kwargs): - return module(*args, **kwargs) - -class ResnetBlock2D(nn.Module): - r""" - A Resnet block. - - Parameters: - in_channels (`int`): The number of channels in the input. - out_channels (`int`, *optional*, default to be `None`): - The number of output channels for the first conv2d layer. - If None, same as `in_channels`. - dropout (`float`, *optional*, defaults to `0.0`): The dropout probability to use. - """ - - def __init__( - self, *, in_channels: int, out_channels: Optional[int] = None, dropout: float = 0.0 - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - - self.nonlinearity = nn.SiLU() - - self.norm1 = torch.nn.GroupNorm( - num_groups=32, num_channels=in_channels, eps=1e-6, affine=True - ) - - self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) - - self.norm2 = torch.nn.GroupNorm( - num_groups=32, num_channels=out_channels, eps=1e-6, affine=True - ) - - self.dropout = torch.nn.Dropout(dropout) - self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) - - self.use_in_shortcut = self.in_channels != out_channels - - self.conv_shortcut = None - if self.use_in_shortcut: - self.conv_shortcut = nn.Conv2d( - in_channels, out_channels, kernel_size=1, stride=1, padding=0 - ) - - def forward(self, input_tensor: torch.Tensor) -> torch.Tensor: - hidden = input_tensor - - hidden = self.norm1(hidden) - hidden = self.nonlinearity(hidden) - hidden = self.conv1(hidden) - - hidden = self.norm2(hidden) - hidden = self.nonlinearity(hidden) - hidden = self.dropout(hidden) - hidden = self.conv2(hidden) - - if self.conv_shortcut is not None: - input_tensor = self.conv_shortcut(input_tensor) - - output_tensor = input_tensor + hidden - - return output_tensor - -class Upsample3D(nn.Module): - """A 3D upsampling layer.""" - - def __init__( - self, - channels: int, - inflation_mode: _inflation_mode_t = "tail", - temporal_up: bool = False, - spatial_up: bool = True, - slicing: bool = False, - ): - super().__init__() - self.channels = channels - self.conv = init_causal_conv3d( - self.channels, self.channels, kernel_size=3, padding=1, inflation_mode=inflation_mode - ) - - self.temporal_up = temporal_up - self.spatial_up = spatial_up - self.temporal_ratio = 2 if temporal_up else 1 - self.spatial_ratio = 2 if spatial_up else 1 - self.slicing = slicing - - upscale_ratio = (self.spatial_ratio**2) * self.temporal_ratio - self.upscale_conv = nn.Conv3d( - self.channels, self.channels * upscale_ratio, kernel_size=1, padding=0 - ) - identity = ( - torch.eye(self.channels).repeat(upscale_ratio, 1).reshape_as(self.upscale_conv.weight) - ) - - self.upscale_conv.weight.data.copy_(identity) - nn.init.zeros_(self.upscale_conv.bias) - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.FloatTensor, - memory_state: MemoryState, - ) -> torch.FloatTensor: - return gradient_checkpointing( - self.custom_forward, - hidden_states, - memory_state, - enabled=self.training and self.gradient_checkpointing, - ) - - def custom_forward( - self, - hidden_states: torch.FloatTensor, - memory_state: MemoryState, - ) -> torch.FloatTensor: - assert hidden_states.shape[1] == self.channels - - if self.slicing: - split_size = hidden_states.size(2) // 2 - hidden_states = list( - hidden_states.split([split_size, hidden_states.size(2) - split_size], dim=2) - ) - else: - hidden_states = [hidden_states] - - for i in range(len(hidden_states)): - hidden_states[i] = self.upscale_conv(hidden_states[i]) - hidden_states[i] = rearrange( - hidden_states[i], - "b (x y z c) f h w -> b c (f z) (h x) (w y)", - x=self.spatial_ratio, - y=self.spatial_ratio, - z=self.temporal_ratio, - ) - - # [Overridden] For causal temporal conv - if self.temporal_up and memory_state != MemoryState.ACTIVE: - hidden_states[0] = remove_head(hidden_states[0]) - - if self.slicing: - hidden_states = self.conv(hidden_states, memory_state=memory_state) - return torch.cat(hidden_states, dim=2) - else: - return self.conv(hidden_states[0], memory_state=memory_state) - - -class Downsample3D(nn.Module): - """A 3D downsampling layer.""" - - def __init__( - self, - channels: int, - inflation_mode: _inflation_mode_t = "tail", - temporal_down: bool = False, - spatial_down: bool = True, - ): - super().__init__() - self.channels = channels - self.temporal_down = temporal_down - self.spatial_down = spatial_down - - self.temporal_ratio = 2 if temporal_down else 1 - self.spatial_ratio = 2 if spatial_down else 1 - - self.temporal_kernel = 3 if temporal_down else 1 - self.spatial_kernel = 3 if spatial_down else 1 - - self.conv = init_causal_conv3d( - self.channels, - self.channels, - kernel_size=(self.temporal_kernel, self.spatial_kernel, self.spatial_kernel), - stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio), - padding=((1 if self.temporal_down else 0), 0, 0), - inflation_mode=inflation_mode, - ) - self.gradient_checkpointing = False - - def forward( - self, - hidden_states: torch.FloatTensor, - memory_state: MemoryState, - ) -> torch.FloatTensor: - return gradient_checkpointing( - self.custom_forward, - hidden_states, - memory_state, - enabled=self.training and self.gradient_checkpointing, - ) - - def custom_forward( - self, - hidden_states: torch.FloatTensor, - memory_state: MemoryState, - ) -> torch.FloatTensor: - - assert hidden_states.shape[1] == self.channels - - if self.spatial_down: - hidden_states = safe_pad_operation(hidden_states, (0, 1, 0, 1), mode="constant", value=0) - - hidden_states = self.conv(hidden_states, memory_state=memory_state) - return hidden_states - - -class ResnetBlock3D(ResnetBlock2D): - def __init__( - self, - *args, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - **kwargs, - ): - super().__init__(*args, **kwargs) - self.conv1 = init_causal_conv3d( - self.in_channels, - self.out_channels, - kernel_size=3, - stride=1, - padding=1, - inflation_mode=inflation_mode, - ) - - self.conv2 = init_causal_conv3d( - self.out_channels, - self.out_channels, - kernel_size=(1, 3, 3) if time_receptive_field == "half" else (3, 3, 3), - stride=1, - padding=(0, 1, 1) if time_receptive_field == "half" else (1, 1, 1), - inflation_mode=inflation_mode, - ) - - if self.use_in_shortcut: - self.conv_shortcut = init_causal_conv3d( - self.in_channels, - self.out_channels, - kernel_size=1, - stride=1, - padding=0, - bias=(self.conv_shortcut.bias is not None), - inflation_mode=inflation_mode, - ) - self.gradient_checkpointing = False - - def forward(self, input_tensor: torch.Tensor, memory_state: MemoryState = MemoryState.UNSET): - return gradient_checkpointing( - self.custom_forward, - input_tensor, - memory_state, - enabled=self.training and self.gradient_checkpointing, - ) - - def custom_forward( - self, input_tensor: torch.Tensor, memory_state: MemoryState = MemoryState.UNSET - ): - assert memory_state != MemoryState.UNSET - hidden_states = input_tensor - - hidden_states = causal_norm_wrapper(self.norm1, hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.conv1(hidden_states, memory_state=memory_state) - - hidden_states = causal_norm_wrapper(self.norm2, hidden_states) - hidden_states = self.nonlinearity(hidden_states) - hidden_states = self.dropout(hidden_states) - hidden_states = self.conv2(hidden_states, memory_state=memory_state) - - if self.conv_shortcut is not None: - input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state) - - output_tensor = input_tensor + hidden_states - - return output_tensor - - -class DownEncoderBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - add_downsample: bool = True, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - temporal_down: bool = True, - spatial_down: bool = True, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - in_channels = in_channels if i == 0 else out_channels - resnets.append( - ResnetBlock3D( - in_channels=in_channels, - out_channels=out_channels, - dropout=dropout, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - self.downsamplers = None - if add_downsample: - # Todo: Refactor this line before V5 Image VAE Training. - self.downsamplers = nn.ModuleList( - [ - Downsample3D( - channels=out_channels, - inflation_mode=inflation_mode, - temporal_down=temporal_down, - spatial_down=spatial_down, - ) - ] - ) - - def forward( - self, hidden_states: torch.FloatTensor, memory_state: MemoryState - ) -> torch.FloatTensor: - for resnet in self.resnets: - hidden_states = resnet(hidden_states, memory_state=memory_state) - - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, memory_state=memory_state) - - return hidden_states - - -class UpDecoderBlock3D(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - dropout: float = 0.0, - num_layers: int = 1, - add_upsample: bool = True, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - temporal_up: bool = True, - spatial_up: bool = True, - slicing: bool = False, - ): - super().__init__() - resnets = [] - - for i in range(num_layers): - input_channels = in_channels if i == 0 else out_channels - - resnets.append( - ResnetBlock3D( - in_channels=input_channels, - out_channels=out_channels, - dropout=dropout, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - ) - - self.resnets = nn.ModuleList(resnets) - - self.upsamplers = None - # Todo: Refactor this line before V5 Image VAE Training. - if add_upsample: - self.upsamplers = nn.ModuleList( - [ - Upsample3D( - channels=out_channels, - inflation_mode=inflation_mode, - temporal_up=temporal_up, - spatial_up=spatial_up, - slicing=slicing, - ) - ] - ) - - def forward( - self, hidden_states: torch.FloatTensor, memory_state: MemoryState - ) -> torch.FloatTensor: - for resnet in self.resnets: - hidden_states = resnet(hidden_states, memory_state=memory_state) - - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, memory_state=memory_state) - - return hidden_states - - -class UNetMidBlock3D(nn.Module): - def __init__( - self, - channels: int, - dropout: float = 0.0, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - ): - super().__init__() - self.resnets = nn.ModuleList( - [ - ResnetBlock3D( - in_channels=channels, - out_channels=channels, - dropout=dropout, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ), - ResnetBlock3D( - in_channels=channels, - out_channels=channels, - dropout=dropout, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ), - ] - ) - - def forward(self, hidden_states: torch.Tensor, memory_state: MemoryState): - for resnet in self.resnets: - hidden_states = resnet(hidden_states, memory_state) - return hidden_states - - -class Encoder3D(nn.Module): - r""" - The `Encoder` layer of a variational autoencoder that encodes - its input into a latent representation. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - block_out_channels: Tuple[int, ...] = (64,), - layers_per_block: int = 2, - double_z: bool = True, - temporal_down_num: int = 2, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), - ): - super().__init__() - self.layers_per_block = layers_per_block - - self.temporal_down_num = temporal_down_num - - self.conv_in = init_causal_conv3d( - in_channels, - block_out_channels[0], - kernel_size=3, - stride=1, - padding=1, - inflation_mode=inflation_mode, - ) - - self.down_blocks = nn.ModuleList([]) - - # down - output_channel = block_out_channels[0] - for i in range(len(block_out_channels)): - input_channel = output_channel - output_channel = block_out_channels[i] - is_final_block = i == len(block_out_channels) - 1 - is_temporal_down_block = i >= len(block_out_channels) - self.temporal_down_num - 1 - # Note: take the last one - - down_block = DownEncoderBlock3D( - num_layers=self.layers_per_block, - in_channels=input_channel, - out_channels=output_channel, - add_downsample=not is_final_block, - temporal_down=is_temporal_down_block, - spatial_down=True, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - self.down_blocks.append(down_block) - - # mid - self.mid_block = UNetMidBlock3D( - channels=block_out_channels[-1], - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - - # out - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[-1], num_groups=32, eps=1e-6 - ) - self.conv_act = nn.SiLU() - - conv_out_channels = 2 * out_channels if double_z else out_channels - self.conv_out = init_causal_conv3d( - block_out_channels[-1], conv_out_channels, 3, padding=1, inflation_mode=inflation_mode - ) - - assert len(selective_checkpointing) == len(self.down_blocks) - self.set_gradient_checkpointing(selective_checkpointing) - - def set_gradient_checkpointing(self, checkpointing_types): - gradient_checkpointing = [] - for down_block, sac_type in zip(self.down_blocks, checkpointing_types): - if sac_type == "coarse": - gradient_checkpointing.append(True) - elif sac_type == "fine": - for n, m in down_block.named_modules(): - if hasattr(m, "gradient_checkpointing"): - m.gradient_checkpointing = True - logger.debug(f"set gradient_checkpointing: {n}") - gradient_checkpointing.append(False) - else: - gradient_checkpointing.append(False) - self.gradient_checkpointing = gradient_checkpointing - logger.info(f"[Encoder3D] gradient_checkpointing: {checkpointing_types}") - - def forward(self, sample: torch.FloatTensor, memory_state: MemoryState) -> torch.FloatTensor: - r"""The forward method of the `Encoder` class.""" - sample = self.conv_in(sample, memory_state=memory_state) - # down - for down_block, sac in zip(self.down_blocks, self.gradient_checkpointing): - sample = gradient_checkpointing( - down_block, - sample, - memory_state=memory_state, - enabled=self.training and sac, - ) - - # middle - sample = self.mid_block(sample, memory_state=memory_state) - - # post-process - sample = causal_norm_wrapper(self.conv_norm_out, sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample, memory_state=memory_state) - - return sample - - -class Decoder3D(nn.Module): - r""" - The `Decoder` layer of a variational autoencoder that - decodes its latent representation into an output sample. - """ - - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - block_out_channels: Tuple[int, ...] = (64,), - layers_per_block: int = 2, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - temporal_up_num: int = 2, - slicing_up_num: int = 0, - selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), - ): - super().__init__() - self.layers_per_block = layers_per_block - self.temporal_up_num = temporal_up_num - - self.conv_in = init_causal_conv3d( - in_channels, - block_out_channels[-1], - kernel_size=3, - stride=1, - padding=1, - inflation_mode=inflation_mode, - ) - - self.up_blocks = nn.ModuleList([]) - - # mid - self.mid_block = UNetMidBlock3D( - channels=block_out_channels[-1], - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - - # up - reversed_block_out_channels = list(reversed(block_out_channels)) - output_channel = reversed_block_out_channels[0] - for i in range(len(reversed_block_out_channels)): - prev_output_channel = output_channel - output_channel = reversed_block_out_channels[i] - - is_final_block = i == len(block_out_channels) - 1 - is_temporal_up_block = i < self.temporal_up_num - is_slicing_up_block = i >= len(block_out_channels) - slicing_up_num - # Note: Keep symmetric - - up_block = UpDecoderBlock3D( - num_layers=self.layers_per_block + 1, - in_channels=prev_output_channel, - out_channels=output_channel, - add_upsample=not is_final_block, - temporal_up=is_temporal_up_block, - slicing=is_slicing_up_block, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - self.up_blocks.append(up_block) - - # out - self.conv_norm_out = nn.GroupNorm( - num_channels=block_out_channels[0], num_groups=32, eps=1e-6 - ) - self.conv_act = nn.SiLU() - self.conv_out = init_causal_conv3d( - block_out_channels[0], out_channels, 3, padding=1, inflation_mode=inflation_mode - ) - - assert len(selective_checkpointing) == len(self.up_blocks) - self.set_gradient_checkpointing(selective_checkpointing) - - def set_gradient_checkpointing(self, checkpointing_types): - gradient_checkpointing = [] - for up_block, sac_type in zip(self.up_blocks, checkpointing_types): - if sac_type == "coarse": - gradient_checkpointing.append(True) - elif sac_type == "fine": - for n, m in up_block.named_modules(): - if hasattr(m, "gradient_checkpointing"): - m.gradient_checkpointing = True - logger.debug(f"set gradient_checkpointing: {n}") - gradient_checkpointing.append(False) - else: - gradient_checkpointing.append(False) - self.gradient_checkpointing = gradient_checkpointing - logger.info(f"[Decoder3D] gradient_checkpointing: {checkpointing_types}") - - def forward(self, sample: torch.FloatTensor, memory_state: MemoryState) -> torch.FloatTensor: - r"""The forward method of the `Decoder` class.""" - - sample = self.conv_in(sample, memory_state=memory_state) - - # middle - sample = self.mid_block(sample, memory_state=memory_state) - - # up - for up_block, sac in zip(self.up_blocks, self.gradient_checkpointing): - sample = gradient_checkpointing( - up_block, - sample, - memory_state=memory_state, - enabled=self.training and sac, - ) - - # post-process - sample = causal_norm_wrapper(self.conv_norm_out, sample) - sample = self.conv_act(sample) - sample = self.conv_out(sample, memory_state=memory_state) - - return sample - - -class VideoAutoencoderKL(nn.Module): - def __init__( - self, - in_channels: int = 3, - out_channels: int = 3, - block_out_channels: Tuple[int] = (64,), - layers_per_block: int = 1, - latent_channels: int = 4, - use_quant_conv: bool = True, - use_post_quant_conv: bool = True, - enc_selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), - dec_selective_checkpointing: Tuple[_selective_checkpointing_t] = ("none",), - temporal_scale_num: int = 0, - slicing_up_num: int = 0, - inflation_mode: _inflation_mode_t = "tail", - time_receptive_field: _receptive_field_t = "half", - slicing_sample_min_size: int = None, - spatial_downsample_factor: int = 16, - temporal_downsample_factor: int = 8, - freeze_encoder: bool = False, - ): - super().__init__() - self.spatial_downsample_factor = spatial_downsample_factor - self.temporal_downsample_factor = temporal_downsample_factor - self.freeze_encoder = freeze_encoder - if slicing_sample_min_size is None: - slicing_sample_min_size = temporal_downsample_factor - self.slicing_sample_min_size = slicing_sample_min_size - self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num) - - # pass init params to Encoder - self.encoder = Encoder3D( - in_channels=in_channels, - out_channels=latent_channels, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - double_z=True, - temporal_down_num=temporal_scale_num, - selective_checkpointing=enc_selective_checkpointing, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - - # pass init params to Decoder - self.decoder = Decoder3D( - in_channels=latent_channels, - out_channels=out_channels, - block_out_channels=block_out_channels, - layers_per_block=layers_per_block, - # [Override] add temporal_up_num parameter - temporal_up_num=temporal_scale_num, - slicing_up_num=slicing_up_num, - selective_checkpointing=dec_selective_checkpointing, - inflation_mode=inflation_mode, - time_receptive_field=time_receptive_field, - ) - - self.quant_conv = ( - init_causal_conv3d( - in_channels=2 * latent_channels, - out_channels=2 * latent_channels, - kernel_size=1, - inflation_mode=inflation_mode, - ) - if use_quant_conv - else None - ) - self.post_quant_conv = ( - init_causal_conv3d( - in_channels=latent_channels, - out_channels=latent_channels, - kernel_size=1, - inflation_mode=inflation_mode, - ) - if use_post_quant_conv - else None - ) - - self.use_slicing = False - - def enable_slicing(self): - self.use_slicing = True - - def disable_slicing(self): - self.use_slicing = False - - def encode(self, x: torch.FloatTensor) -> CausalEncoderOutput: - if x.ndim == 4: - x = x.unsqueeze(2) - h = self.slicing_encode(x) - p = DiagonalGaussianDistribution(h) - z = p.sample() - return CausalEncoderOutput(z, p) - - def decode(self, z: torch.FloatTensor) -> CausalDecoderOutput: - if z.ndim == 4: - z = z.unsqueeze(2) - x = self.slicing_decode(z) - return CausalDecoderOutput(x) - - def _encode(self, x: torch.Tensor, memory_state: MemoryState) -> torch.Tensor: - x = causal_conv_slice_inputs(x, self.slicing_sample_min_size, memory_state=memory_state) - h = self.encoder(x, memory_state=memory_state) - h = self.quant_conv(h, memory_state=memory_state) if self.quant_conv is not None else h - h = causal_conv_gather_outputs(h) - return h - - def _decode(self, z: torch.Tensor, memory_state: MemoryState) -> torch.Tensor: - z = causal_conv_slice_inputs(z, self.slicing_latent_min_size, memory_state=memory_state) - z = ( - self.post_quant_conv(z, memory_state=memory_state) - if self.post_quant_conv is not None - else z - ) - x = self.decoder(z, memory_state=memory_state) - x = causal_conv_gather_outputs(x) - return x - - def slicing_encode(self, x: torch.Tensor) -> torch.Tensor: - sp_size = 1 - if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size * sp_size: - x_slices = x[:, :, 1:].split(split_size=self.slicing_sample_min_size * sp_size, dim=2) - encoded_slices = [ - self._encode( - torch.cat((x[:, :, :1], x_slices[0]), dim=2), - memory_state=MemoryState.INITIALIZING, - ) - ] - for x_idx in range(1, len(x_slices)): - encoded_slices.append( - self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE) - ) - return torch.cat(encoded_slices, dim=2) - else: - return self._encode(x, memory_state=MemoryState.DISABLED) - - def slicing_decode(self, z: torch.Tensor) -> torch.Tensor: - sp_size = 1 - if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size * sp_size: - z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size * sp_size, dim=2) - decoded_slices = [ - self._decode( - torch.cat((z[:, :, :1], z_slices[0]), dim=2), - memory_state=MemoryState.INITIALIZING, - ) - ] - for z_idx in range(1, len(z_slices)): - decoded_slices.append( - self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE) - ) - return torch.cat(decoded_slices, dim=2) - else: - return self._decode(z, memory_state=MemoryState.DISABLED) - - def forward(self, x: torch.FloatTensor) -> CausalAutoencoderOutput: - with torch.no_grad() if self.freeze_encoder else nullcontext(): - z, p = self.encode(x) - x = self.decode(z).sample - return CausalAutoencoderOutput(x, z, p) - - def preprocess(self, x: torch.Tensor): - # x should in [B, C, T, H, W], [B, C, H, W] - assert x.ndim == 4 or x.size(2) % self.temporal_downsample_factor == 1 - return x - - def postprocess(self, x: torch.Tensor): - # x should in [B, C, T, H, W], [B, C, H, W] - return x - - def set_causal_slicing( - self, - *, - split_size: Optional[int], - memory_device: _memory_device_t, - ): - assert ( - split_size is None or memory_device is not None - ), "if split_size is set, memory_device must not be None." - if split_size is not None: - self.enable_slicing() - self.slicing_sample_min_size = split_size - self.slicing_latent_min_size = split_size // self.temporal_downsample_factor - else: - self.disable_slicing() - for module in self.modules(): - if isinstance(module, InflatedCausalConv3d): - module.set_memory_device(memory_device) - - def set_memory_limit(self, conv_max_mem: Optional[float], norm_max_mem: Optional[float]): - set_norm_limit(norm_max_mem) - for m in self.modules(): - if isinstance(m, InflatedCausalConv3d): - m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf")) - - -class VideoAutoencoderKLWrapper(VideoAutoencoderKL): - def __init__( - self, *args, spatial_downsample_factor: int, temporal_downsample_factor: int, **kwargs - ): - self.spatial_downsample_factor = spatial_downsample_factor - self.temporal_downsample_factor = temporal_downsample_factor - super().__init__(*args, **kwargs) - - def forward(self, x) -> CausalAutoencoderOutput: - z, _, p = self.encode(x) - x, _ = self.decode(z) - return CausalAutoencoderOutput(x, z, None, p) - - def encode(self, x) -> CausalEncoderOutput: - if x.ndim == 4: - x = x.unsqueeze(2) - p = super().encode(x).latent_dist - z = p.sample().squeeze(2) - return CausalEncoderOutput(z, None, p) - - def decode(self, z) -> CausalDecoderOutput: - if z.ndim == 4: - z = z.unsqueeze(2) - x = super().decode(z).sample.squeeze(2) - return CausalDecoderOutput(x, None) - - def preprocess(self, x): - # x should in [B, C, T, H, W], [B, C, H, W] - assert x.ndim == 4 or x.size(2) % 4 == 1 - return x - - def postprocess(self, x): - # x should in [B, C, T, H, W], [B, C, H, W] - return x - - def set_causal_slicing( - self, - *, - split_size: Optional[int], - memory_device: Optional[Literal["cpu", "same"]], - ): - assert ( - split_size is None or memory_device is not None - ), "if split_size is set, memory_device must not be None." - if split_size is not None: - self.enable_slicing() - else: - self.disable_slicing() - self.slicing_sample_min_size = split_size - if split_size is not None: - self.slicing_latent_min_size = split_size // self.temporal_downsample_factor - for module in self.modules(): - if isinstance(module, InflatedCausalConv3d): - module.set_memory_device(memory_device) diff --git a/modules/seedvr/test.py b/modules/seedvr/test.py index b8c53f874..0907396b1 100644 --- a/modules/seedvr/test.py +++ b/modules/seedvr/test.py @@ -36,6 +36,8 @@ def upscale_image(model_name:str, image_path:str): runner=runner, images=image_tensor, cfg_scale=cfg, + cfg_rescale=0.0, + steps=1, seed=seed, res_w=resolution, batch_size=1, diff --git a/modules/ui_definitions.py b/modules/ui_definitions.py index c5a3a4bce..2a38277dd 100644 --- a/modules/ui_definitions.py +++ b/modules/ui_definitions.py @@ -631,9 +631,6 @@ def create_settings(cmd_opts): "detailer_unload": OptionInfo(False, "Move detailer model to CPU when complete"), "detailer_augment": OptionInfo(False, "Detailer use model augment"), - "postprocessing_sep_seedvr": OptionInfo("

SeedVR

", "", gr.HTML), - "seedvr_cfg_scale": OptionInfo(3.5, "SeedVR CFG Scale", gr.Slider, {"minimum": 1, "maximum": 15, "step": 1}), - "postprocessing_sep_upscalers": OptionInfo("

Upscaling

", "", gr.HTML), "upscaler_unload": OptionInfo(False, "Unload upscaler after processing"), "upscaler_latent_steps": OptionInfo(20, "Upscaler latent steps", gr.Slider, {"minimum": 4, "maximum": 100, "step": 1}), diff --git a/scripts/postprocessing_pixelart.py b/scripts/postprocessing_pixelart.py index ca8bf36c5..44bd36359 100644 --- a/scripts/postprocessing_pixelart.py +++ b/scripts/postprocessing_pixelart.py @@ -1,6 +1,7 @@ import gradio as gr from modules import scripts_postprocessing, devices + class ScriptPixelArt(scripts_postprocessing.ScriptPostprocessing): name = "PixelArt" order = 30000 @@ -13,6 +14,7 @@ class ScriptPixelArt(scripts_postprocessing.ScriptPostprocessing): with gr.Row(): pixelart_block_size = gr.Slider(minimum=2, maximum=64, step=1, value=8, label="PixelArt block size", elem_id="extras_pixelart_block_size") pixelart_edge_block_size = gr.Slider(minimum=2, maximum=64, step=1, value=4, label="Edge block size", elem_id="extras_pixelart_edge_block_size") + with gr.Row(): pixelart_image_weight = gr.Slider(minimum=0.0, maximum=2.0, step=0.01, value=1.0, label="Edge image weight", elem_id="extras_pixelart_image_weight") pixelart_sharpen_amount = gr.Slider(minimum=0.0, maximum=1.0, step=0.01, value=0.1, label="PixelArt sharpen", elem_id="extras_pixelart_sharpen_amount") return { diff --git a/scripts/postprocessing_seedvr.py b/scripts/postprocessing_seedvr.py new file mode 100644 index 000000000..e20b1ecfa --- /dev/null +++ b/scripts/postprocessing_seedvr.py @@ -0,0 +1,74 @@ +import gradio as gr +from modules import scripts_postprocessing + + +class ScriptSeedVR(scripts_postprocessing.ScriptPostprocessing): + name = "SeedVR" + + def ui(self): + from modules.postprocess.seedvr_model import MODELS_MAP + with gr.Accordion(self.name, open = False, elem_id="postprocess_seedvr_accordion"): + with gr.Row(): + seedvr_enabled = gr.Checkbox(label="Enable SeedVR", value=False, elem_id="extras_seedvr_enabled") + with gr.Row(): + seedvr_selected = gr.Dropdown(label="SeedVR model", choices=list(MODELS_MAP.keys()), value=list(MODELS_MAP.keys())[0], elem_id="extras_seedvr_model") + with gr.Row(): + seedvr_scale = gr.Slider(minimum=1, maximum=16, step=0.1, value=2, label="SeedVR scale", elem_id="extras_seedvr_scale") + seedvr_seed = gr.Number(step=1, value=-1, label="SeedVR seed", elem_id="extras_seedvr_seed") + seedvr_steps = gr.Number(step=1, value=1, minimum=1, maximum=99, label="SeedVR steps", elem_id="extras_seedvr_steps", visible=False) + with gr.Row(): + seedvr_cfg_scale = gr.Slider(minimum=0.0, maximum=15.0, step=0.01, value=3.5, label="SeedVR guidance scale", elem_id="extras_seedvr_cfg_scale") + seedvr_cfg_rescale = gr.Slider(minimum=0.0, maximum=15.0, step=0.01, value=0.0, label="SeedVR guidance rescale", elem_id="extras_seedvr_cfg_rescale") + with gr.Row(): + seedvr_tile_size = gr.Slider(minimum=64, maximum=4096, step=8, value=1024, label="SeedVR tile size", elem_id="extras_seedvr_tile_size") + seedvr_tile_overlap = gr.Slider(minimum=0, maximum=1.0, step=0.01, value=0.25, label="SeedVR tile overlap", elem_id="extras_seedvr_tile_overlap") + return { + "seedvr_enabled": seedvr_enabled, + "seedvr_selected": seedvr_selected, + "seedvr_scale": seedvr_scale, + "seedvr_seed": seedvr_seed, + "seedvr_steps": seedvr_steps, + "seedvr_cfg_scale": seedvr_cfg_scale, + "seedvr_cfg_rescale": seedvr_cfg_rescale, + "seedvr_tile_size": seedvr_tile_size, + "seedvr_tile_overlap": seedvr_tile_overlap, + } + + def process(self, + pp: scripts_postprocessing.PostprocessedImage, + seedvr_enabled: bool, + seedvr_selected: str, + seedvr_scale: int, + seedvr_seed: int, + seedvr_steps: int, + seedvr_cfg_scale: float, + seedvr_cfg_rescale: float, + seedvr_tile_size: int, + seedvr_tile_overlap: float + ): # pylint: disable=arguments-differ + if not seedvr_enabled: + return + from modules import shared, upscaler + from modules.logger import log + image = pp.image + instance: upscaler.UpscalerData = next(iter([x for x in shared.sd_upscalers if x.name == seedvr_selected]), None) + scaler: upscaler.Upscaler = instance.scaler + + log.info(f'Upscaler: type="SeedVR" model="{seedvr_selected}" scale={seedvr_scale} seed={seedvr_seed} steps={seedvr_steps} cfg_scale={seedvr_cfg_scale} cfg_rescale={seedvr_cfg_rescale} tile_size={seedvr_tile_size} tile_overlap={seedvr_tile_overlap}') + + jobid = shared.state.begin('Upscale') + + scaler.scale = float(seedvr_scale) + upscaled = scaler.do_upscale(image, + seedvr_selected, + cfg_scale=seedvr_cfg_scale, + cfg_rescale=seedvr_cfg_rescale, + steps=seedvr_steps, + seed=seedvr_seed, + tile_size=seedvr_tile_size, + tile_overlap=seedvr_tile_overlap, + ) + shared.state.end(jobid) + + pp.image = upscaled + pp.info["SeedVR"] = f"Scale={seedvr_scale} Seed={seedvr_seed} CFG Scale={seedvr_cfg_scale} CFG Rescale={seedvr_cfg_rescale}"