| name | seedvr2-video-restoration |
| title | SeedVR2: One-Step Video Restoration via Diffusion Adversarial Post-Training |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.05301 |
| keywords | ["video-restoration","diffusion-models","adversarial-training","efficient-inference","high-resolution"] |
| description | Achieves single-step video restoration at 1080p resolution with 4x speedup over multi-step diffusion approaches via adversarial training, adaptive window attention, and feature matching loss. |
SeedVR2: One-Step Video Restoration
Core Concept
SeedVR2 tackles computational inefficiency in diffusion-based video restoration by eliminating multi-step iterative refinement. Rather than requiring 64+ denoising steps, the model completes restoration in a single forward pass while maintaining or improving quality over iterative methods. This is achieved through adversarial training against real data, learnable attention mechanisms that adapt to input resolution, and efficient loss functions optimized for high-resolution processing.
Architecture Overview
- Diffusion Transformer Base: Swin-MMDIT architecture with 16B total parameters (generator + discriminator)
- Adaptive Window Attention: Dynamic spatial window sizing that scales to arbitrary resolutions while maintaining computational efficiency
- Causal Video VAE: Temporal compression for efficient processing of multi-frame sequences
- Feature Matching Loss: Efficient alternative to LPIPS for adversarial training on high-resolution video without costly pixel-space decoding
- Progressive Distillation: Training from 64-step teacher model with gradual temporal length increase for stability
- RpGAN Stabilization: Approximate R2 regularization to maintain adversarial training stability across thousands of iterations
Implementation
The following code illustrates the core architectural components:
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Tuple, Optional
class AdaptiveWindowAttention(nn.Module):
"""
Adaptive window sizing for efficient high-resolution video processing.
"""
def __init__(self, dim: int, window_base_size: int = 7):
super().__init__()
self.dim = dim
self.window_base_size = window_base_size
() -> [, ]:
scale_h = h /
scale_w = w /
adaptive_h = (.window_base_size * (scale_h ** ))
adaptive_w = (.window_base_size * (scale_w ** ))
adaptive_h = adaptive_h adaptive_h % == adaptive_h +
adaptive_w = adaptive_w adaptive_w % == adaptive_w +
adaptive_h, adaptive_w
() -> torch.Tensor:
B, T, C, H, W = x.shape
win_h, win_w = .compute_adaptive_window(H, W)
x_flat = x.view(B * T, C, H, W)
x_out = F.avg_pool2d(x_flat, kernel_size=(H // win_h, ))
x_out = F.interpolate(x_out, size=(H, W), mode=)
x_out.view(B, T, C, H, W)
(nn.Module):
():
().__init__()
.feature_extractor = feature_extractor
() -> torch.Tensor:
gen_features = .feature_extractor(generated)
real_features = .feature_extractor(real)
loss = F.mse_loss(gen_features, real_features)
loss
:
():
.generator = generator
.discriminator = discriminator
.gen_optimizer = torch.optim.Adam(generator.parameters(), lr=learning_rate)
.dis_optimizer = torch.optim.Adam(discriminator.parameters(), lr=learning_rate)
.feature_loss = FeatureMatchingLoss(._build_feature_extractor())
() -> nn.Module:
nn.Sequential(
nn.Conv2d(, , kernel_size=, stride=, padding=),
nn.ReLU(),
nn.Conv2d(, , kernel_size=, stride=, padding=),
nn.ReLU(),
)
() -> [, ]:
restored = .generator(degraded_video)
dis_fake = .discriminator(restored)
gen_loss = -dis_fake.mean()
gen_loss += * .feature_loss(restored, real_video)
.gen_optimizer.zero_grad()
gen_loss.backward()
.gen_optimizer.step()
dis_real = .discriminator(real_video)
dis_fake = .discriminator(restored.detach())
dis_loss = F.softplus(dis_fake).mean() + F.softplus(-dis_real).mean()
r2_reg = * (dis_real ** ).mean()
dis_loss = dis_loss + r2_reg
.dis_optimizer.zero_grad()
dis_loss.backward()
.dis_optimizer.step()
(gen_loss), (dis_loss)