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:
+75
@@ -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
|
||||
@@ -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.
|
||||
|
||||
[](https://arxiv.org/abs/2512.23709)
|
||||
[](https://www.patreon.com/Machinedelusions)
|
||||
|
||||

|
||||
|
||||
## 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
@@ -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")
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 310 KiB |
@@ -0,0 +1 @@
|
||||
# FL DiffVSR Core
|
||||
@@ -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
|
||||
@@ -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]
|
||||
@@ -0,0 +1 @@
|
||||
# FL DiffVSR Nodes
|
||||
@@ -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,)
|
||||
@@ -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,)
|
||||
@@ -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 = ""
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
# Stream-DiffVSR adapted source
|
||||
@@ -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
|
||||
@@ -0,0 +1,3 @@
|
||||
from .ddim_scheduler import DDIMScheduler
|
||||
|
||||
__all__ = ["DDIMScheduler"]
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -0,0 +1,3 @@
|
||||
from . import flow_utils
|
||||
|
||||
__all__ = ["flow_utils"]
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user