diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 820a26678..7cc2f614a 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 820a266789ac26c36129cba1fd7b589f7c457c9a +Subproject commit 7cc2f614a536bc34fc9bdd9c0acf50b59165903a diff --git a/modules/postprocess/seedvr_model.py b/modules/postprocess/seedvr_model.py index a85b0fa0f..5da179635 100644 --- a/modules/postprocess/seedvr_model.py +++ b/modules/postprocess/seedvr_model.py @@ -1,9 +1,10 @@ import time +import os import random import numpy as np import torch from PIL import Image -from modules import devices +from modules import devices, timer from modules.shared import opts from modules.upscaler import Upscaler, UpscalerData from modules.image import convert @@ -29,8 +30,6 @@ class UpscalerSeedVR(Upscaler): ] self.model = None self.model_loaded = None - self.tile_size = 1024 - self.tile_overlap = 0.25 self.device = devices.device self.step = 1 self.frames = 0 @@ -38,6 +37,28 @@ class UpscalerSeedVR(Upscaler): self.pbar = None self.task = None self.fps = 24 + self.timer = None + + def set_vae_params(self, vae_memory: float, tile_size: int, tile_overlap: float, vae_tile_encode: bool = True, vae_tile_decode: bool = True): + if vae_memory >= 0.99: + vae_memory = None + self.model.config.vae.memory_limit = {'conv_max_mem': vae_memory, 'norm_max_mem': vae_memory} + self.model.vae.set_memory_limit(**self.model.config.vae.memory_limit) + self.model.vae.tile_sample_min_size = tile_size + self.model.vae.tile_latent_min_size = tile_size // 8 + self.model.vae.tile_overlap_factor = tile_overlap + if vae_tile_encode: + self.model.vae.use_slicing_encode = False + self.model.vae.use_tiling_encode = True + else: + self.model.vae.use_slicing_encode = True + self.model.vae.use_tiling_encode = False + if vae_tile_decode: + self.model.vae.use_slicing_decode = False + self.model.vae.use_tiling_decode = True + else: + self.model.vae.use_slicing_decode = True + self.model.vae.use_tiling_decode = False def load_model(self, path: str): model_name = MODELS_MAP.get(path, None) @@ -52,68 +73,70 @@ class UpscalerSeedVR(Upscaler): device=devices.device, dtype=devices.dtype, ) - self.model_loaded = model_name self.model.dit.device = devices.device self.model.dit.dtype = devices.dtype self.model.vae_encode = self.vae_encode self.model.vae_decode = self.vae_decode - # Patch generation_loop's generation_step() with our wrapper; stash the original once - # so reloads don't re-wrap the wrapper itself (infinite recursion). - if not hasattr(generation, "generation_step_original"): + if not hasattr(generation, "generation_step_original"): # Patch generation_loop's generation_step() with our wrapper; stash the original once so reloads don't re-wrap the wrapper itself (infinite recursion). generation.generation_step_original = generation.generation_step generation.generation_step = self.model_step self.model._internal_dict = { 'dit': self.model.dit, 'vae': self.model.vae, } - t1 = time.time() self.model.dit.config = self.model.config.dit - 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) - + t1 = time.time() log.info(f'Upscaler loaded: name="{self.name}" model="{model_name}" time={t1 - t0:.2f}') def vae_encode(self, samples): latents = [] if len(samples) == 0: return latents - self.pbar.update(self.task, description=f'encode: samples={samples[0].shape if len(samples) > 0 else None} tile={self.model.vae.tile_sample_min_size} overlap={self.model.vae.tile_overlap_factor}') + self.pbar.update(self.task, description=f'encode: images={list(samples[0].shape) if len(samples) > 0 else None}') if self.offload: + t0 = time.time() self.model.dit = self.model.dit.to(device="cpu") self.model.vae = self.model.vae.to(device=self.device) - devices.torch_gc() + self.timer.ts('offload', t0) + devices.torch_gc(fast=True) + t0 = time.time() from einops import rearrange scale = self.model.config.vae.scaling_factor shift = self.model.config.vae.get("shifting_factor", 0.0) batches = [sample.unsqueeze(0) for sample in samples] - for sample in batches: - sample = sample.to(self.device, self.model.vae.dtype) - sample = self.model.vae.preprocess(sample) - latent = self.model.vae.encode(sample).latent - latent = latent.unsqueeze(2) if latent.ndim == 4 else latent - latent = rearrange(latent, "b c ... -> b ... c") - latent = (latent - shift) * scale - latents.append(latent) + with devices.inference_context(): + for sample in batches: + sample = sample.to(self.device, self.model.vae.dtype) + sample = self.model.vae.preprocess(sample) + latent = self.model.vae.encode(sample).latent + latent = latent.unsqueeze(2) if latent.ndim == 4 else latent + latent = rearrange(latent, "b c ... -> b ... c") + latent = (latent - shift) * scale + latents.append(latent.contiguous()) latents = [latent.squeeze(0) for latent in latents] + self.timer.ts('encode', t0) if self.offload: + t0 = time.time() self.model.vae = self.model.vae.to(device="cpu") - devices.torch_gc() + self.timer.ts('offload', t0) + devices.torch_gc(fast=True) return latents def vae_decode(self, latents, target_dtype: torch.dtype = None): - self.pbar.update(self.task, description=f'decode: latents={latents[0].shape if len(latents) > 0 else None} tile={self.model.vae.tile_latent_min_size} overlap={self.model.vae.tile_overlap_factor}') + self.pbar.update(self.task, description=f'decode: latents={list(latents[0].shape) if len(latents) > 0 else None}') samples = [] if len(latents) == 0: return samples from einops import rearrange if self.offload: + t0 = time.time() self.model.dit = self.model.dit.to(device="cpu") self.model.vae = self.model.vae.to(device=self.device) - devices.torch_gc() + self.timer.ts('offload', t0) + devices.torch_gc(fast=True) + t0 = time.time() scale = self.model.config.vae.scaling_factor shift = self.model.config.vae.get("shifting_factor", 0.0) latents = [latent.unsqueeze(0) for latent in latents] @@ -122,29 +145,40 @@ class UpscalerSeedVR(Upscaler): latent = latent.to(self.device, self.model.vae.dtype) latent = latent / scale + shift latent = rearrange(latent, "b ... c -> b c ...") - latent = latent.squeeze(2) + latent = latent.squeeze(2).contiguous() sample = self.model.vae.decode(latent).sample sample = self.model.vae.postprocess(sample) - samples.append(sample) - samples = [sample.squeeze(0) for sample in samples] + samples.append(sample.squeeze(0).contiguous()) + self.timer.ts('decode', t0) if self.offload: + t0 = time.time() self.model.vae = self.model.vae.to(device="cpu") - devices.torch_gc() + self.timer.ts('offload', t0) + devices.torch_gc(fast=True) return samples def model_step(self, *args, **kwargs): + from modules.shared import state + if state.interrupted or state.skipped: + return None from modules.seedvr.src.core import generation if self.offload: + t0 = time.time() self.model.vae = self.model.vae.to(device="cpu") self.model.dit = self.model.dit.to(device=self.device) - devices.torch_gc() + self.timer.ts('offload', t0) + devices.torch_gc(fast=True) + t0 = time.time() with devices.inference_context(): - self.pbar.update(self.task, description=f'inference: step={self.step}') + self.pbar.update(self.task, description=f'inference: batch={self.step}') result = generation.generation_step_original(*args, **kwargs) self.pbar.update(self.task, advance=self.step) + self.timer.ts('step', t0) if self.offload: + t0 = time.time() self.model.dit = self.model.dit.to(device="cpu") - devices.torch_gc() + self.timer.ts('offload', t0) + devices.torch_gc(fast=True) return result def read_image(self, image: str | Image.Image): @@ -170,6 +204,7 @@ class UpscalerSeedVR(Upscaler): return None, None frames = [] width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + self.fps = int(cap.get(cv2.CAP_PROP_FPS)) while True: ret, frame = cap.read() if not ret: @@ -182,12 +217,28 @@ class UpscalerSeedVR(Upscaler): return None, None tensor = torch.from_numpy(np.array(frames)).to(device=devices.device, dtype=devices.dtype) / 255.0 self.frames = tensor.shape[0] - self.fps = int(cap.get(cv2.CAP_PROP_FPS)) return tensor, width except Exception as e: log.error(f'Upscaler: name="SeedVR2" video="{video_path}" {e}') return None, None + def create_video(self, tensor: torch.Tensor, codec: str = 'libx264', codec_opt: str = 'crf:16', interpolate: int = 0): + t0 = time.time() + from modules.video_models.video_save import save_video + pixels = tensor.permute(3, 0, 1, 2).unsqueeze(0) # from (t, h, w, c) to (n, c, t, h, w) + _frames, filename, _thumb = save_video(p=None, + pixels=pixels, + mp4_fps=self.fps, + mp4_thumb=False, + mp4_frames=False, + reclamp=False, + mp4_codec=codec, + mp4_opt=codec_opt, + mp4_interpolate=interpolate, + ) + self.timer.ts('save', t0) + return filename + def do_upscale(self, img: Image.Image | str, selected_file, @@ -200,25 +251,29 @@ class UpscalerSeedVR(Upscaler): tile_overlap: float = 0.25, batch_size: int = 1, batch_overlap: int = 0, - offload: bool = True + offload: bool = True, + interpolate: int = 1, + codec: str = 'libx264', + codec_opt: str = 'crf:16', + vae_memory: float = 0.2, + vae_tile_encode: bool = True, + vae_tile_decode: bool = True, ): + self.timer = timer.Timer() self.offload = offload self.load_model(selected_file) + self.set_vae_params(vae_memory=vae_memory, tile_size=tile_size, tile_overlap=tile_overlap, vae_tile_encode=vae_tile_encode, vae_tile_decode=vae_tile_decode) if self.model is None: return img if not self.offload: self.model.dit = self.model.dit.to(device=devices.device) self.model.vae = self.model.vae.to(device=devices.device) - devices.torch_gc() + devices.torch_gc(fast=True) + self.timer.record('load') 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 if isinstance(img, Image.Image): tensor, width = self.read_image(img) elif isinstance(img, str): @@ -226,6 +281,7 @@ class UpscalerSeedVR(Upscaler): else: log.error(f'Upscaler: name="SeedVR2" image="{img}" unsupported type {type(img)}') return img + self.timer.record('read') if tensor is None or width is None: log.error(f'Upscaler: name="SeedVR2" image="{img}" failed to read') @@ -235,8 +291,11 @@ class UpscalerSeedVR(Upscaler): seed = int(random.randrange(4294967294)) if seed == -1 else int(seed) self.step = 1 if self.frames == 1 else batch_size - batch_overlap - t0 = time.time() - log.info(f'Upscaler: type="{self.name}" model="{selected_file}" scale={self.scale} cfg={cfg_scale}:{cfg_rescale} seed={seed} steps={steps} frames={self.frames} mode={"image" if self.frames == 1 else "video"} tile={self.tile_size}:{self.tile_overlap} batch={batch_size}:{batch_overlap} offload={self.offload}') + mode = "mode=image" if self.frames == 1 else f"mode=video frames={self.frames}" + batch_info = f'batch=(size={batch_size} overlap={batch_overlap})' + vae_info = f'vae=(tiled={vae_tile_encode}/{vae_tile_decode} memory={vae_memory} size={tile_size} overlap={tile_overlap})' + log.info(f'Upscaler: type="{self.name}" model="{selected_file}" {mode} scale={self.scale} cfg={cfg_scale}:{cfg_rescale} seed={seed} steps={steps} offload={self.offload} {batch_info} {vae_info}') + import rich.progress as rp self.pbar = rp.Progress(rp.TextColumn('[cyan]SeedVR:'), rp.BarColumn(), rp.MofNCompleteColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=console) self.task = self.pbar.add_task(total=self.frames, description='starting...') @@ -245,6 +304,7 @@ class UpscalerSeedVR(Upscaler): from modules.seedvr.src.optimization import memory_manager memory_manager.clear_rope_cache(self.model) memory_manager.preinitialize_rope_cache(self.model) + self.timer.record('init') result_tensor = generation.generation_loop( runner=self.model, images=tensor, @@ -256,18 +316,20 @@ class UpscalerSeedVR(Upscaler): seed=seed, res_w=width, device=devices.device, + color_reconstruct=True, ) memory_manager.clear_rope_cache(self.model) self.pbar.update(self.task, completed=self.frames) t1 = time.time() - tiles = getattr(self.model.vae, "tiles", None) self.frames = result_tensor.shape[0] if result_tensor is not None else 0 - log.info(f'Upscaler: type="{self.name}" model="{selected_file}" scale={self.scale} cfg={cfg_scale} seed={seed} tiles={tiles} frames={self.frames} time={t1 - t0:.2f}') + self.timer.add('inference', self.timer.get('step') - self.timer.get('encode') - self.timer.get('decode')) + self.timer.rm('step') - if self.offload: - self.model.dit = self.model.dit.to(device="cpu") - self.model.vae = self.model.vae.to(device="cpu") + t0 = time.time() + self.model.dit = self.model.dit.to(device="cpu") + self.model.vae = self.model.vae.to(device="cpu") + self.timer.ts('offload', t0) if opts.upscaler_unload: self.model.dit = None self.model.vae = None @@ -275,15 +337,14 @@ class UpscalerSeedVR(Upscaler): self.model = None log.debug(f'Upscaler unload: type="{self.name}" model="{selected_file}"') devices.torch_gc(force=True) + self.timer.ts('cleanup', t1) if self.frames == 1: - img = convert.to_pil(result_tensor.squeeze()) - return img + result = convert.to_pil(result_tensor.squeeze()) elif self.frames > 1: - from modules.video_models.video_save import save_video - pixels = result_tensor.permute(3, 0, 1, 2).unsqueeze(0) # from (t, h, w, c) to (n, c, t, h, w) - _frames, filename, _thumb = save_video(p=None, pixels=pixels, mp4_fps=self.fps, mp4_thumb=False, mp4_frames=False, reclamp=False) - return filename + result = self.create_video(result_tensor, codec=codec, codec_opt=codec_opt, interpolate=interpolate) else: log.error(f'Upscaler: name="SeedVR2" model="{selected_file}" no frames generated') - return img + result = img + log.info(f'Upscaler: type="{self.name}" model="{selected_file}" frames={self.frames} {self.timer.summary()}') + return result diff --git a/modules/postprocessing.py b/modules/postprocessing.py index 64509ce05..36667bf67 100644 --- a/modules/postprocessing.py +++ b/modules/postprocessing.py @@ -121,14 +121,17 @@ def run_postprocessing(extras_mode, def process_video(): outputs = [] params = {} - info = '' # TODO process: video add infotext if not video or not isinstance(video, str) or not os.path.isfile(video): log.error(f'Process: mode=video file="{video}" not found') - return outputs, video, info, params + return outputs, video, '', params log.debug(f'Process: video={video} {args}') shared.state.textinfo = video pp = scripts_postprocessing.PostprocessedImage(video=video) scripts_manager.scripts_postproc.run(pp, args) + + from modules.video import get_video_info + params = get_video_info(pp.video) + info = ', '.join([f'{k}: {v}' for k, v in params.items()]) return pp.video, info, params if extras_mode == 3: diff --git a/modules/seedvr/config_3b.yaml b/modules/seedvr/config_3b.yaml index 711d4ec71..c5708e79a 100644 --- a/modules/seedvr/config_3b.yaml +++ b/modules/seedvr/config_3b.yaml @@ -54,7 +54,7 @@ vae: - "modules.seedvr.src.models.video_vae_v3.modules.attn_video_vae" name: "VideoAutoencoderKLWrapper" args: "as_params" - freeze_encoder: False + freeze_encoder: True gradient_checkpoint: True # Disabled to prevent VRAM leaks in inference slicing: split_size: 4 diff --git a/modules/seedvr/config_7b.yaml b/modules/seedvr/config_7b.yaml index 0e5cb146c..416dc01b8 100644 --- a/modules/seedvr/config_7b.yaml +++ b/modules/seedvr/config_7b.yaml @@ -51,7 +51,7 @@ vae: - "modules.seedvr.src.models.video_vae_v3.modules.attn_video_vae" name: "VideoAutoencoderKLWrapper" args: "as_params" - freeze_encoder: False + freeze_encoder: True # gradient_checkpoint: True slicing: split_size: 4 diff --git a/modules/seedvr/src/common/cache.py b/modules/seedvr/src/common/cache.py index 3566852ba..b43802664 100644 --- a/modules/seedvr/src/common/cache.py +++ b/modules/seedvr/src/common/cache.py @@ -21,6 +21,9 @@ class Cache: self.cache[key] = result return result + def clear(self): + self.cache.clear() + def namespace(self, namespace: str): return Cache( disable=self.disable, @@ -31,3 +34,15 @@ class Cache: def get(self, key: str): key = self.prefix + key return self.cache[key] + + def size(self): + num = len(self.cache) + total_size = 0 + for value in self.cache.values(): + if hasattr(value, "element_size") and hasattr(value, "nelement"): + total_size += value.element_size() * value.nelement() + elif isinstance(value, (list, tuple)): + for item in value: + if hasattr(item, "element_size") and hasattr(item, "nelement"): + total_size += item.element_size() * item.nelement() + return num, total_size diff --git a/modules/seedvr/src/core/generation.py b/modules/seedvr/src/core/generation.py index 569250c3b..0b1d42f0f 100644 --- a/modules/seedvr/src/core/generation.py +++ b/modules/seedvr/src/core/generation.py @@ -84,6 +84,7 @@ def generation_step(runner, text_embeds_dict, cond_latents, temporal_overlap, de # Process samples with advanced optimization samples = optimized_video_rearrange(video_tensors) + del video_tensors noises = noises[0].to("cpu") aug_noises = aug_noises[0].to("cpu") cond_latents = cond_latents[0].to("cpu") @@ -106,7 +107,7 @@ def cut_videos(videos): return result -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'): +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', color_reconstruct=True): """ Main generation loop with context-aware temporal processing @@ -159,7 +160,8 @@ def generation_loop(runner, images, cfg_scale=1.0, cfg_rescale=0.0, steps=1, see ]) # Initialize generation state - batch_samples = [] + final_video_images = None + current_idx = 0 # Load text embeddings with adaptive dtype text_embeds = {"texts_pos": [runner.text_pos_embeds], "texts_neg": [runner.text_neg_embeds]} @@ -215,7 +217,8 @@ def generation_loop(runner, images, cfg_scale=1.0, cfg_rescale=0.0, steps=1, see # Normal generation samples = generation_step(runner, text_embeds, cond_latents=cond_latents, temporal_overlap=temporal_overlap, device=device) - #del cond_latents + if samples is None: + return del cond_latents # Post-process samples @@ -223,47 +226,35 @@ def generation_loop(runner, images, cfg_scale=1.0, cfg_rescale=0.0, steps=1, see del samples #del samples if ori_lengths[0] < sample.shape[0]: - sample = sample[:ori_lengths[0]] + sample = sample[:ori_lengths[0]].contiguous() # Apply color correction if available - transformed_video = transformed_video.to(device) - input_video = [optimized_single_video_rearrange(transformed_video)] - del transformed_video - sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)]) - del input_video + if color_reconstruct: + transformed_video = transformed_video.to(device) + input_video = [optimized_single_video_rearrange(transformed_video)] + del transformed_video + sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)]) + del input_video # Convert to final image format sample = optimized_sample_to_image_format(sample) sample = sample.clip(-1, 1).mul_(0.5).add_(0.5) - sample_cpu = sample.to(torch.float16).to("cpu") + sample = sample.detach().to(torch.float16, non_blocking=True).cpu() + if final_video_images is None: + total_frames = len(images) + H, W, C = sample.shape[1], sample.shape[2], sample.shape[3] + final_video_images = torch.empty((total_frames, H, W, C), dtype=torch.float16) + + batch_frames = sample.shape[0] + final_video_images[current_idx:current_idx + batch_frames] = sample + current_idx += batch_frames del sample - batch_samples.append(sample_cpu) - #del sample if progress_callback: progress_callback(batch_count+1, total_batches, current_frames, "Processing batch...") - # 1. Calculer la taille totale finale - total_frames = sum(batch.shape[0] for batch in batch_samples) - if len(batch_samples) > 0: - sample_shape = batch_samples[0].shape - H, W, C = sample_shape[1], sample_shape[2], sample_shape[3] - final_video_images = torch.empty((total_frames, H, W, C), dtype=torch.float16) - block_size = 500 - current_idx = 0 - - for block_start in range(0, len(batch_samples), block_size): - block_end = min(block_start + block_size, len(batch_samples)) - current_block = [] - for i in range(block_start, block_end): - current_block.append(batch_samples[i].to(device)) - block_result = torch.cat(current_block, dim=0) - block_frames = block_result.shape[0] - final_video_images[current_idx:current_idx + block_frames] = block_result.to("cpu") - current_idx += block_frames - del current_block, block_result - else: + if final_video_images is None: print("SeedVR2: No batch_samples to process") final_video_images = torch.empty((0, 0, 0, 0), dtype=torch.float16) diff --git a/modules/seedvr/src/core/infer.py b/modules/seedvr/src/core/infer.py index f18ee433f..eb2400093 100644 --- a/modules/seedvr/src/core/infer.py +++ b/modules/seedvr/src/core/infer.py @@ -7,6 +7,7 @@ from modules.seedvr.src.models.dit_v2 import na if TYPE_CHECKING: from modules.seedvr.src.models.dit_v2.nadit import NaDiT + from modules.seedvr.src.models.video_vae_v3.modules.attn_video_vae import VideoAutoencoderKLWrapper def optimized_channels_to_last(tensor: torch.Tensor) -> torch.Tensor: @@ -49,7 +50,7 @@ class SeedVRPipeline(): self.config = config self.device = device self.dtype = dtype - self.vae = None + self.vae: VideoAutoencoderKLWrapper = None self.dit: NaDiT = None self.sampler = None self.schedule = None @@ -132,6 +133,7 @@ class SeedVRPipeline(): latent = rearrange(latent, "b c ... -> b ... c") #latent = optimized_channels_to_last(latent) latent = (latent - shift) * scale + latent = latent.contiguous() latents.append(latent) # Ungroup back to individual latent with the original order. @@ -174,10 +176,11 @@ class SeedVRPipeline(): latent = latent / scale + shift latent = rearrange(latent, "b ... c -> b c ...") #latent = optimized_channels_to_second(latent) - latent = latent.squeeze(2) + latent = latent.squeeze(2).contiguous() # 🚀 OPTIMISATION 3: Décodage direct SANS autocast (utilise l'autocast externe) sample = self.vae.decode(latent).sample + del latent #sample = self.vae.decode(latent).sample #sample = self.vae.decode(latent).sample @@ -186,6 +189,7 @@ class SeedVRPipeline(): sample = self.vae.postprocess(sample) samples.append(sample) + del sample # Ungroup back to individual sample with the original order. if self.config.vae.grouping: @@ -277,9 +281,9 @@ class SeedVRPipeline(): text_neg_embeds, text_neg_shapes = na.flatten(texts_neg) # Adapter les embeddings texte au dtype cible (compatible avec FP8) - if isinstance(text_pos_embeds, torch.Tensor): + if isinstance(text_pos_embeds, torch.Tensor) and text_pos_embeds.dtype != target_dtype: text_pos_embeds = text_pos_embeds.to(target_dtype) - if isinstance(text_neg_embeds, torch.Tensor): + if isinstance(text_neg_embeds, torch.Tensor) and text_neg_embeds.dtype != target_dtype: text_neg_embeds = text_neg_embeds.to(target_dtype) # Flatten. @@ -289,7 +293,9 @@ class SeedVRPipeline(): # Adapter les latents au dtype cible (compatible avec FP8) latents = latents.to(target_dtype) if latents.dtype != target_dtype else latents latents_cond = latents_cond.to(target_dtype) if latents_cond.dtype != target_dtype else latents_cond - self.dit = self.dit.to(device=self.device, dtype=target_dtype) + current_dit_param = next(self.dit.parameters()) + if current_dit_param.dtype != target_dtype or current_dit_param.device != torch.device(self.device): + self.dit = self.dit.to(device=self.device, dtype=target_dtype) latents = self.sampler.sample( x=latents, @@ -321,6 +327,7 @@ class SeedVRPipeline(): vae_dtype = self.vae.dtype decode_dtype = torch.float16 if (vae_dtype == torch.float16 or target_dtype == torch.float16) else vae_dtype samples = self.vae_decode(latents, target_dtype=decode_dtype) + del latents if samples and len(samples) > 0 and samples[0].dtype != torch.float16: samples = [sample.to(torch.float16, non_blocking=True) for sample in samples] diff --git a/modules/seedvr/src/models/video_vae_v3/modules/__init__.py b/modules/seedvr/src/models/video_vae_v3/modules/__init__.py new file mode 100644 index 000000000..e69de29bb 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 f5effd8b8..78a2d294d 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 @@ -114,14 +114,26 @@ class Upsample3D(Upsample2D): hidden_states = [hidden_states] # ADD BY NUMZ 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, - ) + if self.use_conv and hasattr(self, "upscale_conv") and self.upscale_conv.kernel_size == (1, 1, 1): + hidden_states[i] = hidden_states[i].repeat_interleave(self.temporal_ratio, dim=2) + if self.spatial_ratio != 1: + hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=3) + hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=4) + elif self.use_conv: + 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, + ).contiguous() + else: + if self.temporal_ratio != 1: + hidden_states[i] = hidden_states[i].repeat_interleave(self.temporal_ratio, dim=2) + if self.spatial_ratio != 1: + hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=3) + hidden_states[i] = hidden_states[i].repeat_interleave(self.spatial_ratio, dim=4) # [Overridden] For causal temporal conv if self.temporal_up and memory_state != MemoryState.ACTIVE: @@ -1146,9 +1158,13 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): @apply_forward_hook def encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput: - # h = self.slicing_encode(x) - h = self.tiled_encode(x) - posterior = DiagonalGaussianDistribution(h) + if self.use_slicing_encode: + encoded = self.slicing_encode(x) + elif self.use_tiling_encode: + encoded = self.tiled_encode(x) + else: + encoded = self._encode(x) + posterior = DiagonalGaussianDistribution(encoded) if not return_dict: return (posterior,) @@ -1159,8 +1175,12 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): def decode( self, z: torch.Tensor, return_dict: bool = True ) -> Union[DecoderOutput, torch.Tensor]: - # decoded = self.slicing_decode(z) - decoded = self.tiled_decode(z) + if self.use_slicing_decode: + decoded = self.slicing_decode(z) + elif self.use_tiling_decode: + decoded = self.tiled_decode(z) + else: + decoded = self._decode(z) if not return_dict: return (decoded,) @@ -1170,8 +1190,7 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): def _encode( self, x: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED ) -> torch.Tensor: - _x = x.to(self.device) - _x = causal_conv_slice_inputs(_x, self.slicing_sample_min_size, memory_state=memory_state) + _x = causal_conv_slice_inputs(x.to(self.device), self.slicing_sample_min_size, memory_state=memory_state) h = self.encoder(_x, memory_state=memory_state) if self.quant_conv is not None: output = self.quant_conv(h, memory_state=memory_state) @@ -1243,53 +1262,109 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor)) blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor) row_limit = self.tile_latent_min_size - blend_extent - rows = [] + prev_row = None self.tiles = 0 - for i in range(0, x.shape[3], overlap_size): - row = [] - for j in range(0, x.shape[4], overlap_size): + + row_positions = list(range(0, x.shape[3], overlap_size)) + col_positions = list(range(0, x.shape[4], overlap_size)) + enc = None + output_width = 0 + h_cursor = 0 + + for _row_idx, i in enumerate(row_positions): + row_tiles = [] + for tile_idx, j in enumerate(col_positions): tile = x[:, :, :, i : i + self.tile_sample_min_size, j : j + self.tile_sample_min_size] tile = self._encode(tile) - row.append(tile) + if tile.ndim == 4: + tile = tile.unsqueeze(0) + if prev_row is not None: + tile = self.blend_v(prev_row[tile_idx], tile, blend_extent) + if tile_idx > 0: + tile = self.blend_h(row_tiles[-1], tile, blend_extent) + row_tiles.append(tile) self.tiles += 1 - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=4)) - enc = torch.cat(result_rows, dim=3) - return enc + + cropped_tiles = [tile[:, :, :, :row_limit, :row_limit] for tile in row_tiles] + row_width = 0 + for cropped in cropped_tiles: + row_width += cropped.shape[-1] + if output_width < cropped.shape[-1]: + output_width = cropped.shape[-1] + + if enc is None: + enc = torch.empty( + cropped_tiles[0].shape[0], + cropped_tiles[0].shape[1], + cropped_tiles[0].shape[2], + len(row_positions) * row_limit, + len(col_positions) * row_limit, + dtype=cropped_tiles[0].dtype, + device=cropped_tiles[0].device, + ) + + w_cursor = 0 + for cropped in cropped_tiles: + enc[:, :, :, h_cursor : h_cursor + cropped.shape[-2], w_cursor : w_cursor + cropped.shape[-1]] = cropped + w_cursor += cropped.shape[-1] + + h_cursor += cropped_tiles[0].shape[-2] + prev_row = row_tiles + + return enc[:, :, :, :h_cursor, :w_cursor] def tiled_decode(self, z: torch.Tensor) -> torch.Tensor: overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor)) blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor) row_limit = self.tile_sample_min_size - blend_extent - rows = [] - for i in range(0, z.shape[3], overlap_size): - row = [] - for j in range(0, z.shape[4], overlap_size): + prev_row = None + + row_positions = list(range(0, z.shape[3], overlap_size)) + col_positions = list(range(0, z.shape[4], overlap_size)) + dec = None + output_width = 0 + h_cursor = 0 + + for _row_idx, i in enumerate(row_positions): + row_tiles = [] + for tile_idx, j in enumerate(col_positions): tile = z[:, :, :, i : i + self.tile_latent_min_size, j : j + self.tile_latent_min_size] decoded = self.decoder(tile) - row.append(decoded) - rows.append(row) - result_rows = [] - for i, row in enumerate(rows): - result_row = [] - for j, tile in enumerate(row): - if i > 0: - tile = self.blend_v(rows[i - 1][j], tile, blend_extent) - if j > 0: - tile = self.blend_h(row[j - 1], tile, blend_extent) - result_row.append(tile[:, :, :, :row_limit, :row_limit]) - result_rows.append(torch.cat(result_row, dim=4)) - dec = torch.cat(result_rows, dim=3) - return dec + if decoded.ndim == 4: + decoded = decoded.unsqueeze(0) + if prev_row is not None: + decoded = self.blend_v(prev_row[tile_idx], decoded, blend_extent) + if tile_idx > 0: + decoded = self.blend_h(row_tiles[-1], decoded, blend_extent) + row_tiles.append(decoded) + + cropped_tiles = [tile[:, :, :, :row_limit, :row_limit] for tile in row_tiles] + row_width = 0 + for cropped in cropped_tiles: + row_width += cropped.shape[-1] + if output_width < cropped.shape[-1]: + output_width = cropped.shape[-1] + + if dec is None: + dec = torch.empty( + cropped_tiles[0].shape[0], + cropped_tiles[0].shape[1], + cropped_tiles[0].shape[2], + len(row_positions) * row_limit, + len(col_positions) * row_limit, + dtype=cropped_tiles[0].dtype, + device=cropped_tiles[0].device, + ) + + w_cursor = 0 + for cropped in cropped_tiles: + dec[:, :, :, h_cursor : h_cursor + cropped.shape[-2], w_cursor : w_cursor + cropped.shape[-1]] = cropped + w_cursor += cropped.shape[-1] + + h_cursor += cropped_tiles[0].shape[-2] + prev_row = row_tiles + + return dec[:, :, :, :h_cursor, :w_cursor] def forward( self, x: torch.FloatTensor, mode: Literal["encode", "decode", "all"] = "all", **kwargs @@ -1330,7 +1405,6 @@ 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) diff --git a/modules/seedvr/src/models/video_vae_v3/modules/causal_inflation_lib.py b/modules/seedvr/src/models/video_vae_v3/modules/causal_inflation_lib.py index c6d35f0cb..98b04772a 100644 --- a/modules/seedvr/src/models/video_vae_v3/modules/causal_inflation_lib.py +++ b/modules/seedvr/src/models/video_vae_v3/modules/causal_inflation_lib.py @@ -150,6 +150,8 @@ class InflatedCausalConv3d(Conv3d): assert memory_state != MemoryState.UNSET if memory_state != MemoryState.ACTIVE: self.memory = None + if torch.is_tensor(input) and memory_state == MemoryState.DISABLED: + return self.basic_forward(input, memory_state) if ( math.isinf(self.memory_limit) and torch.is_tensor(input) diff --git a/modules/seedvr/src/optimization/memory_manager.py b/modules/seedvr/src/optimization/memory_manager.py index a89b84812..3464cf229 100644 --- a/modules/seedvr/src/optimization/memory_manager.py +++ b/modules/seedvr/src/optimization/memory_manager.py @@ -84,27 +84,12 @@ def clear_rope_cache(runner) -> None: runner: The model runner containing the cache """ if hasattr(runner, 'cache') and hasattr(runner.cache, 'cache'): - # Count entries before cleanup - len(runner.cache.cache) - - # Free all tensors from cache - for _key, value in runner.cache.cache.items(): - if isinstance(value, (tuple, list)): - for item in value: - if hasattr(item, 'cpu'): - item.cpu() - del item - elif hasattr(value, 'cpu'): - value.cpu() - del value - - # Clear the cache - runner.cache.cache.clear() + runner.cache.clear() if hasattr(runner, 'dit'): - cleared_lru_count = 0 for module in runner.dit.modules(): if isinstance(module, RotaryEmbeddingBase): if hasattr(module.get_axial_freqs, 'cache_clear'): module.get_axial_freqs.cache_clear() - cleared_lru_count += 1 + if hasattr(module, 'cache') and hasattr(module.cache, 'clear'): + module.cache.clear() diff --git a/modules/seedvr/src/optimization/performance.py b/modules/seedvr/src/optimization/performance.py index 99d99fbfd..a4847269f 100644 --- a/modules/seedvr/src/optimization/performance.py +++ b/modules/seedvr/src/optimization/performance.py @@ -54,7 +54,7 @@ def optimized_video_rearrange(video_tensors: List[torch.Tensor]) -> List[torch.T batch_3d = batch_3d.permute(0, 2, 1, 3, 4) # [batch, 1, c, h, w] for i, idx in enumerate(indices_3d): - samples[idx] = batch_3d[i] # [1, c, h, w] + samples[idx] = batch_3d[i].contiguous() # [1, c, h, w] # 🚀 BATCH PROCESSING for 4D videos (c t h w -> t c h w) if videos_4d: @@ -67,13 +67,12 @@ def optimized_video_rearrange(video_tensors: List[torch.Tensor]) -> List[torch.T batch_4d = batch_4d.permute(0, 2, 1, 3, 4) # [batch, t, c, h, w] for i, idx in enumerate(indices_4d): - samples[idx] = batch_4d[i] # [t, c, h, w] + samples[idx] = batch_4d[i].contiguous() # [t, c, h, w] else: # 🔄 FALLBACK: Different shapes, optimized individual processing for i, idx in enumerate(indices_4d): # Use permute instead of rearrange (faster) - samples[idx] = videos_4d[i].permute(1, 0, 2, 3) # c t h w -> t c h w - + samples[idx] = videos_4d[i].permute(1, 0, 2, 3).contiguous() # c t h w -> t c h w return samples diff --git a/modules/seedvr/src/utils/color_fix.py b/modules/seedvr/src/utils/color_fix.py index 7d95cb50d..f7620c79e 100644 --- a/modules/seedvr/src/utils/color_fix.py +++ b/modules/seedvr/src/utils/color_fix.py @@ -2,8 +2,9 @@ import torch from PIL import Image from torch import Tensor from torch.nn import functional as F -from modules.seedvr.src.common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation from torchvision.transforms import ToTensor, ToPILImage +from modules.seedvr.src.common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation + def adain_color_fix(target: Image.Image, source: Image.Image): # Convert images to tensors @@ -118,6 +119,10 @@ def wavelet_reconstruction(content_feat:Tensor, style_feat:Tensor): align_corners=False ) + # align devices so reconstruction does not mix CPU and GPU tensors + if style_feat.device != content_feat.device: + style_feat = style_feat.to(content_feat.device) + # calculate the wavelet decomposition of the content feature content_high_freq, content_low_freq = wavelet_decomposition(content_feat) del content_low_freq diff --git a/modules/timer.py b/modules/timer.py index 01b8dbf94..dcb2b642b 100644 --- a/modules/timer.py +++ b/modules/timer.py @@ -28,6 +28,13 @@ class Timer: self.records[name] = 0 self.records[name] += t + def rm(self, name): + if name in self.records: + del self.records[name] + + def get(self, name): + return self.records.get(name, 0) + def ts(self, name, t): elapsed = time.time() - t self.add(name, elapsed) diff --git a/modules/ui_common.py b/modules/ui_common.py index 0b52a673e..78583ea9f 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -62,14 +62,18 @@ def infotext_to_html(text): res.pop('Negative template', None) runtime = {} - runtime['App'] = res.get('App', '') - res.pop('App', None) - runtime['Version'] = res.get('Version', '') - res.pop('Version', None) - runtime['Pipeline'] = res.get('Pipeline', '') - res.pop('Pipeline', None) - runtime['Operations'] = res.get('Operations', '') - res.pop('Operations', None) + if 'App' in res: + runtime['App'] = res.get('App', '') + res.pop('App', None) + if 'Version' in res: + runtime['Version'] = res.get('Version', '') + res.pop('Version', None) + if 'Pipeline' in res: + runtime['Pipeline'] = res.get('Pipeline', '') + res.pop('Pipeline', None) + if 'Operations' in res: + runtime['Operations'] = res.get('Operations', '') + res.pop('Operations', None) params = [f'{k}: {v}' for k, v in res.items() if v is not None and not k.endswith('-1') and not k.endswith('-2')] params = '| '.join(params) if len(params) > 0 else '' diff --git a/modules/ui_extra_networks_styles.py b/modules/ui_extra_networks_styles.py index 1154b7b7f..793e6d836 100644 --- a/modules/ui_extra_networks_styles.py +++ b/modules/ui_extra_networks_styles.py @@ -77,8 +77,8 @@ class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage): name = getattr(style, 'name', '') if name == '': return item - txt = f'Prompt: {getattr(style, "prompt", "")}' - if len(getattr(style, 'negative_prompt', '')) > 0: + txt = f'Prompt: {getattr(style, "prompt", "") or ""}' + if len(getattr(style, 'negative_prompt', '') or '') > 0: txt += f'\nNegative: {style.negative_prompt}' item = { "type": 'Style', diff --git a/modules/ui_postprocessing.py b/modules/ui_postprocessing.py index e0e6c67cb..83e21d891 100644 --- a/modules/ui_postprocessing.py +++ b/modules/ui_postprocessing.py @@ -95,7 +95,7 @@ def create_ui(): submit.click( _js="submit_postprocessing", - fn=call_queue.wrap_gradio_gpu_call(submit_process, extra_outputs=[None, ''], name='Postprocess'), + fn=call_queue.wrap_gradio_gpu_call(submit_process, extra_outputs=[None, None, ''], name='Postprocess'), inputs=[ tab_index, extras_image, diff --git a/scripts/postprocessing_seedvr.py b/scripts/postprocessing_seedvr.py index 120b8427d..2796a3d46 100644 --- a/scripts/postprocessing_seedvr.py +++ b/scripts/postprocessing_seedvr.py @@ -16,20 +16,36 @@ class ScriptSeedVR(scripts_postprocessing.ScriptPostprocessing): seedvr_offload = gr.Checkbox(label="Offload model", value=True, elem_id="extras_seedvr_offload") 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_steps = gr.Slider(step=1, value=1, minimum=1, maximum=99, label="SeedVR steps", elem_id="extras_seedvr_steps") - with gr.Row(): - seedvr_seed = gr.Number(step=1, value=-1, label="SeedVR seed", elem_id="extras_seedvr_seed") - with gr.Row(): - seedvr_cfg_scale = gr.Slider(minimum=0.0, maximum=15.0, step=0.01, value=1.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") - with gr.Row(): - seedvr_batch_size = gr.Slider(minimum=1, maximum=64, step=1, value=1, label="SeedVR batch size", elem_id="extras_seedvr_batch_size") - seedvr_batch_overlap = gr.Slider(minimum=0, maximum=16, step=1, value=0, label="SeedVR batch overlap", elem_id="extras_seedvr_batch_overlap") + with gr.Accordion('SeedVR advanced', open = False, elem_id="postprocess_seedvr_advanced_accordion"): + with gr.Row(): + seedvr_steps = gr.Slider(step=1, value=1, minimum=1, maximum=99, label="SeedVR steps", elem_id="extras_seedvr_steps") + seedvr_seed = gr.Number(step=1, value=-1, label="SeedVR seed", elem_id="extras_seedvr_seed") + with gr.Row(): + seedvr_cfg_scale = gr.Slider(minimum=0.0, maximum=15.0, step=0.01, value=1.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.Accordion('SeedVR VAE', open = False, elem_id="postprocess_seedvr_vae_accordion"): + with gr.Row(): + seedvr_vae_tile_encode = gr.Checkbox(label="VAE tiled encode", value=True, elem_id="extras_seedvr_vae_tile_encode") + seedvr_vae_tile_decode = gr.Checkbox(label="VAE tiled decode", value=True, elem_id="extras_seedvr_vae_tile_decode") + 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") + with gr.Row(): + seedvr_vae_memory = gr.Slider(minimum=0.1, maximum=1.0, step=0.01, value=1.0, label="SeedVR VAE memory", elem_id="extras_seedvr_vae_memory") + with gr.Accordion('SeedVR video', open = False, elem_id="postprocess_seedvr_video_accordion"): + with gr.Row(): + seedvr_batch_size = gr.Slider(minimum=1, maximum=64, step=1, value=1, label="SeedVR batch size", elem_id="extras_seedvr_batch_size") + seedvr_batch_overlap = gr.Slider(minimum=0, maximum=16, step=1, value=0, label="SeedVR batch overlap", elem_id="extras_seedvr_batch_overlap") + with gr.Row(): + seedvr_interpolate = gr.Slider(label="RIFE interpolate frames", minimum=0, maximum=4, step=1, value=0, elem_id="extras_seedvr_interpolate") + with gr.Row(): + from modules.video_models.video_utils import get_codecs + from modules.ui_common import create_refresh_button + seedvr_codec = gr.Dropdown(label="Video codec", choices=['none', 'libx264'], value='libx264', type='value') + create_refresh_button(seedvr_codec, get_codecs, elem_id="video_mp4_codec_refresh") + seedvr_codec_opt = gr.Textbox(label="Video options", value='crf:16', elem_id="video_mp4_opt") + return { "seedvr_enabled": seedvr_enabled, "seedvr_selected": seedvr_selected, @@ -43,6 +59,12 @@ class ScriptSeedVR(scripts_postprocessing.ScriptPostprocessing): "seedvr_batch_size": seedvr_batch_size, "seedvr_batch_overlap": seedvr_batch_overlap, "seedvr_offload": seedvr_offload, + "seedvr_interpolate": seedvr_interpolate, + "seedvr_codec": seedvr_codec, + "seedvr_codec_opt": seedvr_codec_opt, + "seedvr_vae_memory": seedvr_vae_memory, + "seedvr_vae_tile_encode": seedvr_vae_tile_encode, + "seedvr_vae_tile_decode": seedvr_vae_tile_decode, } def process(self, @@ -58,20 +80,23 @@ class ScriptSeedVR(scripts_postprocessing.ScriptPostprocessing): seedvr_tile_overlap: float, seedvr_batch_size: int, seedvr_batch_overlap: int, - seedvr_offload: bool + seedvr_offload: bool, + seedvr_interpolate: int, + seedvr_codec: str, + seedvr_codec_opt: str, + seedvr_vae_memory: float, + seedvr_vae_tile_encode: bool, + seedvr_vae_tile_decode: bool, ): # pylint: disable=arguments-differ if not seedvr_enabled: return from modules import shared, upscaler - from modules.logger import log _input = pp.image or pp.video if _input is None: return instance: upscaler.UpscalerData = next(iter([x for x in shared.sd_upscalers if x.name == seedvr_selected]), None) scaler: UpscalerSeedVR = 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} batch_size={seedvr_batch_size} batch_overlap={seedvr_batch_overlap}') - jobid = shared.state.begin('Upscale') scaler.scale = float(seedvr_scale) @@ -85,7 +110,13 @@ class ScriptSeedVR(scripts_postprocessing.ScriptPostprocessing): tile_overlap=seedvr_tile_overlap, batch_size=seedvr_batch_size, batch_overlap=seedvr_batch_overlap, - offload=seedvr_offload + offload=seedvr_offload, + interpolate=seedvr_interpolate, + codec=seedvr_codec, + codec_opt=seedvr_codec_opt, + vae_memory=seedvr_vae_memory, + vae_tile_encode=seedvr_vae_tile_encode, + vae_tile_decode=seedvr_vae_tile_decode, ) shared.state.end(jobid)