import torch
from torchvision.transforms import Compose, Lambda, Normalize
from modules.seedvr.src.optimization.performance import optimized_video_rearrange, optimized_single_video_rearrange, optimized_sample_to_image_format
from modules.seedvr.src.common.seed import set_seed
from modules.seedvr.src.data.image.transforms.divisible_crop import DivisibleCrop
from modules.seedvr.src.data.image.transforms.na_resize import NaResize
from modules.seedvr.src.utils.color_fix import wavelet_reconstruction
def generation_step(runner, text_embeds_dict, cond_latents, temporal_overlap, device):
"""
Execute a single generation step with adaptive dtype handling
Args:
runner: SeedVRPipeline instance
text_embeds_dict (dict): Text embeddings for positive and negative prompts
cond_latents (list): Conditional latents for generation
temporal_overlap (int): Number of frames for temporal overlap
Returns:
tuple: (samples, last_latents) for potential temporal continuation
Features:
- Adaptive dtype detection (FP8/FP16/BFloat16)
- Optimal autocast configuration for each model type
- Memory-efficient noise generation and reuse
- Automatic device placement with dtype preservation
- Advanced inference optimization
"""
model_dtype = next(runner.dit.parameters()).dtype
if model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
dtype = torch.bfloat16
elif model_dtype == torch.float16:
dtype = torch.float16
else:
dtype = torch.bfloat16
def _move_to_cuda(x):
"""Move tensors to CUDA with adaptive optimal dtype"""
return [i.to(device, dtype=dtype) for i in x]
with torch.cuda.device(device):
base_noise = torch.randn_like(cond_latents[0], dtype=dtype)
noises = [base_noise]
aug_noises = [base_noise * 0.1 + torch.randn_like(base_noise) * 0.05]
noises, aug_noises, cond_latents = _move_to_cuda(noises), _move_to_cuda(aug_noises), _move_to_cuda(cond_latents)
cond_noise_scale = 0.0
def _add_noise(x, aug_noise):
t = (
torch.tensor([1000.0], device=device, dtype=dtype)
* cond_noise_scale
)
shape = torch.tensor(x.shape[1:], device=device)[None]
t = runner.timestep_transform(t, shape)
x = runner.schedule.forward(x, aug_noise, t)
return x
condition = runner.get_condition(
noises[0],
task="sr",
latent_blur=_add_noise(cond_latents[0], aug_noises[0]),
)
conditions = [condition]
with torch.no_grad():
video_tensors = runner.inference(
noises=noises,
conditions=conditions,
temporal_overlap=temporal_overlap,
**text_embeds_dict,
)
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")
conditions = conditions[0].to("cpu")
condition = condition.to("cpu")
del noises, aug_noises, cond_latents, conditions, condition
return samples
def blend_overlapping_frames(prev_tail, cur_head, overlap):
"""Crossfade the previous batch tail into the current batch head: Hann window from 3 frames, linear below."""
if overlap >= 3:
t = torch.linspace(0.0, 1.0, steps=overlap, dtype=torch.float32)
u = ((t - 1.0 / 3.0) * 3.0).clamp(0.0, 1.0)
w_prev = 0.5 + 0.5 * torch.cos(torch.pi * u)
else:
w_prev = torch.linspace(1.0, 0.0, steps=overlap, dtype=torch.float32)
w_prev = w_prev.view(overlap, 1, 1, 1).to(prev_tail.device)
blended = prev_tail.float() * w_prev + cur_head.float() * (1.0 - w_prev)
return blended.to(prev_tail.dtype)
def cut_videos(videos):
t = videos.size(1)
if t % 4 == 1:
return videos
padding_needed = (4 - (t % 4)) % 4 + 1
last_frame = videos[:, -1:].expand(-1, padding_needed, -1, -1).contiguous()
result = torch.cat([videos, last_frame], dim=1)
return result
def generation_loop(runner,
images,
cfg_scale=1.5,
cfg_rescale=0.0,
steps=1,
seed=-1,
res_w=720,
batch_size=1,
temporal_overlap=0,
progress_callback=None,
device:str='cpu',
color_reconstruct=True,
):
"""
Main generation loop with context-aware temporal processing
Args:
runner: SeedVRPipeline instance
images (torch.Tensor): Input images for upscaling
cfg_scale (float): Classifier-free guidance scale
seed (int): Random seed for reproducibility
res_w (int): Target resolution width
batch_size (int): Batch size for processing
temporal_overlap (int): Frames for temporal continuity
progress_callback (callable): Optional callback for progress reporting
Returns:
torch.Tensor: Generated video frames
Features:
- Context-aware generation with temporal overlap
- Adaptive dtype pipeline (FP8/FP16/BFloat16)
- Memory-optimized batch processing
- Advanced video transformation pipeline
- Intelligent VRAM management throughout process
- Real-time progress reporting
"""
model_dtype = None
model_dtype = next(runner.dit.parameters()).dtype
compute_dtype = model_dtype
runner.config.diffusion.cfg.scale = cfg_scale
runner.config.diffusion.cfg.rescale = cfg_rescale
runner.config.diffusion.timesteps.sampling.steps = steps
runner.configure_diffusion()
set_seed(seed)
video_transform = Compose([
NaResize(
resolution=(res_w),
mode="side",
downsample_only=False,
),
Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),
DivisibleCrop((16, 16)),
Normalize(0.5, 0.5),
Lambda(lambda x: x.permute(1, 0, 2, 3)),
])
final_video_images = None
text_embeds = {"texts_pos": [runner.text_pos_embeds], "texts_neg": [runner.text_neg_embeds]}
step = batch_size - temporal_overlap
if step <= 0:
step = batch_size
temporal_overlap = 0
total_batches = len(range(0, len(images), step))
for batch_count, batch_idx in enumerate(range(0, len(images), step)):
if batch_idx == 0:
start_idx = 0
end_idx = min(batch_size, len(images))
effective_batch_size = end_idx - start_idx
else:
start_idx = batch_idx
end_idx = min(start_idx + batch_size, len(images))
effective_batch_size = end_idx - start_idx
if effective_batch_size <= temporal_overlap:
break
current_frames = end_idx - start_idx
video = images[start_idx:end_idx]
video = video.permute(0, 3, 1, 2).to(device, dtype=compute_dtype)
transformed_video = video_transform(video)
del video
ori_lengths = [transformed_video.size(1)]
t = transformed_video.size(1)
if len(images) >= 5 and t % 4 != 1:
transformed_video = cut_videos(transformed_video)
cond_latents = runner.vae_encode([transformed_video])
samples = generation_step(runner, text_embeds, cond_latents=cond_latents, temporal_overlap=temporal_overlap, device=device)
if samples is None:
return
del cond_latents
sample = samples[0]
del samples
if ori_lengths[0] < sample.shape[0]:
sample = sample[:ori_lengths[0]].contiguous()
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
sample = optimized_sample_to_image_format(sample)
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
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)
write_start = start_idx
if batch_idx > 0 and temporal_overlap > 0:
prev_tail = final_video_images[start_idx:start_idx + temporal_overlap]
final_video_images[start_idx:start_idx + temporal_overlap] = blend_overlapping_frames(prev_tail, sample[:temporal_overlap], temporal_overlap)
sample = sample[temporal_overlap:]
write_start = start_idx + temporal_overlap
final_video_images[write_start:write_start + sample.shape[0]] = sample
del sample
if progress_callback:
progress_callback(batch_count+1, total_batches, current_frames, "Processing batch...")
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)
return final_video_images
def prepare_video_transforms(res_w):
"""
Prepare optimized video transformation pipeline
Args:
res_w (int): Target resolution width
Returns:
Compose: Configured transformation pipeline
Features:
- Resolution-aware upscaling (no downsampling)
- Proper normalization for model compatibility
- Memory-efficient tensor operations
"""
return Compose([
NaResize(
resolution=(res_w),
mode="side",
downsample_only=False,
),
Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),
DivisibleCrop((16, 16)),
Normalize(0.5, 0.5),
Lambda(lambda x: x.permute(1, 0, 2, 3)),
])
def calculate_optimal_batch_params(total_frames, batch_size, temporal_overlap):
"""
Calculate optimal batch processing parameters
Args:
total_frames (int): Total number of frames
batch_size (int): Desired batch size
temporal_overlap (int): Temporal overlap frames
Returns:
dict: Optimized parameters and recommendations
Features:
- 4n+1 constraint optimization
- Padding waste calculation
- Performance recommendations
"""
step = batch_size - temporal_overlap
if step <= 0:
step = batch_size
temporal_overlap = 0
optimal_batches = [x for x in [i for i in range(1, 200) if i % 4 == 1] if x <= total_frames]
best_batch = max(optimal_batches) if optimal_batches else 1
padding_waste = 0
if batch_size not in optimal_batches:
padding_waste = sum(((i // 4) + 1) * 4 + 1 - i for i in range(batch_size, total_frames, batch_size))
return {
'step': step,
'temporal_overlap': temporal_overlap,
'best_batch': best_batch,
'padding_waste': padding_waste,
'is_optimal': batch_size in optimal_batches
}