Initial release - FL DiffVSR video super-resolution nodes

- 4x video upscaling with temporal coherence using Stream-DiffVSR
- Model Loader node with precision/device options and xformers support
- Upscaler node with chunked processing for memory efficiency
- Automatic model download from HuggingFace
- Text-guided upscaling support

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
Fill
2026-01-05 22:09:14 -08:00
co-authored by Claude Opus 4.5
commit 37aa6cf9ce
25 changed files with 2417 additions and 0 deletions
+75
View File
@@ -0,0 +1,75 @@
# Python
__pycache__/
*.py[cod]
*$py.class
*.so
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
*.egg-info/
.installed.cfg
*.egg
# Virtual environments
venv/
ENV/
env/
.venv/
# IDE
.idea/
.vscode/
*.swp
*.swo
*~
.project
.pydevproject
.settings/
# OS
.DS_Store
.DS_Store?
._*
.Spotlight-V100
.Trashes
ehthumbs.db
Thumbs.db
# Logs
*.log
logs/
# Model files (downloaded separately)
*.safetensors
*.bin
*.ckpt
*.pt
*.pth
# Temporary files
*.tmp
*.temp
.cache/
# Jupyter
.ipynb_checkpoints/
*.ipynb
# Testing
.pytest_cache/
.coverage
htmlcov/
# Distribution
*.tar.gz
*.zip
+82
View File
@@ -0,0 +1,82 @@
# FL DiffVSR
Diffusion-based video super-resolution nodes for ComfyUI powered by Stream-DiffVSR. Upscale videos 4x with temporal coherence for smooth, artifact-free results.
[![Stream-DiffVSR](https://img.shields.io/badge/Stream--DiffVSR-Original%20Paper-blue?style=for-the-badge&logo=arxiv&logoColor=white)](https://arxiv.org/abs/2512.23709)
[![Patreon](https://img.shields.io/badge/Patreon-Support%20Me-F96854?style=for-the-badge&logo=patreon&logoColor=white)](https://www.patreon.com/Machinedelusions)
![Workflow Preview](assets/readme.png)
## Features
- **4x Video Upscaling** - Upscale video frames to 4x resolution with high fidelity
- **Temporal Coherence** - Maintains consistency across frames for flicker-free results
- **Diffusion-Based** - Leverages diffusion models for superior detail reconstruction
- **Text Guidance** - Optional prompt support for guided upscaling
- **Memory Efficient** - Chunked processing and xformers support for lower VRAM usage
- **Automatic Downloads** - Models download automatically from HuggingFace on first use
## Nodes
| Node | Description |
|------|-------------|
| **FL DiffVSR Load Model** | Downloads and loads Stream-DiffVSR model from HuggingFace |
| **FL DiffVSR Upscale** | Upscales video frames with temporal coherence |
## Installation
### ComfyUI Manager
Search for "FL DiffVSR" and install.
### Manual
```bash
cd ComfyUI/custom_nodes
git clone https://github.com/filliptm/ComfyUI-FL-DiffVSR.git
cd ComfyUI-FL-DiffVSR
pip install -r requirements.txt
```
## Quick Start
1. Add **FL DiffVSR Load Model** node and configure precision/device settings
2. Connect to **FL DiffVSR Upscale** node
3. Feed your video frames as an IMAGE batch
4. Adjust inference steps (4 recommended for speed/quality balance)
5. Generate upscaled frames
## Parameters
### Model Loader
| Parameter | Options | Description |
|-----------|---------|-------------|
| precision | auto, fp32, fp16, bf16 | Model precision (auto selects fp16 for GPU) |
| device | auto, cuda, cpu | Target device for inference |
| enable_xformers | true/false | Enable memory-efficient attention |
### Upscaler
| Parameter | Default | Description |
|-----------|---------|-------------|
| inference_steps | 4 | Denoising steps (higher = better quality, slower) |
| guidance_scale | 0.0 | CFG scale (0 = no guidance) |
| chunk_size | 8 | Frames per batch (lower = less VRAM) |
| prompt | "" | Optional text guidance |
| negative_prompt | "" | Optional negative prompt |
| seed | -1 | Random seed (-1 for random) |
## Requirements
- Python 3.10+
- 8GB VRAM minimum (16GB+ recommended for larger videos)
- NVIDIA GPU recommended (CPU supported but slow)
## Model
The Stream-DiffVSR model downloads automatically to `ComfyUI/models/stream_diffvsr/` on first use (~2GB).
| Model | Source | Size |
|-------|--------|------|
| Stream-DiffVSR | Jamichsu/Stream-DiffVSR | ~2GB |
## License
Apache 2.0
+30
View File
@@ -0,0 +1,30 @@
"""
FL DiffVSR - ComfyUI Node Pack for Stream-DiffVSR Video Super-Resolution
A standalone ComfyUI node pack for 4x video upscaling with temporal coherence
using diffusion-based super-resolution.
Based on Stream-DiffVSR: https://arxiv.org/abs/2512.23709
Model: Jamichsu/Stream-DiffVSR (HuggingFace)
"""
from .nodes.model_loader import FL_DiffVSR_LoadModel
from .nodes.upscaler import FL_DiffVSR_Upscale
NODE_CLASS_MAPPINGS = {
"FL_DiffVSR_LoadModel": FL_DiffVSR_LoadModel,
"FL_DiffVSR_Upscale": FL_DiffVSR_Upscale,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"FL_DiffVSR_LoadModel": "FL DiffVSR Load Model",
"FL_DiffVSR_Upscale": "FL DiffVSR Upscale",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
print("\n" + "=" * 50)
print("FL DiffVSR - Video Super-Resolution Node Pack")
print("4x Upscaling with Temporal Coherence")
print("=" * 50 + "\n")
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 310 KiB

+1
View File
@@ -0,0 +1 @@
# FL DiffVSR Core
+114
View File
@@ -0,0 +1,114 @@
"""
Model Manager for FL DiffVSR
Handles downloading and managing Stream-DiffVSR models from HuggingFace.
"""
import os
from pathlib import Path
from typing import Optional
import folder_paths
class ModelManager:
"""Manages Stream-DiffVSR model downloading and paths."""
HF_REPO_ID = "Jamichsu/Stream-DiffVSR"
MODEL_FOLDER_NAME = "stream_diffvsr"
def __init__(self):
self.models_dir = self._get_models_dir()
def _get_models_dir(self) -> str:
"""Get the path to the stream_diffvsr models directory."""
models_base = folder_paths.models_dir
models_dir = os.path.join(models_base, self.MODEL_FOLDER_NAME)
os.makedirs(models_dir, exist_ok=True)
return models_dir
def get_model_path(self) -> str:
"""Get the path where models are/will be stored."""
return self.models_dir
def check_models_exist(self) -> bool:
"""Check if all required model files exist."""
required_dirs = [
"unet",
"controlnet",
"vae",
"text_encoder",
"tokenizer",
"scheduler",
]
for dir_name in required_dirs:
dir_path = os.path.join(self.models_dir, dir_name)
if not os.path.exists(dir_path):
return False
# Check for at least one file in each directory
if not any(os.scandir(dir_path)):
return False
return True
def download_models(self, force: bool = False) -> str:
"""
Download all model components from HuggingFace.
Args:
force: If True, download even if models already exist
Returns:
Path to the downloaded models
"""
if self.check_models_exist() and not force:
print(f"Stream-DiffVSR models already exist at {self.models_dir}")
return self.models_dir
try:
from huggingface_hub import snapshot_download
print(f"Downloading Stream-DiffVSR models from {self.HF_REPO_ID}...")
print(f"Target directory: {self.models_dir}")
local_dir = snapshot_download(
repo_id=self.HF_REPO_ID,
local_dir=self.models_dir,
local_dir_use_symlinks=False,
ignore_patterns=[
"*.md",
"*.txt",
".git*",
"*.py",
"*.yml",
"*.yaml",
],
)
print(f"Models downloaded successfully to: {local_dir}")
return local_dir
except ImportError:
raise ImportError(
"huggingface_hub is required to download models. "
"Please install it with: pip install huggingface_hub"
)
except Exception as e:
raise RuntimeError(f"Failed to download Stream-DiffVSR models: {e}")
def get_component_path(self, component: str) -> str:
"""Get the path to a specific model component."""
return os.path.join(self.models_dir, component)
# Singleton instance
_model_manager: Optional[ModelManager] = None
def get_model_manager() -> ModelManager:
"""Get the singleton ModelManager instance."""
global _model_manager
if _model_manager is None:
_model_manager = ModelManager()
return _model_manager
+310
View File
@@ -0,0 +1,310 @@
"""
Pipeline Wrapper for FL DiffVSR
Wraps the Stream-DiffVSR pipeline for ComfyUI integration.
"""
import torch
import torch.nn.functional as F
from PIL import Image
import numpy as np
from typing import Optional, List, Union
from torchvision.models.optical_flow import raft_large, Raft_Large_Weights
class StreamDiffVSRWrapper:
"""
Wrapper around StreamDiffVSRPipeline for ComfyUI integration.
Handles model loading, optical flow, and frame-by-frame processing.
"""
def __init__(
self,
model_path: str,
device: torch.device,
dtype: torch.dtype,
enable_xformers: bool = True,
):
self.model_path = model_path
self.device = device
self.dtype = dtype
# Load pipeline components
self._load_pipeline(enable_xformers)
# Load optical flow model (RAFT)
self._load_flow_model()
# Temporal state
self.prev_frame_rgb = None
self.frame_count = 0
def _load_pipeline(self, enable_xformers: bool):
"""Load the Stream-DiffVSR pipeline."""
from diffusers import UNet2DConditionModel, ControlNetModel
from transformers import CLIPTextModel, CLIPTokenizer
# Import local modules
from ..stream_diffvsr.temporal_autoencoder.autoencoder_tiny import TemporalAutoencoderTiny
from ..stream_diffvsr.pipeline.stream_diffvsr_pipeline import StreamDiffVSRPipeline
from ..stream_diffvsr.scheduler.ddim_scheduler import DDIMScheduler
print("Loading Stream-DiffVSR pipeline components...")
# Load UNet
print(" Loading UNet...")
self.unet = UNet2DConditionModel.from_pretrained(
self.model_path, subfolder="unet", torch_dtype=self.dtype
).to(self.device)
# Load ControlNet
print(" Loading ControlNet...")
self.controlnet = ControlNetModel.from_pretrained(
self.model_path, subfolder="controlnet", torch_dtype=self.dtype
).to(self.device)
# Load Temporal VAE
print(" Loading Temporal VAE...")
self.vae = TemporalAutoencoderTiny.from_pretrained(
self.model_path, subfolder="vae", torch_dtype=self.dtype
).to(self.device)
# Load text encoder and tokenizer
print(" Loading Text Encoder...")
self.text_encoder = CLIPTextModel.from_pretrained(
self.model_path, subfolder="text_encoder", torch_dtype=self.dtype
).to(self.device)
# Load tokenizer - the Stream-DiffVSR model's tokenizer is incomplete (missing merges.txt)
# So we load from openai/clip-vit-large-patch14 which is the base CLIP model
print(" Loading Tokenizer...")
try:
self.tokenizer = CLIPTokenizer.from_pretrained(
self.model_path, subfolder="tokenizer"
)
except (TypeError, OSError) as e:
print(f" Local tokenizer incomplete, loading from openai/clip-vit-large-patch14...")
self.tokenizer = CLIPTokenizer.from_pretrained(
"openai/clip-vit-large-patch14"
)
# Load scheduler
print(" Loading Scheduler...")
self.scheduler = DDIMScheduler.from_pretrained(
self.model_path, subfolder="scheduler"
)
# Create pipeline
self.pipeline = StreamDiffVSRPipeline(
vae=self.vae,
text_encoder=self.text_encoder,
tokenizer=self.tokenizer,
unet=self.unet,
controlnet=self.controlnet,
scheduler=self.scheduler,
safety_checker=None,
feature_extractor=None,
requires_safety_checker=False,
)
# Enable memory optimizations
if enable_xformers:
try:
self.pipeline.enable_xformers_memory_efficient_attention()
print(" xformers memory efficient attention enabled")
except Exception as e:
print(f" Could not enable xformers: {e}")
print("Stream-DiffVSR pipeline loaded successfully!")
def _load_flow_model(self):
"""Load RAFT optical flow model."""
print("Loading RAFT optical flow model...")
self.flow_model = raft_large(weights=Raft_Large_Weights.DEFAULT)
self.flow_model = self.flow_model.to(self.device).eval()
self.flow_model.requires_grad_(False)
print("RAFT model loaded successfully!")
def reset_temporal_state(self):
"""Reset temporal state for new video sequence."""
self.prev_frame_rgb = None
self.frame_count = 0
self.vae.reset_temporal_condition()
def process_frames_chunked(
self,
images: List[Image.Image],
chunk_size: int = 8,
prompt: str = "",
negative_prompt: str = "",
num_inference_steps: int = 4,
guidance_scale: float = 0.0,
generator: Optional[torch.Generator] = None,
progress_callback: Optional[callable] = None,
) -> List[Image.Image]:
"""
Process frames in memory-efficient chunks.
Args:
images: List of PIL Images to upscale
chunk_size: Number of frames per chunk (lower = less VRAM)
prompt: Text prompt for guidance
negative_prompt: Negative text prompt
num_inference_steps: Number of denoising steps
guidance_scale: CFG scale
generator: Random generator for reproducibility
progress_callback: Optional callback for progress updates
Returns:
List of upscaled PIL Images
"""
all_results = []
prev_frame_rgb = None
prev_upscaled_for_flow = None
total_frames = len(images)
# Calculate number of chunks
num_chunks = (total_frames + chunk_size - 1) // chunk_size
print(f" Processing {total_frames} frames in {num_chunks} chunk(s) of up to {chunk_size} frames each")
for chunk_idx in range(num_chunks):
start_idx = chunk_idx * chunk_size
end_idx = min(start_idx + chunk_size, total_frames)
chunk_images = images[start_idx:end_idx]
print(f" Processing chunk {chunk_idx + 1}/{num_chunks} (frames {start_idx}-{end_idx - 1})")
# Reset temporal state but pass previous frame context
self.vae.reset_temporal_condition()
# Process chunk with previous frame state for continuity
output, prev_frame_rgb, prev_upscaled_for_flow = self.pipeline(
prompt=prompt,
images=chunk_images,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
negative_prompt=negative_prompt if negative_prompt else None,
generator=generator,
of_model=self.flow_model,
of_rescale_factor=1,
output_type="pil",
callback=progress_callback,
callback_steps=1,
# Pass previous frame state for chunk continuity
prev_frame_rgb=prev_frame_rgb,
prev_upscaled_for_flow=prev_upscaled_for_flow,
# Pass frame offset for progress callback
frame_offset=start_idx,
)
# Extract results from this chunk
for img_list in output.images:
if isinstance(img_list, list):
all_results.extend(img_list)
else:
all_results.append(img_list)
# Clear VRAM between chunks
torch.cuda.empty_cache()
return all_results
def process_frames(
self,
images: List[Image.Image],
prompt: str = "",
negative_prompt: str = "",
num_inference_steps: int = 4,
guidance_scale: float = 0.0,
generator: Optional[torch.Generator] = None,
progress_callback: Optional[callable] = None,
chunk_size: int = 0,
) -> List[Image.Image]:
"""
Process a list of frames through the pipeline.
Args:
images: List of PIL Images to upscale
prompt: Text prompt for guidance
negative_prompt: Negative text prompt
num_inference_steps: Number of denoising steps
guidance_scale: CFG scale
generator: Random generator for reproducibility
progress_callback: Optional callback for progress updates
chunk_size: Number of frames per chunk (0 = process all at once)
Returns:
List of upscaled PIL Images
"""
# Use chunked processing if chunk_size > 0 and we have more frames than chunk_size
if chunk_size > 0 and len(images) > chunk_size:
return self.process_frames_chunked(
images=images,
chunk_size=chunk_size,
prompt=prompt,
negative_prompt=negative_prompt,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
progress_callback=progress_callback,
)
# Reset temporal state for new sequence
self.reset_temporal_state()
# Run pipeline (original behavior for small batches or chunk_size=0)
output, _, _ = self.pipeline(
prompt=prompt,
images=images,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
negative_prompt=negative_prompt if negative_prompt else None,
generator=generator,
of_model=self.flow_model,
of_rescale_factor=1,
output_type="pil",
callback=progress_callback,
callback_steps=1,
prev_frame_rgb=None,
prev_upscaled_for_flow=None,
frame_offset=0,
)
# Extract results
results = []
for img_list in output.images:
if isinstance(img_list, list):
results.extend(img_list)
else:
results.append(img_list)
return results
def to(self, device: torch.device):
"""Move all models to specified device."""
self.device = device
self.unet = self.unet.to(device)
self.controlnet = self.controlnet.to(device)
self.vae = self.vae.to(device)
self.text_encoder = self.text_encoder.to(device)
self.flow_model = self.flow_model.to(device)
return self
def tensor_to_pil(tensor: torch.Tensor) -> Image.Image:
"""
Convert a ComfyUI tensor [H, W, C] (0-1 range) to PIL Image.
"""
arr = tensor.cpu().numpy()
arr = (arr * 255).clip(0, 255).astype(np.uint8)
return Image.fromarray(arr, mode='RGB')
def pil_to_tensor(pil_image: Image.Image) -> torch.Tensor:
"""
Convert PIL Image to ComfyUI tensor [1, H, W, C] (0-1 range).
"""
arr = np.array(pil_image).astype(np.float32) / 255.0
tensor = torch.from_numpy(arr)
return tensor.unsqueeze(0) # [H, W, C] -> [1, H, W, C]
+1
View File
@@ -0,0 +1 @@
# FL DiffVSR Nodes
+83
View File
@@ -0,0 +1,83 @@
"""
FL DiffVSR Model Loader Node
Downloads and loads Stream-DiffVSR model from HuggingFace.
"""
import torch
from ..core.model_manager import get_model_manager
from ..core.pipeline_wrapper import StreamDiffVSRWrapper
class FL_DiffVSR_LoadModel:
"""
Load Stream-DiffVSR model from HuggingFace.
Downloads model components to ComfyUI/models/stream_diffvsr/ on first use.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"precision": (["auto", "fp32", "fp16", "bf16"], {"default": "auto"}),
"device": (["auto", "cuda", "cpu"], {"default": "auto"}),
},
"optional": {
"enable_xformers": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("FL_DIFFVSR_MODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "load_model"
CATEGORY = "FL DiffVSR"
def load_model(self, precision: str, device: str, enable_xformers: bool = True):
"""Load the Stream-DiffVSR pipeline."""
# Get model manager
model_manager = get_model_manager()
# Download models if needed
if not model_manager.check_models_exist():
print("Stream-DiffVSR models not found. Downloading from HuggingFace...")
model_manager.download_models()
model_path = model_manager.get_model_path()
# Determine device
if device == "auto":
if torch.cuda.is_available():
target_device = torch.device("cuda")
else:
target_device = torch.device("cpu")
else:
target_device = torch.device(device)
# Determine dtype
if precision == "auto":
if target_device.type == "cuda":
dtype = torch.float16
else:
dtype = torch.float32
elif precision == "fp16":
dtype = torch.float16
elif precision == "bf16":
dtype = torch.bfloat16
else:
dtype = torch.float32
# Only enable xformers on CUDA
use_xformers = enable_xformers and target_device.type == "cuda"
# Load pipeline wrapper
wrapper = StreamDiffVSRWrapper(
model_path=model_path,
device=target_device,
dtype=dtype,
enable_xformers=use_xformers,
)
print(f"Stream-DiffVSR model loaded on {target_device} with {dtype}")
return (wrapper,)
+159
View File
@@ -0,0 +1,159 @@
"""
FL DiffVSR Upscaler Node
Upscales video frames using Stream-DiffVSR with temporal coherence.
"""
import torch
from PIL import Image
import numpy as np
import comfy.utils
from ..core.pipeline_wrapper import StreamDiffVSRWrapper, tensor_to_pil, pil_to_tensor
class FL_DiffVSR_Upscale:
"""
Upscale video frames using Stream-DiffVSR.
Processes frames sequentially with temporal coherence.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("FL_DIFFVSR_MODEL",),
"images": ("IMAGE",),
"inference_steps": ("INT", {
"default": 4,
"min": 1,
"max": 50,
"step": 1,
"tooltip": "Number of denoising steps (4 recommended for speed/quality balance)"
}),
"guidance_scale": ("FLOAT", {
"default": 0.0,
"min": 0.0,
"max": 20.0,
"step": 0.5,
"tooltip": "Classifier-free guidance scale (0 = no guidance)"
}),
"chunk_size": ("INT", {
"default": 8,
"min": 1,
"max": 64,
"step": 1,
"tooltip": "Number of frames to process at once (lower = less VRAM, 0 = process all at once)"
}),
},
"optional": {
"prompt": ("STRING", {
"default": "",
"multiline": False,
"tooltip": "Optional text prompt for guidance"
}),
"negative_prompt": ("STRING", {
"default": "",
"multiline": False,
"tooltip": "Optional negative prompt"
}),
"seed": ("INT", {
"default": -1,
"min": -1,
"max": 0xffffffffffffffff,
"tooltip": "-1 for random seed"
}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("upscaled_images",)
FUNCTION = "upscale"
CATEGORY = "FL DiffVSR"
def upscale(
self,
model: StreamDiffVSRWrapper,
images: torch.Tensor,
inference_steps: int,
guidance_scale: float,
chunk_size: int = 8,
prompt: str = "",
negative_prompt: str = "",
seed: int = -1,
):
"""
Upscale images using Stream-DiffVSR pipeline.
Args:
model: StreamDiffVSRWrapper instance
images: Tensor of shape [B, H, W, C] in ComfyUI format (0-1 range)
inference_steps: Number of denoising steps
guidance_scale: CFG scale
chunk_size: Number of frames per chunk (lower = less VRAM)
prompt: Text prompt for guidance
negative_prompt: Negative prompt
seed: Random seed (-1 for random)
Returns:
Upscaled images tensor [B, H*4, W*4, C]
"""
num_frames = images.shape[0]
total_steps = num_frames * inference_steps
print(f"FL DiffVSR: Processing {num_frames} frames with {inference_steps} steps ({total_steps} total steps)...")
if chunk_size > 0:
print(f"FL DiffVSR: Using chunk_size={chunk_size} for VRAM-efficient processing")
# Set seed if specified
generator = None
if seed != -1:
generator = torch.Generator(device=model.device).manual_seed(seed)
# Convert ComfyUI tensors to PIL images
pil_images = []
for i in range(num_frames):
frame = images[i] # [H, W, C]
pil_img = tensor_to_pil(frame)
pil_images.append(pil_img)
# Setup ComfyUI progress bar
pbar = comfy.utils.ProgressBar(total_steps)
current_progress = [0] # Use list to allow modification in closure
def progress_callback(frame_idx, step_idx, total_frames, total_steps_per_frame, latents):
# Calculate overall progress
overall_step = frame_idx * total_steps_per_frame + step_idx + 1
steps_to_update = overall_step - current_progress[0]
if steps_to_update > 0:
pbar.update(steps_to_update)
current_progress[0] = overall_step
# Process through pipeline with chunking support
upscaled_pil = model.process_frames(
images=pil_images,
prompt=prompt,
negative_prompt=negative_prompt,
num_inference_steps=inference_steps,
guidance_scale=guidance_scale,
generator=generator,
progress_callback=progress_callback,
chunk_size=chunk_size,
)
# Ensure progress bar completes
remaining = total_steps - current_progress[0]
if remaining > 0:
pbar.update(remaining)
# Convert back to ComfyUI tensors
upscaled_tensors = []
for pil_img in upscaled_pil:
tensor = pil_to_tensor(pil_img) # [1, H, W, C]
upscaled_tensors.append(tensor)
# Concatenate all frames [B, H*4, W*4, C]
output = torch.cat(upscaled_tensors, dim=0)
print(f"FL DiffVSR: Upscaling complete. Output shape: {list(output.shape)}")
return (output,)
+26
View File
@@ -0,0 +1,26 @@
[project]
name = "comfyui-fl-diffvsr"
description = "FL DiffVSR - Diffusion-based video super-resolution nodes for ComfyUI. Features 4x upscaling with temporal coherence using Stream-DiffVSR for smooth, artifact-free video enhancement. Supports text-guided upscaling, chunked processing for memory efficiency, and automatic model downloading from HuggingFace."
version = "1.0.0"
license = "Apache-2.0"
dependencies = [
"torch>=2.0.0",
"torchvision>=0.15.0",
"diffusers>=0.21.0",
"transformers>=4.30.0",
"safetensors>=0.3.0",
"huggingface_hub>=0.16.0",
"accelerate>=0.20.0",
"Pillow>=9.0.0",
"numpy>=1.20.0",
"xformers>=0.0.20",
"einops>=0.6.0"
]
[project.urls]
Repository = "https://github.com/filliptm/ComfyUI-FL-DiffVSR"
[tool.comfy]
PublisherId = "machinedelusions"
DisplayName = "FL DiffVSR"
Icon = ""
+25
View File
@@ -0,0 +1,25 @@
# FL DiffVSR Requirements
# Core dependencies for Stream-DiffVSR video super-resolution
# Core ML frameworks
torch>=2.0.0
torchvision>=0.15.0
# Diffusers and transformers
diffusers>=0.21.0
transformers>=4.30.0
# Model loading
safetensors>=0.3.0
huggingface_hub>=0.16.0
accelerate>=0.20.0
# Image processing
Pillow>=9.0.0
numpy>=1.20.0
# Optional: Memory efficient attention (recommended)
xformers>=0.0.20
# Tensor operations
einops>=0.6.0
+1
View File
@@ -0,0 +1 @@
# Stream-DiffVSR adapted source
+3
View File
@@ -0,0 +1,3 @@
from .stream_diffvsr_pipeline import StreamDiffVSRPipeline
__all__ = ["StreamDiffVSRPipeline"]
@@ -0,0 +1,582 @@
# Copyright 2023 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import inspect
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
import numpy as np
import PIL.Image
import torch
import torch.nn.functional as F
import torchvision.transforms as T
from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer
from diffusers.image_processor import PipelineImageInput, VaeImageProcessor
from diffusers.loaders import FromSingleFileMixin, LoraLoaderMixin, TextualInversionLoaderMixin
from diffusers.models import AutoencoderKL, ControlNetModel, UNet2DConditionModel
from diffusers.schedulers import KarrasDiffusionSchedulers
from diffusers.utils import (
deprecate,
is_accelerate_available,
is_accelerate_version,
logging,
)
from diffusers.utils.torch_utils import randn_tensor, is_compiled_module
from diffusers.pipelines import DiffusionPipeline
from diffusers.pipelines.controlnet import MultiControlNetModel
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput, StableDiffusionSafetyChecker
from ..util import flow_utils as of
from ..temporal_autoencoder.autoencoder_tiny import TemporalAutoencoderTiny
from ..scheduler.ddim_scheduler import DDIMScheduler
logger = logging.get_logger(__name__)
class StreamDiffVSRPipeline(
DiffusionPipeline, TextualInversionLoaderMixin, LoraLoaderMixin, FromSingleFileMixin
):
"""
Stream-DiffVSR Pipeline for video super-resolution.
Args:
vae: Temporal VAE model for encoding/decoding
text_encoder: CLIP text encoder
tokenizer: CLIP tokenizer
unet: UNet2DConditionModel for denoising
controlnet: ControlNet for temporal conditioning
scheduler: DDIM scheduler
safety_checker: Optional safety checker
feature_extractor: Optional feature extractor
"""
_optional_components = ["safety_checker", "feature_extractor"]
def __init__(
self,
vae: Union[AutoencoderKL, TemporalAutoencoderTiny],
text_encoder: CLIPTextModel,
tokenizer: CLIPTokenizer,
unet: UNet2DConditionModel,
controlnet: Union[ControlNetModel, List[ControlNetModel], Tuple[ControlNetModel], MultiControlNetModel],
scheduler: Union[KarrasDiffusionSchedulers, DDIMScheduler],
safety_checker: StableDiffusionSafetyChecker,
feature_extractor: CLIPImageProcessor,
requires_safety_checker: bool = True,
):
super().__init__()
if safety_checker is None and requires_safety_checker:
logger.warning(
f"You have disabled the safety checker for {self.__class__} by passing `safety_checker=None`."
)
if safety_checker is not None and feature_extractor is None:
raise ValueError(
"Make sure to define a feature extractor when loading {self.__class__} if you want to use the safety checker."
)
if isinstance(controlnet, (list, tuple)):
controlnet = MultiControlNetModel(controlnet)
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
unet=unet,
controlnet=controlnet,
scheduler=scheduler,
safety_checker=safety_checker,
feature_extractor=feature_extractor,
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor, do_convert_rgb=True)
self.control_image_processor = VaeImageProcessor(
vae_scale_factor=self.vae_scale_factor, do_convert_rgb=True, do_normalize=True
)
self.register_to_config(requires_safety_checker=requires_safety_checker)
def enable_vae_slicing(self):
self.vae.enable_slicing()
def disable_vae_slicing(self):
self.vae.disable_slicing()
def enable_vae_tiling(self):
self.vae.enable_tiling()
def disable_vae_tiling(self):
self.vae.disable_tiling()
def enable_model_cpu_offload(self, gpu_id=0):
if is_accelerate_available() and is_accelerate_version(">=", "0.17.0.dev0"):
from accelerate import cpu_offload_with_hook
else:
raise ImportError("`enable_model_cpu_offload` requires `accelerate v0.17.0` or higher.")
device = torch.device(f"cuda:{gpu_id}")
hook = None
for cpu_offloaded_model in [self.text_encoder, self.unet, self.vae]:
_, hook = cpu_offload_with_hook(cpu_offloaded_model, device, prev_module_hook=hook)
if self.safety_checker is not None:
_, hook = cpu_offload_with_hook(self.safety_checker, device, prev_module_hook=hook)
cpu_offload_with_hook(self.controlnet, device)
self.final_offload_hook = hook
def encode_prompt(
self,
prompt,
device,
num_images_per_prompt,
do_classifier_free_guidance,
negative_prompt=None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
lora_scale: Optional[float] = None,
):
if lora_scale is not None and isinstance(self, LoraLoaderMixin):
self._lora_scale = lora_scale
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
if isinstance(self, TextualInversionLoaderMixin):
prompt = self.maybe_convert_prompt(prompt, self.tokenizer)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.tokenizer.model_max_length,
truncation=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
attention_mask = text_inputs.attention_mask.to(device)
else:
attention_mask = None
prompt_embeds = self.text_encoder(
text_input_ids.to(device),
attention_mask=attention_mask,
)
prompt_embeds = prompt_embeds[0]
if self.text_encoder is not None:
prompt_embeds_dtype = self.text_encoder.dtype
elif self.unet is not None:
prompt_embeds_dtype = self.unet.dtype
else:
prompt_embeds_dtype = prompt_embeds.dtype
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
bs_embed, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1)
if do_classifier_free_guidance and negative_prompt_embeds is None:
uncond_tokens: List[str]
if negative_prompt is None:
uncond_tokens = [""] * batch_size
elif isinstance(negative_prompt, str):
uncond_tokens = [negative_prompt]
else:
uncond_tokens = negative_prompt
if isinstance(self, TextualInversionLoaderMixin):
uncond_tokens = self.maybe_convert_prompt(uncond_tokens, self.tokenizer)
max_length = prompt_embeds.shape[1]
uncond_input = self.tokenizer(
uncond_tokens,
padding="max_length",
max_length=max_length,
truncation=True,
return_tensors="pt",
)
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
attention_mask = uncond_input.attention_mask.to(device)
else:
attention_mask = None
negative_prompt_embeds = self.text_encoder(
uncond_input.input_ids.to(device),
attention_mask=attention_mask,
)
negative_prompt_embeds = negative_prompt_embeds[0]
if do_classifier_free_guidance:
seq_len = negative_prompt_embeds.shape[1]
negative_prompt_embeds = negative_prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1)
negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
return prompt_embeds, negative_prompt_embeds
def prepare_extra_step_kwargs(self, generator, eta):
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
extra_step_kwargs = {}
if accepts_eta:
extra_step_kwargs["eta"] = eta
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
if accepts_generator:
extra_step_kwargs["generator"] = generator
return extra_step_kwargs
def prepare_latents(self, batch_size, num_channels_latents, height, width, dtype, device, generator, latents=None):
shape = (batch_size, num_channels_latents, height // self.vae_scale_factor, width // self.vae_scale_factor)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}."
)
if latents is None:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
else:
latents = latents.to(device)
latents = latents * self.scheduler.init_noise_sigma
return latents
def compute_flows(self, of_model, images, rescale_factor=1):
print('Computing forward flows...')
forward_flows = []
for i in range(1, len(images)):
# RAFT optical flow model requires float32 input
prev_image = images[i - 1].float()
cur_image = images[i].float()
fflow = of.get_flow(of_model, cur_image, prev_image, rescale_factor=rescale_factor)
forward_flows.append(fflow)
return forward_flows
def compute_single_flow(self, of_model, prev_image, cur_image, rescale_factor=1):
"""Compute optical flow between two adjacent frames.
To save VRAM, we compute flow at 1/4 resolution and upscale it back.
This is much more memory efficient for high-resolution inputs.
"""
# RAFT optical flow model requires float32 input
prev_image_f32 = prev_image.float()
cur_image_f32 = cur_image.float()
# Downscale images for RAFT to save VRAM (RAFT is very memory hungry at high res)
# Use 1/2 resolution for flow computation (better quality than 1/4)
flow_scale = 2
_, _, h, w = prev_image_f32.shape
small_h, small_w = h // flow_scale, w // flow_scale
prev_small = F.interpolate(prev_image_f32, size=(small_h, small_w), mode='bilinear', align_corners=False)
cur_small = F.interpolate(cur_image_f32, size=(small_h, small_w), mode='bilinear', align_corners=False)
# Compute flow at lower resolution
fflow_small = of.get_flow(of_model, cur_small, prev_small, rescale_factor=rescale_factor)
# Upscale flow back to original resolution and scale the flow values
# Flow is in [B, H, W, 2] format after get_flow
fflow_small_permuted = fflow_small.permute(0, 3, 1, 2) # [B, 2, H, W]
fflow_upscaled = F.interpolate(fflow_small_permuted, size=(h, w), mode='bilinear', align_corners=False)
fflow = fflow_upscaled.permute(0, 2, 3, 1) # [B, H, W, 2]
# Scale flow values to match the resolution change
fflow = fflow * flow_scale
return fflow
@torch.no_grad()
def __call__(
self,
prompt: Union[str, List[str]] = None,
images: PipelineImageInput = None,
height: Optional[int] = None,
width: Optional[int] = None,
num_inference_steps: int = 50,
guidance_scale: float = 7.5,
negative_prompt: Optional[Union[str, List[str]]] = None,
num_images_per_prompt: Optional[int] = 1,
eta: float = 0.0,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.FloatTensor] = None,
prompt_embeds: Optional[torch.FloatTensor] = None,
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
output_type: Optional[str] = "pil",
return_dict: bool = True,
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
callback_steps: int = 1,
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
controlnet_conditioning_scale: Union[float, List[float]] = 1.0,
guess_mode: bool = False,
control_guidance_start: Union[float, List[float]] = 0.0,
control_guidance_end: Union[float, List[float]] = 1.0,
of_model=None,
of_rescale_factor: int = 1,
timesteps_to_be_used: Optional[List[float]] = None,
# New parameters for chunked processing
prev_frame_rgb: Optional[torch.FloatTensor] = None,
prev_upscaled_for_flow: Optional[torch.FloatTensor] = None,
frame_offset: int = 0,
):
"""
Run the Stream-DiffVSR pipeline.
"""
controlnet = self.controlnet._orig_mod if is_compiled_module(self.controlnet) else self.controlnet
# align format for control guidance
if not isinstance(control_guidance_start, list) and isinstance(control_guidance_end, list):
control_guidance_start = len(control_guidance_end) * [control_guidance_start]
elif not isinstance(control_guidance_end, list) and isinstance(control_guidance_start, list):
control_guidance_end = len(control_guidance_start) * [control_guidance_end]
elif not isinstance(control_guidance_start, list) and not isinstance(control_guidance_end, list):
mult = len(controlnet.nets) if isinstance(controlnet, MultiControlNetModel) else 1
control_guidance_start, control_guidance_end = mult * [control_guidance_start], mult * [control_guidance_end]
# Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
device = self._execution_device
do_classifier_free_guidance = guidance_scale > 1.0
if isinstance(controlnet, MultiControlNetModel) and isinstance(controlnet_conditioning_scale, float):
controlnet_conditioning_scale = [controlnet_conditioning_scale] * len(controlnet.nets)
global_pool_conditions = (
controlnet.config.global_pool_conditions
if isinstance(controlnet, ControlNetModel)
else controlnet.nets[0].config.global_pool_conditions
)
guess_mode = guess_mode or global_pool_conditions
# Encode input prompt
text_encoder_lora_scale = (
cross_attention_kwargs.get("scale", None) if cross_attention_kwargs is not None else None
)
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
prompt,
device,
num_images_per_prompt,
do_classifier_free_guidance,
negative_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
lora_scale=text_encoder_lora_scale,
)
# Get model dtype early for consistency
model_dtype = self.unet.dtype
if do_classifier_free_guidance:
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds])
# Ensure prompt embeds are in model dtype
prompt_embeds = prompt_embeds.to(dtype=model_dtype)
# Prepare timesteps
if timesteps_to_be_used is None:
self.scheduler.set_timesteps(num_inference_steps, device=device)
else:
self.scheduler.set_timesteps(timesteps=timesteps_to_be_used, device=device)
timesteps = self.scheduler.timesteps
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
# Create tensor stating which controlnets to keep
controlnet_keep = []
for i in range(len(timesteps)):
keeps = [
1.0 - float(i / len(timesteps) < s or (i + 1) / len(timesteps) > e)
for s, e in zip(control_guidance_start, control_guidance_end)
]
controlnet_keep.append(keeps[0] if isinstance(controlnet, ControlNetModel) else keeps)
interp_mode = 'bilinear' if of_rescale_factor == 1 else 'nearest'
# Initialize state from previous chunk if provided
rgb_for_warpping_to_next_frame = prev_frame_rgb
prev_upscaled = prev_upscaled_for_flow
# Store raw images for per-frame processing
raw_images = images
output_images = []
num_channels_latents = self.vae.config.latent_channels
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
total_frames = len(raw_images)
with self.progress_bar(total=len(timesteps)*total_frames) as progress_bar:
for num_image, raw_image in enumerate(raw_images):
# === PROCESS ONE FRAME AT A TIME (VRAM efficient) ===
# Preprocess current frame only - DO NOT pass height/width to avoid resizing
# The preprocessor should just normalize, not resize the input
image = self.control_image_processor.preprocess(raw_image).to(dtype=model_dtype, device=device)
# Upscale current frame only (4x bicubic for flow/conditioning)
upscaled = F.interpolate(image, scale_factor=4, mode='bicubic').to(dtype=model_dtype)
# Get dimensions from upscaled output (for latent preparation)
frame_height, frame_width = upscaled.shape[-2:]
# Prepare latent for current frame only
latent = self.prepare_latents(
batch_size * num_images_per_prompt,
num_channels_latents,
frame_height,
frame_width,
model_dtype,
device,
generator
)
# Compute flow only between adjacent frames
flow = None
if prev_upscaled is not None:
flow = self.compute_single_flow(of_model, prev_upscaled, upscaled, rescale_factor=of_rescale_factor)
dec_temporal_features = None
warped_prev_est = None
# Compute Temporal Texture Guidance if we have previous frame
if rgb_for_warpping_to_next_frame is not None and flow is not None:
warped_prev_est = of.flow_warp(rgb_for_warpping_to_next_frame, flow, interp_mode=interp_mode)
warped_prev_est = warped_prev_est.to(dtype=model_dtype)
enc_layer_features = self.vae.encode(warped_prev_est, return_features_only=True)
dec_temporal_features = enc_layer_features[::-1]
for i, t in enumerate(timesteps):
# Ensure timestep is on the correct device
t = t.to(device)
latent_model_input = torch.cat([latent] * 2) if do_classifier_free_guidance else latent
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
# Double image for CFG (unconditional + conditional)
image_input = torch.cat([image] * 2) if do_classifier_free_guidance else image
latent_model_input = torch.cat([latent_model_input, image_input], dim=1)
# controlnet(s) inference
if guess_mode and do_classifier_free_guidance:
control_model_input = latent
control_model_input = self.scheduler.scale_model_input(control_model_input, t)
controlnet_prompt_embeds = prompt_embeds.chunk(2)[1]
else:
control_model_input = latent_model_input
controlnet_prompt_embeds = prompt_embeds
if isinstance(controlnet_keep[i], list):
cond_scale = [c * s for c, s in zip(controlnet_conditioning_scale, controlnet_keep[i])]
else:
controlnet_cond_scale = controlnet_conditioning_scale
if isinstance(controlnet_cond_scale, list):
controlnet_cond_scale = controlnet_cond_scale[0]
cond_scale = controlnet_cond_scale * controlnet_keep[i]
# Use ControlNet only if we have warped previous estimate
if warped_prev_est is None:
down_block_res_samples = None
mid_block_res_sample = None
else:
# Double warped_prev_est for CFG
controlnet_cond_input = torch.cat([warped_prev_est] * 2) if do_classifier_free_guidance else warped_prev_est
down_block_res_samples, mid_block_res_sample = self.controlnet(
control_model_input,
t,
encoder_hidden_states=controlnet_prompt_embeds,
controlnet_cond=controlnet_cond_input,
conditioning_scale=cond_scale,
guess_mode=guess_mode,
return_dict=False,
timestep_cond=None
)
if guess_mode and do_classifier_free_guidance:
down_block_res_samples = [torch.cat([torch.zeros_like(d), d]) for d in down_block_res_samples]
mid_block_res_sample = torch.cat([torch.zeros_like(mid_block_res_sample), mid_block_res_sample])
# predict the noise residual
noise_pred = self.unet(
latent_model_input,
t,
encoder_hidden_states=prompt_embeds,
cross_attention_kwargs=cross_attention_kwargs,
down_block_additional_residuals=down_block_res_samples,
mid_block_additional_residual=mid_block_res_sample,
return_dict=False,
)[0]
# perform guidance
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
# compute the previous noisy sample x_t -> x_t-1
step_output = self.scheduler.step(noise_pred, t, latent, **extra_step_kwargs)
latent, x0_est = step_output.prev_sample, step_output.pred_original_sample
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if callback is not None and i % callback_steps == 0:
# Pass frame index (with offset), step index, total frames, total steps
callback(num_image + frame_offset, i, total_frames, len(timesteps), latent)
if not output_type == "latent":
decoded_image = self.vae.decode(latent / self.vae.config.scaling_factor, temporal_features=dec_temporal_features, return_dict=False)[0]
else:
decoded_image = latent
# Update state for next frame
# Clone tensors to avoid reference issues between frames
rgb_for_warpping_to_next_frame = decoded_image.clone()
prev_upscaled = upscaled.clone()
has_nsfw_concept = None
do_denormalize = [True] * decoded_image[0].shape[0]
final_image = self.image_processor.postprocess(decoded_image, output_type=output_type, do_denormalize=do_denormalize)
output_images.append(final_image)
self.vae.reset_temporal_condition()
# Free memory for this frame
del image, latent, upscaled, decoded_image
if flow is not None:
del flow
if warped_prev_est is not None:
del warped_prev_est
if hasattr(self, "final_offload_hook") and self.final_offload_hook is not None:
self.unet.to("cpu")
self.controlnet.to("cpu")
torch.cuda.empty_cache()
if hasattr(self, "final_offload_hook") and self.final_offload_hook is not None:
self.final_offload_hook.offload()
# Return output along with state for next chunk
output = StableDiffusionPipelineOutput(images=output_images, nsfw_content_detected=None)
return output, rgb_for_warpping_to_next_frame, prev_upscaled
+3
View File
@@ -0,0 +1,3 @@
from .ddim_scheduler import DDIMScheduler
__all__ = ["DDIMScheduler"]
+298
View File
@@ -0,0 +1,298 @@
# Copyright 2024 Stanford University Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from dataclasses import dataclass
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput
from diffusers.utils.torch_utils import randn_tensor
from diffusers.schedulers import KarrasDiffusionSchedulers, SchedulerMixin
@dataclass
class DDIMSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
pred_original_sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)` for images):
The predicted denoised sample `(x_{0})` based on the model output from the current timestep.
`pred_original_sample` can be used to preview progress or for guidance.
"""
prev_sample: torch.Tensor
pred_original_sample: Optional[torch.Tensor] = None
def betas_for_alpha_bar(
num_diffusion_timesteps,
max_beta=0.999,
alpha_transform_type="cosine",
):
"""
Create a beta schedule that discretizes the given alpha_t_bar function.
"""
if alpha_transform_type == "cosine":
def alpha_bar_fn(t):
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
elif alpha_transform_type == "exp":
def alpha_bar_fn(t):
return math.exp(t * -12.0)
else:
raise ValueError(f"Unsupported alpha_transform_type: {alpha_transform_type}")
betas = []
for i in range(num_diffusion_timesteps):
t1 = i / num_diffusion_timesteps
t2 = (i + 1) / num_diffusion_timesteps
betas.append(min(1 - alpha_bar_fn(t2) / alpha_bar_fn(t1), max_beta))
return torch.tensor(betas, dtype=torch.float32)
def rescale_zero_terminal_snr(betas):
"""Rescales betas to have zero terminal SNR."""
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
alphas_bar_sqrt = alphas_cumprod.sqrt()
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
alphas_bar_sqrt -= alphas_bar_sqrt_T
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
alphas_bar = alphas_bar_sqrt**2
alphas = alphas_bar[1:] / alphas_bar[:-1]
alphas = torch.cat([alphas_bar[0:1], alphas])
betas = 1 - alphas
return betas
class DDIMScheduler(SchedulerMixin, ConfigMixin):
"""
DDIMScheduler extends the denoising procedure introduced in DDPMs with non-Markovian guidance.
"""
_compatibles = [e.name for e in KarrasDiffusionSchedulers]
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
beta_start: float = 0.0001,
beta_end: float = 0.02,
beta_schedule: str = "linear",
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
clip_sample: bool = True,
set_alpha_to_one: bool = True,
steps_offset: int = 0,
prediction_type: str = "epsilon",
thresholding: bool = False,
dynamic_thresholding_ratio: float = 0.995,
clip_sample_range: float = 1.0,
sample_max_value: float = 1.0,
timestep_spacing: str = "leading",
rescale_betas_zero_snr: bool = False,
):
if trained_betas is not None:
self.betas = torch.tensor(trained_betas, dtype=torch.float32)
elif beta_schedule == "linear":
self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32)
elif beta_schedule == "scaled_linear":
self.betas = torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2
elif beta_schedule == "squaredcos_cap_v2":
self.betas = betas_for_alpha_bar(num_train_timesteps)
else:
raise NotImplementedError(f"{beta_schedule} is not implemented for {self.__class__}")
if rescale_betas_zero_snr:
self.betas = rescale_zero_terminal_snr(self.betas)
self.alphas = 1.0 - self.betas
self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)
self.final_alpha_cumprod = torch.tensor(1.0) if set_alpha_to_one else self.alphas_cumprod[0]
self.init_noise_sigma = 1.0
self.num_inference_steps = None
self.timesteps = torch.from_numpy(np.arange(0, 1000)[::-1].copy().astype(np.int64))
self.timestep_spacing = timestep_spacing
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
return sample
def _get_variance(self, timestep, prev_timestep):
alpha_prod_t = self.alphas_cumprod[timestep]
alpha_prod_t_prev = self.alphas_cumprod[prev_timestep] if prev_timestep >= 0 else self.final_alpha_cumprod
beta_prod_t = 1 - alpha_prod_t
beta_prod_t_prev = 1 - alpha_prod_t_prev
variance = (beta_prod_t_prev / beta_prod_t) * (1 - alpha_prod_t / alpha_prod_t_prev)
return variance
def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor:
dtype = sample.dtype
batch_size, channels, *remaining_dims = sample.shape
if dtype not in (torch.float32, torch.float64):
sample = sample.float()
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
abs_sample = sample.abs()
s = torch.quantile(abs_sample, self.config.dynamic_thresholding_ratio, dim=1)
s = torch.clamp(s, min=1, max=self.config.sample_max_value)
s = s.unsqueeze(1)
sample = torch.clamp(sample, -s, s) / s
sample = sample.reshape(batch_size, channels, *remaining_dims)
sample = sample.to(dtype)
return sample
def set_timesteps(self, num_inference_steps: int = None, device: Union[str, torch.device] = None, timesteps: List[float] = None):
if timesteps is not None:
self.timesteps = torch.tensor(timesteps, device=device)
self.num_inference_steps = len(timesteps)
return
if num_inference_steps > self.config.num_train_timesteps:
raise ValueError(
f"`num_inference_steps`: {num_inference_steps} cannot be larger than `self.config.train_timesteps`:"
f" {self.config.num_train_timesteps}"
)
self.num_inference_steps = num_inference_steps
if self.timestep_spacing == "linspace":
timesteps = (
np.linspace(0, self.config.num_train_timesteps - 1, num_inference_steps)
.round()[::-1]
.copy()
.astype(np.int64)
)
elif self.timestep_spacing == "leading":
step_ratio = self.config.num_train_timesteps // self.num_inference_steps
timesteps = (np.arange(0, num_inference_steps) * step_ratio).round()[::-1].copy().astype(np.int64)
timesteps += self.config.steps_offset
elif self.timestep_spacing == "trailing":
step_ratio = self.config.num_train_timesteps / self.num_inference_steps
timesteps = np.round(np.arange(self.config.num_train_timesteps, 0, -step_ratio)).astype(np.int64)
timesteps -= 1
else:
raise ValueError(
f"{self.timestep_spacing} is not supported. Please make sure to choose one of 'leading' or 'trailing'."
)
self.timesteps = torch.from_numpy(timesteps).to(device)
def step(
self,
model_output: torch.Tensor,
timestep: int,
sample: torch.Tensor,
eta: float = 0.0,
use_clipped_model_output: bool = False,
generator=None,
variance_noise: Optional[torch.Tensor] = None,
return_dict: bool = True,
) -> Union[DDIMSchedulerOutput, Tuple]:
if self.num_inference_steps is None:
raise ValueError(
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
)
prev_timestep = timestep - self.config.num_train_timesteps // self.num_inference_steps
alpha_prod_t = self.alphas_cumprod[timestep]
alpha_prod_t_prev = self.alphas_cumprod[prev_timestep] if prev_timestep >= 0 else self.final_alpha_cumprod
beta_prod_t = 1 - alpha_prod_t
if self.config.prediction_type == "epsilon":
pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)
pred_epsilon = model_output
elif self.config.prediction_type == "sample":
pred_original_sample = model_output
pred_epsilon = (sample - alpha_prod_t ** (0.5) * pred_original_sample) / beta_prod_t ** (0.5)
elif self.config.prediction_type == "v_prediction":
pred_original_sample = (alpha_prod_t**0.5) * sample - (beta_prod_t**0.5) * model_output
pred_epsilon = (alpha_prod_t**0.5) * model_output + (beta_prod_t**0.5) * sample
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`, or"
" `v_prediction`"
)
if self.config.thresholding:
pred_original_sample = self._threshold_sample(pred_original_sample)
elif self.config.clip_sample:
pred_original_sample = pred_original_sample.clamp(
-self.config.clip_sample_range, self.config.clip_sample_range
)
variance = self._get_variance(timestep, prev_timestep)
std_dev_t = eta * variance ** (0.5)
if use_clipped_model_output:
pred_epsilon = (sample - alpha_prod_t ** (0.5) * pred_original_sample) / beta_prod_t ** (0.5)
pred_sample_direction = (1 - alpha_prod_t_prev - std_dev_t**2) ** (0.5) * pred_epsilon
prev_sample = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction
if eta > 0:
if variance_noise is not None and generator is not None:
raise ValueError(
"Cannot pass both generator and variance_noise."
)
if variance_noise is None:
variance_noise = randn_tensor(
model_output.shape, generator=generator, device=model_output.device, dtype=model_output.dtype
)
variance = std_dev_t * variance_noise
prev_sample = prev_sample + variance
if not return_dict:
return (prev_sample,)
return DDIMSchedulerOutput(prev_sample=prev_sample, pred_original_sample=pred_original_sample)
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
self.alphas_cumprod = self.alphas_cumprod.to(device=original_samples.device)
alphas_cumprod = self.alphas_cumprod.to(dtype=original_samples.dtype)
timesteps = timesteps.to(original_samples.device)
sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5
sqrt_alpha_prod = sqrt_alpha_prod.flatten()
while len(sqrt_alpha_prod.shape) < len(original_samples.shape):
sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1)
sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten()
while len(sqrt_one_minus_alpha_prod.shape) < len(original_samples.shape):
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1)
noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
@@ -0,0 +1,3 @@
from .autoencoder_tiny import TemporalAutoencoderTiny
__all__ = ["TemporalAutoencoderTiny"]
@@ -0,0 +1,278 @@
# Copyright 2024 Ollin Boer Bohan and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput
from diffusers.utils.accelerate_utils import apply_forward_hook
from diffusers.models.modeling_utils import ModelMixin
from .vae import DecoderOutput, TemporalDecoderTiny, EncoderTiny
from .models.unets.unet_2d_blocks import TemporalAutoencoderTinyBlock
@dataclass
class TemporalAutoencoderTinyOutput(BaseOutput):
"""
Output of TemporalAutoencoderTiny encoding method.
Args:
latents (`torch.Tensor`): Encoded outputs of the `Encoder`.
"""
latents: torch.Tensor
class TemporalAutoencoderTiny(ModelMixin, ConfigMixin):
"""
A tiny distilled VAE model for encoding images into latents and decoding latent representations into images.
"""
_supports_gradient_checkpointing = True
@register_to_config
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
encoder_block_out_channels: Tuple[int, ...] = (64, 64, 64, 64),
decoder_block_out_channels: Tuple[int, ...] = (64, 64, 64, 64),
act_fn: str = "relu",
upsample_fn: str = "nearest",
latent_channels: int = 4,
upsampling_scaling_factor: int = 2,
num_encoder_blocks: Tuple[int, ...] = (1, 3, 3, 3),
num_decoder_blocks: Tuple[int, ...] = (3, 3, 3, 1),
latent_magnitude: int = 3,
latent_shift: float = 0.5,
force_upcast: bool = False,
scaling_factor: float = 1.0,
shift_factor: float = 0.0,
block_out_channels: Tuple[int, ...] = None, # For compatibility with saved configs
):
super().__init__()
if len(encoder_block_out_channels) != len(num_encoder_blocks):
raise ValueError("`encoder_block_out_channels` should have the same length as `num_encoder_blocks`.")
if len(decoder_block_out_channels) != len(num_decoder_blocks):
raise ValueError("`decoder_block_out_channels` should have the same length as `num_decoder_blocks`.")
self.encoder = EncoderTiny(
in_channels=in_channels,
out_channels=latent_channels,
num_blocks=num_encoder_blocks,
block_out_channels=encoder_block_out_channels,
act_fn=act_fn,
)
self.encoder.requires_grad_(False)
self.decoder = TemporalDecoderTiny(
in_channels=latent_channels,
out_channels=out_channels,
num_blocks=num_decoder_blocks,
block_out_channels=decoder_block_out_channels,
upsampling_scaling_factor=upsampling_scaling_factor,
act_fn=act_fn,
upsample_fn=upsample_fn,
)
self.decoder.requires_grad_(False)
for name, param in self.decoder.named_parameters():
if "alpha" in name or "temporal_processor" in name:
param.requires_grad_(True)
self.latent_magnitude = latent_magnitude
self.latent_shift = latent_shift
self.scaling_factor = scaling_factor
self.use_slicing = False
self.use_tiling = False
self.spatial_scale_factor = 2**out_channels
self.tile_overlap_factor = 0.125
self.tile_sample_min_size = 512
self.tile_latent_min_size = self.tile_sample_min_size // self.spatial_scale_factor
self.register_to_config(block_out_channels=decoder_block_out_channels)
self.register_to_config(force_upcast=False)
def reset_temporal_condition(self):
"""reset temporal memory"""
for module in self.encoder.layers:
if isinstance(module, TemporalAutoencoderTinyBlock):
module.reset_temporal()
for module in self.decoder.layers:
if isinstance(module, TemporalAutoencoderTinyBlock):
module.reset_temporal()
def _set_gradient_checkpointing(self, module, value: bool = False) -> None:
if isinstance(module, (EncoderTiny, TemporalDecoderTiny)):
module.gradient_checkpointing = value
def scale_latents(self, x: torch.Tensor) -> torch.Tensor:
"""raw latents -> [0, 1]"""
return x.div(2 * self.latent_magnitude).add(self.latent_shift).clamp(0, 1)
def unscale_latents(self, x: torch.Tensor) -> torch.Tensor:
"""[0, 1] -> raw latents"""
return x.sub(self.latent_shift).mul(2 * self.latent_magnitude)
def enable_slicing(self) -> None:
self.use_slicing = True
def disable_slicing(self) -> None:
self.use_slicing = False
def enable_tiling(self, use_tiling: bool = True) -> None:
self.use_tiling = use_tiling
def disable_tiling(self) -> None:
self.enable_tiling(False)
def _tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
sf = self.spatial_scale_factor
tile_size = self.tile_sample_min_size
blend_size = int(tile_size * self.tile_overlap_factor)
traverse_size = tile_size - blend_size
ti = range(0, x.shape[-2], traverse_size)
tj = range(0, x.shape[-1], traverse_size)
blend_masks = torch.stack(
torch.meshgrid([torch.arange(tile_size / sf) / (blend_size / sf - 1)] * 2, indexing="ij")
)
blend_masks = blend_masks.clamp(0, 1).to(x.device)
out = torch.zeros(x.shape[0], 4, x.shape[-2] // sf, x.shape[-1] // sf, device=x.device)
for i in ti:
for j in tj:
tile_in = x[..., i : i + tile_size, j : j + tile_size]
tile_out = out[..., i // sf : (i + tile_size) // sf, j // sf : (j + tile_size) // sf]
tile = self.encoder(tile_in)
h, w = tile.shape[-2], tile.shape[-1]
blend_mask_i = torch.ones_like(blend_masks[0]) if i == 0 else blend_masks[0]
blend_mask_j = torch.ones_like(blend_masks[1]) if j == 0 else blend_masks[1]
blend_mask = blend_mask_i * blend_mask_j
tile, blend_mask = tile[..., :h, :w], blend_mask[..., :h, :w]
tile_out.copy_(blend_mask * tile + (1 - blend_mask) * tile_out)
return out
def _tiled_decode(self, x: torch.Tensor) -> torch.Tensor:
sf = self.spatial_scale_factor
tile_size = self.tile_latent_min_size
blend_size = int(tile_size * self.tile_overlap_factor)
traverse_size = tile_size - blend_size
ti = range(0, x.shape[-2], traverse_size)
tj = range(0, x.shape[-1], traverse_size)
blend_masks = torch.stack(
torch.meshgrid([torch.arange(tile_size * sf) / (blend_size * sf - 1)] * 2, indexing="ij")
)
blend_masks = blend_masks.clamp(0, 1).to(x.device)
out = torch.zeros(x.shape[0], 3, x.shape[-2] * sf, x.shape[-1] * sf, device=x.device)
for i in ti:
for j in tj:
tile_in = x[..., i : i + tile_size, j : j + tile_size]
tile_out = out[..., i * sf : (i + tile_size) * sf, j * sf : (j + tile_size) * sf]
tile = self.decoder(tile_in)
h, w = tile.shape[-2], tile.shape[-1]
blend_mask_i = torch.ones_like(blend_masks[0]) if i == 0 else blend_masks[0]
blend_mask_j = torch.ones_like(blend_masks[1]) if j == 0 else blend_masks[1]
blend_mask = (blend_mask_i * blend_mask_j)[..., :h, :w]
tile_out.copy_(blend_mask * tile + (1 - blend_mask) * tile_out)
return out
@apply_forward_hook
def encode(self, x: torch.Tensor, return_dict: bool = True, return_layers_features: bool = True, return_features_only: bool = False) -> Union[TemporalAutoencoderTinyOutput, Tuple[torch.Tensor]]:
layer_features = [] if return_layers_features else None
if self.use_slicing and x.shape[0] > 1:
output = [
self._tiled_encode(x_slice) if self.use_tiling else self.encoder(x_slice)
for x_slice in x.split(1)
]
output = torch.cat(output)
else:
if self.use_tiling:
output = self._tiled_encode(x)
elif return_layers_features:
current_features = x
for module in self.encoder.layers:
current_features = module(current_features)
if isinstance(module, TemporalAutoencoderTinyBlock):
layer_features.append(current_features)
if return_features_only:
return layer_features
output = self.encoder(x)
if not return_dict:
return (output,), layer_features
return TemporalAutoencoderTinyOutput(latents=output)
@apply_forward_hook
def decode(
self, x: torch.Tensor, temporal_features=None, generator: Optional[torch.Generator] = None, return_dict: bool = True
) -> Union[DecoderOutput, Tuple[torch.Tensor]]:
if self.use_slicing and x.shape[0] > 1:
output = [
self._tiled_decode(x_slice) if self.use_tiling else self.decoder(x_slice) for x_slice in x.split(1)
]
output = torch.cat(output)
elif temporal_features is not None:
block_idx = 0
for module in self.decoder.layers:
if isinstance(module, TemporalAutoencoderTinyBlock):
module.prev_features = temporal_features[block_idx]
block_idx += 1
output = self.decoder(x)
else:
output = self._tiled_decode(x) if self.use_tiling else self.decoder(x)
if not return_dict:
return (output,)
return DecoderOutput(sample=output)
def forward(
self,
sample: torch.Tensor,
previous_sample: Optional[torch.Tensor] = None,
return_dict: bool = False,
) -> Union[DecoderOutput, Tuple[torch.Tensor]]:
layer_features = None
if previous_sample is not None:
prev_enc, layer_features = self.encode(previous_sample, return_dict=return_dict)
if layer_features is not None:
temporal_features = layer_features[::-1]
else:
temporal_features = None
dec = self.decode(sample, temporal_features=temporal_features, return_dict=return_dict)[0]
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
@@ -0,0 +1 @@
# Temporal autoencoder models
@@ -0,0 +1,3 @@
from .unet_2d_blocks import TemporalAutoencoderTinyBlock
__all__ = ["TemporalAutoencoderTinyBlock"]
@@ -0,0 +1,98 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import torch
import torch.nn as nn
from diffusers.models.activations import get_activation
class TemporalAutoencoderTinyBlock(nn.Module):
"""
Tiny Autoencoder block used in [`AutoencoderTiny`]. It is a mini residual module consisting of plain conv + ReLU
blocks.
Args:
in_channels (`int`): The number of input channels.
out_channels (`int`): The number of output channels.
act_fn (`str`):
` The activation function to use. Supported values are `"swish"`, `"mish"`, `"gelu"`, and `"relu"`.
Returns:
`torch.Tensor`: A tensor with the same shape as the input tensor, but with the number of channels equal to
`out_channels`.
"""
def __init__(self, in_channels: int, out_channels: int, act_fn: str):
super().__init__()
act_fn = get_activation(act_fn)
self.conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
act_fn,
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
act_fn,
nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),
)
self.skip = (
nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)
if in_channels != out_channels
else nn.Identity()
)
self.fuse = nn.ReLU()
# temporal layers
self.prev_features = None
self.alpha = nn.Parameter(torch.tensor(0.5))
self.temporal_processor = nn.Sequential(
nn.Conv1d(out_channels, out_channels, 3, padding=1),
act_fn,
nn.Conv1d(out_channels, out_channels, 3, padding=1)
)
def forward(self, x):
current_features = self.conv(x)
if self.prev_features is not None:
B, C, H, W = current_features.shape
pool_kernel = (4, 4)
avg_pool = nn.AvgPool2d(kernel_size=pool_kernel, stride=pool_kernel)
current_pooled = avg_pool(current_features)
prev_pooled = avg_pool(self.prev_features)
temporal_input = torch.cat([
current_pooled.view(B, C, -1),
prev_pooled.view(B, C, -1)
], dim=2)
temporal_out = self.temporal_processor(temporal_input)
pool_h, pool_w = current_pooled.shape[2], current_pooled.shape[3]
temporal_out_fuse = self.alpha * temporal_out[:, :, :pool_h * pool_w].view(B, C, pool_h, pool_w) + \
(1 - self.alpha) * temporal_out[:, :, -pool_h * pool_w:].view(B, C, pool_h, pool_w)
temporal_out_fuse = nn.functional.interpolate(
temporal_out_fuse,
size=(H, W),
mode='bilinear',
align_corners=False
)
current_features = current_features + 0.1 * temporal_out_fuse
return self.fuse(current_features + self.skip(x))
def reset_temporal(self):
self.prev_features = None
+138
View File
@@ -0,0 +1,138 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from dataclasses import dataclass
from typing import Optional, Tuple
import torch
import torch.nn as nn
from diffusers.utils import BaseOutput
from diffusers.models.activations import get_activation
from .models.unets.unet_2d_blocks import TemporalAutoencoderTinyBlock
@dataclass
class DecoderOutput(BaseOutput):
"""
Output of decoding method.
Args:
sample (`torch.Tensor` of shape `(batch_size, num_channels, height, width)`):
The decoded output sample from the last layer of the model.
"""
sample: torch.Tensor
commit_loss: Optional[torch.FloatTensor] = None
class EncoderTiny(nn.Module):
"""
The `EncoderTiny` layer is a simpler version of the `Encoder` layer.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
num_blocks: Tuple[int, ...],
block_out_channels: Tuple[int, ...],
act_fn: str,
):
super().__init__()
layers = []
for i, num_block in enumerate(num_blocks):
num_channels = block_out_channels[i]
if i == 0:
layers.append(nn.Conv2d(in_channels, num_channels, kernel_size=3, padding=1))
else:
layers.append(
nn.Conv2d(
num_channels,
num_channels,
kernel_size=3,
padding=1,
stride=2,
bias=False,
)
)
for _ in range(num_block):
layers.append(TemporalAutoencoderTinyBlock(num_channels, num_channels, act_fn))
layers.append(nn.Conv2d(block_out_channels[-1], out_channels, kernel_size=3, padding=1))
self.layers = nn.Sequential(*layers)
self.gradient_checkpointing = False
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.layers(x.add(1).div(2))
return x
class TemporalDecoderTiny(nn.Module):
"""
The `TemporalDecoderTiny` layer is a simpler version of the `Decoder` layer with temporal processing.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
num_blocks: Tuple[int, ...],
block_out_channels: Tuple[int, ...],
upsampling_scaling_factor: int,
act_fn: str,
upsample_fn: str,
):
super().__init__()
layers = [
nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, padding=1),
get_activation(act_fn),
]
for i, num_block in enumerate(num_blocks):
is_final_block = i == (len(num_blocks) - 1)
num_channels = block_out_channels[i]
for _ in range(num_block):
block = TemporalAutoencoderTinyBlock(num_channels, num_channels, act_fn)
layers.append(block)
if not is_final_block:
layers.append(nn.Upsample(scale_factor=upsampling_scaling_factor, mode=upsample_fn))
conv_out_channel = num_channels if not is_final_block else out_channels
layers.append(
nn.Conv2d(
num_channels,
conv_out_channel,
kernel_size=3,
padding=1,
bias=is_final_block,
)
)
self.layers = nn.Sequential(*layers)
self.gradient_checkpointing = False
def forward(self, x: torch.Tensor) -> torch.Tensor:
# Clamp
x = torch.tanh(x / 3) * 3
x = self.layers(x)
# scale image from [0, 1] to [-1, 1] to match diffusers convention
return x.mul(2).sub(1)
+3
View File
@@ -0,0 +1,3 @@
from . import flow_utils
__all__ = ["flow_utils"]
+100
View File
@@ -0,0 +1,100 @@
import torch
import torch.nn.functional as F
def flow_warp(x, flow, interp_mode='bilinear', padding_mode='zeros'):
"""Warp an image or feature map with optical flow
Args:
x (Tensor): size (N, C, H, W)
flow (Tensor): size (N, H, W, 2), normal value
interp_mode (str): 'nearest' or 'bilinear'
padding_mode (str): 'zeros' or 'border' or 'reflection'
Returns:
Tensor: warped image or feature map
"""
if flow.dim() == 4 and flow.shape[1] == 2:
flow = flow.permute(0, 2, 3, 1) # [N, 2, H, W] -> [N, H, W, 2]
assert x.size()[-2:] == flow.size()[1:3]
_, _, H, W = x.size()
# Ensure flow matches input dtype and device
flow = flow.to(dtype=x.dtype, device=x.device)
# mesh grid
grid_y, grid_x = torch.meshgrid(torch.arange(0, H), torch.arange(0, W), indexing='ij')
grid = torch.stack((grid_x, grid_y), 2).float() # W(x), H(y), 2
grid.requires_grad = False
grid = grid.type_as(x)
vgrid = grid + flow
# scale grid to [-1,1]
vgrid_x = 2.0 * vgrid[:, :, :, 0] / max(W - 1, 1) - 1.0
vgrid_y = 2.0 * vgrid[:, :, :, 1] / max(H - 1, 1) - 1.0
vgrid_scaled = torch.stack((vgrid_x, vgrid_y), dim=3)
output = F.grid_sample(x, vgrid_scaled, mode=interp_mode, padding_mode=padding_mode, align_corners=False)
return output
def get_flow(of_model, target, source, rescale_factor=1):
flows = of_model(target, source)
flow = flows[-1]
flow = F.interpolate(flow//rescale_factor, scale_factor=1/rescale_factor, mode='bilinear') if rescale_factor != 1 else flow
flow = flow.permute(0, 2, 3, 1) # permute to B, H, W, 2
return flow
def compute_flow_magnitude(flow):
flow_mag = flow[:, :, :, 0] ** 2 + flow[:, :, :, 1] ** 2
return flow_mag
def compute_flow_gradients(flow):
B = flow.shape[0]
H = flow.shape[1]
W = flow.shape[2]
device = flow.device
flow_x_du = torch.zeros((B, H, W), device=device)
flow_x_dv = torch.zeros((B, H, W), device=device)
flow_y_du = torch.zeros((B, H, W), device=device)
flow_y_dv = torch.zeros((B, H, W), device=device)
flow_x = flow[:, :, :, 0]
flow_y = flow[:, :, :, 1]
flow_x_du[:, :, :-1] = flow_x[:, :, :-1] - flow_x[:, :, 1:]
flow_x_dv[:, :-1, :] = flow_x[:, :-1, :] - flow_x[:, 1:, :]
flow_y_du[:, :, :-1] = flow_y[:, :, :-1] - flow_y[:, :, 1:]
flow_y_dv[:, :-1, :] = flow_y[:, :-1, :] - flow_y[:, 1:, :]
return flow_x_du, flow_x_dv, flow_y_du, flow_y_dv
def detect_occlusion(fw_flow, bw_flow):
# inputs: flow_forward, flow_backward
# return: occlusion mask
tmp = bw_flow
bw_flow = fw_flow
fw_flow = tmp
fw_flow_w = flow_warp(fw_flow.permute(0,3,1,2), bw_flow).permute(0,2,3,1)
fb_flow_sum = fw_flow_w + bw_flow
fb_flow_mag = compute_flow_magnitude(fb_flow_sum)
fw_flow_w_mag = compute_flow_magnitude(fw_flow_w)
bw_flow_mag = compute_flow_magnitude(bw_flow)
mask1 = fb_flow_mag > 0.01 * (fw_flow_w_mag + bw_flow_mag) + 0.5
fx_du, fx_dv, fy_du, fy_dv = compute_flow_gradients(bw_flow)
fx_mag = fx_du ** 2 + fx_dv ** 2
fy_mag = fy_du ** 2 + fy_dv ** 2
mask2 = (fx_mag + fy_mag) > 0.01 * bw_flow_mag + 0.002
mask = torch.logical_or(mask1, mask2)
occlusion = torch.ones((fw_flow.shape[0], fw_flow.shape[1], fw_flow.shape[2]), device=fw_flow.device)
occlusion[mask == 1] = 0
return occlusion
def get_flow_forward_backward(net, current, prev, rescale_factor=1):
flow_forward = get_flow(net, current, prev, rescale_factor=rescale_factor)
flow_backward = get_flow(net, prev, current, rescale_factor=rescale_factor)
return flow_forward, flow_backward