Speed Up 30 to 50%, fix Memory Leak, refacto)

This commit is contained in:
NumZ
2025-06-30 12:46:41 +02:00
parent ad45a539af
commit 2786c05faa
28 changed files with 2770 additions and 546 deletions
+4
View File
@@ -17,3 +17,7 @@ run_*.py
vram_diagnostic.py
test_*.py
VRAM_OPTIMIZATIONS_SUMMARY.md
seedvr2.py
src/core/isolated_generation.py
src/core/subprocess_runner.py
models/video_vae_v3_mine_bad/
+119 -44
View File
@@ -9,12 +9,34 @@ Official release of [SeedVR2](https://github.com/ByteDance-Seed/SeedVR) for Comf
<img src="docs/usage.png">
## 🆙 Todo
## 📋 Quick Access
- Fixed unloading the 3B model when the process is finished (sorry about that, I'm trying to find out what's going on)
- [🆙 Note and futur releases](#-note-and-futur-releases)
- [🚀 Updates](#-updates)
- [🎯 Features](#-features)
- [🔧 Requirements](#-requirements)
- [📦 Installation](#-installation)
- [📖 Usage](#-usage)
- [📊 Benchmarks](#-benchmarks)
- [🔧 Limitations](#-Limitations)
- [🤝 Contributing](#-contributing)
- [🙏 Credits](#-credits)
- [📄 License](#-license)
## 🆙 Note and futur releases
- Improve FP8 integration, we are loosing some FP8 advantages during the process.
- Tile-VAE integration if it works for video, I have test to do or if some dev want help, you are welcome.
- 7B FP8 model seems to have quality issues, use 7BFP16 instead (If FP8 don't give OOM then FP16 will works) I have to review this.
## 🚀 Updates
**2025.06.30**
- 🚀 Speed Up the process and less VRAM used (see new benchmark).
- 🛠️ Fixed leak memory on 3B models.
- ✅ refactored the code for better sharing with the community, feel free to propose pull requests.
**2025.06.24**
- 🚀 Speed up the process until x4 (see new benchmark)
@@ -30,18 +52,18 @@ Official release of [SeedVR2](https://github.com/ByteDance-Seed/SeedVR) for Comf
- 🛠️ Initial push
## Features
## 🎯 Features
- High-quality Upscaling
- Suitable for any video length once the right settings are found
- Model Will Be Download Automatically from [Models](https://huggingface.co/numz/SeedVR2_comfyUI/tree/main)
## Requirements
## 🔧 Requirements
- A Huge VRAM capabilities is better, from my test, even the 3B version need a lot of VRAM at least 18GB.
- Last ComfyUI version with python 3.12.9 (may be works with older versions but I haven't test it)
## Installation
## 📦 Installation
1. Clone this repository into your ComfyUI custom nodes directory:
@@ -86,65 +108,118 @@ python_embeded\python.exe -m pip install -r flash_attn
or can be found here ([MODELS](https://huggingface.co/numz/SeedVR2_comfyUI/tree/main))
## Usage
## 📖 Usage
1. In ComfyUI, locate the **SeedVR2 Video Upscaler** node in the node menu.
<img src="docs/node.png" width="100%">
2. things to know
2. ⚠️ **THINGS TO KNOW !!**
**temporal consistency** : at least a batch_size of 5 is required to activate temporal consistency
**temporal consistency** : at least a **batch_size** of 5 is required to activate temporal consistency. SEEDVR2 need at least 5 frames to calculate it. A higher batch_size give better performances/results but need more than 24GB VRAM.
2. Configure the node parameters:
**VRAM usage** : The input video resolution impacts VRAM consumption during the process. The larger the input video, the more VRAM will consume during the process. So, if you experience OOMs with a batch_size of at least 5, try reducing the input video resolution until it resolves.
Of course, the output resolution also has an impact, so if your hardware doesn't allow it, reduce the output resolution.
3. Configure the node parameters:
- `model`: Select your 3B or 7B model
- `seed`: a seed but it generate another seed from this one
- `new_width`: New desired Width, will keep ration on height
- `cfg_scale`:
- `batch_size`: VERY IMPORTANT!, this model consume a lot of VRAM, All your VRAM, even for the 3B model, so for GPU under 24GB VRAM keep this value Low, good value is "1" without temporal consistency
- `new_resolution`: New desired short edge in px, will keep ratio on other edge
- `cfg_scale`: usually 1.0, I don't think this could have a big impact, need to make more tests.
- `batch_size`: VERY IMPORTANT!, this model consume a lot of VRAM, All your VRAM, even for the 3B model, so for GPU under 24GB VRAM keep this value Low, good value is "1" without temporal consistency, "5" for temporal consistency, but higher is this value better is the result.
- `preserve_vram`: for VRAM < 24GB, If true, It will unload unused models during process, longer but works, otherwise probably OOM with
## Performance
## 📊 Benchmarks
**NVIDIA H100 93GB VRAM** (values in parentheses are from the previous benchmark):
**7B models on NVIDIA H100 93GB VRAM** (values in parentheses are from the previous benchmark):
| nb frames | Resolution | Batch Size | Time fp8 (s) | FPS fp8 | Time fp16 (s) | FPS fp16 |
| --------- | ------------------- | ---------- | ---------------- | ----------- | ---------------- | ----------- |
| 3 | 512×768 → 1080×1620 | 1 | 10.18 (58.10) | 0.29 (0.05) | 10.67 (60.13) | 0.28 (0.05) |
| 15 | 512×768 → 1080×1620 | 5 | 26.71 (135.63) | 0.56 (0.11) | 27.75 (144.18) | 0.54 (0.10) |
| 27 | 512×768 → 1080×1620 | 9 | 33.97 (163.22) | 0.79 (0.17) | 35.08 (177.61) | 0.77 (0.15) |
| 39 | 512×768 → 1080×1620 | 13 | 41.01 (189.36) | 0.95 (0.21) | 42.08 (210.11) | 0.93 (0.19) |
| 51 | 512×768 → 1080×1620 | 17 | 48.12 (215.80) | 1.06 (0.24) | 49.44 (242.64) | 1.03 (0.21) |
| 63 | 512×768 → 1080×1620 | 21 | 55.40 (241.79) | 1.14 (0.26) | 56.70 (275.55) | 1.11 (0.23) |
| 75 | 512×768 → 1080×1620 | 25 | 62.60 (267.93) | 1.20 (0.28) | 63.80 (308.51) | 1.18 (0.24) |
| 123 | 512×768 → 1080×1620 | 41 | 91.38 (373.60) | 1.35 (0.33) | 92.90 (440.01) | 1.32 (0.28) |
| 243 | 512×768 → 1080×1620 | 81 | 164.25 (642.20) | 1.48 (0.38) | 166.09 (780.20) | 1.46 (0.31) |
| 363 | 512×768 → 1080×1620 | 121 | 238.18 (913.61) | 1.52 (0.40) | 239.80 (1114.32) | 1.51 (0.33) |
| 453 | 512×768 → 1080×1620 | 151 | 296.52 (1132.01) | 1.53 (0.40) | 298.65 (1384.86) | 1.52 (0.33) |
| 633 | 512×768 → 1080×1620 | 211 | 406.65 (1541.09) | 1.56 (0.41) | 409.44 (1887.62) | 1.55 (0.34) |
| 903 | 512×768 → 1080×1620 | 301 | OOM (OOM) | OOM (OOM) | OOM (OOM) | OOM (OOM) |
| nb frames | Resolution | Batch Size | execution time fp8 (s) | FPS fp8 | execution time fp16 (s) | FPS fp16 | perf progress since start |
| --------- | ------------------- | ---------- | ---------------------- | ----------- | ----------------------- | ------------------ | ------------------------- |
| 15 | 512×768 → 1080×1620 | 5 | 23.75 (26.71) | 0.63 (0.56) | 24.23 (27.75) | 0.61 (0.54) (0.10) | x6.1 |
| 27 | 512×768 → 1080×1620 | 9 | 27.75 (33.97) | 0.97 (0.79) | 28.48 (35.08) | 0.94 (0.77) (0.15) | x6.2 |
| 39 | 512×768 → 1080×1620 | 13 | 32.02 (41.01) | 1.21 (0.95) | 32.62 (42.08) | 1.19 (0.93) (0.19) | x6.2 |
| 51 | 512×768 → 1080×1620 | 17 | 36.39 (48.12) | 1.40 (1.06) | 37.30 (49.44) | 1.36 (1.03) (0.21) | x6.4 |
| 63 | 512×768 → 1080×1620 | 21 | 40.80 (55.40) | 1.54 (1.14) | 41.32 (56.70) | 1.52 (1.11) (0.23) | x6.6 |
| 75 | 512×768 → 1080×1620 | 25 | 45.37 (62.60) | 1.65 (1.20) | 45.79 (63.80) | 1.63 (1.18) (0.24) | x6.8 |
| 123 | 512×768 → 1080×1620 | 41 | 62.44 (91.38) | 1.96 (1.35) | 62.28 (92.90) | 1.97 (1.32) (0.28) | x7.0 |
| 243 | 512×768 → 1080×1620 | 81 | 106.13 (164.25) | 2.28 (1.48) | 104.68 (166.09) | 2.32 (1.46) (0.31) | x7.4 |
| 363 | 512×768 → 1080×1620 | 121 | 151.01 (238.18) | 2.40 (1.52) | 148.67 (239.80) | 2.44 (1.51) (0.33) | x7.4 |
| 453 | 512×768 → 1080×1620 | 151 | 186.98 (296.52) | 2.42 (1.53) | 184.11 (298.65) | 2.46 (1.52) (0.33) | x7.4 |
| 633 | 512×768 → 1080×1620 | 211 | 253.77 (406.65) | 2.49 (1.56) | 249.43 (409.44) | 2.53 (1.55) (0.34) | x7.4 |
| 903 | 512×768 → 1080×1620 | 301 | OOM (OOM) | (OOM) | OOM (OOM) | (OOM) (OOM) | |
| 149 | 854x480 → 1920x1080 | 149 | | | 450.22 | 0.41 | |
**NVIDIA RTX4090 24GB VRAM** (preserved_vram=off)
| Model | Images | Resolution | Batch Size | Time (seconds) | FPS | Note |
| ------------------------- | ------ | ------------------- | ---------- | -------------- | --- | --- |
| 3B fp8 | 5 | 512x768 → 1080x1620 | 1 | 22.52 | 0.22 | |
| 3B fp16 | 5 | 512x768 → 1080x1620 | 1 | 27.84 | 0.18 | |
| 7B fp8 | 5 | 512x768 → 1080x1620 | 1 | 75.51 | 0.07 | |
| 7B fp16 | 5 | 512x768 → 1080x1620 | 1 | 78.93 | 0.06 | |
| 3B fp8 | 10 | 512x768 → 1080x1620 | 5 | 39.75 | 0.15 | preserve_memory=on|
| 3B fp8 | 20 | 512x768 → 1080x1620 | 1 | 65.40 | 0.31 | |
| 3B fp16 | 20 | 512x768 → 1080x1620 | 1 | 91.12 | 0.22 | |
| 3B fp8 | 20 | 512x768 → 1280x1920 | 1 | 89.10 | 0.22 | |
| 3B fp8 | 20 | 512x768 → 1480x2220 | 1 | 136.08| 0.15 | |
| 3B fp8 | 20 | 512x768 → 1620x2430 | 1 | 191.28 | 0.10 | preserve_memory=on without GPU overload so longer 320sec |
**3B FP8 models on NVIDIA H100 93GB VRAM** (values in parentheses are from the previous benchmark):
## Limitations
| nb frames | Resolution | Batch Size | execution time fp8 (s) | FPS fp8 | execution time fp16 (s) | FPS fp16 |
| --------- | ------------------- | ---------- | ---------------------- | ------- | ----------------------- | -------- |
| 149 | 854x480 → 1920x1080 | 149 | 361.22 | 0.41 | | |
**NVIDIA RTX4090 24GB VRAM**
| Model | nb frames | Resolution | Batch Size | execution time (seconds) | FPS | Note |
| ------- | --------- | ------------------- | ---------- | ------------------------ | ----------- | ---------------------------------------- |
| 3B fp8 | 5 | 512x768 → 1080x1620 | 1 | 14.66 (22.52) | 0.34 (0.22) | |
| 3B fp16 | 5 | 512x768 → 1080x1620 | 1 | 17.02 (27.84) | 0.29 (0.18) | |
| 7B fp8 | 5 | 512x768 → 1080x1620 | 1 | 46.23 (75.51) | 0.11 (0.07) | preserve_memory=on |
| 7B fp16 | 5 | 512x768 → 1080x1620 | 1 | 43.58 (78.93) | 0.11 (0.06) | preserve_memory=on |
| 3B fp8 | 10 | 512x768 → 1080x1620 | 5 | 39.75 | 0.25 | preserve_memory=on |
| 3B fp8 | 100 | 512x768 → 1080x1620 | 5 | 322.77 | 0.31 | preserve_memory=on |
| 3B fp8 | 1000 | 512x768 → 1080x1620 | 5 | 3624.08 | 0.28 | preserve_memory=on |
| 3B fp8 | 20 | 512x768 → 1080x1620 | 1 | 40.71 (65.40) | 0.49 (0.31) | |
| 3B fp16 | 20 | 512x768 → 1080x1620 | 1 | 44.76 (91.12) | 0.45 (0.22) | |
| 3B fp8 | 20 | 512x768 → 1280x1920 | 1 | 61.14 (89.10) | 0.33 (0.22) | |
| 3B fp8 | 20 | 512x768 → 1480x2220 | 1 | 79.66 (136.08) | 0.25 (0.15) | |
| 3B fp8 | 20 | 512x768 → 1620x2430 | 1 | 125.79 (191.28) | 0.16 (0.10) | preserve_memory=off (preserve_memory=on) |
| 3B fp8 | 149 | 854x480 → 1920x1080 | 5 | 782.76 | 0.19 | preserve_memory=on |
## ⚠️ Limitations
- Use a lot of VRAM, it will take all!!
- Processing speed depends on GPU capabilities
## Credits
## 🤝 Contributing
Contributions are welcome! Please feel free to submit a Pull Request. For major changes, please open an issue first to discuss what you would like to change.
Please make sure to update tests as appropriate.
### How to contribute:
1. Fork the repository
2. Create your feature branch (`git checkout -b feature/AmazingFeature`)
3. Commit your changes (`git commit -m 'Add some AmazingFeature'`)
4. Push to the branch (`git push origin feature/AmazingFeature`)
5. Open a Pull Request
### Development Setup:
1. Clone the repository
2. Install dependencies
3. Make your changes
4. Test your changes
5. Submit a pull request
### Code Style:
- Follow the existing code style
- Add comments for complex logic
- Update documentation if needed
- Ensure all tests pass
### Reporting Issues:
When reporting issues, please include:
- Your system specifications
- ComfyUI version
- Python version
- Error messages
- Steps to reproduce the issue
## 🙏 Credits
- Original [SeedVR2](https://github.com/ByteDance-Seed/SeedVR) implementation
+16 -1
View File
@@ -1,5 +1,20 @@
from .seedvr2 import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
"""
SeedVR2 Video Upscaler - Transition progressive vers architecture modulaire
Ce fichier gère la transition entre:
- Ancien code monolithique (seedvr2.py)
- Nouvelle architecture modulaire (src/)
Migration en cours...
"""
# 🆕 TENTATIVE: Nouvelle architecture modulaire
from .src.interfaces.comfyui_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
USING_MODULAR = True
# Export pour ComfyUI
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
# Métadonnées
__version__ = "1.5.0-transition" if not USING_MODULAR else "2.0.0-modular"
+32 -3
View File
@@ -87,11 +87,40 @@ def resolve_inheritance(config: Union[DictConfig, ListConfig]) -> Any:
return config
def import_item(path: str, name: str) -> Any:
def import_item(path: Union[str, List[str]], name: str) -> Any:
"""
Import a python item. Example: import_item("path.to.file", "MyClass") -> MyClass
Import a python item with fallback support.
Args:
path: Single path string or list of paths to try (fallback order)
name: Class/function name to import
Returns:
Imported object
Example:
import_item("path.to.file", "MyClass") -> MyClass
import_item(["path1.to.file", "path2.to.file"], "MyClass") -> MyClass (first working path)
"""
return getattr(importlib.import_module(path), name)
if isinstance(path, str):
# Single path - original behavior
return getattr(importlib.import_module(path), name)
elif isinstance(path, (list, ListConfig)):
# Multiple paths - try each until one works
last_error = None
for single_path in path:
try:
return getattr(importlib.import_module(single_path), name)
except ImportError as e:
last_error = e
continue
# If we get here, none of the paths worked
raise ImportError(f"Could not import '{name}' from any of the paths: {path}. Last error: {last_error}")
else:
raise ValueError(f"Path must be string or list of strings, got: {type(path)}")
def create_object(config: DictConfig) -> Any:
+16 -2
View File
@@ -4,6 +4,13 @@ __object__:
dit:
model:
__object__:
path:
- "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit_v2.nadit"
- "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit_v2.nadit"
- "models.dit_v2.nadit"
name: "NaDiT"
args: "as_params"
vid_in_channels: 33
vid_out_channels: 16
vid_dim: 2560
@@ -40,15 +47,22 @@ ema:
vae:
model:
__object__:
path:
- "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae"
- "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae"
- "models.video_vae_v3.modules.attn_video_vae"
name: "VideoAutoencoderKLWrapper"
args: "as_params"
freeze_encoder: False
gradient_checkpoint: True
gradient_checkpoint: True # Disabled to prevent VRAM leaks in inference
slicing:
split_size: 4
memory_device: same
memory_limit:
conv_max_mem: 0.5
norm_max_mem: 0.5
checkpoint: models/SEEDVR2/ema_vae_fp16.safetensors
checkpoint: ema_vae_fp16.safetensors
scaling_factor: 0.9152
compile: False
grouping: False
+15 -1
View File
@@ -4,6 +4,13 @@ __object__:
dit:
model:
__object__:
path:
- "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit.nadit"
- "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.dit.nadit"
- "models.dit.nadit"
name: "NaDiT"
args: "as_params"
vid_in_channels: 33
vid_out_channels: 16
vid_dim: 3072
@@ -37,6 +44,13 @@ ema:
vae:
model:
__object__:
path:
- "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae"
- "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.models.video_vae_v3.modules.attn_video_vae"
- "models.video_vae_v3.modules.attn_video_vae"
name: "VideoAutoencoderKLWrapper"
args: "as_params"
freeze_encoder: False
# gradient_checkpoint: True
slicing:
@@ -45,7 +59,7 @@ vae:
memory_limit:
conv_max_mem: 0.5
norm_max_mem: 0.5
checkpoint: models/SEEDVR2/ema_vae_fp16.safetensors
checkpoint: ema_vae_fp16.safetensors
scaling_factor: 0.9152
compile: False
grouping: False
+1
View File
@@ -82,6 +82,7 @@ class NaRotaryEmbedding3d(RotaryEmbedding3d):
torch.FloatTensor,
]:
freqs = cache("rope_freqs_3d", lambda: self.get_freqs(shape))
freqs = freqs.to(device=q.device, dtype=q.dtype)
q = rearrange(q, "L h d -> h L d")
k = rearrange(k, "L h d -> h L d")
q = apply_rotary_emb(freqs, q.float()).to(q.dtype)
+2
View File
@@ -19,6 +19,7 @@ from flash_attn import flash_attn_varlen_func
from torch import nn
class TorchAttention(nn.Module):
def tflops(self, args, kwargs, output) -> float:
assert len(args) == 0 or len(args) > 2, "query, key should both provided by args / kwargs"
@@ -44,3 +45,4 @@ class FlashAttentionVarlen(nn.Module):
def forward(self, *args, **kwargs):
kwargs["deterministic"] = torch.are_deterministic_algorithms_enabled()
return flash_attn_varlen_func(*args, **kwargs)
+7
View File
@@ -105,12 +105,19 @@ class NaMMSRTransformerBlock(nn.Module):
}
vid_attn, txt_attn = self.attn_norm(vid, txt)
vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="in", **ada_kwargs)
vid_attn, txt_attn = self.attn(vid_attn, txt_attn, vid_shape, txt_shape, cache)
vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="out", **ada_kwargs)
vid_attn, txt_attn = (vid_attn + vid), (txt_attn + txt)
vid_mlp, txt_mlp = self.mlp_norm(vid_attn, txt_attn)
# ADD BY NUMZ
if vid_mlp.dtype != vid_attn.dtype:
vid_mlp = vid_mlp.to(vid_attn.dtype)
if txt_mlp.dtype != txt_attn.dtype:
txt_mlp = txt_mlp.to(txt_attn.dtype)
# END BY NUMZ
vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="in", **ada_kwargs)
vid_mlp, txt_mlp = self.mlp(vid_mlp, txt_mlp)
vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="out", **ada_kwargs)
+1 -1
View File
@@ -19,7 +19,7 @@ from einops import rearrange
from rotary_embedding_torch import RotaryEmbedding, apply_rotary_emb
from torch import nn
from ...common.cache import Cache
from common.cache import Cache
class RotaryEmbeddingBase(nn.Module):
+51 -33
View File
@@ -11,6 +11,7 @@
from contextlib import nullcontext
import time
from typing import Literal, Optional, Tuple, Union
import diffusers
import torch
@@ -112,6 +113,7 @@ class Upsample3D(Upsample2D):
hidden_states: torch.FloatTensor,
output_size: Optional[int] = None,
memory_state: MemoryState = MemoryState.DISABLED,
preserve_vram: bool = False,
**kwargs,
) -> torch.FloatTensor:
assert hidden_states.shape[1] == self.channels
@@ -130,7 +132,9 @@ class Upsample3D(Upsample2D):
)
else:
hidden_states = [hidden_states]
# ADD BY NUMZ
if preserve_vram:
torch.cuda.empty_cache()
for i in range(len(hidden_states)):
hidden_states[i] = self.upscale_conv(hidden_states[i])
hidden_states[i] = rearrange(
@@ -147,10 +151,12 @@ class Upsample3D(Upsample2D):
if not self.slicing:
hidden_states = hidden_states[0]
# ADD BY NUMZ
if preserve_vram:
torch.cuda.empty_cache()
if self.use_conv:
if self.name == "conv":
hidden_states = self.conv(hidden_states, memory_state=memory_state)
hidden_states = self.conv(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
else:
hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
@@ -294,13 +300,19 @@ class ResnetBlock3D(ResnetBlock2D):
)
def forward(
self, input_tensor, temb, memory_state: MemoryState = MemoryState.DISABLED, **kwargs
self, input_tensor, temb, memory_state: MemoryState = MemoryState.DISABLED, preserve_vram: bool = False, **kwargs
):
hidden_states = input_tensor
hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = causal_norm_wrapper(self.norm1, hidden_states, preserve_vram=preserve_vram)
# ADD BY NUMZ
try:
hidden_states = self.nonlinearity(hidden_states)
except Exception as e:
print("OOM second chance")
torch.cuda.empty_cache()
time.sleep(1)
hidden_states = self.nonlinearity(hidden_states)
if self.upsample is not None:
# upsample_nearest_nhwc fails with large batch sizes.
@@ -314,7 +326,7 @@ class ResnetBlock3D(ResnetBlock2D):
input_tensor = self.downsample(input_tensor, memory_state=memory_state)
hidden_states = self.downsample(hidden_states, memory_state=memory_state)
hidden_states = self.conv1(hidden_states, memory_state=memory_state)
hidden_states = self.conv1(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
if self.time_emb_proj is not None:
if not self.skip_time_act:
@@ -333,10 +345,10 @@ class ResnetBlock3D(ResnetBlock2D):
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states, memory_state=memory_state)
hidden_states = self.conv2(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
if self.conv_shortcut is not None:
input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state)
input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state, preserve_vram=preserve_vram)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
@@ -529,14 +541,15 @@ class UpDecoderBlock3D(UpDecoderBlock2D):
hidden_states: torch.FloatTensor,
temb: Optional[torch.FloatTensor] = None,
memory_state: MemoryState = MemoryState.DISABLED,
preserve_vram: bool = False,
) -> torch.FloatTensor:
for resnet, temporal in zip(self.resnets, self.temporal_modules):
hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state)
hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state, preserve_vram=preserve_vram)
hidden_states = temporal(hidden_states)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states, memory_state=memory_state)
hidden_states = upsampler(hidden_states, memory_state=memory_state, preserve_vram=preserve_vram)
return hidden_states
@@ -791,9 +804,10 @@ class Encoder3D(nn.Module):
sample: torch.FloatTensor,
extra_cond=None,
memory_state: MemoryState = MemoryState.DISABLED,
preserve_vram: bool = False,
) -> torch.FloatTensor:
r"""The forward method of the `Encoder` class."""
sample = self.conv_in(sample, memory_state=memory_state)
sample = self.conv_in(sample, memory_state=memory_state, preserve_vram=preserve_vram)
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
@@ -966,12 +980,14 @@ class Decoder3D(nn.Module):
sample: torch.FloatTensor,
latent_embeds: Optional[torch.FloatTensor] = None,
memory_state: MemoryState = MemoryState.DISABLED,
preserve_vram: bool = False,
) -> torch.FloatTensor:
r"""The forward method of the `Decoder` class."""
sample = self.conv_in(sample, memory_state=memory_state)
upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
#upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
upscale_dtype = sample.dtype
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
@@ -1010,7 +1026,7 @@ class Decoder3D(nn.Module):
# up
for up_block in self.up_blocks:
sample = up_block(sample, latent_embeds, memory_state=memory_state)
sample = up_block(sample, latent_embeds, memory_state=memory_state, preserve_vram=preserve_vram)
# post-process
sample = causal_norm_wrapper(self.conv_norm_out, sample)
@@ -1164,8 +1180,8 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
self.decoder.mid_block.attentions = torch.nn.ModuleList([None])
@apply_forward_hook
def encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
h = self.slicing_encode(x)
def encode(self, x: torch.FloatTensor, return_dict: bool = True, preserve_vram: bool = False) -> AutoencoderKLOutput:
h = self.slicing_encode(x, preserve_vram=preserve_vram)
posterior = DiagonalGaussianDistribution(h)
if not return_dict:
@@ -1175,9 +1191,9 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
@apply_forward_hook
def decode(
self, z: torch.Tensor, return_dict: bool = True
self, z: torch.Tensor, preserve_vram: bool = False, return_dict: bool = True
) -> Union[DecoderOutput, torch.Tensor]:
decoded = self.slicing_decode(z)
decoded = self.slicing_decode(z, preserve_vram=preserve_vram)
if not return_dict:
return (decoded,)
@@ -1185,11 +1201,11 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
return DecoderOutput(sample=decoded)
def _encode(
self, x: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED
self, x: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED, preserve_vram: bool = False
) -> torch.Tensor:
_x = x.to(self.device)
_x = causal_conv_slice_inputs(_x, self.slicing_sample_min_size, memory_state=memory_state)
h = self.encoder(_x, memory_state=memory_state)
h = self.encoder(_x, memory_state=memory_state, preserve_vram=preserve_vram)
if self.quant_conv is not None:
output = self.quant_conv(h, memory_state=memory_state)
else:
@@ -1198,17 +1214,17 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
return output.to(x.device)
def _decode(
self, z: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED
self, z: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED, preserve_vram: bool = False
) -> torch.Tensor:
_z = z.to(self.device)
_z = causal_conv_slice_inputs(_z, self.slicing_latent_min_size, memory_state=memory_state)
if self.post_quant_conv is not None:
_z = self.post_quant_conv(_z, memory_state=memory_state)
output = self.decoder(_z, memory_state=memory_state)
output = self.decoder(_z, memory_state=memory_state, preserve_vram=preserve_vram)
output = causal_conv_gather_outputs(output)
return output.to(z.device)
def slicing_encode(self, x: torch.Tensor) -> torch.Tensor:
def slicing_encode(self, x: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor:
sp_size = get_sequence_parallel_world_size()
if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size * sp_size:
x_slices = x[:, :, 1:].split(split_size=self.slicing_sample_min_size * sp_size, dim=2)
@@ -1216,17 +1232,18 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
self._encode(
torch.cat((x[:, :, :1], x_slices[0]), dim=2),
memory_state=MemoryState.INITIALIZING,
preserve_vram=preserve_vram
)
]
for x_idx in range(1, len(x_slices)):
encoded_slices.append(
self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE)
self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE, preserve_vram=preserve_vram)
)
return torch.cat(encoded_slices, dim=2)
else:
return self._encode(x)
return self._encode(x, preserve_vram=preserve_vram)
def slicing_decode(self, z: torch.Tensor) -> torch.Tensor:
def slicing_decode(self, z: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor:
sp_size = get_sequence_parallel_world_size()
if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size * sp_size:
z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size * sp_size, dim=2)
@@ -1234,15 +1251,16 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL):
self._decode(
torch.cat((z[:, :, :1], z_slices[0]), dim=2),
memory_state=MemoryState.INITIALIZING,
preserve_vram=preserve_vram
)
]
for z_idx in range(1, len(z_slices)):
decoded_slices.append(
self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE)
self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE, preserve_vram=preserve_vram)
)
return torch.cat(decoded_slices, dim=2)
else:
return self._decode(z)
return self._decode(z, preserve_vram=preserve_vram)
def tiled_encode(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
raise NotImplementedError
@@ -1298,17 +1316,17 @@ class VideoAutoencoderKLWrapper(VideoAutoencoderKL):
x = self.decode(z).sample
return CausalAutoencoderOutput(x, z, p)
def encode(self, x: torch.FloatTensor) -> CausalEncoderOutput:
def encode(self, x: torch.FloatTensor, preserve_vram: bool = False) -> CausalEncoderOutput:
if x.ndim == 4:
x = x.unsqueeze(2)
p = super().encode(x).latent_dist
p = super().encode(x, preserve_vram=preserve_vram).latent_dist
z = p.sample().squeeze(2)
return CausalEncoderOutput(z, p)
def decode(self, z: torch.FloatTensor) -> CausalDecoderOutput:
def decode(self, z: torch.FloatTensor, preserve_vram: bool = False) -> CausalDecoderOutput:
if z.ndim == 4:
z = z.unsqueeze(2)
x = super().decode(z).sample.squeeze(2)
x = super().decode(z, preserve_vram).sample.squeeze(2)
return CausalDecoderOutput(x)
def preprocess(self, x: torch.Tensor):
@@ -14,6 +14,7 @@
import math
from contextlib import contextmanager
import time
from typing import List, Optional, Union
import torch
import torch.nn.functional as F
@@ -86,6 +87,7 @@ class InflatedCausalConv3d(Conv3d):
split_dim=3,
padding=(0, 0, 0, 0, 0, 0),
prev_cache=None,
preserve_vram = False,
):
# Compatible with no limit.
if math.isinf(self.memory_limit):
@@ -117,7 +119,8 @@ class InflatedCausalConv3d(Conv3d):
x = list(x.split(split_sizes, dim=split_dim))
if prev_cache is not None:
prev_cache = list(prev_cache.split(split_sizes, dim=split_dim))
if preserve_vram:
torch.cuda.empty_cache()
# Loop Fwd.
cache = None
for idx in range(len(x)):
@@ -155,17 +158,31 @@ class InflatedCausalConv3d(Conv3d):
split_dim=split_dim + 1,
padding=padding,
prev_cache=cache,
preserve_vram=preserve_vram
)
# Update cache.
cache = next_cache
return torch.cat(x, split_dim)
# ADD BY NUMZ
if preserve_vram:
torch.cuda.empty_cache()
#print("empty cache 1")
#time.sleep(2)
try:
output = torch.cat(x, split_dim)
except Exception as e:
print("OOM second chance")
torch.cuda.empty_cache()
time.sleep(2)
output = torch.cat(x, split_dim)
return output
def forward(
self,
input: Union[Tensor, List[Tensor]],
memory_state: MemoryState = MemoryState.UNSET,
preserve_vram: bool = False,
) -> Tensor:
assert memory_state != MemoryState.UNSET
if memory_state != MemoryState.ACTIVE:
@@ -176,7 +193,7 @@ class InflatedCausalConv3d(Conv3d):
and get_sequence_parallel_group() is None
):
return self.basic_forward(input, memory_state)
return self.slicing_forward(input, memory_state)
return self.slicing_forward(input, memory_state, preserve_vram)
def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET):
mem_size = self.stride[0] - self.kernel_size[0]
@@ -203,6 +220,7 @@ class InflatedCausalConv3d(Conv3d):
self,
input: Union[Tensor, List[Tensor]],
memory_state: MemoryState = MemoryState.UNSET,
preserve_vram: bool = False,
) -> Tensor:
squeeze_out = False
if torch.is_tensor(input):
@@ -249,6 +267,7 @@ class InflatedCausalConv3d(Conv3d):
input[i],
padding=padding,
prev_cache=cache,
preserve_vram=preserve_vram
)
# Update cache.
@@ -303,7 +322,7 @@ def init_causal_conv3d(
return InflatedCausalConv3d(*args, inflation_mode=inflation_mode, **kwargs)
def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor:
def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor, preserve_vram: bool = False) -> torch.Tensor:
input_dtype = x.dtype
if isinstance(norm_layer, (nn.LayerNorm, RMSNorm)):
if x.ndim == 4:
@@ -332,9 +351,25 @@ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor:
weights = norm_layer.weight.chunk(num_chunks, dim=0)
biases = norm_layer.bias.chunk(num_chunks, dim=0)
for i, (w, b) in enumerate(zip(weights, biases)):
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
try:
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
except Exception as e:
print("OOM Second Chance : Group Norm")
torch.cuda.empty_cache()
time.sleep(2)
x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
x[i] = x[i].to(input_dtype)
x = torch.cat(x, dim=1)
# ADD BY NUMZ
if preserve_vram:
torch.cuda.empty_cache()
# ADD BY NUMZ
try:
x = torch.cat(x, dim=1)
except Exception as e:
print("OOM Second Chance : Cat")
torch.cuda.empty_cache()
time.sleep(2)
x = torch.cat(x, dim=1)
else:
x = norm_layer(x)
x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
-368
View File
@@ -1,368 +0,0 @@
# // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
# //
# // 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 os
import random
import threading
from abc import ABC
from concurrent.futures import ThreadPoolExecutor, as_completed
from dataclasses import dataclass
from functools import partial
from itertools import chain
from typing import Any, Dict, List, Optional, Tuple, Union
import pyarrow as pa
import pyarrow.parquet as pq
from omegaconf import DictConfig
from common.distributed import get_global_rank, get_world_size
from common.fs import copy, exists, listdir, mkdir, remove
from common.partition import partition_by_groups
from common.persistence.utils import get_local_path
from data.common.parquet_sampler import (
IdentityParquetSampler,
ParquetSampler,
create_parquet_sampler,
)
from data.common.utils import filter_parquets, get_parquet_metadata
# Function to save a Parquet file and copy it to a target path
def save_and_copy(
pa_table,
local_path: str,
target_path: str,
row_group_size: int,
executor: ThreadPoolExecutor,
do_async: bool = False,
futures: List[Tuple[threading.Thread, str]] = [],
):
# Function to handle completion of the future
def _make_on_complete(local_path):
def _on_complete(future):
target_path = future.result()
remove(local_path)
# del future
print(f"Target path saved: {target_path}")
return _on_complete
# Function to write Parquet table and copy it
def _fn(pa_table, local_path, target_path, row_group_size):
pq.write_table(
pa_table,
local_path,
row_group_size=row_group_size,
)
mkdir(os.path.dirname(target_path))
copy(local_path, target_path)
return target_path
# Submit the task to the executor
future = executor.submit(_fn, pa_table, local_path, target_path, row_group_size)
future.add_done_callback(_make_on_complete(local_path))
futures.append(future)
# If not asynchronous, wait for all futures to complete
if not do_async:
for future in as_completed(futures):
try:
future.result()
except Exception as exc:
print(f"Generated an exception: {exc}")
executor.shutdown(wait=True)
@dataclass
class FileListOutput:
existing_files: List[str]
source_files: List[Any]
target_files: List[str]
@dataclass
class PersistedParquet:
path: str
# Method to save the Parquet file
def save(
self,
row_group_size: int,
executor: ThreadPoolExecutor,
pa_table: Optional[pa.Table] = None,
data_dict: Optional[Dict[str, List[Union[str, bytes]]]] = None,
is_last_file=False,
futures: List[threading.Thread] = [],
):
assert (pa_table is None) != (data_dict is None)
local_path = get_local_path(self.path)
if not pa_table:
schema_dict = self.generate_schema_from_dict(data_dict)
pa_table = pa.Table.from_pydict(data_dict, schema=schema_dict)
save_and_copy(
pa_table,
local_path=local_path,
target_path=self.path,
row_group_size=row_group_size,
executor=executor,
do_async=not is_last_file,
futures=futures,
)
# Method to generate schema from a dictionary
def generate_schema_from_dict(
self,
data_dict: Dict[str, List[Union[str, bytes]]],
):
schema_dict = {}
for key, value in data_dict.items():
if isinstance(value[0], str):
schema_dict[key] = pa.string()
elif isinstance(value[0], bytes):
schema_dict[key] = pa.binary()
else:
raise ValueError(f"Unsupported data type for key '{key}': {type(value)}")
return pa.schema(schema_dict)
# Base class for managing Parquet files
class ParquetManager(ABC):
"""
Base class for the DumpingManager and RepackingManager.
"""
def __init__(
self,
task: Optional[DictConfig] = None,
target_dir: str = ".",
):
self.task = task
self.target_dir = target_dir.rstrip("/")
self.executor = ThreadPoolExecutor(max_workers=4)
self.futures = []
# Method to get list of Parquet files from source path
def get_parquet_files(
self,
source_path: str,
parquet_sampler: ParquetSampler = IdentityParquetSampler(),
path_mode: str = "dir",
):
# Helper function to flatten nested lists
def _flatten(paths):
if isinstance(paths, list):
if any(isinstance(i, list) for i in paths):
return list(chain(*paths))
else:
return paths
else:
return [paths]
file_paths = _flatten(source_path)
if path_mode == "dir":
file_paths = map(listdir, file_paths)
if isinstance(parquet_sampler.size, float):
file_paths = map(filter_parquets, file_paths)
file_paths = map(parquet_sampler, file_paths)
file_paths = list(chain(*file_paths))
else:
file_paths = chain(*file_paths)
file_paths = parquet_sampler(filter_parquets(file_paths))
return file_paths
# Method to save a Parquet file
def save_parquet(
self,
*,
file_name: str,
row_group_size: int,
pa_table: Optional[pa.Table] = None,
data_dict: Optional[Dict[str, List[Union[str, bytes]]]] = None,
override: bool = True,
is_last_file: bool = False,
):
persist = self._get_parquet(file_name)
if override or not exists(persist.path):
persist.save(
pa_table=pa_table,
data_dict=data_dict,
executor=self.executor,
row_group_size=row_group_size,
is_last_file=is_last_file,
futures=self.futures,
)
# Method to get a PersistedParquet object
def _get_parquet(self, file_name: str) -> PersistedParquet:
return PersistedParquet(file_name)
# Class to manage dumping of Parquet files
class DumpingManager(ParquetManager):
"""
Dumping manager handles parquet saving and resuming.
"""
def __init__(
self,
task: DictConfig,
target_dir: str,
):
super().__init__(task=task, target_dir=target_dir)
# Method to generate saving path
def generate_saving_path(self, file_path: str, rsplit: int):
part_list = file_path.rsplit("/", rsplit)
result_folder = "/".join(
[self.target_dir] + [f"epoch_{self.task.epoch}"] + part_list[-rsplit:-1]
)
result_file = "/".join([result_folder, part_list[-1]])
return result_folder, result_file
# Method to configure task paths
def configure_task_path(self, source_path: str, rsplit: int, path_mode: str = "dir"):
file_paths = self.get_parquet_files(
source_path=source_path,
path_mode=path_mode,
)
# Shuffle file paths
random.Random(0).shuffle(file_paths)
# Partition the file paths based on task configuration
full_source_files = partition_by_groups(file_paths, self.task.total_count)[self.task.index]
full_source_files = partition_by_groups(full_source_files, get_world_size())[
get_global_rank()
]
if not full_source_files:
return FileListOutput([], [], [])
generate_saving_path = partial(self.generate_saving_path, rsplit=rsplit)
full_paths = map(generate_saving_path, full_source_files)
full_target_folders, full_target_files = map(list, zip(*full_paths))
full_target_folders = set(full_target_folders)
existing_file_paths = map(
lambda folder: listdir(folder) if exists(folder) else [], full_target_folders
)
existing_file_paths = chain(*existing_file_paths)
self.existing_files = list(
filter(
lambda path: path.endswith(".parquet") and path in full_target_files,
existing_file_paths,
)
)
filtered_pairs = list(
filter(
lambda pair: pair[1] not in self.existing_files,
zip(full_source_files, full_target_files),
)
)
if filtered_pairs:
filtered_source_files, filtered_target_files = map(list, zip(*filtered_pairs))
else:
filtered_source_files, filtered_target_files = [], []
# Skip existing file paths if specified
skip_exists = self.task.skip_exists
self.source_files = filtered_source_files if skip_exists else full_source_files
self.target_files = filtered_target_files if skip_exists else full_target_files
return FileListOutput(self.existing_files, self.source_files, self.target_files)
class RepackingManager(ParquetManager):
"""
Repacking manager handles parquet spliting and saving.
"""
def __init__(
self,
task: DictConfig,
target_dir: str,
repackaging: DictConfig,
):
super().__init__(task=task, target_dir=target_dir)
self.repackaging = repackaging
# Configure the task paths for repacking
def configure_task_path(
self,
source_path: str,
parquet_sampler: Optional[DictConfig] = None,
path_mode: str = "dir",
):
parquet_sampler = create_parquet_sampler(config=parquet_sampler)
file_paths = self.get_parquet_files(
source_path=source_path,
parquet_sampler=parquet_sampler,
path_mode=path_mode,
)
random.Random(0).shuffle(file_paths)
target_dir = self.target_dir
size = abs(parquet_sampler.size)
if self.task:
# Partition the file paths based on task configuration
file_paths = partition_by_groups(file_paths, self.task.total_count)[self.task.index]
target_dir = os.path.join(target_dir, f"{self.task.total_count}_{self.task.index}")
if size > 1:
size = len(
partition_by_groups(range(size), self.task.total_count)[self.task.index]
)
# Get metadata for each Parquet file
metadatas = get_parquet_metadata(file_paths, self.repackaging.num_processes)
# Create a list of (file_path, row) tuples for each row in the files
target_items = [
(file_path, row)
for file_path, metadata in zip(file_paths, metadatas)
for row in range(metadata.num_rows)
]
# Shuffle the target items
random.Random(0).shuffle(target_items)
if size > 1:
target_items = target_items[:size]
# Partition the items into groups for each target file
items_per_file = partition_by_groups(target_items, self.repackaging.num_files)
# Generate target file paths
target_files = [
os.path.join(target_dir, f"{str(i).zfill(5)}.parquet")
for i in range(self.repackaging.num_files)
]
existing_file_paths = listdir(target_dir) if exists(target_dir) else []
self.existing_files = list(
filter(
lambda path: path.endswith(".parquet"),
existing_file_paths,
)
)
self.source_files = items_per_file
self.target_files = target_files
return FileListOutput(self.existing_files, self.source_files, self.target_files)
+22 -12
View File
@@ -34,7 +34,6 @@ except ImportError:
from .data.image.transforms.divisible_crop import DivisibleCrop
from .data.image.transforms.na_resize import NaResize
from .data.video.transforms.rearrange import Rearrange
script_directory = os.path.dirname(os.path.abspath(__file__))
if os.path.exists(os.path.join(script_directory, "./projects/video_diffusion_sr/color_fix.py")):
from .projects.video_diffusion_sr.color_fix import wavelet_reconstruction
@@ -591,12 +590,12 @@ def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, tempo
print(f"🔄 INFERENCE time: {time.time() - t} seconds")
# Traitement des échantillons avec OPTIMISATION 🚀
t = time.time()
#t = time.time()
samples = optimized_video_rearrange(video_tensors)
last_latents = samples[-temporal_overlap:]
print(f"🔄 sample size: {len(samples)}")
print(f"🔄 Samples shape: {samples[0].shape}")
print(f"🚀 OPTIMIZED REARRANGE time: {time.time() - t} seconds")
#print(f"🔄 sample size: {len(samples)}")
#print(f"🔄 Samples shape: {samples[0].shape}")
#print(f"🚀 OPTIMIZED REARRANGE time: {time.time() - t} seconds")
# Nettoyage agressif des tenseurs intermédiaires
#del video_tensors, noises, aug_noises, cond_latents, conditions
@@ -857,6 +856,9 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
tps_vae = time.time()
with torch.autocast("cuda", autocast_dtype, enabled=True):
cond_latents = runner.vae_encode([transformed_video])
tps = time.time()
transformed_video = transformed_video.to("cpu")
print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
print(f"🔄 Cond latents shape: {cond_latents[0].shape}, time: {time.time() - tps_vae} seconds")
#text_embeds = {"texts_pos": [text_pos_embeds], "texts_neg": [text_neg_embeds]}
@@ -911,13 +913,17 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
if temporal_overlap>0 and not is_first_batch and sample.shape[0] > effective_batch_size - temporal_overlap:
sample = sample[temporal_overlap:] # Supprimer les frames de chevauchement en sortie
# 🚀 OPTIMISATION: Utiliser PyTorch natif au lieu de rearrange (2-5x plus rapide)
tps = time.time()
transformed_video = transformed_video.to(device)
print(f"🔄 Transformed video to device time: {time.time() - tps} seconds")
tps = time.time()
input_video = [optimized_single_video_rearrange(transformed_video)]
#print(f"🔄 Optimized single video rearrange time: {time.time() - t} seconds")
t = time.time()
#t = time.time()
if use_colorfix:
sample = wavelet_reconstruction(sample, input_video[0][: sample.size(0)])
print(f"🔄 Wavelet reconstruction time: {time.time() - t} seconds")
t = time.time()
#print(f"🔄 Wavelet reconstruction time: {time.time() - t} seconds")
#t = time.time()
# 🚀 OPTIMISATION: Remplacer rearrange par fonction optimisée
sample = optimized_sample_to_image_format(sample)
#print(f"🔄 Optimized sample format time: {time.time() - t} seconds")
@@ -925,12 +931,16 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
sample = sample.to("cpu")
batch_samples.append(sample)
print(f"🔄 Batch samples time: {time.time() - t} seconds")
#print(f"🔄 Batch samples time: {time.time() - t} seconds")
#t = time.time()
# Nettoyage ultra-agressif après chaque batch
print(f"🔄 Time batch: {time.time() - tps_loop} seconds")
#input_video = input_video[0].to("cpu")
tps = time.time()
video = video.to("cpu")
transformed_video = transformed_video.to("cpu")
print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
print(f"🔄 Time batch: {time.time() - tps_loop} seconds")
del samples, sample, input_video, video, transformed_video
@@ -1071,12 +1081,12 @@ class FP8CompatibleDiT(torch.nn.Module):
#self._force_nadit_bfloat16()
else:
print("🎯 Detected NaDiT 7B FP16")
self._force_nadit_bfloat16()
#self._force_nadit_bfloat16()
elif self.is_fp8_model and is_nadit_v2_3b:
# Pour NaDiT v2 3B FP8: Convertir TOUT le modèle en BFloat16
print("🎯 Detected NaDiT v2 3B FP8 - Converting all parameters to BFloat16")
self._force_nadit_bfloat16()
#self._force_nadit_bfloat16()
#else:
# Pour les autres modèles (3B FP16, etc.): Forcer seulement RoPE en BFloat16
#print("🎯 Standard model - Converting only RoPE to BFloat16")
+139
View File
@@ -0,0 +1,139 @@
"""
SeedVR2 Video Upscaler - Modular Architecture
Refactored from monolithic seedvr2.py for better maintainability
Author: Refactored codebase
Version: 2.0.0 - Modular
Available Modules:
- utils: Download and path utilities
- optimization: Memory, performance, and compatibility optimizations
- core: Model management and generation pipeline (NEW)
- processing: Video and tensor processing (coming next)
- interfaces: ComfyUI integration
"""
# Track which modules are available for progressive migration
MODULES_AVAILABLE = {
'downloads': True, # ✅ Module 1 - Downloads and model management
'memory_manager': True, # ✅ Module 2 - Memory optimization
'performance': True, # ✅ Module 3 - Performance optimizations
'compatibility': True, # ✅ Module 4 - FP8/FP16 compatibility
'model_manager': True, # ✅ Module 5 - Model configuration and loading
'generation': True, # ✅ Module 6 - Generation loop and inference
'video_transforms': True, # ✅ Module 7 - Video processing and transforms
'comfyui_node': True, # ✅ Module 8 - ComfyUI node interface (COMPLETE!)
'infer': True, # ✅ Module 9 - Infer
}
# Core imports (always available)
import os
import sys
# Add current directory to path for fallback imports
current_dir = os.path.dirname(os.path.abspath(__file__))
parent_dir = os.path.dirname(current_dir)
if parent_dir not in sys.path:
sys.path.insert(0, parent_dir)
# Progressive import system with fallback
# ===== MODULE 1: Downloads =====
if MODULES_AVAILABLE['downloads']:
from src.utils.downloads import (
download_weight,
get_base_cache_dir
)
# ===== MODULE 2: Memory Manager =====
if MODULES_AVAILABLE['memory_manager']:
from src.optimization.memory_manager import (
get_vram_usage,
clear_vram_cache,
reset_vram_peak,
preinitialize_rope_cache,
clear_rope_cache,
)
# ===== MODULE 3: Performance =====
if MODULES_AVAILABLE['performance']:
from src.optimization.performance import (
optimized_video_rearrange,
optimized_single_video_rearrange,
optimized_sample_to_image_format,
temporal_latent_blending,
)
# ===== MODULE 4: Compatibility =====
if MODULES_AVAILABLE['compatibility']:
from src.optimization.compatibility import (
FP8CompatibleDiT,
apply_fp8_compatibility_hooks,
remove_compatibility_hooks,
)
# ===== MODULE 5: Model Manager =====
if MODULES_AVAILABLE['model_manager']:
from src.core.model_manager import (
configure_runner,
load_quantized_state_dict,
configure_dit_model_inference,
configure_vae_model_inference,
)
# ===== MODULE 6: Generation =====
if MODULES_AVAILABLE['generation']:
from src.core.generation import (
generation_step,
generation_loop,
load_text_embeddings,
calculate_optimal_batch_params,
prepare_video_transforms
)
# ===== MODULE 7: Video Transforms =====
if MODULES_AVAILABLE['infer']:
from src.core.infer import VideoDiffusionInfer
# ===== MODULE 8: ComfyUI Node =====
if MODULES_AVAILABLE['comfyui_node']:
from src.interfaces.comfyui_node import (
SeedVR2,
NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS
)
# Export all available functions
__all__ = [
# Utils
'download_weight', 'get_base_cache_dir',
# Memory Management
'get_vram_usage', 'clear_vram_cache', 'reset_vram_peak',
'preinitialize_rope_cache', 'clear_rope_cache',
# Performance & Video Processing
'optimized_video_rearrange', 'optimized_single_video_rearrange', 'optimized_sample_to_image_format',
'temporal_latent_blending',
'validate_video_format', 'ensure_4n_plus_1_format', 'calculate_padding_requirements', 'apply_wavelet_reconstruction', 'temporal_consistency_check',
# Compatibility
'FP8CompatibleDiT', 'apply_fp8_compatibility_hooks', 'remove_compatibility_hooks',
# Core Model & Generation & Infer
'configure_runner', 'load_quantized_state_dict', 'configure_dit_model_inference', 'configure_vae_model_inference',
'generation_step', 'generation_loop', 'load_text_embeddings', 'calculate_optimal_batch_params',
'prepare_video_transforms', 'VideoDiffusionInfer',
# ComfyUI Interface
'SeedVR2', 'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'format_execution_results',
# Progress tracking
'get_refactoring_progress'
]
+46
View File
@@ -0,0 +1,46 @@
"""
Core Module for SeedVR2
Contains the main business logic and model management functionality:
- Model configuration and loading
- Architecture detection and memory estimation
- Runner creation and management
- Generation pipeline and logic
"""
from .model_manager import (
configure_runner,
load_quantized_state_dict,
configure_dit_model_inference,
configure_vae_model_inference,
)
from .generation import (
generation_step,
generation_loop,
cut_videos,
prepare_video_transforms,
load_text_embeddings,
calculate_optimal_batch_params
)
from .infer import VideoDiffusionInfer
__all__ = [
# Model management
'configure_runner',
'load_quantized_state_dict',
'configure_dit_model_inference',
'configure_vae_model_inference',
# Generation logic
'generation_step',
'generation_loop',
'cut_videos',
'prepare_video_transforms',
'load_text_embeddings',
'calculate_optimal_batch_params',
# Infer
'VideoDiffusionInfer'
]
+546
View File
@@ -0,0 +1,546 @@
"""
Generation Logic Module for SeedVR2
This module handles the main generation pipeline including:
- Single generation steps with adaptive dtype handling
- Complete generation loop with temporal awareness
- Context-aware batch processing with overlapping
- Video preprocessing and post-processing
- Optimized memory management during generation
Key Features:
- Native FP8 pipeline support for 2x speedup and 50% VRAM reduction
- Context-aware generation with temporal overlap for smooth transitions
- Adaptive dtype detection and optimal autocast configuration
- Intelligent batch processing with memory optimization
- Advanced video format handling (4n+1 constraint)
"""
import os
import torch
import time
import gc
from torchvision.transforms import Compose, Lambda, Normalize
# Import required modules
from src.optimization.memory_manager import reset_vram_peak
from src.optimization.performance import (
optimized_video_rearrange, optimized_single_video_rearrange,
optimized_sample_to_image_format, temporal_latent_blending
)
from common.seed import set_seed
# Get script directory for embeddings
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
# Import transforms and color fix
from data.image.transforms.divisible_crop import DivisibleCrop
from data.image.transforms.na_resize import NaResize
from src.utils.color_fix import wavelet_reconstruction
def generation_step(runner, text_embeds_dict, preserve_vram, cond_latents, temporal_overlap):
"""
Execute a single generation step with adaptive dtype handling
Args:
runner: VideoDiffusionInfer instance
text_embeds_dict (dict): Text embeddings for positive and negative prompts
preserve_vram (bool): Whether to enable VRAM optimization
cond_latents (list): Conditional latents for generation
temporal_overlap (int): Number of frames for temporal overlap
Returns:
tuple: (samples, last_latents) for potential temporal continuation
Features:
- Adaptive dtype detection (FP8/FP16/BFloat16)
- Optimal autocast configuration for each model type
- Memory-efficient noise generation and reuse
- Automatic device placement with dtype preservation
- Advanced inference optimization
"""
device = "cuda" if torch.cuda.is_available() else "cpu"
# Adaptive dtype detection for optimal performance
model_dtype = next(runner.dit.parameters()).dtype
# Configure dtypes according to model architecture
if model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
# FP8 native: use BFloat16 for intermediate calculations (optimal compatibility)
dtype = torch.bfloat16
autocast_dtype = torch.bfloat16
elif model_dtype == torch.float16:
dtype = torch.float16
autocast_dtype = torch.float16
else:
dtype = torch.bfloat16
autocast_dtype = torch.bfloat16
def _move_to_cuda(x):
"""Move tensors to CUDA with adaptive optimal dtype"""
return [i.to(device, dtype=dtype) for i in x]
# Memory optimization: Generate noise once and reuse to save VRAM
with torch.cuda.device(device):
base_noise = torch.randn_like(cond_latents[0], dtype=dtype)
noises = [base_noise]
aug_noises = [base_noise * 0.1 + torch.randn_like(base_noise) * 0.05]
# Move tensors with adaptive dtype (optimized for FP8/FP16/BFloat16)
noises, aug_noises, cond_latents = _move_to_cuda(noises), _move_to_cuda(aug_noises), _move_to_cuda(cond_latents)
cond_noise_scale = 0.0
def _add_noise(x, aug_noise):
# Use adaptive optimal dtype
t = (
torch.tensor([1000.0], device=device, dtype=dtype)
* cond_noise_scale
)
shape = torch.tensor(x.shape[1:], device=device)[None]
t = runner.timestep_transform(t, shape)
x = runner.schedule.forward(x, aug_noise, t)
return x
# Generate conditions with memory optimization
condition = runner.get_condition(
noises[0],
task="sr",
latent_blur=_add_noise(cond_latents[0], aug_noises[0]),
)
conditions = [condition]
t = time.time()
# Use adaptive autocast for optimal performance
with torch.no_grad():
with torch.autocast("cuda", autocast_dtype, enabled=True):
video_tensors = runner.inference(
noises=noises,
conditions=conditions,
preserve_vram=preserve_vram, # Memory offload optimization
temporal_overlap=temporal_overlap,
**text_embeds_dict,
)
print(f"🔄 INFERENCE time: {time.time() - t} seconds")
# Process samples with advanced optimization
samples = optimized_video_rearrange(video_tensors)
#last_latents = samples[-temporal_overlap:] if temporal_overlap > 0 else samples[-1:]
noises = noises[0].to("cpu")
aug_noises = aug_noises[0].to("cpu")
cond_latents = cond_latents[0].to("cpu")
conditions = conditions[0].to("cpu")
condition = condition.to("cpu")
return samples #, last_latents
def cut_videos(videos):
"""
Correct video cutting respecting the constraint: frames % 4 == 1
Args:
videos (torch.Tensor): Video tensor to format
Returns:
torch.Tensor: Properly formatted video tensor
Features:
- Ensures frames % 4 == 1 constraint for model compatibility
- Intelligent padding with last frame repetition
- Memory-efficient tensor operations
"""
t = videos.size(1)
if t % 4 == 1:
return videos
# Calculate next valid number (4n + 1)
padding_needed = (4 - (t % 4)) % 4 + 1
# Apply padding to reach 4n+1 format
last_frame = videos[:, -1:].expand(-1, padding_needed, -1, -1).contiguous()
result = torch.cat([videos, last_frame], dim=1)
return result
def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_size=90, preserve_vram=False, temporal_overlap=0, debug=False):
"""
Main generation loop with context-aware temporal processing
Args:
runner: VideoDiffusionInfer instance
images (torch.Tensor): Input images for upscaling
cfg_scale (float): Classifier-free guidance scale
seed (int): Random seed for reproducibility
res_w (int): Target resolution width
batch_size (int): Batch size for processing
preserve_vram (str/bool): VRAM preservation mode
temporal_overlap (int): Frames for temporal continuity
Returns:
torch.Tensor: Generated video frames
Features:
- Context-aware generation with temporal overlap
- Adaptive dtype pipeline (FP8/FP16/BFloat16)
- Memory-optimized batch processing
- Advanced video transformation pipeline
- Intelligent VRAM management throughout process
"""
device = "cuda" if torch.cuda.is_available() else "cpu"
# Adaptive model dtype detection for maximum performance
model_dtype = None
try:
# Get real dtype of loaded DiT model
model_dtype = next(runner.dit.parameters()).dtype
# Adapt dtypes according to model
if model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
# For FP8, use BFloat16 for intermediate calculations (compatible)
compute_dtype = torch.bfloat16
autocast_dtype = torch.bfloat16
vae_dtype = torch.bfloat16 # VAE stays BFloat16 for compatibility
elif model_dtype == torch.float16:
compute_dtype = torch.float16
autocast_dtype = torch.float16
vae_dtype = torch.float16
else: # BFloat16 or others
compute_dtype = torch.bfloat16
autocast_dtype = torch.bfloat16
vae_dtype = torch.bfloat16
except Exception as e:
print(f"⚠️ Could not detect model dtype: {e}, falling back to BFloat16")
model_dtype = torch.bfloat16
compute_dtype = torch.bfloat16
autocast_dtype = torch.bfloat16
vae_dtype = torch.bfloat16
# Optimization tips for users
if torch.cuda.is_available():
total_frames = len(images)
optimal_batches = [x for x in [i for i in range(1, 200) if i % 4 == 1] if x <= total_frames]
if optimal_batches:
best_batch = max(optimal_batches)
if best_batch != batch_size:
print(f"\n💡 TIP: For {total_frames} frames, use batch_size={best_batch} to avoid padding")
if batch_size not in optimal_batches:
padding_waste = sum(((i // 4) + 1) * 4 + 1 - i for i in range(batch_size, total_frames, batch_size))
print(f" Currently: ~{padding_waste} wasted padding frames")
# Configure classifier-free guidance
runner.config.diffusion.cfg.scale = cfg_scale
runner.config.diffusion.cfg.rescale = 0.0
# Configure sampling steps
runner.config.diffusion.timesteps.sampling.steps = 1
runner.configure_diffusion()
# Set random seed
set_seed(seed)
# Advanced video transformation pipeline
video_transform = Compose([
NaResize(
resolution=(res_w),
mode="side",
# Upsample image, model only trained for high res
downsample_only=False,
),
Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),
DivisibleCrop((16, 16)),
Normalize(0.5, 0.5),
Lambda(lambda x: x.permute(1, 0, 2, 3)), # t c h w -> c t h w (faster than Rearrange)
])
# Initialize generation state
batch_samples = []
final_tensor = None
# Load text embeddings with adaptive dtype
text_pos_embeds = torch.load(os.path.join(script_directory, 'pos_emb.pt')).to(device, dtype=compute_dtype)
text_neg_embeds = torch.load(os.path.join(script_directory, 'neg_emb.pt')).to(device, dtype=compute_dtype)
text_embeds = {"texts_pos": [text_pos_embeds], "texts_neg": [text_neg_embeds]}
# Memory optimization
reset_vram_peak()
# Calculate processing parameters
step = batch_size - temporal_overlap
if step <= 0:
step = batch_size
temporal_overlap = 0
# Move images to CPU for memory efficiency
#t = time.time()
#images = images.to("cpu")
#print(f"🔄 Images to CPU time: {time.time() - t} seconds")
try:
# Main processing loop with context awareness
for batch_idx in range(0, len(images), step):
# Calculate batch indices with overlap
if batch_idx == 0:
# First batch: no overlap
start_idx = 0
end_idx = min(batch_size, len(images))
effective_batch_size = end_idx - start_idx
is_first_batch = True
else:
# Subsequent batches: temporal overlap
start_idx = batch_idx
end_idx = min(start_idx + batch_size, len(images))
effective_batch_size = end_idx - start_idx
is_first_batch = False
if effective_batch_size <= temporal_overlap:
break # Not enough new frames, stop
tps_loop = time.time()
batch_number = (batch_idx // step + 1) if step > 0 else 1
print(f"\n🎬 Batch {batch_number}: frames {start_idx}-{end_idx-1}")
# Process current batch
video = images[start_idx:end_idx]
if debug:
print(f"🔄 video Compute dtype: {compute_dtype}")
# Use adaptive computation dtype
video = video.permute(0, 3, 1, 2).to(device, dtype=compute_dtype)
# Apply video transformations with memory optimization
transformed_video = video_transform(video)
del video
#video = video.to("cpu")
#del video
ori_lengths = [transformed_video.size(1)]
# Handle correct format: frames % 4 == 1
t = transformed_video.size(1)
print(f"📹 Sequence of {t} frames")
if len(images) >= 5 and t % 4 != 1:
if debug:
print(f"🔄 Transformed video shape before cut: {transformed_video.shape}")
transformed_video = cut_videos(transformed_video)
if debug:
print(f"🔄 Transformed video shape: {transformed_video.shape}")
# Context-aware temporal strategy
# First batch: standard complete diffusion
tps_vae = time.time()
runner.vae.to(device)
if debug:
print(f"🔄 VAE to GPU time: {time.time() - tps_vae} seconds")
tps_vae = time.time()
if debug:
print(f"🔄 VAE dtype: {autocast_dtype}")
with torch.autocast("cuda", autocast_dtype, enabled=True):
cond_latents = runner.vae_encode([transformed_video])
if debug:
print(f"🔄 VAE encode time: {time.time() - tps_vae} seconds")
#tps = time.time()
#transformed_video = transformed_video.to("cpu")
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
if debug:
print(f"🔄 Cond latents shape: {cond_latents[0].shape}, time: {time.time() - tps_vae} seconds")
# Normal generation
samples = generation_step(runner, text_embeds, preserve_vram, cond_latents=cond_latents, temporal_overlap=temporal_overlap)
#del cond_latents
del cond_latents
# Post-process samples
sample = samples[0]
del samples
# 🔧 DIAGNOSTIC: Vérifier les valeurs après generation_step
'''
if debug:
print(f"🔍 DIAGNOSTIC - Après generation_step:")
print(f" 📊 Sample shape: {sample.shape}")
print(f" 📊 Sample dtype: {sample.dtype}")
print(f" 📊 Sample device: {sample.device}")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
print(f" 📊 Sample std: {sample.std():.6f}")
print(f" 📊 Non-zero count: {(sample != 0).sum()}/{sample.numel()}")
'''
#del samples
if ori_lengths[0] < sample.shape[0]:
sample = sample[:ori_lengths[0]]
#if temporal_overlap > 0 and not is_first_batch and sample.shape[0] > effective_batch_size - temporal_overlap:
# sample = sample[temporal_overlap:] # Remove overlap frames from output
# Apply color correction if available
tps = time.time()
transformed_video = transformed_video.to(device)
if debug:
print(f"🔄 Transformed video to device time: {time.time() - tps} seconds")
input_video = [optimized_single_video_rearrange(transformed_video)]
del transformed_video
#transformed_video = transformed_video.to("cpu")
#del transformed_video
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)])
del input_video
'''
if debug:
# 🔧 DIAGNOSTIC: Vérifier les valeurs après wavelet_reconstruction
print(f"🔍 DIAGNOSTIC - Après wavelet_reconstruction:")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
print(f" 📊 Non-zero count: {(sample != 0).sum()}/{sample.numel()}")
'''
#del input_video
# Convert to final image format
sample = optimized_sample_to_image_format(sample)
'''
if debug:
# 🔧 DIAGNOSTIC: Vérifier les valeurs après optimized_sample_to_image_format
print(f"🔍 DIAGNOSTIC - Après optimized_sample_to_image_format:")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
'''
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
'''
if debug:
# 🔧 DIAGNOSTIC: Vérifier les valeurs finales
print(f"🔍 DIAGNOSTIC - Valeurs finales:")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
print(f" 🎯 Est-ce que l'image est noire? {sample.max() < 0.01}")
'''
sample_cpu = sample.to("cpu")
del sample
batch_samples.append(sample_cpu)
#del sample
# Aggressive cleanup after each batch
tps = time.time()
#transformed_video = transformed_video.to("cpu")
#print(f"🔄 Transformed video to cpu time: {time.time() - tps} seconds")
if debug:
print(f"🔄 Time batch: {time.time() - tps_loop} seconds")
if preserve_vram:
torch.cuda.empty_cache()
#del transformed_video
#clear_vram_cache()
finally:
# Final cleanup of embeddings
text_pos_embeds = text_pos_embeds.to("cpu")
text_neg_embeds = text_neg_embeds.to("cpu")
#del text_pos_embeds, text_neg_embeds
#clear_vram_cache()
for i in range(len(batch_samples)):
batch_samples[i] = batch_samples[i].to(device)
# Concatenate all batch results
final_video_images = torch.cat(batch_samples, dim=0)
final_video_images = final_video_images.to("cpu")
# Critical correction: Convert to Float16 for ComfyUI compatibility
if final_video_images.dtype != torch.float16:
final_video_images = final_video_images.to(torch.float16)
# Cleanup batch_samples
#del batch_samples
return final_video_images
def prepare_video_transforms(res_w):
"""
Prepare optimized video transformation pipeline
Args:
res_w (int): Target resolution width
Returns:
Compose: Configured transformation pipeline
Features:
- Resolution-aware upscaling (no downsampling)
- Proper normalization for model compatibility
- Memory-efficient tensor operations
"""
return Compose([
NaResize(
resolution=(res_w),
mode="side",
downsample_only=False, # Model trained for high resolution
),
Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),
DivisibleCrop((16, 16)),
Normalize(0.5, 0.5),
Lambda(lambda x: x.permute(1, 0, 2, 3)), # t c h w -> c t h w
])
def load_text_embeddings(script_directory, device, dtype):
"""
Load and prepare text embeddings for generation
Args:
script_directory (str): Script directory path
device (str): Target device
dtype (torch.dtype): Target dtype
Returns:
dict: Text embeddings dictionary
Features:
- Adaptive dtype handling
- Device-optimized loading
- Memory-efficient embedding preparation
"""
text_pos_embeds = torch.load(os.path.join(script_directory, 'pos_emb.pt')).to(device, dtype=dtype)
text_neg_embeds = torch.load(os.path.join(script_directory, 'neg_emb.pt')).to(device, dtype=dtype)
return {"texts_pos": [text_pos_embeds], "texts_neg": [text_neg_embeds]}
def calculate_optimal_batch_params(total_frames, batch_size, temporal_overlap):
"""
Calculate optimal batch processing parameters
Args:
total_frames (int): Total number of frames
batch_size (int): Desired batch size
temporal_overlap (int): Temporal overlap frames
Returns:
dict: Optimized parameters and recommendations
Features:
- 4n+1 constraint optimization
- Padding waste calculation
- Performance recommendations
"""
step = batch_size - temporal_overlap
if step <= 0:
step = batch_size
temporal_overlap = 0
# Find optimal batch sizes (4n+1 constraint)
optimal_batches = [x for x in [i for i in range(1, 200) if i % 4 == 1] if x <= total_frames]
best_batch = max(optimal_batches) if optimal_batches else 1
# Calculate potential padding waste
padding_waste = 0
if batch_size not in optimal_batches:
padding_waste = sum(((i // 4) + 1) * 4 + 1 - i for i in range(batch_size, total_frames, batch_size))
return {
'step': step,
'temporal_overlap': temporal_overlap,
'best_batch': best_batch,
'padding_waste': padding_waste,
'is_optimal': batch_size in optimal_batches
}
@@ -19,26 +19,29 @@ import gc
from einops import rearrange
from omegaconf import DictConfig, ListConfig
from torch import Tensor
from src.optimization.memory_manager import clear_vram_cache
from models.video_vae_v3.modules.types import MemoryState
from ...common.config import create_object
from ...common.decorators import log_on_entry, log_runtime
from ...common.diffusion import (
from common.config import create_object
from common.decorators import log_on_entry, log_runtime
from common.diffusion import (
classifier_free_guidance_dispatcher,
create_sampler_from_config,
create_sampling_timesteps_from_config,
create_schedule_from_config,
)
from ...common.distributed import (
from common.distributed import (
get_device,
get_global_rank,
)
from ...common.distributed.meta_init_utils import (
from common.distributed.meta_init_utils import (
meta_non_persistent_buffer_init_fn,
)
# from common.fs import download
from ...models.dit_v2 import na
from models.dit_v2 import na
def optimized_channels_to_last(tensor):
"""🚀 Optimized replacement for rearrange(tensor, 'b c ... -> b ... c')
@@ -73,10 +76,9 @@ def optimized_channels_to_second(tensor):
return tensor.permute(*dims)
class VideoDiffusionInfer():
def __init__(self, config: DictConfig):
# print(config)
def __init__(self, config: DictConfig, debug: bool = False):
self.config = config
self.debug = debug
def get_condition(self, latent: Tensor, latent_blur: Tensor, task: str) -> Tensor:
t, h, w, c = latent.shape
cond = torch.zeros([t, h, w, c + 1], device=latent.device, dtype=latent.dtype)
@@ -102,7 +104,7 @@ class VideoDiffusionInfer():
cond[:, ..., -1:] = 1.0
return cond
raise NotImplementedError
'''
@log_on_entry
@log_runtime
def configure_dit_model(self, device="cpu", checkpoint=None):
@@ -154,7 +156,7 @@ class VideoDiffusionInfer():
self.vae.set_causal_slicing(**self.config.vae.slicing)
# ------------------------------ Diffusion ------------------------------ #
'''
def configure_diffusion(self):
self.schedule = create_schedule_from_config(
config=self.config.diffusion.schedule,
@@ -174,7 +176,7 @@ class VideoDiffusionInfer():
# -------------------------------- Helper ------------------------------- #
@torch.no_grad()
def vae_encode(self, samples: List[Tensor]) -> List[Tensor]:
def vae_encode(self, samples: List[Tensor], preserve_vram: bool = False) -> List[Tensor]:
use_sample = self.config.vae.get("use_sample", True)
latents = []
if len(samples) > 0:
@@ -201,6 +203,7 @@ class VideoDiffusionInfer():
sample = self.vae.preprocess(sample)
if use_sample:
latent = self.vae.encode(sample).latent
#latent = self.vae.encode(sample, preserve_vram).latent
else:
# Deterministic vae encode, only used for i2v inference (optionally)
latent = self.vae.encode(sample).posterior.mode().squeeze(2)
@@ -217,9 +220,10 @@ class VideoDiffusionInfer():
latents = [latent.squeeze(0) for latent in latents]
return latents
@torch.no_grad()
def vae_decode(self, latents: List[Tensor], target_dtype: torch.dtype = None) -> List[Tensor]:
def vae_decode(self, latents: List[Tensor], target_dtype: torch.dtype = None, preserve_vram: bool = False) -> List[Tensor]:
"""🚀 VAE decode optimisé - décodage direct sans chunking, compatible avec autocast externe"""
samples = []
if len(latents) > 0:
@@ -234,11 +238,15 @@ class VideoDiffusionInfer():
if isinstance(shift, ListConfig):
shift = torch.tensor(shift, device=device, dtype=dtype)
# 🚀 OPTIMISATION 1: Group latents intelligemment pour batch processing
if self.config.vae.grouping:
latents, indices = na.pack(latents)
else:
latents = [latent.unsqueeze(0) for latent in latents]
if self.debug:
print(f"🔄 shape of latents: {latents[0].shape}")
#print(f"🔄 GROUPING time: {time.time() - t} seconds")
t = time.time()
# 🚀 OPTIMISATION 2: Traitement batch optimisé avec dtype adaptatif
@@ -253,8 +261,9 @@ class VideoDiffusionInfer():
latent = latent.squeeze(2)
# 🚀 OPTIMISATION 3: Décodage direct SANS autocast (utilise l'autocast externe)
with torch.autocast("cuda", torch.float16, enabled=True):
sample = self.vae.decode(latent).sample
#with torch.autocast("cuda", torch.float16, enabled=True):
sample = self.vae.decode(latent, preserve_vram).sample
#sample = self.vae.decode(latent).sample
#sample = self.vae.decode(latent).sample
# 🚀 OPTIMISATION 4: Post-processing conditionnel
@@ -264,9 +273,11 @@ class VideoDiffusionInfer():
samples.append(sample)
# 🚀 OPTIMISATION 5: Nettoyage sélectif
if i % 2 == 0 or i == len(latents) - 1:
torch.cuda.empty_cache()
print(f"🔄 DECODE time: {time.time() - t} seconds")
#if i % 2 == 0 or i == len(latents) - 1:
#torch.cuda.empty_cache()
if self.debug:
print(f"🔄 DECODE time: {time.time() - t} seconds")
#t = time.time()
# Ungroup back to individual sample with the original order.
if self.config.vae.grouping:
@@ -329,13 +340,6 @@ class VideoDiffusionInfer():
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
def clear_vram_cache(self):
"""Nettoyer le cache VRAM"""
if torch.cuda.is_available():
torch.cuda.empty_cache()
import gc
gc.collect()
@torch.no_grad()
def inference(
self,
@@ -344,7 +348,7 @@ class VideoDiffusionInfer():
texts_pos: Union[List[str], List[Tensor], List[Tuple[Tensor]]],
texts_neg: Union[List[str], List[Tensor], List[Tuple[Tensor]]],
cfg_scale: Optional[float] = None,
dit_offload: bool = False,
preserve_vram: bool = False,
temporal_overlap: int = 0,
) -> List[Tensor]:
assert len(noises) == len(conditions) == len(texts_pos) == len(texts_neg)
@@ -363,19 +367,21 @@ class VideoDiffusionInfer():
# 🚀 OPTIMISATION: Détecter le dtype du modèle pour performance optimale
model_dtype = next(self.dit.parameters()).dtype
if self.debug:
print(f"🎯 model_dtype: {model_dtype}")
# Adapter les dtypes selon le modèle
if model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
# FP8 natif: utiliser BFloat16 pour les calculs intermédiaires (compatible)
target_dtype = torch.bfloat16
target_dtype = torch.float16
#print(f"🚀 FP8 model detected: using BFloat16 for intermediate calculations")
elif model_dtype == torch.float16:
target_dtype = torch.float16
target_dtype = torch.bfloat16
#print(f"🎯 FP16 model: using FP16 pipeline")
else:
target_dtype = torch.bfloat16
#print(f"🎯 BFloat16 model: using BFloat16 pipeline")
if self.debug:
print(f"🎯 target_dtype: {target_dtype}")
# Text embeddings.
assert type(texts_pos[0]) is type(texts_neg[0])
if isinstance(texts_pos[0], str):
@@ -410,17 +416,20 @@ class VideoDiffusionInfer():
latents = latents.to(target_dtype) if latents.dtype != target_dtype else latents
latents_cond = latents_cond.to(target_dtype) if latents_cond.dtype != target_dtype else latents_cond
# Enter eval mode.
# was_training = self.dit.training
# self.dit.eval()
if preserve_vram:
if conditions[0].shape[0] > 1:
t = time.time()
self.vae = self.vae.to("cpu")
if self.debug:
print(f"🔄 VAE to CPU time: {time.time() - t} seconds")
t = time.time()
self.dit = self.dit.to(get_device())
if self.debug:
print(f"🔄 Dit to GPU time: {time.time() - t} seconds")
# Sampling avec optimisations VRAM et autocast adaptatif
# print(f"🔄 Starting EulerSampler sampling...")
t = time.time()
self.dit.to(get_device())
print(f"🔄 Dit to GPU time: {time.time() - t} seconds")
t = time.time()
with torch.autocast("cuda", target_dtype, enabled=True):
latents = self.sampler.sample(
x=latents,
@@ -448,47 +457,43 @@ class VideoDiffusionInfer():
rescale=self.config.diffusion.cfg.rescale,
),
)
print(f"🔄 INFERENCE time: {time.time() - t} seconds")
t = time.time()
if dit_offload:
self.dit.to("cpu")
print(f"🔄 Dit to CPU time: {time.time() - t} seconds")
if self.debug:
print(f"🔄 INFERENCE time: {time.time() - t} seconds")
# Exit eval mode.
#self.dit.train(was_training)
# 🚀 PIPELINE OPTIMISÉ: Réduire les transferts et nettoyages
t = time.time()
# Unflatten.
latents = na.unflatten(latents, latents_shapes)
print(f"🔄 UNFLATTEN time: {time.time() - t} seconds")
# 🚀 OPTIMISATION: Préparer VAE en parallèle si pas d'offloading
#if not dit_offload:
#t = time.time()
#self.vae.to(get_device())
#print(f"🔄 VAE to GPU time: {time.time() - t} seconds")
#else:
#t = time.time()
#if dit_offload:
#self.dit.to(get_device())
#print(f"🔄 DIT to GPU time: {time.time() - t} seconds")
#t = time.time()
#self.vae.to(get_device())
#print(f"🔄 VAE to GPU time: {time.time() - t} seconds")
# 🚀 VAE decode optimisé ULTRA-RAPIDE - Élimination des transferts coûteux
#t = time.time()
#print(f"🔄 UNFLATTEN time: {time.time() - t} seconds")
# 🎯 Pré-calcul des dtypes (une seule fois)
vae_dtype = getattr(torch, self.config.vae.dtype)
decode_dtype = torch.float16 if (vae_dtype == torch.float16 or target_dtype == torch.float16) else vae_dtype
if self.debug:
print(f"🎯 decode_dtype: {decode_dtype}")
if preserve_vram:
t = time.time()
self.dit = self.dit.to("cpu")
latents_cond = latents_cond.to("cpu")
latents_shapes = latents_shapes.to("cpu")
if latents[0].shape[0] > 1:
clear_vram_cache()
if self.debug:
print(f"🔄 Dit to CPU time: {time.time() - t} seconds")
if latents[0].shape[0] > 1:
t = time.time()
self.vae = self.vae.to(get_device())
if self.debug:
print(f"🔄 VAE to GPU time: {time.time() - t} seconds")
#with torch.autocast("cuda", decode_dtype, enabled=True):
samples = self.vae_decode(latents, target_dtype=decode_dtype, preserve_vram=preserve_vram)
with torch.autocast("cuda", decode_dtype, enabled=True):
samples = self.vae_decode(latents, target_dtype=decode_dtype)
print(f"🔄 Samples shape: {samples[0].shape}")
if self.debug:
print(f"🔄 Samples shape: {samples[0].shape}")
#print(f"🔄 ULTRA-FAST VAE DECODE time: {time.time() - t} seconds")
#t = time.time()
#self.dit.to(get_device())
@@ -497,7 +502,8 @@ class VideoDiffusionInfer():
#t = time.time()
# 🚀 CORRECTION CRITIQUE: Conversion batch Float16 pour ComfyUI (plus rapide)
if samples and len(samples) > 0 and samples[0].dtype != torch.float16:
print(f"🔧 Converting {len(samples)} samples from {samples[0].dtype} to Float16")
if self.debug:
print(f"🔧 Converting {len(samples)} samples from {samples[0].dtype} to Float16")
samples = [sample.to(torch.float16, non_blocking=True) for sample in samples]
#print(f"🚀 Conversion batch Float16 time: {time.time() - t} seconds")
@@ -514,4 +520,4 @@ class VideoDiffusionInfer():
#print(f"🔄 FINAL CLEANUP time: {time.time() - t} seconds")
return samples
return samples
+358
View File
@@ -0,0 +1,358 @@
"""
Model Management Module for SeedVR2
This module handles all model-related operations including:
- Model configuration and path resolution
- Model loading with format detection (SafeTensors, PyTorch)
- DiT and VAE model setup and inference configuration
- State dict management with native FP8 support
- Universal compatibility wrappers
Key Features:
- Dynamic import path resolution for different ComfyUI environments
- Native FP8 model support with optimal performance
- Automatic compatibility mode for model architectures
- Memory-efficient model loading and configuration
"""
import os
import time
import torch
from omegaconf import DictConfig, OmegaConf
# Import SafeTensors with fallback
try:
from safetensors.torch import load_file as load_safetensors_file
SAFETENSORS_AVAILABLE = True
except ImportError:
print("⚠️ SafeTensors not available, recommended install: pip install safetensors")
SAFETENSORS_AVAILABLE = False
from src.optimization.memory_manager import get_basic_vram_info, clear_vram_cache
from src.optimization.compatibility import FP8CompatibleDiT
from src.optimization.memory_manager import preinitialize_rope_cache
from common.config import load_config, create_object
from src.core.infer import VideoDiffusionInfer
# NOUVEAU: Import des opérations ComfyUI pour FP8
import comfy.ops as ops
import comfy.model_management as model_management
# Get script directory for config paths
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False):
"""
Configure and create a VideoDiffusionInfer runner for the specified model
Args:
model (str): Model filename (e.g., "seedvr2_ema_3b_fp16.safetensors")
base_cache_dir (str): Base directory containing model files
Returns:
VideoDiffusionInfer: Configured runner instance ready for inference
Features:
- Dynamic config loading based on model type (3B vs 7B)
- Automatic import path resolution for different environments
- VAE configuration with proper parameter handling
- Memory optimization and RoPE cache pre-initialization
"""
t = time.time()
vram_info = get_basic_vram_info()
if debug:
print(f"🔄 RUNNER : VRAM INFO: {vram_info}")
# Select config based on model type
if "7b" in model:
config_path = os.path.join(script_directory, './configs_7b', 'main.yaml')
model_weight = "7b_fp8" if "fp8" in model else "7b_fp16"
else:
config_path = os.path.join(script_directory, './configs_3b', 'main.yaml')
model_weight = "3b_fp8" if "fp8" in model else "3b_fp16"
config = load_config(config_path)
if debug:
print(f"🔄 RUNNER : CONFIG LOAD TIME: {time.time() - t} seconds")
# DiT model configuration is now handled directly in the YAML config files
# No need for dynamic path resolution here anymore!
# Load and configure VAE with additional parameters
vae_config_path = os.path.join(script_directory, 'models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml')
t = time.time()
vae_config = OmegaConf.load(vae_config_path)
if debug:
print(f"🔄 RUNNER : VAE CONFIG LOAD TIME: {time.time() - t} seconds")
t = time.time()
# Configure VAE parameters
spatial_downsample_factor = vae_config.get('spatial_downsample_factor', 8)
temporal_downsample_factor = vae_config.get('temporal_downsample_factor', 4)
vae_config.spatial_downsample_factor = spatial_downsample_factor
vae_config.temporal_downsample_factor = temporal_downsample_factor
if debug:
print(f"🔄 RUNNER : VAE CONFIG SET TIME: {time.time() - t} seconds")
# Merge additional VAE config with main config (preserving __object__ from main config)
t = time.time()
config.vae.model = OmegaConf.merge(config.vae.model, vae_config)
if debug:
print(f"🔄 RUNNER : VAE CONFIG MERGE TIME: {time.time() - t} seconds")
t = time.time()
# Create runner
runner = VideoDiffusionInfer(config, debug)
OmegaConf.set_readonly(runner.config, False)
if debug:
print(f"🔄 RUNNER : RUNNER VIDEO DIFFUSION INFER TIME: {time.time() - t} seconds")
# Set device
device = "cuda" if torch.cuda.is_available() else "cpu"
# Configure models
checkpoint_path = os.path.join(base_cache_dir, f'./{model}')
t = time.time()
runner = configure_dit_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug)
if debug:
print(f"🔄 RUNNER : DIT MODEL INFERENCE TIME: {time.time() - t} seconds")
t = time.time()
checkpoint_path = os.path.join(base_cache_dir, f'./{config.vae.checkpoint}')
runner = configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram, model_weight, vram_info, debug)
if debug:
print(f"🔄 RUNNER : VAE MODEL INFERENCE TIME: {time.time() - t} seconds")
t = time.time()
if hasattr(runner.vae, "set_memory_limit"):
runner.vae.set_memory_limit(**runner.config.vae.memory_limit)
if debug:
print(f"🔄 RUNNER : VAE MEMORY LIMIT TIME: {time.time() - t} seconds")
# Pre-initialize RoPE cache for optimal performance
t = time.time()
preinitialize_rope_cache(runner)
if debug:
print(f"🔄 RUNNER : ROPE CACHE PREINITIALIZE TIME: {time.time() - t} seconds")
#clear_vram_cache()
return runner
def load_quantized_state_dict(checkpoint_path, device="cpu", keep_native_fp8=True):
"""
Load state dict from SafeTensors or PyTorch with optimal FP8 native support
Args:
checkpoint_path (str): Path to model checkpoint (.safetensors or .pth)
device (str): Target device for loading
keep_native_fp8 (bool): Whether to preserve native FP8 format for performance
Returns:
dict: State dictionary with optimal dtype handling
Features:
- Automatic format detection (SafeTensors vs PyTorch)
- Native FP8 preservation for 2x speedup and 50% VRAM reduction
- Intelligent dtype conversion when needed for compatibility
- Memory-mapped loading for large models
"""
if checkpoint_path.endswith('.safetensors'):
if not SAFETENSORS_AVAILABLE:
raise ImportError("SafeTensors required to load this model. Install with: pip install safetensors")
state = load_safetensors_file(checkpoint_path, device=device)
elif checkpoint_path.endswith('.pth'):
state = torch.load(checkpoint_path, map_location=device, mmap=True)
else:
raise ValueError(f"Unsupported format. Expected .safetensors or .pth, got: {checkpoint_path}")
# FP8 optimization: Keep native format for maximum performance
fp8_detected = False
fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) if hasattr(torch, 'float8_e4m3fn') else ()
if fp8_types:
# Check if model contains FP8 tensors
for key, tensor in state.items():
if hasattr(tensor, 'dtype') and tensor.dtype in fp8_types:
fp8_detected = True
break
if fp8_detected:
if keep_native_fp8:
# Keep native FP8 format for optimal performance
# Benefits: ~50% less VRAM, ~2x faster inference
return state
else:
# Convert FP8 → BFloat16 for compatibility
converted_state = {}
converted_count = 0
for key, tensor in state.items():
if hasattr(tensor, 'dtype') and tensor.dtype in fp8_types:
converted_state[key] = tensor.to(torch.bfloat16)
converted_count += 1
else:
converted_state[key] = tensor
return converted_state
return state
def configure_dit_model_inference(runner, device, checkpoint, config, preserve_vram=False, model_weight=None, vram_info=None, debug=False):
"""
Configure DiT model for inference without distributed decorators
Args:
runner: VideoDiffusionInfer instance
device (str): Target device
checkpoint (str): Path to model checkpoint
config: Model configuration
Features:
- Automatic format detection and optimal loading
- Native FP8 support with universal compatibility wrapper
- Gradient checkpointing configuration
- Intelligent dtype handling for all model architectures
"""
# Create dit model
t = time.time()
loading_device = "cpu" if preserve_vram else device
with torch.device(device):
runner.dit = create_object(config.dit.model)
# Passer les opérations au modèle
if debug:
print(f"🔄 CONFIG DIT : MODEL CREATE TIME: {time.time() - t} seconds device: {device}")
t = time.time()
runner.dit.set_gradient_checkpointing(config.dit.gradient_checkpoint)
# Detect and log model format
print(f"🚀 Loading model_weight: {model_weight}")
t = time.time()
state_loading_device = "cpu" if "7b" in model_weight and vram_info['total_gb'] < 25 else device
state = load_quantized_state_dict(checkpoint, state_loading_device, keep_native_fp8=True)
if debug:
print(f"🔄 CONFIG DIT : DiT load state dict time: {time.time() - t} seconds")
t = time.time()
runner.dit.load_state_dict(state, strict=True, assign=True)
if 'state' in locals():
del state
if debug:
print(f"🔄 CONFIG DIT : DiT load time: {time.time() - t} seconds")
#state.to("cpu")
#runner.dit = runner.dit.to(device)
# Apply universal compatibility wrapper to ALL models
# This ensures RoPE compatibility and optimal performance across all architectures
t = time.time()
runner.dit = FP8CompatibleDiT(runner.dit)
if debug:
print(f"🔄 CONFIG DIT : FP8CompatibleDiT time: {time.time() - t} seconds")
# Move DiT to CPU to prevent VRAM leaks (especially for 3B model with complex RoPE)
if preserve_vram:
if debug:
print(f"🔄 CONFIG DIT : dit to cpu cause preserve_vram: {preserve_vram}")
runner.dit = runner.dit.to("cpu")
if "7b" in model_weight:
clear_vram_cache()
else:
if state_loading_device == "cpu":
runner.dit.to(device)
return runner
def configure_vae_model_inference(runner, device, checkpoint_path, config, preserve_vram=False, model_weight=None, vram_info=None, debug=False):
"""
Configure VAE model for inference without distributed decorators
Args:
runner: VideoDiffusionInfer instance
config: Model configuration
device (str): Target device
Features:
- Dynamic path resolution for VAE checkpoints
- SafeTensors and PyTorch format support
- FP8 and FP16 VAE handling
- Causal slicing configuration
"""
# Create vae model
dtype = getattr(torch, config.vae.dtype)
t = time.time()
loading_device = "cpu" if preserve_vram else device
with torch.device(device):
runner.vae = create_object(config.vae.model)
if debug:
print(f"🔄 CONFIG VAE : MODEL CREATE TIME: {time.time() - t} seconds device: {device} dtype: {dtype}")
t = time.time()
runner.vae.requires_grad_(False).eval()
if debug:
print(f"🔄 CONFIG VAE : MODEL REQUIRES GRAD TIME: {time.time() - t} seconds device: {device} dtype: {dtype}")
t = time.time()
#runner.vae.to(device=loading_device, dtype=dtype)
#if debug:
# print(f"🔄 CONFIG VAE : TO CPU TIME: {time.time() - t} seconds device: {device} dtype: {dtype}")
# Resolve VAE checkpoint path dynamically
'''
checkpoint_path = config.vae.checkpoint
possible_paths = [
checkpoint_path, # Original path
os.path.join("ComfyUI", checkpoint_path), # With ComfyUI prefix
os.path.join(script_directory, checkpoint_path), # Relative to script directory
os.path.join(script_directory, "..", "..", checkpoint_path), # From ComfyUI root
]
t = time.time()
vae_checkpoint_path = None
for path in possible_paths:
if os.path.exists(path):
vae_checkpoint_path = path
if debug:
print(f"🔄 CONFIG VAE : Found VAE checkpoint at: {vae_checkpoint_path}")
break
if debug:
print(f"🔄 CONFIG VAE : VAE CHECKPOINT PATH TIME: {time.time() - t} seconds")
if vae_checkpoint_path is None:
raise FileNotFoundError(f"VAE checkpoint not found. Tried paths: {possible_paths}")
'''
# Load VAE with format detection
t = time.time()
state_loading_device = "cpu" if "7b" in model_weight and vram_info['total_gb'] < 25 else device
print(f"🚀 Loading VAE SafeTensors: {checkpoint_path}")
# Use optimized loading for all SafeTensors formats
if "fp8_e4m3fn" in checkpoint_path:
state = load_quantized_state_dict(checkpoint_path, state_loading_device, keep_native_fp8=True)
else:
# For FP16 SafeTensors, disable native FP8
state = load_quantized_state_dict(checkpoint_path, state_loading_device, keep_native_fp8=False)
if debug:
print(f"🔄 CONFIG VAE : VAE LOAD TIME: {time.time() - t} seconds")
t = time.time()
runner.vae.load_state_dict(state)
if state_loading_device == "cpu":
runner.vae.to(device)
if 'state' in locals():
del state
if debug:
print(f"🔄 CONFIG VAE : VAE LOAD STATE DICT TIME: {time.time() - t} seconds")
# Set causal slicing if available
t = time.time()
if hasattr(runner.vae, "set_causal_slicing") and hasattr(config.vae, "slicing"):
runner.vae.set_causal_slicing(**config.vae.slicing)
if debug:
print(f"🔄 CONFIG VAE : VAE SET CAUSAL SLICING TIME: {time.time() - t} seconds")
return runner
#runner.vae.to("cpu")
+27
View File
@@ -0,0 +1,27 @@
"""
Interfaces package for SeedVR2
Contains user interface integrations (ComfyUI, etc.)
"""
# ComfyUI Interfaces Module
# Handles ComfyUI node integration and user interface
from .comfyui_node import (
# Main ComfyUI node class
SeedVR2,
# ComfyUI mappings
NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS,
)
__all__ = [
# Core node interface
'SeedVR2',
# ComfyUI mappings
'NODE_CLASS_MAPPINGS',
'NODE_DISPLAY_NAME_MAPPINGS',
]
+222
View File
@@ -0,0 +1,222 @@
# ComfyUI Node Interface
# Clean interface for SeedVR2 VideoUpscaler integration with ComfyUI
# Extracted from original seedvr2.py lines 1731-1812
from datetime import datetime
import os
import gc
import time
import torch
from typing import Tuple, Dict, Any
from src.utils.downloads import download_weight, get_base_cache_dir
from src.core.model_manager import configure_runner
from src.core.generation import generation_loop
from src.optimization.memory_manager import clear_rope_lru_caches, fast_model_cleanup
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
class SeedVR2:
"""
SeedVR2 Video Upscaler ComfyUI Node
High-quality video upscaling using diffusion models with support for:
- Multiple model variants (3B/7B, FP16/FP8)
- Adaptive VRAM management
- Advanced dtype compatibility
- Optimized inference pipeline
"""
def __init__(self):
"""Initialize SeedVR2 node"""
self.runner = None
self.text_pos_embeds = None
self.text_neg_embeds = None
self.current_model_name = ""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Dict[str, Any]]:
"""
Define ComfyUI input parameter types and constraints
Returns:
Dictionary defining input parameters, types, and validation
"""
return {
"required": {
"images": ("IMAGE", ),
"model": ([
"seedvr2_ema_3b_fp16.safetensors",
"seedvr2_ema_3b_fp8_e4m3fn.safetensors",
"seedvr2_ema_7b_fp16.safetensors",
"seedvr2_ema_7b_fp8_e4m3fn.safetensors",
], {
"default": "seedvr2_ema_3b_fp8_e4m3fn.safetensors"
}),
"seed": ("INT", {
"default": 100,
"min": 0,
"max": 5000,
"step": 1,
"tooltip": "Random seed for generation reproducibility"
}),
"new_width": ("INT", {
"default": 1280,
"min": 1,
"max": 4320,
"step": 1,
"tooltip": "Target width for upscaled video"
}),
"cfg_scale": ("FLOAT", {
"default": 1.0,
"min": 0.01,
"max": 2.0,
"step": 0.01,
"tooltip": "Classifier-free guidance scale"
}),
"batch_size": ("INT", {
"default": 5,
"min": 1,
"max": 2048,
"step": 4,
"tooltip": "Number of frames to process per batch (recommend 4n+1 format)"
}),
"preserve_vram": ("BOOLEAN", {"default": True})
},
}
# Define return types for ComfyUI
RETURN_NAMES = ("image", )
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
CATEGORY = "SEEDVR2"
def execute(self, images: torch.Tensor, model: str, seed: int, new_width: int,
cfg_scale: float, batch_size: int, preserve_vram: bool) -> Tuple[torch.Tensor]:
"""Execute SeedVR2 video upscaling"""
temporal_overlap = 0
print(f"🔄 Preparing model: {model}")
download_weight(model)
debug = False
try:
return self._internal_execute(images, model, seed, new_width, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug)
except Exception as e:
self.cleanup(force_ram_cleanup=True)
raise e
def cleanup(self, force_ram_cleanup: bool = True):
"""Fast cleanup with minimal logging"""
if self.runner:
# Clear cache
if hasattr(self.runner, 'cache') and hasattr(self.runner.cache, 'cache'):
for key, value in list(self.runner.cache.cache.items()):
if hasattr(value, 'cpu'):
value.cpu()
if hasattr(value, 'detach'):
value.detach()
del value
self.runner.cache.cache.clear()
# Clear DiT model
if hasattr(self.runner, 'dit') and self.runner.dit is not None:
clear_rope_lru_caches(self.runner.dit)
fast_model_cleanup(self.runner.dit)
del self.runner.dit
self.runner.dit = None
# Clear VAE model
if hasattr(self.runner, 'vae') and self.runner.vae is not None:
#from src.optimization.memory_manager import fast_model_cleanup
fast_model_cleanup(self.runner.vae)
del self.runner.vae
self.runner.vae = None
# Clear other components
for component in ['sampler', 'sampling_timesteps', 'schedule', 'config']:
if hasattr(self.runner, component):
setattr(self.runner, component, None)
del self.runner
self.runner = None
# Clear embeddings
if self.text_pos_embeds is not None:
if hasattr(self.text_pos_embeds, 'cpu'):
self.text_pos_embeds.cpu()
del self.text_pos_embeds
self.text_pos_embeds = None
if self.text_neg_embeds is not None:
if hasattr(self.text_neg_embeds, 'cpu'):
self.text_neg_embeds.cpu()
del self.text_neg_embeds
self.text_neg_embeds = None
self.current_model_name = ""
# Fast RAM cleanup
if force_ram_cleanup:
from src.optimization.memory_manager import fast_ram_cleanup
fast_ram_cleanup()
def _internal_execute(self, images, model, seed, new_width, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug):
"""Internal execution logic"""
total_start_time = time.time()
# Configure runner
if debug:
print("🔄 Configuring inference runner...")
runner_start = time.time()
self.runner = configure_runner(model, get_base_cache_dir(), preserve_vram, debug)
if debug:
print(f"🔄 Runner configuration time: {time.time() - runner_start:.2f}s")
if debug:
print("🚀 Starting video upscaling generation...")
# Execute generation
sample = generation_loop(
self.runner, images, cfg_scale, seed, new_width,
batch_size, preserve_vram, temporal_overlap, debug
)
print(f"✅ Video upscaling completed successfully!")
# Cleanup
print(f"🔄 Total execution time: {time.time() - total_start_time:.2f}s")
self.cleanup(force_ram_cleanup=True)
return (sample,)
def __del__(self):
"""Destructor"""
try:
self.cleanup(force_ram_cleanup=True)
except:
pass
# ComfyUI Node Mappings
NODE_CLASS_MAPPINGS = {
"SeedVR2": SeedVR2,
}
# Human-readable node display names
NODE_DISPLAY_NAME_MAPPINGS = {
"SeedVR2": "SeedVR2 Video Upscaler",
}
# Export version and metadata
__version__ = "2.0.0-modular"
__author__ = "SeedVR2 Team"
__description__ = "High-quality video upscaling using advanced diffusion models"
# Additional exports for introspection
__all__ = [
'SeedVR2',
'NODE_CLASS_MAPPINGS',
'NODE_DISPLAY_NAME_MAPPINGS'
]
+48
View File
@@ -0,0 +1,48 @@
"""
Optimization package for SeedVR2
Contains memory management, performance optimizations, and compatibility layers
"""
# Memory management functions
from .memory_manager import (
get_vram_usage,
clear_vram_cache,
reset_vram_peak,
preinitialize_rope_cache,
clear_rope_cache,
)
# Performance optimization functions
from .performance import (
optimized_video_rearrange,
optimized_single_video_rearrange,
optimized_sample_to_image_format,
temporal_latent_blending,
)
# Compatibility functions and classes
from .compatibility import (
FP8CompatibleDiT,
apply_fp8_compatibility_hooks,
remove_compatibility_hooks,
)
__all__ = [
# Memory management
"get_vram_usage",
"clear_vram_cache",
"reset_vram_peak",
"preinitialize_rope_cache",
"clear_rope_cache",
# Performance optimization
"optimized_video_rearrange",
"optimized_single_video_rearrange",
"optimized_sample_to_image_format",
"temporal_latent_blending",
# Compatibility
"FP8CompatibleDiT",
"apply_fp8_compatibility_hooks",
"remove_compatibility_hooks",
]
+483
View File
@@ -0,0 +1,483 @@
"""
Compatibility module for SeedVR2
Contains FP8/FP16 compatibility layers and wrappers for different model architectures
Extracted from: seedvr2.py (lines 1045-1630)
"""
import time
import torch
from typing import List, Tuple, Union, Any, Optional
class FP8CompatibleDiT(torch.nn.Module):
"""
Wrapper for DiT models with automatic compatibility management + advanced optimizations
- FP8: Keeps native FP8 parameters, converts inputs/outputs
- FP16: Uses native FP16
- RoPE: ALWAYS forced to BFloat16 for maximum compatibility
- Flash Attention: Automatic optimization of attention layers
"""
def __init__(self, dit_model):
super().__init__()
self.dit_model = dit_model
self.model_dtype = self._detect_model_dtype()
self.is_fp8_model = self.model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
self.is_fp16_model = self.model_dtype == torch.float16
# Detect model type
is_nadit_7b = self._is_nadit_model() # NaDiT 7B (dit/nadit)
is_nadit_v2_3b = self._is_nadit_v2_model() # NaDiT v2 3B (dit_v2/nadit)
if is_nadit_7b:
# 🎯 CRITICAL FIX: ALL NaDiT 7B models (FP8 AND FP16) require BFloat16 conversion
# 7B architecture has dtype compatibility issues regardless of storage format
if self.is_fp8_model:
print("🎯 Detected NaDiT 7B FP8 - Converting all parameters to BFloat16")
self._force_nadit_bfloat16()
else:
print("🎯 Detected NaDiT 7B FP16")
elif self.is_fp8_model and is_nadit_v2_3b:
# For NaDiT v2 3B FP8: Convert ALL model to BFloat16
print("🎯 Detected NaDiT v2 3B FP8 - Converting all parameters to BFloat16")
self._force_nadit_bfloat16()
# 🚀 FLASH ATTENTION OPTIMIZATION (Phase 2)
self._apply_flash_attention_optimization()
def _detect_model_dtype(self) -> torch.dtype:
"""Detect main model dtype"""
try:
return next(self.dit_model.parameters()).dtype
except:
return torch.bfloat16
def _is_nadit_model(self) -> bool:
"""Detect if this is a NaDiT model (7B) with precise logic"""
# 🎯 PRIMARY METHOD: Check emb_scale attribute (specific to 7B)
# This is the most reliable criterion to distinguish 7B vs 3B
if hasattr(self.dit_model, 'emb_scale'):
return True
# 🎯 SECONDARY METHOD: Check module path for NaDiT 7B (dit/nadit, not dit_v2)
model_module = str(self.dit_model.__class__.__module__).lower()
if 'dit.nadit' in model_module and 'dit_v2' not in model_module:
return True
return False
def _is_nadit_v2_model(self) -> bool:
"""Detect if this is a NaDiT v2 model (3B) with precise logic"""
# 🎯 PRIMARY METHOD: Check module path for NaDiT v2 (dit_v2/nadit)
model_module = str(self.dit_model.__class__.__module__).lower()
if 'dit_v2' in model_module:
return True
# 🎯 SECONDARY METHOD: Check specific 3B structure
# NaDiT v2 3B has vid_in, txt_in, emb_in but NO emb_scale
if (hasattr(self.dit_model, 'vid_in') and
hasattr(self.dit_model, 'txt_in') and
hasattr(self.dit_model, 'emb_in') and
not hasattr(self.dit_model, 'emb_scale')): # Absence of emb_scale = 3B
return True
return False
def _force_rope_bfloat16(self) -> None:
"""🎯 Force ALL RoPE modules to BFloat16 for maximum compatibility"""
rope_count = 0
for name, module in self.dit_model.named_modules():
# Identify RoPE modules by name or type
if any(keyword in name.lower() for keyword in ['rope', 'rotary', 'embedding']):
# Convert all parameters of this module to BFloat16
for param_name, param in module.named_parameters():
if param.dtype != torch.bfloat16:
param.data = param.data.to(torch.bfloat16)
rope_count += 1
# Also convert buffers (non-trainable parameters)
for buffer_name, buffer in module.named_buffers():
if buffer.dtype != torch.bfloat16:
buffer.data = buffer.data.to(torch.bfloat16)
rope_count += 1
def _force_nadit_bfloat16(self) -> None:
"""🎯 Force ALL NaDiT parameters to BFloat16 to avoid promotion errors"""
print("🔧 Converting ALL NaDiT parameters to BFloat16 for type compatibility...")
t = time.time()
converted_count = 0
original_dtype = None
# Convert ALL parameters to BFloat16 (FP8, FP16, etc.)
for name, param in self.dit_model.named_parameters():
if original_dtype is None:
original_dtype = param.dtype
if param.dtype != torch.bfloat16:
param.data = param.data.to(torch.bfloat16)
converted_count += 1
# Also convert buffers
for name, buffer in self.dit_model.named_buffers():
if buffer.dtype != torch.bfloat16:
buffer.data = buffer.data.to(torch.bfloat16)
converted_count += 1
print(f" ✅ Converted {converted_count} parameters/buffers from {original_dtype} to BFloat16")
# Update detected dtype
self.model_dtype = torch.bfloat16
self.is_fp8_model = False # Model is no longer FP8 after conversion
def _apply_flash_attention_optimization(self) -> None:
"""🚀 FLASH ATTENTION OPTIMIZATION - 30-50% speedup of attention layers"""
attention_layers_optimized = 0
flash_attention_available = self._check_flash_attention_support()
for name, module in self.dit_model.named_modules():
# Identify all attention layers
if self._is_attention_layer(name, module):
# Apply optimization based on availability
if self._optimize_attention_layer(name, module, flash_attention_available):
attention_layers_optimized += 1
if not flash_attention_available:
print(" ℹ️ Flash Attention not available, using PyTorch SDPA as fallback")
def _check_flash_attention_support(self) -> bool:
"""Check if Flash Attention is available"""
try:
# Check PyTorch SDPA (includes Flash Attention on H100/A100)
if hasattr(torch.nn.functional, 'scaled_dot_product_attention'):
return True
# Check flash-attn package
import flash_attn
return True
except ImportError:
return False
def _is_attention_layer(self, name: str, module: torch.nn.Module) -> bool:
"""Identify if a module is an attention layer"""
attention_keywords = [
'attention', 'attn', 'self_attn', 'cross_attn', 'mhattn', 'multihead',
'transformer_block', 'dit_block'
]
# Check by name
if any(keyword in name.lower() for keyword in attention_keywords):
return True
# Check by module type
module_type = type(module).__name__.lower()
if any(keyword in module_type for keyword in attention_keywords):
return True
# Check by attributes (modules with q, k, v projections)
if hasattr(module, 'q_proj') or hasattr(module, 'qkv') or hasattr(module, 'to_q'):
return True
return False
def _optimize_attention_layer(self, name: str, module: torch.nn.Module, flash_attention_available: bool) -> bool:
"""Optimize a specific attention layer"""
try:
# Save original forward method
if not hasattr(module, '_original_forward'):
module._original_forward = module.forward
# Create new optimized forward method
if flash_attention_available:
optimized_forward = self._create_flash_attention_forward(module, name)
else:
optimized_forward = self._create_sdpa_forward(module, name)
# Replace forward method
module.forward = optimized_forward
return True
except Exception as e:
print(f" ⚠️ Failed to optimize attention layer '{name}': {e}")
return False
def _create_flash_attention_forward(self, module: torch.nn.Module, layer_name: str):
"""Create optimized forward with Flash Attention"""
original_forward = module._original_forward
def flash_attention_forward(*args, **kwargs):
try:
# Try to use Flash Attention via SDPA
return self._sdpa_attention_forward(original_forward, module, *args, **kwargs)
except Exception as e:
# Fallback to original implementation
print(f" ⚠️ Flash Attention failed for {layer_name}, using original: {e}")
return original_forward(*args, **kwargs)
return flash_attention_forward
def _create_sdpa_forward(self, module: torch.nn.Module, layer_name: str):
"""Create optimized forward with PyTorch SDPA"""
original_forward = module._original_forward
def sdpa_forward(*args, **kwargs):
try:
return self._sdpa_attention_forward(original_forward, module, *args, **kwargs)
except Exception as e:
# Fallback to original implementation
return original_forward(*args, **kwargs)
return sdpa_forward
def _sdpa_attention_forward(self, original_forward, module: torch.nn.Module, *args, **kwargs):
"""Optimized forward pass using SDPA (Scaled Dot Product Attention)"""
# Detect if we can intercept and optimize this layer
if len(args) >= 1 and isinstance(args[0], torch.Tensor):
input_tensor = args[0]
# Check dimensions to ensure it's standard attention
if len(input_tensor.shape) >= 3: # [batch, seq_len, hidden_dim] or similar
try:
return self._optimized_attention_computation(module, input_tensor, *args[1:], **kwargs)
except:
pass
# Fallback to original implementation
return original_forward(*args, **kwargs)
def _optimized_attention_computation(self, module: torch.nn.Module, input_tensor: torch.Tensor, *args, **kwargs):
"""Optimized attention computation with SDPA"""
# Try to detect standard attention format
batch_size, seq_len = input_tensor.shape[:2]
# Check if module has standard Q, K, V projections
if hasattr(module, 'qkv') or (hasattr(module, 'q_proj') and hasattr(module, 'k_proj') and hasattr(module, 'v_proj')):
return self._compute_sdpa_attention(module, input_tensor, *args, **kwargs)
# If no standard format detected, use original
return module._original_forward(input_tensor, *args, **kwargs)
def _compute_sdpa_attention(self, module: torch.nn.Module, x: torch.Tensor, *args, **kwargs):
"""Optimized SDPA computation for standard attention modules"""
try:
# Case 1: Module with combined QKV projection
if hasattr(module, 'qkv'):
qkv = module.qkv(x)
# Reshape to separate Q, K, V
batch_size, seq_len, _ = qkv.shape
qkv = qkv.reshape(batch_size, seq_len, 3, -1)
q, k, v = qkv.unbind(dim=2)
# Case 2: Separate Q, K, V projections
elif hasattr(module, 'q_proj') and hasattr(module, 'k_proj') and hasattr(module, 'v_proj'):
q = module.q_proj(x)
k = module.k_proj(x)
v = module.v_proj(x)
else:
# Unsupported format, use original
return module._original_forward(x, *args, **kwargs)
# Detect number of heads
head_dim = getattr(module, 'head_dim', None)
num_heads = getattr(module, 'num_heads', None)
if head_dim is None or num_heads is None:
# Try to guess from dimensions
hidden_dim = q.shape[-1]
if hasattr(module, 'num_heads'):
num_heads = module.num_heads
head_dim = hidden_dim // num_heads
else:
# Reasonable defaults
head_dim = 64
num_heads = hidden_dim // head_dim
# Reshape for multi-head attention
batch_size, seq_len = q.shape[:2]
q = q.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
k = k.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
v = v.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
# Use optimized SDPA
with torch.backends.cuda.sdp_kernel(
enable_flash=True,
enable_math=True,
enable_mem_efficient=True
):
attn_output = torch.nn.functional.scaled_dot_product_attention(
q, k, v,
dropout_p=0.0,
is_causal=False
)
# Reshape back
attn_output = attn_output.transpose(1, 2).contiguous().view(
batch_size, seq_len, num_heads * head_dim
)
# Output projection if it exists
if hasattr(module, 'out_proj') or hasattr(module, 'o_proj'):
proj = getattr(module, 'out_proj', None) or getattr(module, 'o_proj', None)
attn_output = proj(attn_output)
return attn_output
except Exception as e:
# In case of error, use original implementation
return module._original_forward(x, *args, **kwargs)
def forward(self, *args, **kwargs):
"""Forward pass with intelligent type management according to architecture"""
is_nadit_7b = self._is_nadit_model()
is_nadit_v2_3b = self._is_nadit_v2_model()
# Input conversion according to architecture
if is_nadit_7b or is_nadit_v2_3b:
# For NaDiT models (7B and v2 3B): Everything to BFloat16
converted_args = []
for arg in args:
if isinstance(arg, torch.Tensor):
if arg.dtype in (torch.float32, torch.float8_e4m3fn, torch.float8_e5m2):
converted_args.append(arg.to(torch.bfloat16))
else:
converted_args.append(arg)
else:
converted_args.append(arg)
converted_kwargs = {}
for key, value in kwargs.items():
if isinstance(value, torch.Tensor):
if value.dtype in (torch.float32, torch.float8_e4m3fn, torch.float8_e5m2):
converted_kwargs[key] = value.to(torch.bfloat16)
else:
converted_kwargs[key] = value
else:
converted_kwargs[key] = value
args = tuple(converted_args)
kwargs = converted_kwargs
else:
# For standard models: Conversion according to model dtype
if self.is_fp8_model:
# Convert FP8 → BFloat16 for calculations
converted_args = []
for arg in args:
if isinstance(arg, torch.Tensor) and arg.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
converted_args.append(arg.to(torch.bfloat16))
else:
converted_args.append(arg)
converted_kwargs = {}
for key, value in kwargs.items():
if isinstance(value, torch.Tensor) and value.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
converted_kwargs[key] = value.to(torch.bfloat16)
else:
converted_kwargs[key] = value
args = tuple(converted_args)
kwargs = converted_kwargs
elif self.is_fp16_model:
# Convert Float32 → FP16 for FP16 models
converted_args = []
for arg in args:
if isinstance(arg, torch.Tensor) and arg.dtype == torch.float32:
converted_args.append(arg.to(torch.float16))
else:
converted_args.append(arg)
converted_kwargs = {}
for key, value in kwargs.items():
if isinstance(value, torch.Tensor) and value.dtype == torch.float32:
converted_kwargs[key] = value.to(torch.float16)
else:
converted_kwargs[key] = value
args = tuple(converted_args)
kwargs = converted_kwargs
try:
return self.dit_model(*args, **kwargs)
except Exception as e:
print(f"❌ Error in forward pass: {e}")
print(f" Model type: NaDiT 7B={is_nadit_7b}, NaDiT v2 3B={is_nadit_v2_3b}")
print(f" Args dtypes: {[arg.dtype if isinstance(arg, torch.Tensor) else type(arg) for arg in args]}")
print(f" Kwargs dtypes: {[(k, v.dtype if isinstance(v, torch.Tensor) else type(v)) for k, v in kwargs.items()]}")
raise
def __getattr__(self, name):
"""Redirect all other attributes to original model"""
if name in ['dit_model', 'model_dtype', 'is_fp8_model', 'is_fp16_model', '_forward_count']:
return super().__getattr__(name)
return getattr(self.dit_model, name)
def __setattr__(self, name, value):
"""Redirect assignments to original model except for our attributes"""
if name in ['dit_model', 'model_dtype', 'is_fp8_model', 'is_fp16_model', '_forward_count']:
super().__setattr__(name, value)
else:
if hasattr(self, 'dit_model'):
setattr(self.dit_model, name, value)
else:
super().__setattr__(name, value)
def apply_fp8_compatibility_hooks(model: torch.nn.Module) -> List[Tuple[str, Any]]:
"""
Hook system to intercept problematic FP8 modules
Alternative if the wrapper is not sufficient.
Args:
model: Model to apply hooks to
Returns:
List of (module_name, hook) tuples for cleanup
"""
def create_fp8_safe_hook(original_dtype: torch.dtype):
def hook_fn(module, input, output):
# Convert FP8 output → BFloat16 if necessary for compatibility
if isinstance(output, torch.Tensor) and output.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
# Temporarily keep in BFloat16 to avoid downstream errors
return output.to(torch.bfloat16)
elif isinstance(output, (tuple, list)):
# Handle multiple outputs
converted_output = []
for item in output:
if isinstance(item, torch.Tensor) and item.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
converted_output.append(item.to(torch.bfloat16))
else:
converted_output.append(item)
return type(output)(converted_output)
return output
return hook_fn
# Apply hooks to critical modules
problematic_modules = []
for name, module in model.named_modules():
# Identify RoPE and attention modules that cause FP8 problems
if any(keyword in name.lower() for keyword in ['rope', 'rotary', 'attention', 'mmattn']):
if hasattr(module, 'register_forward_hook'):
hook = module.register_forward_hook(create_fp8_safe_hook(torch.float8_e4m3fn))
problematic_modules.append((name, hook))
print(f"🔧 Applied FP8 compatibility hooks to {len(problematic_modules)} modules")
return problematic_modules
def remove_compatibility_hooks(hooks: List[Tuple[str, Any]]) -> None:
"""
Remove previously applied compatibility hooks
Args:
hooks: List of (module_name, hook) tuples from apply_fp8_compatibility_hooks
"""
removed_count = 0
for name, hook in hooks:
try:
hook.remove()
removed_count += 1
except Exception as e:
print(f"⚠️ Failed to remove hook from {name}: {e}")
print(f"🧹 Removed {removed_count}/{len(hooks)} compatibility hooks")
+251
View File
@@ -0,0 +1,251 @@
"""
Memory management module for SeedVR2
Handles VRAM usage, cache management, and memory optimization
Extracted from: seedvr2.py (lines 373-405, 607-626, 1016-1044)
"""
import os
import torch
import gc
import time
from typing import Tuple, Optional
from common.cache import Cache
from models.dit_v2.rope import RotaryEmbeddingBase
def get_basic_vram_info():
"""🔍 Méthode basique avec PyTorch natif"""
if not torch.cuda.is_available():
return {"error": "CUDA not available"}
# Mémoire libre et totale (en bytes)
free_memory, total_memory = torch.cuda.mem_get_info()
# Conversion en GB
free_gb = free_memory / (1024**3)
total_gb = total_memory / (1024**3)
return {
"free_gb": free_gb,
"total_gb": total_gb
}
# Utilisation
vram_info = get_basic_vram_info()
print(f"VRAM libre: {vram_info['free_gb']:.2f} GB")
def get_vram_usage() -> Tuple[float, float, float]:
"""
Get current VRAM usage (allocated, reserved, peak)
Returns:
tuple: (allocated_gb, reserved_gb, max_allocated_gb)
Returns (0, 0, 0) if CUDA not available
"""
if torch.cuda.is_available():
allocated = torch.cuda.memory_allocated() / (1024**3)
reserved = torch.cuda.memory_reserved() / (1024**3)
max_allocated = torch.cuda.max_memory_allocated() / (1024**3)
return allocated, reserved, max_allocated
return 0, 0, 0
def clear_vram_cache() -> None:
"""
Clear VRAM cache and run garbage collection
"""
print("🧹 Clearing VRAM cache...")
if torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
def reset_vram_peak() -> None:
"""
Reset VRAM peak counter for new tracking
"""
if torch.cuda.is_available():
torch.cuda.reset_peak_memory_stats()
def preinitialize_rope_cache(runner) -> None:
"""
🚀 Pre-initialize RoPE cache to avoid OOM at first launch
Args:
runner: The model runner containing DiT and VAE models
"""
try:
# Create dummy tensors to simulate common shapes
# Format: [batch, channels, frames, height, width] for vid_shape
# Format: [batch, seq_len] for txt_shape
common_shapes = [
# Common video resolutions
(torch.tensor([[1, 3, 3]], dtype=torch.long), torch.tensor([[77]], dtype=torch.long)), # 1 frame, 77 tokens
(torch.tensor([[4, 3, 3]], dtype=torch.long), torch.tensor([[77]], dtype=torch.long)), # 4 frames
(torch.tensor([[5, 3, 3]], dtype=torch.long), torch.tensor([[77]], dtype=torch.long)), # 5 frames (4n+1 format)
(torch.tensor([[1, 4, 4]], dtype=torch.long), torch.tensor([[77]], dtype=torch.long)), # Higher resolution
]
# Create mock cache for pre-initialization
temp_cache = Cache()
# Access RoPE modules in DiT (recursive search)
def find_rope_modules(module):
rope_modules = []
for name, child in module.named_modules():
if hasattr(child, 'get_freqs') and callable(getattr(child, 'get_freqs')):
rope_modules.append((name, child))
return rope_modules
rope_modules = find_rope_modules(runner.dit)
# Pre-calculate for each RoPE module found
for name, rope_module in rope_modules:
# Temporarily move module to CPU if necessary
original_device = next(rope_module.parameters()).device if list(rope_module.parameters()) else torch.device('cpu')
rope_module.to('cpu')
try:
for vid_shape, txt_shape in common_shapes:
cache_key = f"720pswin_by_size_bysize_{tuple(vid_shape[0].tolist())}_sd3.mmrope_freqs_3d"
def compute_freqs():
try:
# Calculate with reduced dimensions to avoid OOM
with torch.no_grad():
# Detect RoPE module type
module_type = type(rope_module).__name__
if module_type == 'NaRotaryEmbedding3d':
# NaRotaryEmbedding3d: only takes shape (vid_shape)
return rope_module.get_freqs(vid_shape.cpu())
else:
# Standard RoPE: takes vid_shape and txt_shape
return rope_module.get_freqs(vid_shape.cpu(), txt_shape.cpu())
except Exception as e:
print(f" ⚠️ Failed for {cache_key}: {e}")
# Return empty tensors as fallback
time.sleep(1)
clear_vram_cache()
return torch.zeros(1, 64)
# Store in cache
temp_cache(cache_key, compute_freqs)
except Exception as e:
print(f" ❌ Error in module {name}: {e}")
finally:
# Restore to original device
rope_module.to(original_device)
# Copy temporary cache to runner cache
if hasattr(runner, 'cache'):
runner.cache.cache.update(temp_cache.cache)
else:
runner.cache = temp_cache
except Exception as e:
print(f" ⚠️ Error during RoPE pre-init: {e}")
print(" 🔄 Model will work but could have OOM at first launch")
def clear_rope_cache(runner) -> None:
"""
🧹 Clear RoPE cache to free VRAM
Args:
runner: The model runner containing the cache
"""
print("🧹 Cleaning RoPE cache...")
if hasattr(runner, 'cache') and hasattr(runner.cache, 'cache'):
# Count entries before cleanup
cache_size = len(runner.cache.cache)
# Free all tensors from cache
for key, value in runner.cache.cache.items():
if isinstance(value, (tuple, list)):
for item in value:
if hasattr(item, 'cpu'):
item.cpu()
del item
elif hasattr(value, 'cpu'):
value.cpu()
del value
# Clear the cache
runner.cache.cache.clear()
print(f" ✅ RoPE cache cleared ({cache_size} entries removed)")
if hasattr(runner, 'dit'):
cleared_lru_count = 0
for module in runner.dit.modules():
if isinstance(module, RotaryEmbeddingBase):
if hasattr(module.get_axial_freqs, 'cache_clear'):
module.get_axial_freqs.cache_clear()
cleared_lru_count += 1
if cleared_lru_count > 0:
print(f" ✅ Cleared {cleared_lru_count} LRU caches from RoPE modules.")
# Aggressive VRAM cleanup
# clear_vram_cache()
#torch.cuda.empty_cache()
#clear_vram_cache()
print("🎯 RoPE cache cleanup completed!")
def clear_rope_lru_caches(model) -> int:
"""Clear ALL LRU caches from RoPE modules"""
cleared_count = 0
for name, module in model.named_modules():
if hasattr(module, 'get_axial_freqs') and hasattr(module.get_axial_freqs, 'cache_clear'):
module.get_axial_freqs.cache_clear()
cleared_count += 1
return cleared_count
def fast_model_cleanup(model):
"""Fast model cleanup without logs"""
if model is None:
return
# Move to CPU
model.to("cpu")
# Clear parameters and buffers recursively
def clear_recursive(m):
for child in m.children():
clear_recursive(child)
for param in m.parameters():
if param is not None:
param.data = param.data.cpu()
param.grad = None
for buffer in m.buffers():
if buffer is not None:
buffer.data = buffer.data.cpu()
clear_recursive(model)
def fast_ram_cleanup():
"""Fast RAM cleanup without excessive logging"""
# Garbage collection
gc.collect()
# Clear CUDA cache
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
# Clear PyTorch internal caches
try:
torch._C._clear_cache()
except:
pass
+160
View File
@@ -0,0 +1,160 @@
"""
Performance optimization module for SeedVR2
Contains optimized tensor operations and video processing functions
Extracted from: seedvr2.py (lines 1633-1730)
"""
import torch
from typing import List, Union
def optimized_video_rearrange(video_tensors: List[torch.Tensor]) -> List[torch.Tensor]:
"""
🚀 OPTIMIZED version of video rearrangement
Replaces slow loops with vectorized operations
Transforms:
- 3D: c h w -> t c h w (with t=1)
- 4D: c t h w -> t c h w
Expected gains: 5-10x faster than naive loops
Args:
video_tensors: List of video tensors to rearrange
Returns:
List of rearranged tensors in t c h w format
"""
if not video_tensors:
return []
# 🔍 Analyze dimensions to optimize processing
videos_3d = []
videos_4d = []
indices_3d = []
indices_4d = []
for i, video in enumerate(video_tensors):
if video.ndim == 3:
videos_3d.append(video)
indices_3d.append(i)
else: # ndim == 4
videos_4d.append(video)
indices_4d.append(i)
# 🎯 Prepare final result
samples = [None] * len(video_tensors)
# 🚀 BATCH PROCESSING for 3D videos (c h w -> 1 c h w)
if videos_3d:
# Method 1: Stack + permute (faster than rearrange)
# c h w -> c 1 h w -> 1 c h w
batch_3d = torch.stack([v.unsqueeze(1) for v in videos_3d]) # [batch, c, 1, h, w]
batch_3d = batch_3d.permute(0, 2, 1, 3, 4) # [batch, 1, c, h, w]
for i, idx in enumerate(indices_3d):
samples[idx] = batch_3d[i] # [1, c, h, w]
# 🚀 BATCH PROCESSING for 4D videos (c t h w -> t c h w)
if videos_4d:
# Check if all 4D videos have the same shape for maximum optimization
shapes = [v.shape for v in videos_4d]
if len(set(shapes)) == 1:
# 🎯 MAXIMUM OPTIMIZATION: All shapes identical
# Stack + permute in single operation
batch_4d = torch.stack(videos_4d) # [batch, c, t, h, w]
batch_4d = batch_4d.permute(0, 2, 1, 3, 4) # [batch, t, c, h, w]
for i, idx in enumerate(indices_4d):
samples[idx] = batch_4d[i] # [t, c, h, w]
else:
# 🔄 FALLBACK: Different shapes, optimized individual processing
for i, idx in enumerate(indices_4d):
# Use permute instead of rearrange (faster)
samples[idx] = videos_4d[i].permute(1, 0, 2, 3) # c t h w -> t c h w
return samples
def optimized_single_video_rearrange(video: torch.Tensor) -> torch.Tensor:
"""
🚀 OPTIMIZED version for single video tensor
Replaces rearrange() with native PyTorch operations
Transforms:
- 3D: c h w -> 1 c h w (add temporal dimension)
- 4D: c t h w -> t c h w (permute dimensions)
Expected gains: 2-5x faster than rearrange()
Args:
video: Input video tensor
Returns:
Rearranged tensor with temporal dimension first
"""
if video.ndim == 3:
# c h w -> 1 c h w (add temporal dimension t=1)
return video.unsqueeze(0)
else: # ndim == 4
# c t h w -> t c h w (permute channels and temporal)
return video.permute(1, 0, 2, 3)
def optimized_sample_to_image_format(sample: torch.Tensor) -> torch.Tensor:
"""
🚀 OPTIMIZED version to convert sample to image format
Replaces rearrange() with native PyTorch operations
Transforms:
- 3D: c h w -> 1 h w c (add temporal dimension + permute to image format)
- 4D: t c h w -> t h w c (permute to image format)
Expected gains: 2-5x faster than rearrange()
Args:
sample: Input sample tensor
Returns:
Tensor in image format (channels last)
"""
if sample.ndim == 3:
# c h w -> 1 h w c (add temporal dimension then permute)
return sample.unsqueeze(0).permute(0, 2, 3, 1)
else: # ndim == 4
# t c h w -> t h w c (permute channels to last)
return sample.permute(0, 2, 3, 1)
def temporal_latent_blending(latents1: torch.Tensor, latents2: torch.Tensor, blend_frames: int) -> torch.Tensor:
"""
🎨 Temporal blending in latent space to avoid discontinuities
Args:
latents1: Latents from previous batch (end frames)
latents2: Latents from current batch (start frames)
blend_frames: Number of frames to blend
Returns:
Blended latents for smooth transition
"""
if latents1.shape[0] != latents2.shape[0]:
# Adjust dimensions if necessary
min_frames = min(latents1.shape[0], latents2.shape[0])
latents1 = latents1[:min_frames]
latents2 = latents2[:min_frames]
# Create linear blending weights
# Frame 0: 100% latents1, 0% latents2
# Frame n: 0% latents1, 100% latents2
weights1 = torch.linspace(1.0, 0.0, blend_frames).view(-1, 1, 1, 1).to(latents1.device)
weights2 = torch.linspace(0.0, 1.0, blend_frames).view(-1, 1, 1, 1).to(latents2.device)
# Apply blending
blended_latents = weights1 * latents1 + weights2 * latents2
return blended_latents
+10
View File
@@ -0,0 +1,10 @@
"""
Utilities package for SeedVR2
Contains general utility functions like downloads, config loading, etc.
"""
from .downloads import download_weight
__all__ = [
"download_weight",
]
@@ -2,7 +2,7 @@ import torch
from PIL import Image
from torch import Tensor
from torch.nn import functional as F
from ...common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation
from common.half_precision_fixes import safe_pad_operation, safe_interpolate_operation
from torchvision.transforms import ToTensor, ToPILImage
def adain_color_fix(target: Image, source: Image):
@@ -99,6 +99,8 @@ def wavelet_decomposition(image: Tensor, levels=5):
return high_freq, low_freq
def wavelet_reconstruction(content_feat:Tensor, style_feat:Tensor):
"""
Apply wavelet decomposition, so that the content will have the same color as the style.
+70
View File
@@ -0,0 +1,70 @@
"""
Downloads utility module for SeedVR2
Handles model and VAE downloads from HuggingFace Hub
Extracted from: seedvr2.py (line 968-1015)
"""
import os
from huggingface_hub import hf_hub_download
import folder_paths
# Configuration des chemins
base_cache_dir = os.path.join(folder_paths.models_dir, "SEEDVR2")
# S'assurer que le dossier de cache existe
folder_paths.add_model_folder_path("seedvr2", os.path.join(folder_paths.models_dir, "SEEDVR2"))
def download_weight(model):
"""
Télécharge un modèle SeedVR2 et son VAE associé depuis HuggingFace Hub
Args:
model (str): Nom du fichier modèle à télécharger
(ex: "seedvr2_ema_3b_fp16.safetensors")
Gère automatiquement:
- Téléchargement du modèle principal
- Téléchargement du VAE avec fallbacks:
1. ema_vae_fp16.safetensors (priorité)
2. ema_vae_fp8_e4m3fn.safetensors (fallback)
3. ema_vae.pth (legacy fallback)
"""
model_path = os.path.join(base_cache_dir, model)
vae_fp16_path = os.path.join(base_cache_dir, "ema_vae_fp16.safetensors")
# Configuration HuggingFace
repo_id = "numz/SeedVR2_comfyUI"
# 🚀 Téléchargement du modèle principal
if not os.path.exists(model_path):
print(f"📥 Downloading model: {model}")
hf_hub_download(repo_id=repo_id, filename=model, local_dir=base_cache_dir)
print(f"✅ Downloaded: {model}")
# 🚀 Téléchargement du VAE avec stratégie de fallback
if not os.path.exists(vae_fp16_path):
print("📥 Downloading FP16 VAE SafeTensors...")
try:
hf_hub_download(
repo_id=repo_id,
filename="ema_vae_fp16.safetensors",
local_dir=base_cache_dir
)
print("✅ Downloaded: ema_vae_fp16.safetensors (FP16 SafeTensors)")
except Exception as e:
print(f"⚠️ FP16 SafeTensors VAE not available: {e}")
return
def get_base_cache_dir():
"""
Retourne le répertoire de cache base pour les modèles SeedVR2
Returns:
str: Chemin du répertoire de cache
"""
return base_cache_dir