mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
Submodule extensions-builtin/sdnext-modernui updated: 820a266789...7cc2f614a5
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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 ''
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user