seedvr enhancements

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2026-07-21 15:59:45 +02:00
parent 315719f9e5
commit 53e69c7ec2
19 changed files with 385 additions and 201 deletions
+116 -55
View File
@@ -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
+5 -2
View File
@@ -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:
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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
+15
View File
@@ -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
+23 -32
View File
@@ -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)
+12 -5
View File
@@ -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]
@@ -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)
@@ -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)
@@ -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()
@@ -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
+6 -1
View File
@@ -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
+7
View File
@@ -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)
+12 -8
View File
@@ -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 ''
+2 -2
View File
@@ -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',
+1 -1
View File
@@ -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,
+49 -18
View File
@@ -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)