Compare commits
74
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4058e92fb4 | ||
|
|
de6f92b896 | ||
|
|
aaeee0b3ac | ||
|
|
ee2c06890b | ||
|
|
7186408be9 | ||
|
|
878524bc28 | ||
|
|
44da53f2a9 | ||
|
|
a59b7b04d9 | ||
|
|
c8f4a6378f | ||
|
|
a486a15bf4 | ||
|
|
207ce0e9b5 | ||
|
|
bc818a0928 | ||
|
|
8295fabda1 | ||
|
|
c81eabd50f | ||
|
|
81b404b17d | ||
|
|
693a0f361d | ||
|
|
c963128781 | ||
|
|
d7143bb00d | ||
|
|
a652992506 | ||
|
|
8955d15f4e | ||
|
|
411302db13 | ||
|
|
a3e53f0e97 | ||
|
|
7420c7d163 | ||
|
|
ca958fa4cc | ||
|
|
65ed588570 | ||
|
|
14adfe2edc | ||
|
|
1c99f157ec | ||
|
|
3e3e0f6ccb | ||
|
|
493d8c659a | ||
|
|
b1f7acc4b5 | ||
|
|
3d6a69b154 | ||
|
|
6198c6a640 | ||
|
|
a3a3eea735 | ||
|
|
fb9e18ed20 | ||
|
|
e6b71b531b | ||
|
|
2bf1f5510b | ||
|
|
5a9f2d2f05 | ||
|
|
b5bdf51697 | ||
|
|
5bad51acf2 | ||
|
|
6cdd5dca30 | ||
|
|
0c0fb63c56 | ||
|
|
91ef7567ac | ||
|
|
09f2c64bbe | ||
|
|
fc5e823875 | ||
|
|
0c910981e5 | ||
|
|
5d809f19d1 | ||
|
|
8da919967f | ||
|
|
21ffa0d96a | ||
|
|
d6a1e875f0 | ||
|
|
c025ea80e5 | ||
|
|
98653a2261 | ||
|
|
9589938583 | ||
|
|
13fa3d83c7 | ||
|
|
db154cc313 | ||
|
|
8fb2708ff0 | ||
|
|
36dab19ff9 | ||
|
|
3c1179b4c2 | ||
|
|
8b60ac2964 | ||
|
|
0b3de5e5eb | ||
|
|
c627dc6421 | ||
|
|
117d23fcc0 | ||
|
|
4932793d00 | ||
|
|
152e6a6c77 | ||
|
|
3fd5531f4f | ||
|
|
26ab6c84c3 | ||
|
|
65c0733a48 | ||
|
|
4f40eef184 | ||
|
|
1d5bf58b34 | ||
|
|
4ad1281377 | ||
|
|
ee5da0838e | ||
|
|
4006b3a5ee | ||
|
|
6c731260f0 | ||
|
|
1de2eb2afd | ||
|
|
4b9782e3f4 |
@@ -59,5 +59,8 @@ docs/source/inference/examples/
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
|
||||
# Local scripts (keep local but don't track in git)
|
||||
local_scripts/
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
@@ -8,7 +8,7 @@ It features a clean, consistent API that works across popular video models, maki
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
|
||||
@@ -279,11 +279,11 @@ def main(args):
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--batch_size', type=int, default=2, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=32, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
|
||||
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -8,12 +8,14 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
|
||||
```
|
||||
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
@@ -117,4 +117,4 @@ If you're planning to contribute to FastVideo please see the following page:
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
|
||||
@@ -12,7 +12,7 @@ This guide explains how to implement a custom diffusion pipeline in FastVideo, l
|
||||
4. **Register Your Pipeline** - Make it discoverable by the framework
|
||||
5. **Configure Your Pipeline** - (Coming soon)
|
||||
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
## Step 1: Pipeline Modules
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ fastvideo generate --help
|
||||
### Hardware Configuration
|
||||
|
||||
- `--num-gpus {NUM_GPUS}`: Number of GPUs to use
|
||||
- `--tp-size {TP_SIZE}`: Tensor parallelism size (Typically should match the number of GPUs)
|
||||
- `--tp-size {TP_SIZE}`: Tensor parallelism size (only for the encoder, should not be larger than 1 if text encoder offload is enabled, as layerwise offload + prefetch is faster)
|
||||
- `--sp-size {SP_SIZE}`: Sequence parallelism size (Typically should match the number of GPUs)
|
||||
|
||||
#### Video Configuration
|
||||
@@ -68,7 +68,7 @@ Example configuration file (config.json):
|
||||
"output_path": "outputs/",
|
||||
"num_gpus": 2,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"num_frames": 45,
|
||||
"height": 720,
|
||||
"width": 1280,
|
||||
@@ -102,7 +102,7 @@ prompt: "A beautiful woman in a red dress walking down a street"
|
||||
output_path: "outputs/"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
tp_size: 2
|
||||
tp_size: 1
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
|
||||
@@ -121,4 +121,4 @@ If the generated video doesn't match your prompt:
|
||||
- Learn about using [Optimizations](#inference-optimizations)
|
||||
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
|
||||
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
@@ -30,7 +30,7 @@ training_args=(
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 8
|
||||
--tp_size 8
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
@@ -66,7 +66,7 @@ training_args=(
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
@@ -26,6 +26,24 @@
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 98
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-013.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 148
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -30,7 +30,7 @@ training_args=(
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
@@ -63,7 +63,7 @@ training_args=(
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 4
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
)
|
||||
|
||||
@@ -6,7 +6,8 @@ from typing import Any
|
||||
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
|
||||
from fastvideo.v1.configs.sample.wan import (Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam)
|
||||
@@ -23,6 +24,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -94,3 +94,20 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
|
||||
-13.02252404
|
||||
]))
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
@@ -23,7 +23,6 @@ If you only need to use the distributed environment without model parallelism,
|
||||
you can skip the model parallel initialization and destruction steps.
|
||||
"""
|
||||
import contextlib
|
||||
import gc
|
||||
import os
|
||||
import pickle
|
||||
import weakref
|
||||
@@ -1016,15 +1015,6 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
if shutdown_ray:
|
||||
import ray # Lazy import Ray
|
||||
ray.shutdown()
|
||||
gc.collect()
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
if not current_platform.is_cpu():
|
||||
torch.cuda.empty_cache()
|
||||
try:
|
||||
torch._C._host_emptyCache()
|
||||
except AttributeError:
|
||||
logger.warning(
|
||||
"torch._C._host_emptyCache() only available in Pytorch >=2.5")
|
||||
|
||||
|
||||
def in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
|
||||
|
||||
@@ -6,7 +6,6 @@ This module provides a consolidated interface for generating videos using
|
||||
diffusion models.
|
||||
"""
|
||||
|
||||
import gc
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
@@ -277,5 +276,3 @@ class VideoGenerator:
|
||||
"""
|
||||
self.executor.shutdown()
|
||||
del self.executor
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -6,7 +6,7 @@ import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import field
|
||||
from typing import Any
|
||||
from typing import Any, List
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -78,6 +78,8 @@ class FastVideoArgs:
|
||||
|
||||
# Stage verification
|
||||
enable_stage_verification: bool = True
|
||||
|
||||
denoising_step_list: List[int] | None = field(default=None)
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
@@ -254,6 +256,12 @@ class FastVideoArgs:
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
|
||||
parser.add_argument("--denoising-step-list",
|
||||
type=parse_int_list,
|
||||
default=FastVideoArgs.denoising_step_list,
|
||||
help="Comma-separated list of denoising steps (e.g., '1000,757,522')",
|
||||
)
|
||||
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -292,7 +300,7 @@ class FastVideoArgs:
|
||||
assert self.sp_size != -1, "sp_size must be set for training"
|
||||
|
||||
if self.tp_size == -1:
|
||||
self.tp_size = self.num_gpus
|
||||
self.tp_size = 1
|
||||
if self.sp_size == -1:
|
||||
self.sp_size = self.num_gpus
|
||||
if self.hsdp_shard_dim == -1:
|
||||
@@ -305,11 +313,6 @@ class FastVideoArgs:
|
||||
if self.num_gpus < max(self.tp_size, self.sp_size):
|
||||
self.num_gpus = max(self.tp_size, self.sp_size)
|
||||
|
||||
if self.tp_size != self.sp_size:
|
||||
raise ValueError(
|
||||
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
|
||||
)
|
||||
|
||||
if self.enable_torch_compile and self.num_gpus > 1:
|
||||
logger.warning(
|
||||
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
|
||||
@@ -425,6 +428,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
learning_rate: float = 0.0
|
||||
scale_lr: bool = False
|
||||
lr_scheduler: str = "constant"
|
||||
lr_step_rules: str | None = None
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
enable_gradient_checkpointing_type: str | None = None
|
||||
@@ -460,6 +464,17 @@ class TrainingArgs(FastVideoArgs):
|
||||
# VSA training decay parameters
|
||||
VSA_decay_rate: float = 0.01 # decay rate -> 0.02
|
||||
VSA_decay_interval_steps: int = 1 # decay interval steps -> 50
|
||||
|
||||
# distillation args
|
||||
student_critic_update_ratio: int = 5
|
||||
critic_learning_rate: float = 1e-5
|
||||
critic_lr_scheduler: str = "constant"
|
||||
critic_lr_step_rules: str | None = None
|
||||
min_step_ratio: float = 0.2
|
||||
max_step_ratio: float = 0.98
|
||||
teacher_guidance_scale: float = 3.5
|
||||
simulate_student_forward: bool = False
|
||||
num_teacher_noisy_ground_truth_steps: int = 0
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -619,6 +634,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
default="constant",
|
||||
help="Learning rate scheduler type")
|
||||
parser.add_argument("--lr-step-rules",
|
||||
type=str,
|
||||
help="Learning rate step rules")
|
||||
parser.add_argument("--lr-warmup-steps",
|
||||
type=int,
|
||||
default=10,
|
||||
@@ -726,5 +744,48 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=int,
|
||||
default=TrainingArgs.VSA_decay_interval_steps,
|
||||
help="VSA decay interval steps")
|
||||
|
||||
# Distillation arguments
|
||||
parser.add_argument("--student-critic-update-ratio",
|
||||
type=int,
|
||||
default=TrainingArgs.student_critic_update_ratio,
|
||||
help="Ratio of student updates to critic updates.")
|
||||
parser.add_argument("--critic-learning-rate",
|
||||
type=float,
|
||||
default=TrainingArgs.critic_learning_rate,
|
||||
help="Learning rate for critic")
|
||||
parser.add_argument("--critic-lr-scheduler",
|
||||
type=str,
|
||||
default=TrainingArgs.critic_lr_scheduler,
|
||||
help="Learning rate scheduler type for critic")
|
||||
parser.add_argument("--critic-lr-step-rules",
|
||||
type=str,
|
||||
help="Learning rate step rules for critic")
|
||||
parser.add_argument("--min-step-ratio",
|
||||
type=float,
|
||||
default=TrainingArgs.min_step_ratio,
|
||||
help="Minimum step ratio")
|
||||
parser.add_argument("--max-step-ratio",
|
||||
type=float,
|
||||
default=TrainingArgs.max_step_ratio,
|
||||
help="Maximum step ratio")
|
||||
parser.add_argument("--teacher-guidance-scale",
|
||||
type=float,
|
||||
default=TrainingArgs.teacher_guidance_scale,
|
||||
help="Teacher guidance scale")
|
||||
parser.add_argument("--simulate-student-forward",
|
||||
action=StoreBoolean,
|
||||
default=TrainingArgs.simulate_student_forward,
|
||||
help="Whether to simulate student forward")
|
||||
parser.add_argument("--num-teacher-noisy-ground-truth-steps",
|
||||
type=int,
|
||||
default=TrainingArgs.num_teacher_noisy_ground_truth_steps,
|
||||
help="Number of steps to use noisy ground truth for teacher")
|
||||
|
||||
return parser
|
||||
|
||||
def parse_int_list(value: str) -> List[int]:
|
||||
"""Parse a comma-separated string of integers into a list."""
|
||||
if not value:
|
||||
return []
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
@@ -72,6 +72,8 @@ class ComponentLoader(ABC):
|
||||
module_loaders = {
|
||||
"scheduler": (SchedulerLoader, "diffusers"),
|
||||
"transformer": (TransformerLoader, "diffusers"),
|
||||
"teacher_transformer": (TransformerLoader, "diffusers"),
|
||||
"critic_transformer": (TransformerLoader, "diffusers"),
|
||||
"vae": (VAELoader, "diffusers"),
|
||||
"text_encoder": (TextEncoderLoader, "transformers"),
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
|
||||
@@ -18,15 +18,14 @@
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from typing import Any, Optional, Tuple, Union, List
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
from diffusers.utils import BaseOutput, is_scipy_available, logging
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.schedulers.base import BaseScheduler
|
||||
|
||||
@@ -34,7 +33,7 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
@@ -46,8 +45,7 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
@@ -57,16 +55,37 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
shift (`float`, defaults to 1.0):
|
||||
The shift value for the timestep schedule.
|
||||
reverse (`bool`, defaults to `True`):
|
||||
Whether to reverse the timestep schedule.
|
||||
use_dynamic_shifting (`bool`, defaults to False):
|
||||
Whether to apply timestep shifting on-the-fly based on the image resolution.
|
||||
base_shift (`float`, defaults to 0.5):
|
||||
Value to stabilize image generation. Increasing `base_shift` reduces variation and image is more consistent
|
||||
with desired output.
|
||||
max_shift (`float`, defaults to 1.15):
|
||||
Value change allowed to latent vectors. Increasing `max_shift` encourages more variation and image may be
|
||||
more exaggerated or stylized.
|
||||
base_image_seq_len (`int`, defaults to 256):
|
||||
The base image sequence length.
|
||||
max_image_seq_len (`int`, defaults to 4096):
|
||||
The maximum image sequence length.
|
||||
invert_sigmas (`bool`, defaults to False):
|
||||
Whether to invert the sigmas.
|
||||
shift_terminal (`float`, defaults to None):
|
||||
The end value of the shifted timestep schedule.
|
||||
use_karras_sigmas (`bool`, defaults to False):
|
||||
Whether to use Karras sigmas for step sizes in the noise schedule during sampling.
|
||||
use_exponential_sigmas (`bool`, defaults to False):
|
||||
Whether to use exponential sigmas for step sizes in the noise schedule during sampling.
|
||||
use_beta_sigmas (`bool`, defaults to False):
|
||||
Whether to use beta sigmas for step sizes in the noise schedule during sampling.
|
||||
time_shift_type (`str`, defaults to "exponential"):
|
||||
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
|
||||
stochastic_sampling (`bool`, defaults to False):
|
||||
Whether to use stochastic sampling.
|
||||
"""
|
||||
|
||||
_compatibles: list[Any] = []
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
@@ -74,31 +93,53 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
shift: float = 1.0,
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: int | None = None,
|
||||
**kwargs,
|
||||
use_dynamic_shifting: bool = False,
|
||||
base_shift: Optional[float] = 0.5,
|
||||
max_shift: Optional[float] = 1.15,
|
||||
base_image_seq_len: Optional[int] = 256,
|
||||
max_image_seq_len: Optional[int] = 4096,
|
||||
invert_sigmas: bool = False,
|
||||
shift_terminal: Optional[float] = None,
|
||||
use_karras_sigmas: Optional[bool] = False,
|
||||
use_exponential_sigmas: Optional[bool] = False,
|
||||
use_beta_sigmas: Optional[bool] = False,
|
||||
time_shift_type: str = "exponential",
|
||||
stochastic_sampling: bool = False,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
if not reverse:
|
||||
sigmas = sigmas.flip(0)
|
||||
|
||||
self.sigmas = sigmas
|
||||
# the value fed to model
|
||||
self.timesteps = (sigmas[:-1] *
|
||||
num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
self._step_index: int | None = None
|
||||
self._begin_index = 0
|
||||
|
||||
self.supported_solver = ["euler"]
|
||||
if solver not in self.supported_solver:
|
||||
if self.config.use_beta_sigmas and not is_scipy_available():
|
||||
raise ImportError("Make sure to install scipy if you want to use beta sigmas.")
|
||||
if sum([self.config.use_beta_sigmas, self.config.use_exponential_sigmas, self.config.use_karras_sigmas]) > 1:
|
||||
raise ValueError(
|
||||
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
|
||||
)
|
||||
if time_shift_type not in {"exponential", "linear"}:
|
||||
raise ValueError("`time_shift_type` must either be 'exponential' or 'linear'.")
|
||||
|
||||
BaseScheduler.__init__(self)
|
||||
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
if not use_dynamic_shifting:
|
||||
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
|
||||
self.timesteps = sigmas * num_train_timesteps
|
||||
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self._shift = shift
|
||||
|
||||
self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@property
|
||||
def shift(self):
|
||||
"""
|
||||
The value used for shifting.
|
||||
"""
|
||||
return self._shift
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
@@ -125,44 +166,190 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def set_shift(self, shift: float):
|
||||
self._shift = shift
|
||||
|
||||
def scale_noise(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
noise: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Forward process in flow-matching
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
|
||||
|
||||
if sample.device.type == "mps" and torch.is_floating_point(timestep):
|
||||
# mps does not support float64
|
||||
schedule_timesteps = self.timesteps.to(sample.device, dtype=torch.float32)
|
||||
timestep = timestep.to(sample.device, dtype=torch.float32)
|
||||
else:
|
||||
schedule_timesteps = self.timesteps.to(sample.device)
|
||||
timestep = timestep.to(sample.device)
|
||||
|
||||
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
|
||||
if self.begin_index is None:
|
||||
step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep]
|
||||
elif self.step_index is not None:
|
||||
# add_noise is called after first denoising step (for inpainting)
|
||||
step_indices = [self.step_index] * timestep.shape[0]
|
||||
else:
|
||||
# add noise is called before first denoising step to create initial latent(img2img)
|
||||
step_indices = [self.begin_index] * timestep.shape[0]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < len(sample.shape):
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
sample = sigma * noise + (1.0 - sigma) * sample
|
||||
|
||||
return sample
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
|
||||
if self.config.time_shift_type == "exponential":
|
||||
return self._time_shift_exponential(mu, sigma, t)
|
||||
elif self.config.time_shift_type == "linear":
|
||||
return self._time_shift_linear(mu, sigma, t)
|
||||
|
||||
def stretch_shift_to_terminal(self, t: torch.Tensor) -> torch.Tensor:
|
||||
r"""
|
||||
Stretches and shifts the timestep schedule to ensure it terminates at the configured `shift_terminal` config
|
||||
value.
|
||||
|
||||
Reference:
|
||||
https://github.com/Lightricks/LTX-Video/blob/a01a171f8fe3d99dce2728d60a73fecf4d4238ae/ltx_video/schedulers/rf.py#L51
|
||||
|
||||
Args:
|
||||
t (`torch.Tensor`):
|
||||
A tensor of timesteps to be stretched and shifted.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A tensor of adjusted timesteps such that the final value equals `self.config.shift_terminal`.
|
||||
"""
|
||||
one_minus_z = 1 - t
|
||||
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
|
||||
stretched_t = 1 - (one_minus_z / scale_factor)
|
||||
return stretched_t
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
device: str | torch.device = None,
|
||||
n_tokens: int = 0,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
mu: Optional[float] = None,
|
||||
timesteps: Optional[List[float]] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
num_inference_steps (`int`, *optional*):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom values for sigmas to be used for each diffusion step. If `None`, the sigmas are computed
|
||||
automatically.
|
||||
mu (`float`, *optional*):
|
||||
Determines the amount of shifting applied to sigmas when performing resolution-dependent timestep
|
||||
shifting.
|
||||
timesteps (`List[float]`, *optional*):
|
||||
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
|
||||
automatically.
|
||||
"""
|
||||
if self.config.use_dynamic_shifting and mu is None:
|
||||
raise ValueError("`mu` must be passed when `use_dynamic_shifting` is set to be `True`")
|
||||
|
||||
if sigmas is not None and timesteps is not None:
|
||||
if len(sigmas) != len(timesteps):
|
||||
raise ValueError("`sigmas` and `timesteps` should have the same length")
|
||||
|
||||
if num_inference_steps is not None:
|
||||
if (sigmas is not None and len(sigmas) != num_inference_steps) or (
|
||||
timesteps is not None and len(timesteps) != num_inference_steps
|
||||
):
|
||||
raise ValueError(
|
||||
"`sigmas` and `timesteps` should have the same length as num_inference_steps, if `num_inference_steps` is provided"
|
||||
)
|
||||
else:
|
||||
num_inference_steps = len(sigmas) if sigmas is not None else len(timesteps)
|
||||
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
|
||||
sigmas = self.sd3_time_shift(sigmas)
|
||||
# 1. Prepare default sigmas
|
||||
is_timesteps_provided = timesteps is not None
|
||||
|
||||
if not self.config.reverse:
|
||||
sigmas = 1 - sigmas
|
||||
if is_timesteps_provided:
|
||||
timesteps = np.array(timesteps).astype(np.float32)
|
||||
|
||||
self.sigmas = sigmas
|
||||
if not getattr(self.config, "timesteps_scale", True):
|
||||
self.timesteps = sigmas[:-1] # for stepvideo
|
||||
if sigmas is None:
|
||||
if timesteps is None:
|
||||
timesteps = np.linspace(
|
||||
self._sigma_to_t(self.sigma_max), self._sigma_to_t(self.sigma_min), num_inference_steps
|
||||
)
|
||||
sigmas = timesteps / self.config.num_train_timesteps
|
||||
else:
|
||||
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
|
||||
dtype=torch.float32, device=device)
|
||||
# Reset step index
|
||||
self._step_index = None
|
||||
sigmas = np.array(sigmas).astype(np.float32)
|
||||
num_inference_steps = len(sigmas)
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
|
||||
# 2. Perform timestep shifting. Either no shifting is applied, or resolution-dependent shifting of
|
||||
# "exponential" or "linear" type is applied
|
||||
if self.config.use_dynamic_shifting:
|
||||
sigmas = self.time_shift(mu, 1.0, sigmas)
|
||||
else:
|
||||
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
|
||||
|
||||
# 3. If required, stretch the sigmas schedule to terminate at the configured `shift_terminal` value
|
||||
if self.config.shift_terminal:
|
||||
sigmas = self.stretch_shift_to_terminal(sigmas)
|
||||
|
||||
# 4. If required, convert sigmas to one of karras, exponential, or beta sigma schedules
|
||||
if self.config.use_karras_sigmas:
|
||||
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
elif self.config.use_exponential_sigmas:
|
||||
sigmas = self._convert_to_exponential(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
elif self.config.use_beta_sigmas:
|
||||
sigmas = self._convert_to_beta(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
|
||||
# 5. Convert sigmas and timesteps to tensors and move to specified device
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
|
||||
if not is_timesteps_provided:
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
else:
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32, device=device)
|
||||
|
||||
# 6. Append the terminal sigma value.
|
||||
# If a model requires inverted sigma schedule for denoising but timesteps without inversion, the
|
||||
# `invert_sigmas` flag can be set to `True`. This case is only required in Mochi
|
||||
if self.config.invert_sigmas:
|
||||
sigmas = 1.0 - sigmas
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
|
||||
else:
|
||||
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||
|
||||
self.timesteps = timesteps
|
||||
self.sigmas = sigmas
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
@@ -174,17 +361,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
idx: int = indices[pos].item()
|
||||
return indices[pos].item()
|
||||
|
||||
return idx
|
||||
|
||||
def set_shift(self, shift: float) -> None:
|
||||
self.config.shift = shift
|
||||
|
||||
def set_timesteps_scale(self, timesteps_scale: bool) -> None:
|
||||
self.config.timesteps_scale = timesteps_scale
|
||||
|
||||
def _init_step_index(self, timestep) -> None:
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
@@ -192,22 +371,19 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: int | None = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def sd3_time_shift(self, t: torch.Tensor):
|
||||
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: float | torch.FloatTensor,
|
||||
sample: torch.FloatTensor,
|
||||
s_churn: float = 0.0,
|
||||
s_tmin: float = 0.0,
|
||||
s_tmax: float = float("inf"),
|
||||
s_noise: float = 1.0,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
per_token_timesteps: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
**kwargs,
|
||||
) -> FlowMatchDiscreteSchedulerOutput | tuple:
|
||||
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
@@ -219,25 +395,38 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
s_churn (`float`):
|
||||
s_tmin (`float`):
|
||||
s_tmax (`float`):
|
||||
s_noise (`float`, defaults to 1.0):
|
||||
Scaling factor for noise added to the sample.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
per_token_timesteps (`torch.Tensor`, *optional*):
|
||||
The timesteps for each token in the sample.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
tuple.
|
||||
Whether or not to return a
|
||||
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or tuple.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`,
|
||||
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] is returned,
|
||||
otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if isinstance(timestep, (int | torch.IntTensor | torch.LongTensor)):
|
||||
raise ValueError((
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."), )
|
||||
if (
|
||||
isinstance(timestep, int)
|
||||
or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)
|
||||
):
|
||||
raise ValueError(
|
||||
(
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `FlowMatchEulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."
|
||||
),
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
@@ -245,24 +434,454 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
assert self.step_index is not None
|
||||
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
|
||||
if per_token_timesteps is not None:
|
||||
per_token_sigmas = per_token_timesteps / self.config.num_train_timesteps
|
||||
|
||||
if self.config.solver == "euler":
|
||||
prev_sample = sample + model_output.to(torch.float32) * dt
|
||||
sigmas = self.sigmas[:, None, None]
|
||||
lower_mask = sigmas < per_token_sigmas[None] - 1e-6
|
||||
lower_sigmas = lower_mask * sigmas
|
||||
lower_sigmas, _ = lower_sigmas.max(dim=0)
|
||||
|
||||
current_sigma = per_token_sigmas[..., None]
|
||||
next_sigma = lower_sigmas[..., None]
|
||||
dt = current_sigma - next_sigma
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
sigma_idx = self.step_index
|
||||
sigma = self.sigmas[sigma_idx]
|
||||
sigma_next = self.sigmas[sigma_idx + 1]
|
||||
|
||||
current_sigma = sigma
|
||||
next_sigma = sigma_next
|
||||
dt = sigma_next - sigma
|
||||
|
||||
if self.config.stochastic_sampling:
|
||||
x0 = sample - current_sigma * model_output
|
||||
noise = torch.randn_like(sample)
|
||||
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
|
||||
else:
|
||||
prev_sample = sample + dt * model_output
|
||||
|
||||
# upon completion increase step index by one
|
||||
assert self._step_index is not None
|
||||
self._step_index += 1
|
||||
if per_token_timesteps is None:
|
||||
# Cast sample back to model compatible dtype
|
||||
prev_sample = prev_sample.to(model_output.dtype)
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, )
|
||||
return (prev_sample,)
|
||||
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
|
||||
def _convert_to_karras(self, in_sigmas: torch.Tensor, num_inference_steps) -> torch.Tensor:
|
||||
"""Constructs the noise schedule of Karras et al. (2022)."""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
rho = 7.0 # 7.0 is the value used in the paper
|
||||
ramp = np.linspace(0, 1, num_inference_steps)
|
||||
min_inv_rho = sigma_min ** (1 / rho)
|
||||
max_inv_rho = sigma_max ** (1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||
return sigmas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
|
||||
def _convert_to_exponential(self, in_sigmas: torch.Tensor, num_inference_steps: int) -> torch.Tensor:
|
||||
"""Constructs an exponential noise schedule."""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
sigmas = np.exp(np.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps))
|
||||
return sigmas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
|
||||
def _convert_to_beta(
|
||||
self, in_sigmas: torch.Tensor, num_inference_steps: int, alpha: float = 0.6, beta: float = 0.6
|
||||
) -> torch.Tensor:
|
||||
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
sigmas = np.array(
|
||||
[
|
||||
sigma_min + (ppf * (sigma_max - sigma_min))
|
||||
for ppf in [
|
||||
scipy.stats.beta.ppf(timestep, alpha, beta)
|
||||
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
|
||||
]
|
||||
]
|
||||
)
|
||||
return sigmas
|
||||
|
||||
def _time_shift_exponential(self, mu, sigma, t):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||
|
||||
def _time_shift_linear(self, mu, sigma, t):
|
||||
return mu / (mu + (1 / t - 1) ** sigma)
|
||||
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B, C, H, W]
|
||||
- noise: the noise with shape [B, C, H, W]
|
||||
- timestep: the timestep with shape [B]
|
||||
Output: the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timesteps.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
|
||||
class FlowMatchScheduler():
|
||||
order = 1
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.inverse_timesteps = inverse_timesteps
|
||||
self.extra_one_step = extra_one_step
|
||||
self.reverse_sigmas = reverse_sigmas
|
||||
self.set_timesteps(num_inference_steps)
|
||||
|
||||
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False):
|
||||
sigma_start = self.sigma_min + \
|
||||
(self.sigma_max - self.sigma_min) * denoising_strength
|
||||
if self.extra_one_step:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
|
||||
else:
|
||||
self.sigmas = torch.linspace(
|
||||
sigma_start, self.sigma_min, num_inference_steps)
|
||||
if self.inverse_timesteps:
|
||||
self.sigmas = torch.flip(self.sigmas, dims=[0])
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
if self.reverse_sigmas:
|
||||
self.sigmas = 1 - self.sigmas
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
if training:
|
||||
x = self.timesteps
|
||||
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
|
||||
num_inference_steps) ** 2)
|
||||
y_shifted = y - y.min()
|
||||
bsmntw_weighing = y_shifted * \
|
||||
(num_inference_steps / y_shifted.sum())
|
||||
self.linear_timesteps_weights = bsmntw_weighing
|
||||
|
||||
def step(self, model_output, timestep, sample, to_final=False):
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
|
||||
sigma_ = 1 if (
|
||||
self.inverse_timesteps or self.reverse_sigmas) else 0
|
||||
else:
|
||||
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
prev_sample = sample + model_output * (sigma_ - sigma)
|
||||
return prev_sample
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B, C, H, W]
|
||||
- noise: the noise with shape [B, C, H, W]
|
||||
- timestep: the timestep with shape [B]
|
||||
Output: the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) # [21, 1, 1, 1]
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def training_target(self, sample, noise, timestep):
|
||||
target = noise - sample
|
||||
return target
|
||||
|
||||
def training_weight(self, timestep):
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps - timestep.to(self.timesteps.device)).abs())
|
||||
weights = self.linear_timesteps_weights[timestep_id]
|
||||
return weights
|
||||
|
||||
# class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
# """
|
||||
# Euler scheduler.
|
||||
|
||||
# This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
# methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
# Args:
|
||||
# num_train_timesteps (`int`, defaults to 1000):
|
||||
# The number of diffusion steps to train the model.
|
||||
# timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
# The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
# Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
# shift (`float`, defaults to 1.0):
|
||||
# The shift value for the timestep schedule.
|
||||
# reverse (`bool`, defaults to `True`):
|
||||
# Whether to reverse the timestep schedule.
|
||||
# """
|
||||
|
||||
# _compatibles: list[Any] = []
|
||||
# order = 1
|
||||
|
||||
# @register_to_config
|
||||
# def __init__(
|
||||
# self,
|
||||
# num_train_timesteps: int = 1000,
|
||||
# shift: float = 1.0,
|
||||
# reverse: bool = True,
|
||||
# solver: str = "euler",
|
||||
# n_tokens: Optional[int] = None,
|
||||
# **kwargs,
|
||||
# ):
|
||||
# sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
# if not reverse:
|
||||
# sigmas = sigmas.flip(0)
|
||||
|
||||
# self.sigmas = sigmas
|
||||
# # the value fed to model
|
||||
# self.timesteps = (sigmas[:-1] *
|
||||
# num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
# self._step_index: int | None = None
|
||||
# self._begin_index = 0
|
||||
|
||||
# self.supported_solver = ["euler"]
|
||||
# if solver not in self.supported_solver:
|
||||
# raise ValueError(
|
||||
# f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
# )
|
||||
|
||||
# BaseScheduler.__init__(self)
|
||||
|
||||
# @property
|
||||
# def step_index(self):
|
||||
# """
|
||||
# The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
# """
|
||||
# return self._step_index
|
||||
|
||||
# @property
|
||||
# def begin_index(self):
|
||||
# """
|
||||
# The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
# """
|
||||
# return self._begin_index
|
||||
|
||||
# # Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||
# def set_begin_index(self, begin_index: int = 0):
|
||||
# """
|
||||
# Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
# Args:
|
||||
# begin_index (`int`):
|
||||
# The begin index for the scheduler.
|
||||
# """
|
||||
# self._begin_index = begin_index
|
||||
|
||||
# def _sigma_to_t(self, sigma):
|
||||
# return sigma * self.config.num_train_timesteps
|
||||
|
||||
# def set_timesteps(
|
||||
# self,
|
||||
# num_inference_steps: int,
|
||||
# device: Union[str, torch.device] = None,
|
||||
# n_tokens: int = 0,
|
||||
# ):
|
||||
# """
|
||||
# Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
# Args:
|
||||
# num_inference_steps (`int`):
|
||||
# The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
# device (`str` or `torch.device`, *optional*):
|
||||
# The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
# n_tokens (`int`, *optional*):
|
||||
# Number of tokens in the input sequence.
|
||||
# """
|
||||
# self.num_inference_steps = num_inference_steps
|
||||
|
||||
# sigmas = torch.linspace(1, 0, num_inference_steps + 1)
|
||||
# sigmas = self.sd3_time_shift(sigmas)
|
||||
|
||||
# if not self.config.reverse:
|
||||
# sigmas = 1 - sigmas
|
||||
|
||||
# self.sigmas = sigmas
|
||||
# if not getattr(self.config, "timesteps_scale", True):
|
||||
# self.timesteps = sigmas[:-1] # for stepvideo
|
||||
# else:
|
||||
# self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
|
||||
# dtype=torch.float32, device=device)
|
||||
# # Reset step index
|
||||
# self._step_index = None
|
||||
|
||||
# def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
|
||||
# if schedule_timesteps is None:
|
||||
# schedule_timesteps = self.timesteps
|
||||
|
||||
# indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# # The sigma index that is taken for the **very** first `step`
|
||||
# # is always the second index (or the last index if there is only 1)
|
||||
# # This way we can ensure we don't accidentally skip a sigma in
|
||||
# # case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
# pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
# idx: int = indices[pos].item()
|
||||
|
||||
# return idx
|
||||
|
||||
# def set_shift(self, shift: float) -> None:
|
||||
# self.config.shift = shift
|
||||
|
||||
# def set_timesteps_scale(self, timesteps_scale: bool) -> None:
|
||||
# self.config.timesteps_scale = timesteps_scale
|
||||
|
||||
# def _init_step_index(self, timestep) -> None:
|
||||
# if self.begin_index is None:
|
||||
# if isinstance(timestep, torch.Tensor):
|
||||
# timestep = timestep.to(self.timesteps.device)
|
||||
# self._step_index = self.index_for_timestep(timestep)
|
||||
# else:
|
||||
# self._step_index = self._begin_index
|
||||
|
||||
# def scale_model_input(self,
|
||||
# sample: torch.Tensor,
|
||||
# timestep: Optional[int] = None) -> torch.Tensor:
|
||||
# return sample
|
||||
|
||||
# def sd3_time_shift(self, t: torch.Tensor):
|
||||
# return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
|
||||
|
||||
# def step(
|
||||
# self,
|
||||
# model_output: torch.FloatTensor,
|
||||
# timestep: Union[float, torch.FloatTensor],
|
||||
# sample: torch.FloatTensor,
|
||||
# return_dict: bool = True,
|
||||
# **kwargs,
|
||||
# ) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
|
||||
# """
|
||||
# Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
# process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
# Args:
|
||||
# model_output (`torch.FloatTensor`):
|
||||
# The direct output from learned diffusion model.
|
||||
# timestep (`float`):
|
||||
# The current discrete timestep in the diffusion chain.
|
||||
# sample (`torch.FloatTensor`):
|
||||
# A current instance of a sample created by the diffusion process.
|
||||
# generator (`torch.Generator`, *optional*):
|
||||
# A random number generator.
|
||||
# n_tokens (`int`, *optional*):
|
||||
# Number of tokens in the input sequence.
|
||||
# return_dict (`bool`):
|
||||
# Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
# tuple.
|
||||
|
||||
# Returns:
|
||||
# [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
# If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
# returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
# """
|
||||
|
||||
# if isinstance(timestep, (int, torch.IntTensor, torch.LongTensor)):
|
||||
# raise ValueError((
|
||||
# "Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
# " `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
# " one of the `scheduler.timesteps` as a timestep."), )
|
||||
|
||||
# if self.step_index is None:
|
||||
# self._init_step_index(timestep)
|
||||
|
||||
# # Upcast to avoid precision issues when computing prev_sample
|
||||
# sample = sample.to(torch.float32)
|
||||
|
||||
# assert self.step_index is not None
|
||||
# dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
|
||||
|
||||
# if self.config.solver == "euler":
|
||||
# prev_sample = sample + model_output.to(torch.float32) * dt
|
||||
# else:
|
||||
# raise ValueError(
|
||||
# f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
# )
|
||||
|
||||
# # upon completion increase step index by one
|
||||
# assert self._step_index is not None
|
||||
# self._step_index += 1
|
||||
|
||||
# if not return_dict:
|
||||
# return (prev_sample, )
|
||||
|
||||
# return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
# def __len__(self):
|
||||
# return self.config.num_train_timesteps
|
||||
|
||||
@@ -772,47 +772,5 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
"""
|
||||
return sample
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.add_noise
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||
sigmas = self.sigmas.to(device=original_samples.device,
|
||||
dtype=original_samples.dtype)
|
||||
if original_samples.device.type == "mps" and torch.is_floating_point(
|
||||
timesteps):
|
||||
# mps does not support float64
|
||||
schedule_timesteps = self.timesteps.to(original_samples.device,
|
||||
dtype=torch.float32)
|
||||
timesteps = timesteps.to(original_samples.device,
|
||||
dtype=torch.float32)
|
||||
else:
|
||||
schedule_timesteps = self.timesteps.to(original_samples.device)
|
||||
timesteps = timesteps.to(original_samples.device)
|
||||
|
||||
# begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index
|
||||
if self.begin_index is None:
|
||||
step_indices = [
|
||||
self.index_for_timestep(t, schedule_timesteps)
|
||||
for t in timesteps
|
||||
]
|
||||
elif self.step_index is not None:
|
||||
# add_noise is called after first denoising step (for inpainting)
|
||||
step_indices = [self.step_index] * timesteps.shape[0]
|
||||
else:
|
||||
# add noise is called before first denoising step to create initial latent(img2img)
|
||||
step_indices = [self.begin_index] * timesteps.shape[0]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < len(original_samples.shape):
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
|
||||
noisy_samples = alpha_t * original_samples + sigma_t * noise
|
||||
return noisy_samples
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
|
||||
@@ -239,7 +239,6 @@ class ParallelTiledVAE(ABC):
|
||||
|
||||
results = torch.cat(local_results, dim=0).contiguous()
|
||||
del local_results
|
||||
torch.cuda.empty_cache()
|
||||
# first gather size to pad the results
|
||||
local_size = torch.tensor([results.size(0)],
|
||||
device=results.device,
|
||||
@@ -253,7 +252,7 @@ class ParallelTiledVAE(ABC):
|
||||
padded_results = torch.zeros(max_size, device=results.device)
|
||||
padded_results[:results.size(0)] = results
|
||||
del results
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Gather all results
|
||||
gathered_dim_metadata = [None] * world_size
|
||||
gathered_results = torch.zeros_like(padded_results).repeat(
|
||||
|
||||
@@ -222,12 +222,13 @@ class ComposedPipelineBase(ABC):
|
||||
assert len(
|
||||
model_index
|
||||
) > 1, "model_index.json must contain at least one pipeline module"
|
||||
|
||||
|
||||
for module_name in self.required_config_modules:
|
||||
if module_name not in model_index:
|
||||
raise ValueError(
|
||||
f"model_index.json must contain a {module_name} module")
|
||||
|
||||
logger.warning(
|
||||
f"model_index.json does not contain a {module_name} module, adding {module_name} to model_index")
|
||||
if 'transformer' in module_name:
|
||||
model_index[module_name] = model_index['transformer']
|
||||
# all the component models used by the pipeline
|
||||
required_modules = self.required_config_modules
|
||||
logger.info("Loading required modules: %s", required_modules)
|
||||
@@ -242,7 +243,11 @@ class ComposedPipelineBase(ABC):
|
||||
logger.info("Using module %s already provided", module_name)
|
||||
modules[module_name] = loaded_modules[module_name]
|
||||
continue
|
||||
component_model_path = os.path.join(self.model_path, module_name)
|
||||
if 'transformer' in module_name:
|
||||
loading_module_name = module_name.split("_")[-1]
|
||||
else:
|
||||
loading_module_name = module_name
|
||||
component_model_path = os.path.join(self.model_path, loading_module_name)
|
||||
module = PipelineComponentLoader.load_module(
|
||||
module_name=module_name,
|
||||
component_model_path=component_model_path,
|
||||
|
||||
@@ -18,6 +18,12 @@ from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
|
||||
WanTeaCacheParams)
|
||||
|
||||
__all__ = [
|
||||
"ForwardBatch",
|
||||
"TrainingBatch",
|
||||
"AttentionMetadata",
|
||||
"VideoSparseAttentionMetadata",
|
||||
]
|
||||
|
||||
@dataclass
|
||||
class ForwardBatch:
|
||||
@@ -146,9 +152,11 @@ class ForwardBatch:
|
||||
class TrainingBatch:
|
||||
current_timestep: int = 0
|
||||
current_vsa_sparsity: float = 0.0
|
||||
|
||||
|
||||
# Dataloader batch outputs
|
||||
latents: torch.Tensor | None = None
|
||||
noise_latents: torch.Tensor | None = None
|
||||
encoder_hidden_states: torch.Tensor | None = None
|
||||
encoder_attention_mask: torch.Tensor | None = None
|
||||
# i2v
|
||||
@@ -156,6 +164,7 @@ class TrainingBatch:
|
||||
image_embeds: torch.Tensor | None = None
|
||||
image_latents: torch.Tensor | None = None
|
||||
infos: list[dict[str, Any]] | None = None
|
||||
mask_lat_size: torch.Tensor | None = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: torch.Tensor | None = None
|
||||
@@ -163,6 +172,7 @@ class TrainingBatch:
|
||||
sigmas: torch.Tensor | None = None
|
||||
noise: torch.Tensor | None = None
|
||||
|
||||
attn_metadata_vsa: AttentionMetadata | None = None
|
||||
attn_metadata: AttentionMetadata | None = None
|
||||
|
||||
# input kwargs
|
||||
@@ -174,3 +184,19 @@ class TrainingBatch:
|
||||
# Training outputs
|
||||
total_loss: float | None = None
|
||||
grad_norm: float | None = None
|
||||
|
||||
# Distillation-specific attributes
|
||||
encoder_hidden_states_neg: torch.Tensor | None = None
|
||||
encoder_attention_mask_neg: torch.Tensor | None = None
|
||||
conditional_dict: dict[str, Any] | None = None
|
||||
unconditional_dict: dict[str, Any] | None = None
|
||||
|
||||
# Distillation losses
|
||||
student_loss: float = 0.0
|
||||
critic_loss: float = 0.0
|
||||
|
||||
# Training control
|
||||
dmd_log_dict: dict[str, Any] = field(default_factory=dict)
|
||||
critic_log_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DmdDenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.v1.pipelines.stages.image_encoding import ImageEncodingStage
|
||||
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
|
||||
@@ -28,6 +29,7 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
|
||||
@@ -103,6 +103,7 @@ class DecodingStage(PipelineStage):
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import inspect, copy
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
@@ -27,6 +27,7 @@ from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import (
|
||||
@@ -117,6 +118,7 @@ class DenoisingStage(PipelineStage):
|
||||
batch.image_latent = image_latent
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if timesteps is None:
|
||||
raise ValueError("Timesteps must be provided")
|
||||
@@ -176,7 +178,7 @@ class DenoisingStage(PipelineStage):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
|
||||
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
@@ -541,3 +543,233 @@ class DenoisingStage(PipelineStage):
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
|
||||
class DmdDenoisingStage(DenoisingStage):
|
||||
"""
|
||||
Denoising stage for DMD.
|
||||
"""
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__(transformer, scheduler)
|
||||
self.scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=8.0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run the denoising loop.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
{
|
||||
"generator": batch.generator,
|
||||
"eta": batch.eta
|
||||
},
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
# TODO(will): make the precision configurable for inference
|
||||
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Handle sequence parallelism if enabled
|
||||
sp_world_size, rank_in_sp_group = get_sp_world_size(
|
||||
), get_sp_parallel_rank()
|
||||
sp_group = sp_world_size > 1
|
||||
if sp_group:
|
||||
latents = rearrange(batch.latents,
|
||||
"b t (n s) h w -> b t n s h w",
|
||||
n=sp_world_size).contiguous()
|
||||
latents = latents[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.latents = latents
|
||||
if batch.image_latent is not None:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b t (n s) h w -> b t n s h w",
|
||||
n=sp_world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if timesteps is None:
|
||||
raise ValueError("Timesteps must be provided")
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
num_warmup_steps = len(
|
||||
timesteps) - num_inference_steps * self.scheduler.order
|
||||
|
||||
# Prepare image latents and embeddings for I2V generation
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"mask_strategy": dict_to_3d_list(
|
||||
None, t_max=50, l_max=60, h_max=24)
|
||||
},
|
||||
)
|
||||
|
||||
pos_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
neg_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_2": batch.clip_embedding_neg,
|
||||
"encoder_attention_mask": batch.negative_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
# Prepare STA parameters
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
# Get latents and embeddings
|
||||
latents = batch.latents
|
||||
# TODO(yongqi) hard code prepare latents
|
||||
latents = torch.randn(latents.permute(0, 2, 1, 3, 4).shape, dtype=torch.bfloat16, device="cuda", generator=torch.Generator(device="cuda").manual_seed(42))
|
||||
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
timesteps = torch.tensor(
|
||||
fastvideo_args.denoising_step_list, dtype=torch.long, device=get_local_torch_device())
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
# Expand latents for I2V
|
||||
noise_latents = copy.deepcopy(latents)
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
|
||||
if batch.image_latent is not None:
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent.permute(0, 2, 1, 3, 4)],
|
||||
dim=2).to(target_dtype)
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) *
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
if (st_attn_available
|
||||
and self.attn_backend == SlidingTileAttentionBackend
|
||||
) or (vsa_available and self.attn_backend
|
||||
== VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
# TODO(will): clean this up
|
||||
attn_metadata = self.attn_metadata_builder.build(
|
||||
current_timestep=i,
|
||||
forward_batch=batch,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
# TODO(will): finalize the interface. vLLM uses this to
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
batch.is_cfg_negative = False
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
pred_noise = self.transformer(
|
||||
latent_model_input.permute(0, 2, 1, 3, 4),
|
||||
prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
|
||||
t_shape = pred_noise.shape[1]
|
||||
timestep = t_expand.expand(1, t_shape)
|
||||
from fastvideo.v1.training.training_utils import DiffusionWrapper
|
||||
pred_video = DiffusionWrapper._convert_flow_pred_to_x0(
|
||||
flow_pred=pred_noise.flatten(0, 1),
|
||||
xt=noise_latents.flatten(0, 1),
|
||||
timestep=timestep.flatten(0, 1),
|
||||
scheduler=self.scheduler
|
||||
).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
if i < len(timesteps) - 1:
|
||||
next_timestep = timesteps[i + 1] * torch.ones(
|
||||
pred_video.shape[:2], dtype=torch.long, device=pred_video.device)
|
||||
latents = self.scheduler.add_noise(
|
||||
pred_video.flatten(0, 1),
|
||||
torch.randn_like(pred_video.flatten(0, 1)),
|
||||
next_timestep.flatten(0, 1)
|
||||
).unflatten(0, pred_video.shape[:2])
|
||||
else:
|
||||
latents = pred_video.permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
(i + 1) % self.scheduler.order == 0
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
if sp_group:
|
||||
latents = sequence_model_parallel_all_gather(latents, dim=2)
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
return batch
|
||||
|
||||
@@ -136,7 +136,6 @@ class EncodingStage(PipelineStage):
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
self.vae.to("cpu")
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
@@ -5,8 +5,6 @@ Image encoding stages for I2V diffusion pipelines.
|
||||
This module contains implementations of image encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
@@ -68,7 +66,6 @@ class ImageEncodingStage(PipelineStage):
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
import numpy as np
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DmdDenoisingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanDmdPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
Wan video diffusion pipeline with LoRA support.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DmdDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = WanDmdPipeline
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
|
||||
EncodingStage, InputValidationStage, LatentPreparationStage,
|
||||
TextEncodingStage, TimestepPreparationStage)
|
||||
# isort: on
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
|
||||
"image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=ImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DmdDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoDmdPipeline
|
||||
@@ -105,7 +105,7 @@ def run_training():
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--tp_size", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--tp_size", 1,
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
|
||||
@@ -24,7 +24,7 @@ FastHunyuan-diffusers: {
|
||||
"flow_shift": 17,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": true,
|
||||
"fps": 24
|
||||
}
|
||||
@@ -41,7 +41,7 @@ Wan2.1-T2V-1.3B-Diffusers: {
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
@@ -60,7 +60,7 @@ Wan2.1-I2V-14B-480P-Diffusers: {
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
|
||||
@@ -33,7 +33,7 @@ HUNYUAN_PARAMS = {
|
||||
"flow_shift": 17,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
}
|
||||
@@ -50,7 +50,7 @@ WAN_T2V_PARAMS = {
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
@@ -69,7 +69,7 @@ WAN_I2V_PARAMS = {
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
@@ -238,7 +238,7 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
min_acceptable_ssim = 0.97
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim}"
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN", "TORCH_SDPA"])
|
||||
@@ -337,5 +337,5 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
if not success:
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
min_acceptable_ssim = 0.95
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim}"
|
||||
min_acceptable_ssim = 0.93
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} for {model_id} with backend {ATTENTION_BACKEND}"
|
||||
|
||||
@@ -43,7 +43,7 @@ def run_worker():
|
||||
"--num_latent_t", "4",
|
||||
"--num_gpus", "4",
|
||||
"--sp_size", "4",
|
||||
"--tp_size", "4",
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", "4",
|
||||
"--train_sp_batch_size", "1",
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from .training_pipeline import TrainingPipeline
|
||||
from .wan_training_pipeline import WanTrainingPipeline
|
||||
from .distillation_pipeline import DistillationPipeline
|
||||
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
|
||||
@@ -0,0 +1,998 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import gc
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from abc import abstractmethod
|
||||
from collections import deque
|
||||
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
from diffusers.optimization import get_scheduler
|
||||
from einops import rearrange
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
|
||||
get_local_torch_device, get_world_group)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs,TrainingArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.v1.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
|
||||
normalize_dit_input, save_checkpoint, shard_latents_across_sp, prepare_for_saving)
|
||||
from fastvideo.v1.utils import set_random_seed, is_vsa_available
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler, FlowMatchScheduler
|
||||
from fastvideo.v1.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.v1.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadata)
|
||||
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.v1.training.training_utils import DiffusionWrapper
|
||||
import wandb # isort: skip
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class DistillationPipeline(TrainingPipeline):
|
||||
"""
|
||||
A distillation pipeline for training a student model using teacher model guidance.
|
||||
Inherits from TrainingPipeline to reuse training infrastructure.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae", "teacher_transformer", "critic_transformer"]
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
|
||||
Dict[str, Any]]]
|
||||
current_epoch: int = 0
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize the distillation training pipeline with multiple models."""
|
||||
logger.info("Initializing distillation training pipeline...")
|
||||
|
||||
# 1. Call parent initialization first
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
# self.noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=self.timestep_shift)
|
||||
self.noise_scheduler = FlowMatchScheduler(
|
||||
shift=8.0, sigma_min=0.0, extra_one_step=True
|
||||
)
|
||||
self.noise_scheduler.set_timesteps(1000, training=True)
|
||||
|
||||
# 2. Distillation-specific initialization
|
||||
# The parent class already sets self.transformer as the student model
|
||||
self.student_transformer = DiffusionWrapper(self.transformer, self.noise_scheduler)
|
||||
self.teacher_transformer = DiffusionWrapper(self.get_module("teacher_transformer"), self.noise_scheduler)
|
||||
self.critic_transformer = DiffusionWrapper(self.get_module("critic_transformer"), self.noise_scheduler)
|
||||
|
||||
self.teacher_transformer.requires_grad_(False)
|
||||
self.teacher_transformer.eval()
|
||||
self.critic_transformer.requires_grad_(True)
|
||||
self.critic_transformer.train()
|
||||
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
self.critic_transformer = apply_activation_checkpointing(
|
||||
self.critic_transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
# Initialize optimizers
|
||||
critic_params = list(filter(lambda p: p.requires_grad, self.critic_transformer.parameters()))
|
||||
self.critic_transformer_optimizer = torch.optim.AdamW(
|
||||
critic_params,
|
||||
lr=training_args.critic_learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
if training_args.critic_lr_scheduler == "piecewise_constant":
|
||||
assert training_args.critic_lr_step_rules is not None, "critic lr step rules is required when using piecewise_constant lr scheduler"
|
||||
|
||||
self.critic_lr_scheduler = get_scheduler(
|
||||
training_args.critic_lr_scheduler,
|
||||
step_rules=training_args.critic_lr_step_rules,
|
||||
optimizer=self.critic_transformer_optimizer,
|
||||
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
|
||||
num_training_steps=training_args.max_train_steps * self.world_size,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
logger.info("Distillation optimizers initialized: student and critic")
|
||||
|
||||
self.student_critic_update_ratio = self.training_args.student_critic_update_ratio
|
||||
logger.info(f"Distillation pipeline initialized with student_critic_update_ratio={self.student_critic_update_ratio}")
|
||||
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.denoising_step_list, dtype=torch.long, device=get_local_torch_device())
|
||||
logger.info(f"Distillation student model to {len(self.denoising_step_list)} denoising steps")
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
# TODO(yongqi): hardcode for bidirectional distillation
|
||||
self.distill_task_type = "bidirectional_video"
|
||||
self.denoising_loss_type = 'flow'
|
||||
# TODO(yongqi): hardcode for causal distillation
|
||||
self.num_frame_per_block = 3
|
||||
|
||||
self.min_step = int(self.training_args.min_step_ratio * self.num_train_timestep)
|
||||
self.max_step = int(self.training_args.max_step_ratio * self.num_train_timestep)
|
||||
|
||||
self.teacher_guidance_scale = self.training_args.teacher_guidance_scale
|
||||
self.denoising_loss_func = FlowPredLoss()
|
||||
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize validation pipeline - must be implemented by subclasses."""
|
||||
raise NotImplementedError(
|
||||
"Distillation pipelines must implement this method")
|
||||
|
||||
def _prepare_distillation(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Prepare training environment for distillation."""
|
||||
self.student_transformer.requires_grad_(True)
|
||||
self.student_transformer.train()
|
||||
self.critic_transformer.requires_grad_(True)
|
||||
self.critic_transformer.train()
|
||||
|
||||
return training_batch
|
||||
|
||||
def _process_timestep(self, timestep: torch.Tensor, type: str) -> torch.Tensor:
|
||||
"""
|
||||
Pre-process the randomly generated timestep based on the generator's task type.
|
||||
Input:
|
||||
- timestep: [batch_size, num_frame] tensor containing the randomly generated timestep.
|
||||
- type: a string indicating the type of the current model (image, bidirectional_video, or causal_video).
|
||||
Output Behavior:
|
||||
- image: check that the second dimension (num_frame) is 1.
|
||||
- bidirectional_video: broadcast the timestep to be the same for all frames.
|
||||
- causal_video: broadcast the timestep to be the same for all frames **in a block**.
|
||||
"""
|
||||
if type == "image":
|
||||
assert timestep.shape[1] == 1
|
||||
return timestep
|
||||
elif type == "bidirectional_video":
|
||||
for index in range(timestep.shape[0]):
|
||||
timestep[index] = timestep[index, 0]
|
||||
return timestep
|
||||
elif type == "causal_video":
|
||||
# make the noise level the same within every motion block
|
||||
timestep = timestep.reshape(
|
||||
timestep.shape[0], -1, self.num_frame_per_block)
|
||||
timestep[:, :, 1:] = timestep[:, :, 0:1]
|
||||
timestep = timestep.reshape(timestep.shape[0], -1)
|
||||
return timestep
|
||||
else:
|
||||
raise NotImplementedError("Unsupported model type {}".format(type))
|
||||
|
||||
def _student_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Forward pass through student transformer and compute student losses."""
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
simulated_noisy_input = []
|
||||
for timestep in self.denoising_step_list:
|
||||
# Use cross-codebase generator for reproducible noise generation
|
||||
noise = torch.randn(
|
||||
self.video_latent_shape, device=self.device, dtype=dtype)
|
||||
|
||||
noisy_timestep = timestep * torch.ones(
|
||||
self.video_latent_shape[:2], device=self.device, dtype=torch.long)
|
||||
|
||||
if timestep != 0:
|
||||
noisy_video = self.noise_scheduler.add_noise(
|
||||
latents.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
noisy_timestep.flatten(0, 1)
|
||||
).unflatten(0, self.video_latent_shape[:2])
|
||||
else:
|
||||
noisy_video = latents
|
||||
|
||||
simulated_noisy_input.append(noisy_video)
|
||||
|
||||
simulated_noisy_input = torch.stack(simulated_noisy_input, dim=1)
|
||||
|
||||
# Step 2: Randomly sample a timestep and pick the corresponding input
|
||||
# Use cross-codebase generator for reproducible index generation
|
||||
index = torch.randint(0, len(self.denoising_step_list), [
|
||||
self.video_latent_shape[0], self.video_latent_shape[1]], device=self.device, dtype=torch.long)
|
||||
|
||||
index = self._process_timestep(index, type=self.distill_task_type)
|
||||
|
||||
# select the corresponding timestep's noisy input from the stacked tensor [B, T, F, C, H, W]
|
||||
|
||||
noisy_input = torch.gather(
|
||||
simulated_noisy_input, dim=1,
|
||||
index=index.reshape(index.shape[0], 1, index.shape[1], 1, 1, 1).expand(
|
||||
-1, -1, -1, *self.video_latent_shape[2:])
|
||||
).squeeze(1)
|
||||
|
||||
timestep = self.denoising_step_list[index]
|
||||
|
||||
training_batch = self._build_input_kwargs(noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_video = self.student_transformer(training_batch, timestep)
|
||||
|
||||
pred_video = pred_video.type_as(noisy_input)
|
||||
|
||||
return pred_video, timestep.float().detach()
|
||||
|
||||
def _multi_step_simulation_student_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Forward pass through student transformer matching inference procedure."""
|
||||
from fastvideo.v1.training.training_utils import DiffusionWrapper
|
||||
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
|
||||
# Step 1: Randomly sample a target timestep index from denoising_step_list
|
||||
target_timestep_idx = torch.randint(0, len(self.denoising_step_list), [
|
||||
self.video_latent_shape[0], self.video_latent_shape[1]], device=self.device, dtype=torch.long)
|
||||
|
||||
target_timestep_idx = self._process_timestep(target_timestep_idx, type=self.distill_task_type)
|
||||
target_timestep = self.denoising_step_list[target_timestep_idx]
|
||||
|
||||
# Step 2: Simulate the multi-step inference process up to the target timestep
|
||||
# Start from pure noise like in inference
|
||||
current_latents = torch.randn(self.video_latent_shape, device=self.device, dtype=dtype)
|
||||
|
||||
# Only run intermediate steps if target_timestep_idx > 0
|
||||
max_target_idx = target_timestep_idx.max().item()
|
||||
if max_target_idx > 0:
|
||||
# Run student model for all steps before the target timestep
|
||||
with torch.no_grad():
|
||||
for step_idx in range(max_target_idx):
|
||||
current_timestep = self.denoising_step_list[step_idx]
|
||||
logger.info(f"target_timestep: {target_timestep}, current_timestep: {current_timestep}")
|
||||
current_timestep_tensor = current_timestep * torch.ones(
|
||||
self.video_latent_shape[:2], device=self.device, dtype=torch.long)
|
||||
|
||||
# Run student model to get flow prediction
|
||||
training_batch_temp = self._build_input_kwargs(
|
||||
current_latents, current_timestep_tensor, training_batch.conditional_dict, training_batch)
|
||||
pred_flow = self.student_transformer.model(**training_batch_temp.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# Convert flow prediction to x0 prediction
|
||||
pred_clean = DiffusionWrapper._convert_flow_pred_to_x0(
|
||||
flow_pred=pred_flow.flatten(0, 1),
|
||||
xt=current_latents.flatten(0, 1),
|
||||
timestep=current_timestep_tensor.flatten(0, 1),
|
||||
scheduler=self.noise_scheduler
|
||||
).unflatten(0, self.video_latent_shape[:2])
|
||||
|
||||
# Add noise for the next timestep
|
||||
next_timestep = self.denoising_step_list[step_idx + 1]
|
||||
next_timestep_tensor = next_timestep * torch.ones(
|
||||
self.video_latent_shape[:2], device=self.device, dtype=torch.long)
|
||||
current_latents = self.noise_scheduler.add_noise(
|
||||
pred_clean.flatten(0, 1),
|
||||
torch.randn_like(pred_clean.flatten(0, 1)),
|
||||
next_timestep_tensor.flatten(0, 1)
|
||||
).unflatten(0, self.video_latent_shape[:2])
|
||||
|
||||
# Step 3: Use the simulated noisy input for the final training step
|
||||
# For timestep index 0, this is pure noise
|
||||
# For timestep index k > 0, this is the result after k denoising steps + noise at target level
|
||||
noisy_input = current_latents
|
||||
|
||||
# Step 4: Final student prediction (this is what we train on)
|
||||
training_batch = self._build_input_kwargs(noisy_input, target_timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_video = self.student_transformer(training_batch, target_timestep)
|
||||
|
||||
pred_video = pred_video.type_as(noisy_input)
|
||||
|
||||
return pred_video, target_timestep.float().detach()
|
||||
|
||||
|
||||
def _compute_kl_grad(
|
||||
self,
|
||||
noisy_video: torch.Tensor,
|
||||
estimated_clean_video: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
training_batch: TrainingBatch,
|
||||
normalization: bool = True
|
||||
) -> Tuple[torch.Tensor, dict]:
|
||||
assert self.training_args is not None
|
||||
|
||||
# critic_transformer forward
|
||||
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_fake_video = self.critic_transformer(training_batch, timestep)
|
||||
|
||||
if self.current_trainstep > self.training_args.num_teacher_noisy_ground_truth_steps:
|
||||
teacher_noisy_input = noisy_video
|
||||
teacher_timestep = timestep
|
||||
else:
|
||||
teacher_timestep = timestep
|
||||
logger.info(f"Using noisy ground truth for teacher with timestep {teacher_timestep}")
|
||||
batch_size, num_frame = self.video_latent_shape[:2]
|
||||
noisy_ground_truth = self.noise_scheduler.add_noise(
|
||||
training_batch.latents.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
teacher_timestep.flatten(0, 1)
|
||||
).detach().unflatten(0, (batch_size, num_frame))
|
||||
|
||||
teacher_noisy_input = noisy_ground_truth
|
||||
|
||||
# teacher_transformer cond forward
|
||||
training_batch = self._build_input_kwargs(teacher_noisy_input, teacher_timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_real_video_cond = self.teacher_transformer(training_batch, teacher_timestep)
|
||||
|
||||
# teacher_transformer uncond forward
|
||||
training_batch = self._build_input_kwargs(teacher_noisy_input, teacher_timestep, training_batch.unconditional_dict, training_batch)
|
||||
pred_real_video_uncond = self.teacher_transformer(training_batch, teacher_timestep)
|
||||
|
||||
pred_real_video = pred_real_video_cond + (
|
||||
pred_real_video_cond - pred_real_video_uncond
|
||||
) * self.teacher_guidance_scale
|
||||
|
||||
grad = (pred_fake_video - pred_real_video)
|
||||
|
||||
if normalization:
|
||||
p_real = (estimated_clean_video - pred_real_video)
|
||||
normalizer = torch.abs(p_real).mean(dim=[1, 2, 3, 4], keepdim=True)
|
||||
grad = grad / normalizer
|
||||
grad = torch.nan_to_num(grad)
|
||||
|
||||
return grad, {
|
||||
"dmdtrain_latents": estimated_clean_video.detach(),
|
||||
"dmdtrain_noisy_latent": noisy_video.detach(),
|
||||
"dmdtrain_pred_real_video": pred_real_video.detach(),
|
||||
"dmdtrain_pred_fake_video": pred_fake_video.detach(),
|
||||
"dmdtrain_gradient_norm": torch.mean(torch.abs(grad)).detach(),
|
||||
"timestep": timestep.float().detach()
|
||||
}
|
||||
|
||||
def _compute_dmd_loss(self, pred_video: torch.Tensor, training_batch: TrainingBatch) -> Tuple[torch.Tensor, dict]:
|
||||
"""Compute DMD (Diffusion Model Distillation) loss."""
|
||||
|
||||
original_latent = pred_video
|
||||
batch_size, num_frame = self.video_latent_shape[:2]
|
||||
with torch.no_grad():
|
||||
# Use cross-codebase generator for reproducible timestep generation
|
||||
timestep = torch.randint(
|
||||
0,
|
||||
self.num_train_timestep,
|
||||
[batch_size, num_frame],
|
||||
device=self.device,
|
||||
dtype=torch.long
|
||||
)
|
||||
|
||||
timestep = self._process_timestep(
|
||||
timestep, type=self.distill_task_type)
|
||||
|
||||
if self.timestep_shift > 1:
|
||||
timestep = self.timestep_shift * \
|
||||
(timestep / self.num_train_timestep) / \
|
||||
(1 + (self.timestep_shift - 1) * (timestep / self.num_train_timestep)) * self.num_train_timestep
|
||||
|
||||
timestep = timestep.clamp(self.min_step, self.max_step)
|
||||
|
||||
# Use cross-codebase generator for reproducible noise generation
|
||||
noise = torch.randn_like(pred_video)
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
pred_video.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
timestep.flatten(0, 1)
|
||||
).detach().unflatten(0, (batch_size, num_frame))
|
||||
|
||||
grad, dmd_log_dict = self._compute_kl_grad(
|
||||
noisy_video=noisy_latent,
|
||||
estimated_clean_video=original_latent,
|
||||
noise=noise,
|
||||
timestep=timestep,
|
||||
training_batch=training_batch
|
||||
)
|
||||
|
||||
dmd_loss = 0.5 * F.mse_loss(original_latent.double(
|
||||
), (original_latent.double() - grad.double()).detach(), reduction="mean")
|
||||
|
||||
return dmd_loss, dmd_log_dict
|
||||
|
||||
def _student_forward_and_compute_dmd_loss(self, training_batch: TrainingBatch) -> Tuple[TrainingBatch, torch.Tensor, dict]:
|
||||
"""Forward pass through student transformer and compute student losses."""
|
||||
assert self.training_args is not None
|
||||
assert training_batch.conditional_dict is not None
|
||||
assert training_batch.unconditional_dict is not None
|
||||
assert training_batch.latents is not None
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
|
||||
if self.training_args.simulate_student_forward:
|
||||
pred_video, timestep_dmd = self._multi_step_simulation_student_forward(training_batch)
|
||||
else:
|
||||
pred_video, timestep_dmd = self._student_forward(training_batch)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata):
|
||||
dmd_loss, dmd_log_dict = self._compute_dmd_loss(
|
||||
pred_video=pred_video,
|
||||
training_batch=training_batch
|
||||
)
|
||||
|
||||
dmd_log_dict['dmd_timestep_stu'] = timestep_dmd
|
||||
|
||||
return training_batch, dmd_loss, dmd_log_dict
|
||||
|
||||
def _critic_forward_and_compute_loss(self, training_batch: TrainingBatch) -> Tuple[TrainingBatch, torch.Tensor, dict]:
|
||||
assert self.training_args is not None
|
||||
assert training_batch.conditional_dict is not None
|
||||
assert training_batch.unconditional_dict is not None
|
||||
assert training_batch.latents is not None
|
||||
|
||||
with torch.no_grad():
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
|
||||
if self.training_args.simulate_student_forward:
|
||||
generated_video, timestep_gen = self._multi_step_simulation_student_forward(training_batch)
|
||||
else:
|
||||
generated_video, timestep_gen = self._student_forward(training_batch)
|
||||
|
||||
critic_timestep = torch.randint(
|
||||
0,
|
||||
self.num_train_timestep,
|
||||
self.video_latent_shape[:2],
|
||||
device=self.device,
|
||||
dtype=torch.long
|
||||
)
|
||||
critic_timestep = self._process_timestep(
|
||||
critic_timestep, type=self.distill_task_type)
|
||||
|
||||
# TODO: Add timestep warping
|
||||
if self.timestep_shift > 1:
|
||||
critic_timestep = self.timestep_shift * \
|
||||
(critic_timestep / self.num_train_timestep) / (1 + (self.timestep_shift - 1) * (critic_timestep / self.num_train_timestep)) * self.num_train_timestep
|
||||
|
||||
critic_timestep = critic_timestep.clamp(self.min_step, self.max_step)
|
||||
|
||||
# Use cross-codebase generator for reproducible noise generation
|
||||
critic_noise = torch.randn_like(generated_video)
|
||||
|
||||
noisy_generated_video = self.noise_scheduler.add_noise(
|
||||
generated_video.flatten(0, 1),
|
||||
critic_noise.flatten(0, 1),
|
||||
critic_timestep.flatten(0, 1)
|
||||
).unflatten(0, self.video_latent_shape[:2])
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata):
|
||||
training_batch = self._build_input_kwargs(noisy_generated_video, critic_timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_fake_video = self.critic_transformer(training_batch, critic_timestep)
|
||||
|
||||
# # Step 3: Compute the denoising loss for the fake critic
|
||||
pred_fake_video_noise = DiffusionWrapper._convert_x0_to_flow_pred(
|
||||
x0_pred=pred_fake_video.flatten(0, 1),
|
||||
xt=noisy_generated_video.flatten(0, 1),
|
||||
timestep=critic_timestep.flatten(0, 1),
|
||||
scheduler=self.noise_scheduler
|
||||
)
|
||||
|
||||
denoising_loss = self.denoising_loss_func(
|
||||
x=generated_video.flatten(0, 1),
|
||||
noise=critic_noise.flatten(0, 1),
|
||||
flow_pred=pred_fake_video_noise
|
||||
)
|
||||
|
||||
critic_log_dict = {
|
||||
"critictrain_latent": generated_video.detach(),
|
||||
"critictrain_noisy_latent": noisy_generated_video.detach(),
|
||||
"critictrain_pred_video": pred_fake_video.detach(),
|
||||
"critic_timestep": critic_timestep.float().detach(),
|
||||
"critic_timestep_stu": timestep_gen.float().detach()
|
||||
}
|
||||
|
||||
return training_batch, denoising_loss, critic_log_dict
|
||||
|
||||
def _clip_grad_norm(self, training_batch: TrainingBatch, transformer) -> TrainingBatch:
|
||||
assert self.training_args is not None
|
||||
max_grad_norm = self.training_args.max_grad_norm
|
||||
|
||||
# TODO(will): perhaps move this into transformer api so that we can do
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
# Clip gradients for both student and critic models
|
||||
model_parts = [transformer]
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
foreach=None,
|
||||
)
|
||||
assert grad_norm is not float('nan') or grad_norm is not float(
|
||||
'inf')
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
training_batch.grad_norm = grad_norm
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
super()._prepare_dit_inputs(training_batch)
|
||||
conditional_dict = {
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
|
||||
training_batch.conditional_dict = conditional_dict
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
assert training_batch.latents is not None
|
||||
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
self.video_latent_shape = training_batch.latents.shape # [B, C, T, H, W]
|
||||
|
||||
|
||||
return training_batch
|
||||
|
||||
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Train one step with alternating student and critic updates."""
|
||||
assert self.training_args is not None
|
||||
|
||||
training_batch = self._prepare_distillation(training_batch)
|
||||
TRAIN_STUDENT = self.current_trainstep % self.student_critic_update_ratio == 0
|
||||
# for _ in range(self.training_args.gradient_accumulation_steps):
|
||||
training_batch = self._get_next_batch(training_batch)
|
||||
|
||||
training_batch = self._normalize_dit_input(training_batch)
|
||||
training_batch = self._prepare_dit_inputs(training_batch)
|
||||
|
||||
training_batch = self._build_attention_metadata(training_batch)
|
||||
|
||||
import copy
|
||||
training_batch.attn_metadata_vsa = copy.deepcopy(training_batch.attn_metadata)
|
||||
if training_batch.attn_metadata is not None:
|
||||
training_batch.attn_metadata.VSA_sparsity = 0.0
|
||||
|
||||
if TRAIN_STUDENT:
|
||||
training_batch, dmd_loss, dmd_log_dict = self._student_forward_and_compute_dmd_loss(training_batch)
|
||||
training_batch.dmd_log_dict = dmd_log_dict
|
||||
self.optimizer.zero_grad()
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
|
||||
dmd_loss.backward()
|
||||
training_batch = self._clip_grad_norm(training_batch, self.student_transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
avg_dmd_loss = dmd_loss.detach().clone()
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_dmd_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
|
||||
training_batch.student_loss = avg_dmd_loss.item()
|
||||
|
||||
training_batch, critic_loss, critic_log_dict = self._critic_forward_and_compute_loss(training_batch)
|
||||
training_batch.critic_log_dict = critic_log_dict
|
||||
self.critic_transformer_optimizer.zero_grad()
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata):
|
||||
critic_loss.backward()
|
||||
training_batch = self._clip_grad_norm(training_batch, self.critic_transformer)
|
||||
self.critic_transformer_optimizer.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
|
||||
avg_critic_loss = critic_loss.detach().clone()
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_critic_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
|
||||
# Record loss values for logging
|
||||
|
||||
training_batch.critic_loss = avg_critic_loss.item()
|
||||
|
||||
training_batch.total_loss = training_batch.student_loss + training_batch.critic_loss
|
||||
|
||||
return training_batch
|
||||
|
||||
def _resume_from_checkpoint(self) -> None: #TODO(yongqi)
|
||||
"""Resume training from checkpoint with distillation models."""
|
||||
assert self.training_args is not None
|
||||
logger.info("Loading distillation checkpoint from %s",
|
||||
self.training_args.resume_from_checkpoint)
|
||||
|
||||
resumed_step = load_checkpoint(
|
||||
self.student_transformer.model, self.global_rank,
|
||||
self.training_args.resume_from_checkpoint, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
|
||||
# TODO: Add checkpoint loading for critic and teacher models
|
||||
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
logger.info("Successfully resumed from step %s", resumed_step)
|
||||
else:
|
||||
logger.warning("Failed to load checkpoint, starting from step 0")
|
||||
self.init_steps = -1
|
||||
|
||||
def _log_training_info(self) -> None:
|
||||
"""Log distillation-specific training information."""
|
||||
# First call parent class method to get basic training info
|
||||
super()._log_training_info()
|
||||
|
||||
# Then add distillation-specific information
|
||||
logger.info("Distillation-specific settings:")
|
||||
logger.info(" Student/Critic update ratio: %s", self.student_critic_update_ratio)
|
||||
assert isinstance(self.training_args, TrainingArgs)
|
||||
logger.info(" Max gradient norm: %s", self.training_args.max_grad_norm)
|
||||
assert self.teacher_transformer is not None
|
||||
logger.info(" Teacher transformer parameters: %s B",
|
||||
sum(p.numel() for p in self.teacher_transformer.parameters()) / 1e9)
|
||||
assert self.critic_transformer is not None
|
||||
logger.info(" Critic transformer parameters: %s B",
|
||||
sum(p.numel() for p in self.critic_transformer.parameters()) / 1e9)
|
||||
|
||||
def add_visualization(self, generator_log_dict: Dict[str, Any], critic_log_dict: Dict[str, Any], training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging."""
|
||||
wandb_loss_dict = {}
|
||||
|
||||
# Clear GPU cache before VAE decoding to prevent OOM
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# # Use consistent decoding approach - use decode_stage for all
|
||||
# decode_stage = self.validation_pipeline._stages[-1]
|
||||
|
||||
# Process critic training data
|
||||
critic_latents_name = ['critictrain_latent', 'critictrain_noisy_latent', 'critictrain_pred_video']
|
||||
# critic_latents_name = ['critictrain_pred_video']
|
||||
for latent_key in critic_latents_name:
|
||||
latents = critic_log_dict[latent_key]
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
# decoded_latent = decode_stage(ForwardBatch(data_type="video", latents=latents), training_args)
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
wandb_loss_dict[latent_key] = prepare_for_saving(video)
|
||||
# Clean up references
|
||||
del video, latents
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Process DMD training data if available - use decode_stage instead of self.vae.decode
|
||||
|
||||
dmd_latents_name = ['dmdtrain_pred_fake_video', 'dmdtrain_pred_real_video', 'dmdtrain_latents', 'dmdtrain_noisy_latent']
|
||||
for latent_key in dmd_latents_name:
|
||||
latents = generator_log_dict[latent_key]
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
# decoded_latent = decode_stage(ForwardBatch(data_type="video", latents=latents), training_args)
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
wandb_loss_dict[latent_key] = prepare_for_saving(video)
|
||||
# Clean up references
|
||||
del video, latents
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Log to wandb
|
||||
if self.global_rank == 0:
|
||||
wandb.log(wandb_loss_dict, step=step)
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
assert training_args is not None
|
||||
training_args.inference_mode = True
|
||||
training_args.use_cpu_offload = False
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
raise ValueError("Validation pipeline is not set")
|
||||
|
||||
logger.info("Starting validation")
|
||||
|
||||
# Create sampling parameters if not provided
|
||||
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
|
||||
|
||||
# Set deterministic seed for validation
|
||||
# set_random_seed(self.seed)
|
||||
logger.info("Using validation seed: %s", self.seed)
|
||||
|
||||
# Prepare validation prompts
|
||||
logger.info('rank: %s: fastvideo_args.validation_dataset_file: %s',
|
||||
self.global_rank,
|
||||
training_args.validation_dataset_file,
|
||||
local_main_process_only=False)
|
||||
validation_dataset = ValidationDataset(
|
||||
training_args.validation_dataset_file)
|
||||
validation_dataloader = DataLoader(validation_dataset,
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
transformer.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
validation_steps = [step for step in validation_steps if step > 0]
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
# Process each validation prompt for each validation step
|
||||
for num_inference_steps in validation_steps:
|
||||
logger.info("rank: %s: num_inference_steps: %s",
|
||||
self.global_rank,
|
||||
num_inference_steps,
|
||||
local_main_process_only=False)
|
||||
step_videos: list[np.ndarray] = []
|
||||
step_captions: list[str] = []
|
||||
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(sampling_param,
|
||||
training_args,
|
||||
validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage(batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
|
||||
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
self.global_rank,
|
||||
self.rank_in_sp_group,
|
||||
batch.prompt,
|
||||
local_main_process_only=False)
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# Run validation inference
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in video:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
step_videos.append(frames)
|
||||
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
|
||||
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
|
||||
# results to global rank 0
|
||||
if self.rank_in_sp_group == 0:
|
||||
if self.global_rank == 0:
|
||||
# Global rank 0 collects results from all sp_group leaders
|
||||
all_videos = step_videos # Start with own results
|
||||
all_captions = step_captions
|
||||
|
||||
# Receive from other sp_group leaders
|
||||
for sp_group_idx in range(1, num_sp_groups):
|
||||
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
|
||||
recv_videos = world_group.recv_object(src=src_rank)
|
||||
recv_captions = world_group.recv_object(src=src_rank)
|
||||
all_videos.extend(recv_videos)
|
||||
all_captions.extend(recv_captions)
|
||||
|
||||
video_filenames = []
|
||||
for i, (video, caption) in enumerate(
|
||||
zip(all_videos, all_captions, strict=True)):
|
||||
os.makedirs(training_args.output_dir, exist_ok=True)
|
||||
filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
|
||||
)
|
||||
imageio.mimsave(filename, video, fps=sampling_param.fps)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
f"validation_videos_{num_inference_steps}_steps": [
|
||||
wandb.Video(filename, caption=caption)
|
||||
for filename, caption in zip(
|
||||
video_filenames, all_captions, strict=True)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
|
||||
# Save all prompts from all cards to txt file
|
||||
prompt_filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_prompts.txt"
|
||||
)
|
||||
with open(prompt_filename, 'w', encoding='utf-8') as f:
|
||||
for i, caption in enumerate(all_captions):
|
||||
f.write(f"Video_{i}: {caption}\n")
|
||||
logger.info(f"Saved {len(all_captions)} prompts to {prompt_filename}")
|
||||
else:
|
||||
# Other sp_group leaders send their results to global rank 0
|
||||
world_group.send_object(step_videos, dst=0)
|
||||
world_group.send_object(step_captions, dst=0)
|
||||
|
||||
# Re-enable gradients for training
|
||||
transformer.train()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def train(self) -> None:
|
||||
"""Main training loop with distillation-specific logging."""
|
||||
assert self.training_args is not None
|
||||
|
||||
assert self.training_args.seed is not None, "seed must be set"
|
||||
seed = self.training_args.seed
|
||||
set_random_seed(seed + self.global_rank)
|
||||
|
||||
self.noise_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(seed)
|
||||
|
||||
self.validation_generator = torch.Generator(device=get_local_torch_device()).manual_seed(42)
|
||||
|
||||
logger.info("Initialized random seeds with seed: %s", seed)
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
self._log_training_info()
|
||||
self._log_validation(self.student_transformer, self.training_args, 0)
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, self.training_args.max_train_steps),
|
||||
initial=self.init_steps,
|
||||
desc="Steps",
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
|
||||
for step in range(self.init_steps+1,
|
||||
self.training_args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
current_vsa_sparsity = self.training_args.VSA_sparsity if vsa_available else 0.0
|
||||
|
||||
training_batch = TrainingBatch()
|
||||
self.current_trainstep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
total_loss = training_batch.total_loss
|
||||
student_loss = training_batch.student_loss
|
||||
critic_loss = training_batch.critic_loss
|
||||
grad_norm = training_batch.grad_norm
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"total_loss": f"{total_loss:.4f}",
|
||||
"student_loss": f"{student_loss:.4f}",
|
||||
"critic_loss": f"{critic_loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
if self.global_rank == 0:
|
||||
# Prepare logging data
|
||||
log_data = {
|
||||
"train_total_loss": total_loss,
|
||||
"train_student_loss": student_loss,
|
||||
"train_critic_loss": critic_loss,
|
||||
"learning_rate": self.lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
}
|
||||
|
||||
# Add DMD training metrics if available
|
||||
if hasattr(training_batch, 'dmd_log_dict') and training_batch.dmd_log_dict:
|
||||
dmd_metrics = {
|
||||
"dmd_gradient_norm": training_batch.dmd_log_dict.get("dmdtrain_gradient_norm", 0.0),
|
||||
"dmd_timestep": training_batch.dmd_log_dict.get("timestep", 0.0).mean().item(),
|
||||
"dmd_timestep_stu": training_batch.dmd_log_dict.get("dmd_timestep_stu", 0.0).mean().item()
|
||||
}
|
||||
log_data.update(dmd_metrics)
|
||||
|
||||
# Add critic training metrics if available
|
||||
if hasattr(training_batch, 'critic_log_dict') and training_batch.critic_log_dict:
|
||||
critic_metrics = {
|
||||
"critic_timestep": training_batch.critic_log_dict.get("critic_timestep", 0.0).mean().item(),
|
||||
"critic_timestep_stu": training_batch.critic_log_dict.get("critic_timestep_stu", 0.0).mean().item(),
|
||||
}
|
||||
log_data.update(critic_metrics)
|
||||
wandb.log(log_data, step=step)
|
||||
|
||||
# if step % self.training_args.checkpointing_steps == 0:
|
||||
# print("rank", self.global_rank, "save checkpoint at step", step)
|
||||
# save_checkpoint(self.transformer, self.global_rank, #TODO(yongqi)
|
||||
# self.training_args.output_dir, step,
|
||||
# self.optimizer, self.train_dataloader,
|
||||
# self.lr_scheduler, self.noise_random_generator)
|
||||
# if self.transformer:
|
||||
# self.transformer.train()
|
||||
# self.sp_group.barrier()
|
||||
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage before validation: %s MB",
|
||||
gpu_memory_usage)
|
||||
self.add_visualization(training_batch.dmd_log_dict, training_batch.critic_log_dict, self.training_args, step)
|
||||
self._log_validation(self.student_transformer, self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage after validation: %s MB",
|
||||
gpu_memory_usage)
|
||||
|
||||
wandb.finish()
|
||||
save_checkpoint(self.student_transformer.model, self.global_rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
class FlowPredLoss():
|
||||
def __call__(
|
||||
self, x: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
flow_pred: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
return torch.mean((flow_pred - (noise - x)) ** 2)
|
||||
@@ -1,5 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import gc
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
@@ -110,8 +109,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
self.init_steps = 0
|
||||
logger.info("optimizer: %s", self.optimizer)
|
||||
|
||||
|
||||
if training_args.lr_scheduler == "piecewise_constant":
|
||||
assert training_args.lr_step_rules is not None, "lr step rules is required when using piecewise_constant lr scheduler"
|
||||
|
||||
self.lr_scheduler = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
step_rules=training_args.lr_step_rules,
|
||||
optimizer=self.optimizer,
|
||||
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
|
||||
num_training_steps=training_args.max_train_steps * self.world_size,
|
||||
@@ -222,7 +226,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
# indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
indices = (u * self.noise_scheduler.num_train_timesteps).long()
|
||||
timesteps = self.noise_scheduler.timesteps[indices].to(
|
||||
device=training_batch.latents.device)
|
||||
if self.training_args.sp_size > 1:
|
||||
@@ -254,7 +259,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
assert training_batch.timesteps is not None
|
||||
patch_size = self.training_args.pipeline_config.dit_config.patch_size
|
||||
current_vsa_sparsity = training_batch.current_vsa_sparsity
|
||||
|
||||
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
dit_seq_shape = [
|
||||
latents.shape[2] * self.sp_world_size // patch_size[0],
|
||||
@@ -320,6 +325,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
current_timestep=training_batch.current_timestep,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
model_pred = self.transformer(**input_kwargs)
|
||||
|
||||
if self.training_args.precondition_outputs:
|
||||
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
|
||||
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
|
||||
@@ -706,5 +712,4 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -3,13 +3,16 @@ import json
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dcp
|
||||
from torchvision.utils import make_grid
|
||||
from einops import rearrange
|
||||
from safetensors.torch import save_file
|
||||
import wandb
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
@@ -18,6 +21,8 @@ from fastvideo.v1.training.checkpointing_utils import (ModelWrapper,
|
||||
OptimizerWrapper,
|
||||
RandomStateWrapper,
|
||||
SchedulerWrapper)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from abc import ABC
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -550,3 +555,80 @@ def convert_custom_format_to_diffusers_format(state_dict: dict[str, Any],
|
||||
new_state_dict[training_key] = v
|
||||
|
||||
return new_state_dict
|
||||
|
||||
def prepare_for_saving(tensor: torch.Tensor, fps: int = 16, caption: str | None = None) -> wandb.Image | wandb.Video:
|
||||
if tensor.ndim == 4:
|
||||
# Assuming it's an image and has shape [batch_size, 3, height, width]
|
||||
tensor = make_grid(tensor, 4, padding=0, normalize=False)
|
||||
return wandb.Image((tensor * 255).numpy().astype(np.uint8), caption=caption)
|
||||
elif tensor.ndim == 5:
|
||||
# Assuming it's a video and has shape [batch_size, num_frames, 3, height, width]
|
||||
return wandb.Video((tensor * 255).numpy().astype(np.uint8), fps=fps, format="webm", caption=caption)
|
||||
else:
|
||||
raise ValueError("Unsupported tensor shape for saving. Expected 4D (image) or 5D (video) tensor.")
|
||||
|
||||
class DiffusionWrapper(torch.nn.Module, ABC):
|
||||
def __init__(self, transformer, scheduler):
|
||||
super().__init__()
|
||||
self.model = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
def forward(self, training_batch: TrainingBatch, timestep: torch.Tensor):
|
||||
pred_noise = self.model(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
pred_video = self._convert_flow_pred_to_x0(
|
||||
flow_pred=pred_noise.flatten(0, 1),
|
||||
xt=training_batch.noise_latents.flatten(0, 1),
|
||||
timestep=timestep.flatten(0, 1),
|
||||
scheduler=self.scheduler
|
||||
).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
return pred_video
|
||||
|
||||
@staticmethod
|
||||
def _convert_x0_to_flow_pred(x0_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor, scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert x0 prediction to flow matching's prediction.
|
||||
x0_pred: the x0 prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = (x_t - x_0) / sigma_t
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = x0_pred.dtype
|
||||
x0_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(x0_pred.device), [x0_pred, xt,
|
||||
scheduler.sigmas,
|
||||
scheduler.timesteps]
|
||||
)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
flow_pred = (xt - x0_pred) / sigma_t
|
||||
return flow_pred.to(original_dtype)
|
||||
|
||||
@staticmethod
|
||||
def _convert_flow_pred_to_x0(flow_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor, scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert flow matching's prediction to x0 prediction.
|
||||
flow_pred: the prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = noise - x0
|
||||
x_t = (1-sigma_t) * x0 + sigma_t * noise
|
||||
we have x0 = x_t - sigma_t * pred
|
||||
see derivations https://chatgpt.com/share/67bf8589-3d04-8008-bc6e-4cf1a24e2d0e
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = flow_pred.dtype
|
||||
flow_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(flow_pred.device), [flow_pred, xt,
|
||||
scheduler.sigmas,
|
||||
scheduler.timesteps]
|
||||
)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
x0_pred = xt - sigma_t * flow_pred
|
||||
return x0_pred.to(original_dtype)
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
import torch
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.v1.pipelines.wan.wan_dmd_pipeline import WanDmdPipeline
|
||||
from fastvideo.v1.training.distillation_pipeline import DistillationPipeline
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
|
||||
TrainingBatch)
|
||||
|
||||
from fastvideo.v1.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanDistillationPipeline(DistillationPipeline):
|
||||
"""
|
||||
A distillation pipeline for Wan that uses a single transformer model.
|
||||
The main transformer serves as the student model, and copies are made for teacher and critic.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae", "teacher_transformer", "critic_transformer"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize Wan-specific scheduler."""
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.use_cpu_offload = False
|
||||
args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
validation_pipeline = WanDmdPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
inference_mode=True,
|
||||
loaded_modules={"transformer": self.get_module("transformer")},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
def _build_input_kwargs(self, noise_input: torch.Tensor, timestep: torch.Tensor, text_dict: dict[str, torch.Tensor],
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"timestep": timestep[0][:1],
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
training_batch.noise_latents = noise_input
|
||||
return training_batch
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting Wan distillation pipeline...")
|
||||
|
||||
# Create pipeline with original args
|
||||
pipeline = WanDistillationPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
|
||||
args = pipeline.training_args
|
||||
|
||||
# Start training
|
||||
pipeline.train()
|
||||
logger.info("Wan distillation pipeline completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.v1.fastvideo_args import TrainingArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.use_cpu_offload = False
|
||||
main(args)
|
||||
@@ -0,0 +1,231 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.dataset.dataloader.schema import (
|
||||
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.v1.pipelines.wan.wan_i2v_dmd_pipeline import WanImageToVideoDmdPipeline
|
||||
from fastvideo.v1.training.distillation_pipeline import DistillationPipeline
|
||||
from fastvideo.v1.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanI2VDistillationPipeline(DistillationPipeline):
|
||||
"""
|
||||
A distillation pipeline for Wan that uses a single transformer model.
|
||||
The main transformer serves as the student model, and copies are made for teacher and critic.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae", "teacher_transformer", "critic_transformer"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize Wan-specific scheduler."""
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_i2v
|
||||
self.validation_dataset_schema = pyarrow_schema_i2v_validation
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.use_cpu_offload = False
|
||||
# args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
|
||||
validation_pipeline = WanImageToVideoDmdPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
inference_mode=True,
|
||||
loaded_modules={"transformer": self.get_module("transformer")},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
use_cpu_offload=True)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
assert self.training_args is not None
|
||||
assert self.train_dataloader is not None
|
||||
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
latents = batch['vae_latent']
|
||||
latents = latents[:, :, :self.training_args.num_latent_t]
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
clip_features = batch['clip_feature']
|
||||
image_latents = batch['first_frame_latent']
|
||||
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
|
||||
pil_image = batch['pil_image']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.preprocessed_image = pil_image.to(
|
||||
get_local_torch_device())
|
||||
training_batch.image_embeds = clip_features.to(get_local_torch_device())
|
||||
training_batch.image_latents = image_latents.to(
|
||||
get_local_torch_device())
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
|
||||
def _prepare_validation_batch(self, sampling_param: SamplingParam,
|
||||
training_args: TrainingArgs,
|
||||
validation_batch: dict[str, Any],
|
||||
num_inference_steps: int) -> ForwardBatch:
|
||||
sampling_param.prompt = validation_batch['prompt']
|
||||
sampling_param.height = training_args.num_height
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.image_path = validation_batch['video_path']
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
VSA_sparsity=training_args.VSA_sparsity,
|
||||
)
|
||||
|
||||
return batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
|
||||
assert self.training_args is not None
|
||||
assert training_batch.latents is not None
|
||||
assert training_batch.encoder_hidden_states is not None
|
||||
assert training_batch.encoder_attention_mask is not None
|
||||
assert self.noise_random_generator is not None
|
||||
assert training_batch.image_latents is not None
|
||||
|
||||
# First, call parent method to prepare noise, timesteps, etc. for video latents
|
||||
training_batch = super()._prepare_dit_inputs(training_batch)
|
||||
|
||||
assert isinstance(training_batch.image_latents, torch.Tensor)
|
||||
image_latents = training_batch.image_latents.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
|
||||
temporal_compression_ratio = 4
|
||||
num_frames = (self.training_args.num_latent_t -
|
||||
1) * temporal_compression_ratio + 1
|
||||
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
|
||||
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, 1:] = 0
|
||||
|
||||
first_frame_mask = mask_lat_size[:, :, :1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
|
||||
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
|
||||
dim=2)
|
||||
mask_lat_size = mask_lat_size.view(batch_size, -1,
|
||||
temporal_compression_ratio,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(
|
||||
image_latents.device).to(dtype=torch.bfloat16)
|
||||
|
||||
image_latents = torch.cat(
|
||||
[mask_lat_size, image_latents],
|
||||
dim=1)
|
||||
|
||||
training_batch.image_latents = image_latents
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_input_kwargs(self, noise_input: torch.Tensor, timestep: torch.Tensor, text_dict: dict[str, torch.Tensor],
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
assert training_batch.image_embeds is not None
|
||||
assert training_batch.image_latents is not None
|
||||
|
||||
# Image Embeds for conditioning
|
||||
image_embeds = training_batch.image_embeds
|
||||
assert torch.isnan(image_embeds).sum() == 0
|
||||
image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
noisy_model_input = torch.cat(
|
||||
[noise_input, training_batch.image_latents.permute(0, 2, 1, 3, 4)], dim=2)
|
||||
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states": noisy_model_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"timestep": timestep[0][:1],
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
training_batch.noise_latents = noise_input
|
||||
|
||||
return training_batch
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting Wan distillation pipeline...")
|
||||
|
||||
# Create pipeline with original args
|
||||
pipeline = WanI2VDistillationPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
|
||||
args = pipeline.training_args
|
||||
|
||||
# Start training
|
||||
pipeline.train()
|
||||
logger.info("Wan distillation pipeline completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.v1.fastvideo_args import TrainingArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.use_cpu_offload = False
|
||||
main(args)
|
||||
@@ -1,7 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import contextlib
|
||||
import faulthandler
|
||||
import gc
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import signal
|
||||
@@ -69,8 +68,6 @@ class Worker:
|
||||
torch.cuda.set_device(self.device)
|
||||
|
||||
# _check_if_gpu_supports_dtype(self.model_config.dtype)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
@@ -102,9 +99,6 @@ class Worker:
|
||||
if hasattr(self, 'pipeline') and self.pipeline is not None:
|
||||
# Clean up pipeline resources if needed
|
||||
pass
|
||||
# Release CUDA resources
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Destroy the distributed environment
|
||||
cleanup_dist_env_and_memory(shutdown_ray=False)
|
||||
@@ -133,8 +127,6 @@ class Worker:
|
||||
|
||||
# Handle regular RPC calls
|
||||
if method_name == 'execute_forward':
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
forward_batch = recv_rpc['kwargs']['forward_batch']
|
||||
fastvideo_args = recv_rpc['kwargs']['fastvideo_args']
|
||||
output_batch = self.execute_forward(forward_batch,
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY=
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=mini_i2v_dataset/crush-smol_preprocessed/combined_parquet_dataset
|
||||
VALIDATION_DIR=mini_i2v_dataset/crush-smol_preprocessed/validation_parquet_dataset
|
||||
NUM_GPUS=8
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_preprocessed_path "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim 8 \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--max_train_steps 30000 \
|
||||
--learning_rate 2e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 61 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '999,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
@@ -20,7 +20,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
--train_batch_size=4 \
|
||||
--num_latent_t 20 \
|
||||
--sp_size 4 \
|
||||
--tp_size 4 \
|
||||
--tp_size 1 \
|
||||
--hsdp_replicate_dim 1 \
|
||||
--hsdp_shard_dim 4 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/val/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
NUM_GPUS=8
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 8 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 29 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
|
||||
# validation_preprocessed_path
|
||||
@@ -0,0 +1,66 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/test_8/
|
||||
NUM_GPUS=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-14B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-14B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_preprocessed_path "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 20 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 768 \
|
||||
--num_width 1280 \
|
||||
--num_frames 77 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
--VSA_sparsity 0.0 \
|
||||
@@ -0,0 +1,66 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
NUM_GPUS=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path data/Wan2.1-T2V-1.3B-Diffusers-VT \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path data/Wan2.1-T2V-1.3B-Diffusers-VT \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 61 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
--VSA_sparsity 0.9 \
|
||||
@@ -0,0 +1,66 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/val/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
NUM_GPUS=8
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_distillation_pipeline.py \
|
||||
--model_path Wan2.1-I2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan2.1-I2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 8 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune_i2v" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 29 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "6.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
|
||||
@@ -3,13 +3,10 @@
|
||||
num_gpus=4
|
||||
export MODEL_BASE=FastVideo/FastHunyuan-Diffusers
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
|
||||
@@ -4,13 +4,10 @@ num_gpus=4
|
||||
export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
|
||||
@@ -5,13 +5,10 @@ export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_hunyuan.json
|
||||
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size ${num_gpus} \
|
||||
--tp-size ${num_gpus} \
|
||||
--tp-size 1 \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
--num-frames 117 \
|
||||
|
||||
@@ -4,13 +4,10 @@ num_gpus=2
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--height 480 \
|
||||
--width 832 \
|
||||
|
||||
@@ -4,13 +4,10 @@ num_gpus=2
|
||||
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
|
||||
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
|
||||
@@ -4,17 +4,14 @@ num_gpus=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# change model path to local dir if you want to inference using your checkpoint
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--tp-size 1 \
|
||||
--num-gpus $num_gpus \
|
||||
--height 448 \
|
||||
--width 832 \
|
||||
--height 768 \
|
||||
--width 1280\
|
||||
--num-frames 77 \
|
||||
--num-inference-steps 50 \
|
||||
--fps 16 \
|
||||
|
||||
@@ -4,9 +4,6 @@ num_gpus=2
|
||||
export FASTVIDEO_ATTENTION_BACKEND=
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers
|
||||
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
|
||||
Reference in New Issue
Block a user