Compare commits

..
74 changed files with 703 additions and 3528 deletions
+17
View File
@@ -58,6 +58,8 @@ steps:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
@@ -65,6 +67,21 @@ steps:
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/tests/lora/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/transformers/**"
- "fastvideo/v1/pipelines/**"
- "fastvideo/v1/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "pyproject.toml"
+4
View File
@@ -97,6 +97,10 @@ case "$TEST_TYPE" in
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
;;
"inference_lora")
log "Running LoRA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_lora_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
+1 -1
View File
@@ -13,4 +13,4 @@
]
}
]
}
}
+1 -1
View File
@@ -372,4 +372,4 @@ jobs:
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
run: python .github/scripts/runpod_cleanup.py
-3
View File
@@ -59,8 +59,5 @@ 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
+2 -2
View File
@@ -60,7 +60,7 @@ repos:
rev: v1.15.0
hooks:
- id: mypy
args: [--python-version, '3.10', --follow-imports, "skip", ]
args: [--python-version, '3.10', --follow-imports, "skip" ]
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
- repo: local
hooks:
@@ -69,7 +69,7 @@ repos:
entry: bash
args:
- -c
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep -v "^fastvideo/v1/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
language: system
always_run: true
pass_filenames: false
+5 -5
View File
@@ -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=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('--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('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
args = parser.parse_args()
main(args)
+13 -3
View File
@@ -6,7 +6,7 @@ def main():
# Initialize VideoGenerator with the Wan model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=2,
num_gpus=1,
lora_path="benjamin-paine/steamboat-willie-1.3b",
lora_nickname="steamboat"
)
@@ -16,6 +16,7 @@ def main():
"num_frames": 81,
"guidance_scale": 5.0,
"num_inference_steps": 32,
"seed": 42,
}
# Generate video with LoRA style
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
@@ -29,8 +30,17 @@ def main():
negative_prompt=negative_prompt,
**kwargs
)
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
del generator
# Until FSDP resharding bug is fixed, multi-lora requires reloading the model
# see https://github.com/pytorch/pytorch/issues/157209
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
lora_path="motimalu/wan-flat-color-1.3b-v2",
lora_nickname="flat_color"
)
# generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
video = generator.generate_video(
@@ -24,6 +24,7 @@ training_args=(
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -1,13 +1,11 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
@@ -60,6 +58,7 @@ training_args=(
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -24,6 +24,7 @@ training_args=(
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -1,13 +1,11 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
@@ -60,6 +58,7 @@ training_args=(
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -26,24 +26,6 @@
"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
}
]
}
@@ -24,6 +24,7 @@ training_args=(
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -1,13 +1,11 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=t2v_output/t2v_%j.out
#SBATCH --error=t2v_output/t2v_%j.err
@@ -57,6 +55,7 @@ training_args=(
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
+1 -1
View File
@@ -62,6 +62,7 @@ SystemEnv = namedtuple(
DEFAULT_CONDA_PATTERNS = {
"torch",
"numpy",
"mypy"
"cudatoolkit",
"soumith",
"mkl",
@@ -80,7 +81,6 @@ DEFAULT_CONDA_PATTERNS = {
DEFAULT_PIP_PATTERNS = {
"torch",
"numpy",
"mypy",
"flake8",
"triton",
"optree",
+1 -1
View File
@@ -4,7 +4,7 @@
"use_cpu_offload": false,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_precision": "fp32",
"vae_tiling": true,
"vae_sp": true,
"vae_config": {
+3 -3
View File
@@ -11,9 +11,9 @@ from fastvideo.v1.platforms import AttentionBackendEnum
class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_reverse_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
@@ -31,7 +31,7 @@ class HunyuanVideoArchConfig(DiTArchConfig):
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
_param_names_mapping: dict = field(
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
@@ -146,8 +146,8 @@ class HunyuanVideoArchConfig(DiTArchConfig):
r"final_layer.linear.\1",
})
# Reverse mapping for saving checkpoints: training -> diffusers
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: int = 2
patch_size_t: int = 1
@@ -10,7 +10,7 @@ class StepVideoArchConfig(DiTArchConfig):
default_factory=lambda:
[lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit()])
_param_names_mapping: dict = field(
param_names_mapping: dict = field(
default_factory=lambda: {
# transformer block
r"^transformer_blocks\.(\d+)\.norm1\.(weight|bias)$":
+4 -4
View File
@@ -12,7 +12,7 @@ def is_blocks(n: str, m) -> bool:
class WanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
@@ -52,12 +52,12 @@ class WanVideoArchConfig(DiTArchConfig):
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Reverse mapping for saving checkpoints: training -> diffusers
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
_lora_param_names_mapping: dict = field(
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
+2 -2
View File
@@ -62,11 +62,11 @@ class PipelineConfig:
image_encoder_precision: str = "fp32"
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp32", )
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
default_factory=lambda: ("fp32", ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor],
+1 -4
View File
@@ -6,8 +6,7 @@ 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 (Wan2_1_Fun_1_3B_InP_SamplingParam,
WanI2V_14B_480P_SamplingParam,
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam)
@@ -24,8 +23,6 @@ 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
}
-17
View File
@@ -94,20 +94,3 @@ 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
+8 -3
View File
@@ -70,7 +70,7 @@ class VideoGenerator:
"""
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
return cls.from_fastvideo_args(fastvideo_args)
@@ -109,6 +109,7 @@ class VideoGenerator:
prompt: The prompt to use for generation
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
output_path: Path to save the video (overrides the one in fastvideo_args)
output_video_name: Name of the video file to save. Default is the first 100 characters of the prompt.
save_video: Whether to save the video to disk
return_frames: Whether to return the raw frames
num_inference_steps: Number of denoising steps (overrides fastvideo_args)
@@ -228,6 +229,7 @@ class VideoGenerator:
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
output_video_name=kwargs.get("output_video_name", prompt[:100]),
)
# Run inference
@@ -251,7 +253,8 @@ class VideoGenerator:
output_path = batch.output_path
if output_path:
os.makedirs(output_path, exist_ok=True)
video_path = os.path.join(output_path, f"{prompt[:100]}.mp4")
video_path = os.path.join(output_path,
f"{batch.output_video_name}.mp4")
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", video_path)
else:
@@ -267,7 +270,9 @@ class VideoGenerator:
"generation_time": gen_time
}
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
self.executor.set_lora_adapter(lora_nickname, lora_path)
def shutdown(self):
+2 -68
View File
@@ -6,7 +6,7 @@ import argparse
import dataclasses
from contextlib import contextmanager
from dataclasses import field
from typing import Any, List
from typing import Any
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.v1.logger import init_logger
@@ -78,8 +78,6 @@ class FastVideoArgs:
# Stage verification
enable_stage_verification: bool = True
denoising_step_list: List[int] | None = field(default=None)
@property
def training_mode(self) -> bool:
@@ -256,12 +254,6 @@ 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)
@@ -288,7 +280,7 @@ class FastVideoArgs:
return cls(**kwargs) # type: ignore
@classmethod
def from_kwargs(cls, kwargs: dict[str, Any]) -> "FastVideoArgs":
def from_kwargs(cls, **kwargs: Any) -> "FastVideoArgs":
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs)
@@ -428,7 +420,6 @@ 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
@@ -464,17 +455,6 @@ 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":
@@ -634,9 +614,6 @@ 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,
@@ -744,48 +721,5 @@ 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(",")]
+45 -36
View File
@@ -3,9 +3,12 @@
import torch
from torch import nn
from torch.distributed.tensor import DTensor, distribute_tensor
from torch.distributed._composable.fsdp import (CPUOffloadPolicy, OffloadPolicy,
fully_shard)
from torch.distributed.tensor import DTensor
from fastvideo.v1.distributed import (get_tp_rank, split_tensor_along_last_dim,
from fastvideo.v1.distributed import (get_local_torch_device, get_tp_rank,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.linear import (ColumnParallelLinear, LinearBase,
@@ -13,6 +16,7 @@ from fastvideo.v1.layers.linear import (ColumnParallelLinear, LinearBase,
QKVParallelLinear, ReplicatedLinear,
RowParallelLinear)
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.utils import get_mixed_precision_state
class BaseLayerWithLoRA(nn.Module):
@@ -26,12 +30,11 @@ class BaseLayerWithLoRA(nn.Module):
self.lora_A: torch.Tensor = None
self.lora_B: torch.Tensor = None
self.merged: bool = False
self.weight = base_layer.weight
self.cpu_weight = base_layer.weight.to("cpu")
self.unmerge_count = 0
# indicates adapter weights don't contain this layer
# (which shouldn't normally happen, but we want to separate it from the case of erroneous merging)
self.disable_lora: bool = False
self.lora_path: str | None = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.base_layer.forward(x)
@@ -45,12 +48,14 @@ class BaseLayerWithLoRA(nn.Module):
def set_lora_weights(self,
A: torch.Tensor,
B: torch.Tensor,
training_mode: bool = False) -> None:
training_mode: bool = False,
lora_path: str | None = None) -> None:
self.lora_A = A # share storage with weights in the pipeline
self.lora_B = B
self.disable_lora = False
if not training_mode:
self.merge_lora_weights()
self.lora_path = lora_path
@torch.no_grad()
def merge_lora_weights(self) -> None:
@@ -58,27 +63,44 @@ class BaseLayerWithLoRA(nn.Module):
return
if self.merged:
raise ValueError(
"LoRA weights already merged. Please unmerge them first.")
self.unmerge_lora_weights()
assert self.lora_A is not None and self.lora_B is not None, "LoRA weights not set. Please set them first."
if isinstance(self.base_layer.weight, DTensor):
mesh = self.base_layer.weight.data.device_mesh
placements = self.base_layer.weight.data.placements
# Using offload param is on CPU, so current_device is for "CPU -> GPU -> merge -> CPU"
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}").full_tensor()
data += (self.slice_lora_b_weights(self.lora_B)
@ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh,
placements=placements).to(current_device))
get_local_torch_device()).full_tensor()
data += (self.slice_lora_b_weights(self.lora_B).to(data)
@ self.slice_lora_a_weights(self.lora_A).to(data))
# Must re-register updated weights for FSDP to recognize them
self.base_layer.weight = nn.Parameter(data.to(current_device))
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
self.base_layer.bias = nn.Parameter(
self.base_layer.bias.to(
get_local_torch_device(),
non_blocking=True).full_tensor().to(current_device))
offload_policy = CPUOffloadPolicy() if "cpu" in str(
current_device) else OffloadPolicy()
# see https://github.com/pytorch/torchtune/pull/2714/files#diff-909ee7ef184b0d834c40a1980ca4149afc38612ec7a4b344d8e2fc27641758c9R69-R79
# After the 1st forward, self.base_layer becomes a FSDP module and needs to be resharded
if hasattr(self.base_layer, "unshard"):
self.base_layer.unshard()
mp_policy = get_mixed_precision_state().mp_policy
fully_shard(self.base_layer,
mesh=mesh,
mp_policy=mp_policy,
offload_policy=offload_policy)
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.to(
f"cuda:{torch.cuda.current_device()}")
data = self.base_layer.weight.data.to(get_local_torch_device())
data += \
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight = nn.Parameter(data.to(current_device))
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
self.base_layer.weight.data = data.to(current_device,
non_blocking=True)
self.merged = True
@torch.no_grad()
@@ -90,28 +112,15 @@ class BaseLayerWithLoRA(nn.Module):
raise ValueError(
"LoRA weights not merged. Please merge them first before unmerging."
)
self.unmerge_count += 1
# Avoid precision loss
if self.unmerge_count % 3 == 0:
# To avoid precision loss we do not subtract the LoRA weights here
if isinstance(self.base_layer.weight, DTensor):
device = self.base_layer.weight.data.device
self.base_layer.weight = nn.Parameter(self.cpu_weight.to(device))
else:
self.base_layer.weight.data = self.cpu_weight.data.to(
self.base_layer.weight)
if isinstance(self.base_layer.weight, DTensor):
mesh = self.base_layer.weight.data.device_mesh
placement = self.base_layer.weight.data.placements
device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}").full_tensor()
data -= self.slice_lora_b_weights(
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
self.base_layer.weight = nn.Parameter(
distribute_tensor(data, mesh, placements=placement).to(device))
else:
self.base_layer.weight.data -= \
self.slice_lora_b_weights(self.lora_B) @\
self.slice_lora_a_weights(self.lora_A)
self.merged = False
+6 -6
View File
@@ -13,8 +13,8 @@ from fastvideo.v1.platforms import AttentionBackendEnum
class BaseDiT(nn.Module, ABC):
_fsdp_shard_conditions: list = []
_compile_conditions: list = []
_param_names_mapping: dict
_reverse_param_names_mapping: dict
param_names_mapping: dict
reverse_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
num_channels_latents: int
@@ -24,7 +24,7 @@ class BaseDiT(nn.Module, ABC):
def __init_subclass__(cls) -> None:
required_class_attrs = [
"_fsdp_shard_conditions", "_param_names_mapping",
"_fsdp_shard_conditions", "param_names_mapping",
"_compile_conditions"
]
super().__init_subclass__()
@@ -78,9 +78,9 @@ class CachableDiT(BaseDiT):
"""
# These are required class attributes that should be overridden by concrete implementations
_fsdp_shard_conditions = []
_param_names_mapping = {}
_reverse_param_names_mapping = {}
_lora_param_names_mapping: dict = {}
param_names_mapping = {}
reverse_param_names_mapping = {}
lora_param_names_mapping: dict = {}
# Ensure these instance attributes are properly defined in subclasses
hidden_size: int
num_attention_heads: int
+4 -4
View File
@@ -441,10 +441,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
_compile_conditions = HunyuanVideoConfig()._compile_conditions
_supported_attention_backends = HunyuanVideoConfig(
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
_reverse_param_names_mapping = HunyuanVideoConfig(
)._reverse_param_names_mapping
_lora_param_names_mapping = HunyuanVideoConfig()._lora_param_names_mapping
param_names_mapping = HunyuanVideoConfig().param_names_mapping
reverse_param_names_mapping = HunyuanVideoConfig(
).reverse_param_names_mapping
lora_param_names_mapping = HunyuanVideoConfig().lora_param_names_mapping
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
+7 -5
View File
@@ -457,11 +457,13 @@ class StepVideoTransformerBlock(nn.Module):
class StepVideoModel(BaseDiT):
# (Optional) Keep the same attribute for compatibility with splitting, etc.
_fsdp_shard_conditions = StepVideoConfig()._fsdp_shard_conditions
_param_names_mapping = StepVideoConfig()._param_names_mapping
_reverse_param_names_mapping = StepVideoConfig(
)._reverse_param_names_mapping
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
_fsdp_shard_conditions = [
lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit(),
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
param_names_mapping = StepVideoConfig().param_names_mapping
reverse_param_names_mapping = StepVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = StepVideoConfig().lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
+3 -3
View File
@@ -515,9 +515,9 @@ class WanTransformer3DModel(CachableDiT):
_compile_conditions = WanVideoConfig()._compile_conditions
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
_reverse_param_names_mapping = WanVideoConfig()._reverse_param_names_mapping
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
param_names_mapping = WanVideoConfig().param_names_mapping
reverse_param_names_mapping = WanVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = WanVideoConfig().lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
@@ -72,8 +72,6 @@ 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"),
@@ -431,7 +429,6 @@ class TransformerLoader(ComponentLoader):
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
cpu_offload=fastvideo_args.use_cpu_offload,
fsdp_inference=fastvideo_args.use_fsdp_inference,
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
+18 -40
View File
@@ -5,7 +5,6 @@
# Copyright 2025 The FastVideo Authors.
import contextlib
from collections import defaultdict
from collections.abc import Callable, Generator
from itertools import chain
from typing import Any
@@ -19,7 +18,8 @@ from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.utils import get_param_names_mapping
from fastvideo.v1.models.loader.utils import (get_param_names_mapping,
hf_to_custom_state_dict)
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.utils import set_mixed_precision_policy
@@ -62,7 +62,6 @@ def maybe_load_fsdp_model(
device: torch.device,
hsdp_replicate_dim: int,
hsdp_shard_dim: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
@@ -81,12 +80,14 @@ def maybe_load_fsdp_model(
output_dtype,
cast_forward_inputs=False)
set_mixed_precision_policy(master_dtype=default_dtype,
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
output_dtype=output_dtype)
set_mixed_precision_policy(
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
output_dtype=output_dtype,
mp_policy=mp_policy,
)
with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(param_dtype), torch.device("meta"):
model = model_cls(**init_params)
world_size = hsdp_replicate_dim * hsdp_shard_dim
if not training_mode and not fsdp_inference:
@@ -106,9 +107,8 @@ def maybe_load_fsdp_model(
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory)
weight_iterator = safetensors_weights_iterator(weight_dir_list,
to_cpu=cpu_offload)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
load_model_from_full_model_state_dict(
model,
weight_iterator,
@@ -233,36 +233,14 @@ def load_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sd = model.state_dict()
# Find new params
used_keys = set()
sharded_sd = {}
to_merge_params: defaultdict[str, dict[Any, Any]] = defaultdict(dict)
reverse_param_names_mapping = {}
assert param_names_mapping is not None
for source_param_name, full_tensor in full_sd_iterator:
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
used_keys.add(target_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
# cat at output dim according to the merge_index order
sorted_tensors = [
to_merge_params[target_param_name][i]
for i in range(num_params_to_merge)
]
full_tensor = torch.cat(sorted_tensors, dim=0)
del to_merge_params[target_param_name]
else:
continue
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
full_sd_iterator, param_names_mapping) # type: ignore
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
raise ValueError(
f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict"
f"Parameter {target_param_name} not found in custom model state dict. The hf to custom mapping may be incorrect."
)
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
@@ -279,10 +257,10 @@ def load_model_from_full_model_state_dict(
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
model._reverse_param_names_mapping = reverse_param_names_mapping
unused_keys = set(meta_sd.keys()) - used_keys
model.reverse_param_names_mapping = reverse_param_names_mapping
unused_keys = set(meta_sd.keys()) - set(sharded_sd.keys())
if unused_keys:
logger.warning("Found new parameters in meta state dict: %s",
logger.warning("Found unloaded parameters in meta state dict: %s",
unused_keys)
# List of allowed parameter name patterns
+45 -3
View File
@@ -2,7 +2,8 @@
"""Utilities for selecting and loading models."""
import contextlib
import re
from collections.abc import Callable
from collections import defaultdict
from collections.abc import Callable, Iterator
from typing import Any
import torch
@@ -35,7 +36,6 @@ def get_param_names_mapping(
"""
def mapping_fn(name: str) -> tuple[str, Any, Any]:
# Try to match and transform the name using the regex patterns in mapping_dict
for pattern, replacement in mapping_dict.items():
match = re.match(pattern, name)
@@ -52,4 +52,46 @@ def get_param_names_mapping(
# If no pattern matches, return the original name
return name, None, None
return mapping_fn
return mapping_fn
def hf_to_custom_state_dict(
hf_param_sd: dict[str, torch.Tensor] | Iterator[tuple[str, torch.Tensor]],
param_names_mapping: Callable[[str], tuple[str, Any, Any]]
) -> tuple[dict[str, torch.Tensor], dict[str, tuple[str, Any, Any]]]:
"""
Converts a Hugging Face parameter state dictionary to a custom parameter state dictionary.
Args:
hf_param_sd (Dict[str, torch.Tensor]): The Hugging Face parameter state dictionary
param_names_mapping (Callable[[str], tuple[str, Any, Any]]): A function that maps parameter names from source to target format
Returns:
custom_param_sd (Dict[str, torch.Tensor]): The custom formatted parameter state dict
reverse_param_names_mapping (Dict[str, Tuple[str, Any, Any]]): Maps back from custom to hf
"""
custom_param_sd = {}
to_merge_params = defaultdict(dict) # type: ignore
reverse_param_names_mapping = {}
if isinstance(hf_param_sd, dict):
hf_param_sd = hf_param_sd.items() # type: ignore
for source_param_name, full_tensor in hf_param_sd: # type: ignore
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
# cat at output dim according to the merge_index order
sorted_tensors = [
to_merge_params[target_param_name][i]
for i in range(num_params_to_merge)
]
full_tensor = torch.cat(sorted_tensors, dim=0)
del to_merge_params[target_param_name]
else:
continue
custom_param_sd[target_param_name] = full_tensor
return custom_param_sd, reverse_param_names_mapping
@@ -18,14 +18,15 @@
# Modified from diffusers==0.29.2
#
# ==============================================================================
import math
from dataclasses import dataclass
from typing import Any, Optional, Tuple, Union, List
import numpy as np
from typing import Any
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, is_scipy_available, logging
from diffusers.utils import BaseOutput
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.base import BaseScheduler
@@ -33,7 +34,7 @@ logger = init_logger(__name__)
@dataclass
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
@@ -45,7 +46,8 @@ class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
Euler scheduler.
@@ -55,37 +57,16 @@ class FlowMatchEulerDiscreteScheduler(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.
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.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
"""
_compatibles = []
_compatibles: list[Any] = []
order = 1
@register_to_config
@@ -93,53 +74,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
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,
reverse: bool = True,
solver: str = "euler",
n_tokens: int | None = None,
**kwargs,
):
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:
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(
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
if time_shift_type not in {"exponential", "linear"}:
raise ValueError("`time_shift_type` must either be 'exponential' or 'linear'.")
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
BaseScheduler.__init__(self)
@property
def step_index(self):
@@ -166,190 +125,44 @@ class FlowMatchEulerDiscreteScheduler(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: Optional[int] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[float] = None,
timesteps: Optional[List[float]] = None,
num_inference_steps: int,
device: 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`, *optional*):
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.
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.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
"""
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
# 1. Prepare default sigmas
is_timesteps_provided = timesteps is not None
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
if is_timesteps_provided:
timesteps = np.array(timesteps).astype(np.float32)
if not self.config.reverse:
sigmas = 1 - sigmas
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:
sigmas = np.array(sigmas).astype(np.float32)
num_inference_steps = len(sigmas)
# 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
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
self._begin_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
@@ -361,9 +174,17 @@ class FlowMatchEulerDiscreteScheduler(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
return indices[pos].item()
idx: int = indices[pos].item()
def _init_step_index(self, timestep):
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)
@@ -371,19 +192,22 @@ class FlowMatchEulerDiscreteScheduler(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,
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]:
**kwargs,
) -> 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).
@@ -395,38 +219,25 @@ class FlowMatchEulerDiscreteScheduler(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.
per_token_timesteps (`torch.Tensor`, *optional*):
The timesteps for each token in the sample.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
return_dict (`bool`):
Whether or not to return a
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or tuple.
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~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.
[`~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)
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 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)
@@ -434,454 +245,24 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
if per_token_timesteps is not None:
per_token_sigmas = per_token_timesteps / self.config.num_train_timesteps
assert self.step_index is not None
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
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
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
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
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 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 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
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_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,5 +772,47 @@ 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
@@ -113,7 +113,7 @@ class ComposedPipelineBase(ABC):
if args is None or args.inference_mode:
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
@@ -222,13 +222,12 @@ 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:
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']
raise ValueError(
f"model_index.json must contain a {module_name} module")
# all the component models used by the pipeline
required_modules = self.required_config_modules
logger.info("Loading required modules: %s", required_modules)
@@ -243,11 +242,7 @@ class ComposedPipelineBase(ABC):
logger.info("Using module %s already provided", module_name)
modules[module_name] = loaded_modules[module_name]
continue
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)
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
+16 -9
View File
@@ -7,6 +7,7 @@ import torch
import torch.distributed as dist
from safetensors.torch import load_file
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.layers.lora.linear import (BaseLayerWithLoRA, get_lora_layer,
replace_submodule)
@@ -29,13 +30,13 @@ class LoRAPipeline(ComposedPipelineBase):
lora_layers: dict[str, BaseLayerWithLoRA] = {}
fastvideo_args: FastVideoArgs
exclude_lora_layers: list[str] = []
device: torch.device = torch.device(f"cuda:{torch.cuda.current_device()}")
device: torch.device | None = None
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.exclude_lora_layers = self.modules[
"transformer"].config.arch_config.exclude_lora_layers
self.device = get_local_torch_device()
self.convert_to_lora_layers()
if self.fastvideo_args.pipeline_config.lora_path is not None:
self.set_lora_adapter(
@@ -53,7 +54,6 @@ class LoRAPipeline(ComposedPipelineBase):
"""
Converts the transformer to a LoRA transformer.
"""
for name, layer in self.modules["transformer"].named_modules():
if not self.is_target_layer(name):
continue
@@ -85,16 +85,18 @@ class LoRAPipeline(ComposedPipelineBase):
raise ValueError(
f"Adapter {lora_nickname} not found in the pipeline. Please provide lora_path to load it."
)
adapter_updated = False
rank = dist.get_rank()
if lora_path is not None:
lora_local_path = maybe_download_lora(lora_path)
lora_state_dict = load_file(lora_local_path)
lora_state_dict = load_file(lora_local_path,
device=str(self.device))
# Map the hf layer names to our custom layer names
param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"]._param_names_mapping)
self.modules["transformer"].param_names_mapping)
lora_param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"]._lora_param_names_mapping)
self.modules["transformer"].lora_param_names_mapping)
to_merge_params: defaultdict[Hashable,
dict[Any, Any]] = defaultdict(dict)
@@ -119,6 +121,11 @@ class LoRAPipeline(ComposedPipelineBase):
del to_merge_params[target_name]
else:
continue
if target_name in self.lora_adapters[lora_nickname]:
raise ValueError(
f"Target name {target_name} already exists in lora_adapters[{lora_nickname}]"
)
self.lora_adapters[lora_nickname][target_name] = weight.to(
self.device)
adapter_updated = True
@@ -134,12 +141,11 @@ class LoRAPipeline(ComposedPipelineBase):
lora_B_name = name + ".lora_B"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
if layer.merged:
layer.unmerge_lora_weights()
layer.set_lora_weights(
self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name],
training_mode=self.fastvideo_args.training_mode)
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
else:
if rank == 0:
@@ -149,4 +155,5 @@ class LoRAPipeline(ComposedPipelineBase):
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
self.cur_adapter_name = lora_nickname
+1 -27
View File
@@ -18,12 +18,6 @@ from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
WanTeaCacheParams)
__all__ = [
"ForwardBatch",
"TrainingBatch",
"AttentionMetadata",
"VideoSparseAttentionMetadata",
]
@dataclass
class ForwardBatch:
@@ -52,7 +46,7 @@ class ForwardBatch:
negative_prompt: str | list[str] | None = None
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
# Primary encoder embeddings
prompt_embeds: list[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: list[torch.Tensor] | None = None
@@ -152,11 +146,9 @@ 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
@@ -164,7 +156,6 @@ 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
@@ -172,7 +163,6 @@ 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
@@ -184,19 +174,3 @@ 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,7 +10,6 @@ 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
@@ -29,7 +28,6 @@ __all__ = [
"LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
@@ -103,7 +103,6 @@ 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
+2 -234
View File
@@ -3,7 +3,7 @@
Denoising stage for diffusion pipelines.
"""
import inspect, copy
import inspect
from collections.abc import Iterable
from typing import Any
@@ -27,7 +27,6 @@ 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 (
@@ -118,7 +117,6 @@ 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")
@@ -178,7 +176,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:
@@ -543,233 +541,3 @@ 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
@@ -81,7 +81,6 @@ class TextEncodingStage(PipelineStage):
output_hidden_states=True,
)
prompt_embeds = postprocess_func(outputs)
batch.prompt_embeds.append(prompt_embeds)
if batch.prompt_attention_mask is not None:
batch.prompt_attention_mask.append(attention_mask)
@@ -14,7 +14,6 @@ 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__)
@@ -1,70 +0,0 @@
# 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
@@ -1,79 +0,0 @@
# 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
View File
@@ -0,0 +1,195 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import pytest
from fastvideo import VideoGenerator
from fastvideo.v1.logger import init_logger
from fastvideo.v1.tests.utils import compute_video_ssim_torchvision, write_ssim_results
from diffusers import DiffusionPipeline
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines import build_pipeline
from fastvideo.v1.models.loader.utils import hf_to_custom_state_dict, get_param_names_mapping
from torch.testing import assert_close
from torch.distributed.tensor import DTensor
from fastvideo.v1.worker import MultiprocExecutor
import torch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29500"
# Base parameters for LoRA inference tests
WAN_LORA_PARAMS = {
"num_gpus": 1,
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 32,
"guidance_scale": 5.0,
"flow_shift": 3.0,
"seed": 42,
"fps": 24,
"neg_prompt": "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
"text-encoder-precision": ("fp32",),
"use_cpu_offload": True,
}
# LoRA configurations for testing
LORA_CONFIGS = [
{
"lora_path": "benjamin-paine/steamboat-willie-1.3b",
"lora_nickname": "steamboat",
"prompt": "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image.",
"negative_prompt": "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
"ssim_threshold": 0.79
},
# {
# "lora_path": "motimalu/wan-flat-color-1.3b-v2",
# "lora_nickname": "flat_color",
# "prompt": "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid.",
# "negative_prompt": "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
# "ssim_threshold": 0.79
# }
]
MODEL_TO_PARAMS = {
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WAN_LORA_PARAMS,
}
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_merge_lora_weights(model_id):
lora_config = LORA_CONFIGS[0] # test only one
hf_pipe = DiffusionPipeline.from_pretrained(model_id, torch_dtype=torch.bfloat16)
hf_pipe.enable_model_cpu_offload()
lora_nickname = lora_config["lora_nickname"]
lora_path = lora_config["lora_path"]
args = FastVideoArgs.from_kwargs(
model_path=model_id,
use_cpu_offload=True,
dit_precision="bf16",
)
pipe = build_pipeline(args)
pipe.set_lora_adapter(lora_nickname, lora_path)
custom_transformer = pipe.modules["transformer"]
custom_state_dict = custom_transformer.state_dict()
hf_pipe.load_lora_weights(lora_path, adapter_name=lora_nickname)
for name, layer in hf_pipe.transformer.named_modules():
if hasattr(layer, "unmerge"):
layer.unmerge()
layer.merge(adapter_names=[lora_nickname])
hf_transformer = hf_pipe.transformer
param_names_mapping = get_param_names_mapping(custom_transformer.param_names_mapping)
hf_state_dict, _ = hf_to_custom_state_dict(hf_transformer.state_dict(), param_names_mapping)
for key in hf_state_dict.keys():
if "base_layer" not in key:
continue
hf_param = hf_state_dict[key]
custom_param = custom_state_dict[key].to_local() if isinstance(custom_state_dict[key], DTensor) else custom_state_dict[key]
assert_close(hf_param, custom_param, atol=7e-4, rtol=7e-4)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
"""
Test that runs LoRA inference with LoRA switching and compares the output
to reference videos using SSIM.
"""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path.join(script_dir, 'generated_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
os.makedirs(output_dir, exist_ok=True)
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
num_inference_steps = BASE_PARAMS["num_inference_steps"]
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"flow_shift": BASE_PARAMS["flow_shift"],
"use_cpu_offload": BASE_PARAMS["use_cpu_offload"],
}
if "text-encoder-precision" in BASE_PARAMS:
init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
generation_kwargs = {
"num_inference_steps": num_inference_steps,
"output_path": output_dir,
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"guidance_scale": BASE_PARAMS["guidance_scale"],
"seed": BASE_PARAMS["seed"],
"fps": BASE_PARAMS["fps"],
"save_video": True,
}
generator = VideoGenerator.from_pretrained(model_path=BASE_PARAMS["model_path"], **init_kwargs)
for lora_config in LORA_CONFIGS:
lora_nickname = lora_config["lora_nickname"]
lora_path = lora_config["lora_path"]
prompt = lora_config["prompt"]
generation_kwargs["negative_prompt"] = lora_config["negative_prompt"]
generator.set_lora_adapter(lora_nickname=lora_nickname, lora_path=lora_path)
output_video_name = f"{lora_path.split('/')[-1]}_{prompt[:50]}"
generation_kwargs["output_path"] = output_dir
generation_kwargs["output_video_name"] = output_video_name
generator.generate_video(prompt, **generation_kwargs)
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, 'L40S_reference_videos', model_id.split('/')[-1], ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
# Find the matching reference video for the switched LoRA
reference_video_name = None
for filename in os.listdir(reference_folder):
# Check if the filename starts with the expected output_video_name and ends with .mp4
if filename.startswith(output_video_name) and filename.endswith('.mp4'):
reference_video_name = filename # Remove .mp4 extension to match the logic below
break
if not reference_video_name:
logger.error(f"Reference video not found for adapter: {lora_path} with prompt: {prompt[:50]} and backend: {ATTENTION_BACKEND}")
raise FileNotFoundError(f"Reference video missing for adapter {lora_path}")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name + ".mp4")
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
)
ssim_values = compute_video_ssim_torchvision(reference_video_path,
generated_video_path,
use_ms_ssim=True)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
logger.info(f"Writing SSIM results to directory: {output_dir}")
success = write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path, num_inference_steps,
prompt)
if not success:
logger.error("Failed to write SSIM results to file")
min_acceptable_ssim = lora_config["ssim_threshold"]
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} for adapter {lora_config['lora_path']}"
+5
View File
@@ -97,3 +97,8 @@ def run_precision_tests_STA():
@app.function(gpu="H100:1", image=image, timeout=900)
def run_precision_tests_VSA():
run_test("python csrc/attn/tests/test_block_sparse.py")
@app.function(gpu="L40S:1", image=image, timeout=3600)
def run_inference_lora_tests():
run_test("pytest ./fastvideo/v1/tests/inference/lora/test_lora_inference_similarity.py -vs")
@@ -7,7 +7,7 @@ import pytest
from fastvideo import VideoGenerator
from fastvideo.v1.logger import init_logger
from fastvideo.v1.tests.ssim.compute_ssim import compute_video_ssim_torchvision
from fastvideo.v1.tests.utils import compute_video_ssim_torchvision, write_ssim_results
from fastvideo.v1.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
@@ -99,44 +99,6 @@ I2V_IMAGE_PATHS = [
]
def write_ssim_results(output_dir, ssim_values, reference_path, generated_path,
num_inference_steps, prompt):
"""
Write SSIM results to a JSON file in the same directory as the generated videos.
"""
try:
logger.info(
f"Attempting to write SSIM results to directory: {output_dir}")
if not os.path.exists(output_dir):
os.makedirs(output_dir, exist_ok=True)
mean_ssim, min_ssim, max_ssim = ssim_values
result = {
"mean_ssim": mean_ssim,
"min_ssim": min_ssim,
"max_ssim": max_ssim,
"reference_video": reference_path,
"generated_video": generated_path,
"parameters": {
"num_inference_steps": num_inference_steps,
"prompt": prompt
}
}
test_name = f"steps{num_inference_steps}_{prompt[:100]}"
result_file = os.path.join(output_dir, f"{test_name}_ssim.json")
logger.info(f"Writing JSON results to: {result_file}")
with open(result_file, 'w') as f:
json.dump(result, f, indent=2)
logger.info(f"SSIM results written to {result_file}")
return True
except Exception as e:
logger.error(f"ERROR writing SSIM results: {str(e)}")
return False
@pytest.mark.parametrize("prompt", I2V_TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN", "TORCH_SDPA"])
@@ -1,15 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import json
from fastvideo.v1.logger import init_logger
import numpy as np
import torch
from pytorch_msssim import ms_ssim, ssim
from torchvision.io import read_video
logger = init_logger(__name__)
def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
"""
Compute SSIM between two videos.
Args:
video1_path: Path to the first video.
video2_path: Path to the second video.
use_ms_ssim: Whether to use Multi-Scale Structural Similarity(MS-SSIM) instead of SSIM.
"""
print(f"Computing SSIM between {video1_path} and {video2_path}...")
if not os.path.exists(video1_path):
raise FileNotFoundError(f"Video1 not found: {video1_path}")
if not os.path.exists(video2_path):
raise FileNotFoundError(f"Video2 not found: {video2_path}")
frames1, _, _ = read_video(video1_path,
pts_unit='sec',
@@ -65,7 +81,26 @@ def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
def compare_folders(reference_folder, generated_folder, use_ms_ssim=True):
"""
Compare videos with the same filename between reference_folder and generated_folder
Example usage:
results = compare_folders(reference_folder, generated_folder,
args.use_ms_ssim)
for video_name, ssim_value in results.items():
if ssim_value is not None:
print(
f"{video_name}: {ssim_value[0]:.4f}, Min SSIM: {ssim_value[1]:.4f}, Max SSIM: {ssim_value[2]:.4f}"
)
else:
print(f"{video_name}: Error during comparison")
valid_ssims = [v for v in results.values() if v is not None]
if valid_ssims:
avg_ssim = np.mean([v[0] for v in valid_ssims])
print(f"\nAverage SSIM across all videos: {avg_ssim:.4f}")
else:
print("\nNo valid SSIM values to average")
"""
reference_videos = [
f for f in os.listdir(reference_folder) if f.endswith('.mp4')
]
@@ -92,54 +127,40 @@ def compare_folders(reference_folder, generated_folder, use_ms_ssim=True):
return results
def write_ssim_results(output_dir, ssim_values, reference_path, generated_path,
num_inference_steps, prompt):
"""
Write SSIM results to a JSON file in the same directory as the generated videos.
"""
try:
logger.info(
f"Attempting to write SSIM results to directory: {output_dir}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description='Compare videos using SSIM/MS-SSIM metrics')
parser.add_argument('--reference',
'-r',
type=str,
help='Path to reference videos directory')
parser.add_argument('--generated',
'-g',
type=str,
help='Path to generated videos directory')
parser.add_argument('--use-ms-ssim',
action='store_true',
help='Use MS-SSIM instead of SSIM')
args = parser.parse_args()
if not os.path.exists(output_dir):
os.makedirs(output_dir, exist_ok=True)
script_dir = os.path.dirname(os.path.abspath(__file__))
mean_ssim, min_ssim, max_ssim = ssim_values
reference_folder = args.reference if args.reference else os.path.join(
script_dir, 'reference_videos')
generated_folder = args.generated if args.generated else os.path.join(
script_dir, 'generated_videos')
result = {
"mean_ssim": mean_ssim,
"min_ssim": min_ssim,
"max_ssim": max_ssim,
"reference_video": reference_path,
"generated_video": generated_path,
"parameters": {
"num_inference_steps": num_inference_steps,
"prompt": prompt
}
}
if not os.path.exists(reference_folder):
print(f"ERROR: Reference folder {reference_folder} does not exist!")
exit(1)
test_name = f"steps{num_inference_steps}_{prompt[:100]}"
result_file = os.path.join(output_dir, f"{test_name}_ssim.json")
logger.info(f"Writing JSON results to: {result_file}")
with open(result_file, 'w') as f:
json.dump(result, f, indent=2)
if not os.path.exists(generated_folder):
print(f"ERROR: Generated folder {generated_folder} does not exist!")
exit(1)
print(f"Comparing videos between {reference_folder} and {generated_folder}")
results = compare_folders(reference_folder, generated_folder,
args.use_ms_ssim)
print("\n===== SSIM Results Summary =====")
for video_name, ssim_value in results.items():
if ssim_value is not None:
print(
f"{video_name}: {ssim_value[0]:.4f}, Min SSIM: {ssim_value[1]:.4f}, Max SSIM: {ssim_value[2]:.4f}"
)
else:
print(f"{video_name}: Error during comparison")
valid_ssims = [v for v in results.values() if v is not None]
if valid_ssims:
avg_ssim = np.mean([v[0] for v in valid_ssims])
print(f"\nAverage SSIM across all videos: {avg_ssim:.4f}")
else:
print("\nNo valid SSIM values to average")
logger.info(f"SSIM results written to {result_file}")
return True
except Exception as e:
logger.error(f"ERROR writing SSIM results: {str(e)}")
return False
+1 -2
View File
@@ -1,5 +1,4 @@
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
from .distillation_pipeline import DistillationPipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
@@ -1,998 +0,0 @@
# 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)
+11 -21
View File
@@ -85,6 +85,14 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
assert self.transformer is not None
self.set_schemas()
# Set random seeds for deterministic training
set_random_seed(self.seed)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.transformer.requires_grad_(True)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
@@ -109,13 +117,8 @@ 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,
@@ -226,8 +229,7 @@ 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.num_train_timesteps).long()
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=training_batch.latents.device)
if self.training_args.sp_size > 1:
@@ -259,7 +261,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],
@@ -325,7 +327,6 @@ 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
@@ -430,12 +431,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
local_main_process_only=False)
assert self.training_args is not None
# Set random seeds for deterministic training
set_random_seed(self.seed)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
if self.training_args.resume_from_checkpoint:
@@ -581,7 +576,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
generator=self.validation_random_generator,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
@@ -604,10 +599,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# 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,
@@ -712,4 +703,3 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
+16 -96
View File
@@ -3,16 +3,14 @@ import json
import math
import os
import time
from typing import Any, Dict
from collections.abc import Iterator
from typing import Any
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)
@@ -21,8 +19,6 @@ 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__)
@@ -167,9 +163,9 @@ def save_checkpoint(transformer,
weight_path,
local_main_process_only=False)
# Convert fastvideo custom format to diffusers format and save
diffusers_state_dict = convert_custom_format_to_diffusers_format(
cpu_state, transformer)
# Convert training format to diffusers format and save
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, transformer.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
logger.info("rank: %s, consolidated checkpoint saved to %s",
@@ -492,24 +488,25 @@ def _has_foreach_support(tensors: list[torch.Tensor],
t is None or type(t) in [torch.Tensor] for t in tensors)
def convert_custom_format_to_diffusers_format(state_dict: dict[str, Any],
transformer) -> dict[str, Any]:
def custom_to_hf_state_dict(
state_dict: dict[str, Any] | Iterator[tuple[str, torch.Tensor]],
reverse_param_names_mapping: dict[str, tuple[str, int,
int]]) -> dict[str, Any]:
"""
Convert fastvideo custom format state dict to diffusers format using reverse_param_names_mapping.
Convert fastvideo's custom model format to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in training format
transformer: Transformer model object with _reverse_param_names_mapping
state_dict: State dict in fastvideo's custom format
reverse_param_names_mapping: Reverse mapping from fastvideo's custom format to diffusers format
Returns:
State dict in diffusers format
"""
assert len(
reverse_param_names_mapping) > 0, "reverse_param_names_mapping is empty"
if isinstance(state_dict, Iterator):
state_dict = dict(state_dict)
new_state_dict = {}
# Get the reverse mapping from the transformer
reverse_param_names_mapping = transformer._reverse_param_names_mapping
assert reverse_param_names_mapping != {}, "reverse_param_names_mapping is empty"
# Group parameters that need to be split (merged parameters)
merge_groups: dict[str, list[tuple[str, int, int]]] = {}
@@ -555,80 +552,3 @@ 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)
@@ -1,95 +0,0 @@
# 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)
@@ -1,231 +0,0 @@
# 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)
+9 -6
View File
@@ -29,6 +29,7 @@ from diffusers.loaders.lora_base import (
_best_guess_weight_name) # watch out for potetential removal from diffusers
from huggingface_hub import snapshot_download
from remote_pdb import RemotePdb
from torch.distributed.fsdp import MixedPrecisionPolicy
import fastvideo.v1.envs as envs
from fastvideo.v1.logger import init_logger
@@ -684,11 +685,11 @@ def remote_breakpoint() -> None:
@dataclass
class MixedPrecisionState:
master_dtype: torch.dtype | None = None
param_dtype: torch.dtype | None = None
reduce_dtype: torch.dtype | None = None
output_dtype: torch.dtype | None = None
compute_dtype: torch.dtype | None = None
mp_policy: MixedPrecisionPolicy | None = None
# Thread-local storage for mixed precision state
@@ -702,10 +703,12 @@ def get_mixed_precision_state() -> MixedPrecisionState:
return cast(MixedPrecisionState, _mixed_precision_state.state)
def set_mixed_precision_policy(master_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
output_dtype: torch.dtype | None = None):
def set_mixed_precision_policy(
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
output_dtype: torch.dtype | None = None,
mp_policy: MixedPrecisionPolicy | None = None,
):
"""Set mixed precision policy globally.
Args:
@@ -714,10 +717,10 @@ def set_mixed_precision_policy(master_dtype: torch.dtype,
output_dtype: Optional output dtype
"""
state = MixedPrecisionState(
master_dtype=master_dtype,
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
output_dtype=output_dtype,
mp_policy=mp_policy,
)
_mixed_precision_state.state = state
+5
View File
@@ -0,0 +1,5 @@
from .executor import Executor
from .gpu_worker import run_worker_process
from .multiproc_executor import MultiprocExecutor
__all__ = ["Executor", "run_worker_process", "MultiprocExecutor"]
+3 -1
View File
@@ -49,7 +49,9 @@ class Executor(ABC):
return cast(ForwardBatch, outputs[0]["output_batch"])
@abstractmethod
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
"""
Set the LoRA adapter for the workers.
"""
+10 -1
View File
@@ -87,7 +87,9 @@ class Worker:
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
return cast(ForwardBatch, output_batch)
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
self.pipeline.set_lora_adapter(lora_nickname, lora_path)
def shutdown(self) -> dict[str, Any]:
@@ -132,6 +134,13 @@ class Worker:
output_batch = self.execute_forward(forward_batch,
fastvideo_args)
self.pipe.send({"output_batch": output_batch.output.cpu()})
elif method_name == 'set_lora_adapter':
lora_nickname = recv_rpc['kwargs']['lora_nickname']
lora_path = recv_rpc['kwargs']['lora_path']
self.set_lora_adapter(lora_nickname, lora_path)
logger.info("Worker %d set LoRA adapter %s with path %s",
self.rank, lora_nickname, lora_path)
self.pipe.send({"status": "lora_adapter_set"})
else:
# Handle other methods dynamically if needed
args = recv_rpc.get('args', ())
+12 -6
View File
@@ -75,12 +75,18 @@ class MultiprocExecutor(Executor):
})
return cast(ForwardBatch, responses[0]["output_batch"])
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
self.collective_rpc("set_lora_adapter",
kwargs={
"lora_nickname": lora_nickname,
"lora_path": lora_path
})
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
responses = self.collective_rpc("set_lora_adapter",
kwargs={
"lora_nickname": lora_nickname,
"lora_path": lora_path
})
for i, response in enumerate(responses):
if response["status"] != "lora_adapter_set":
raise RuntimeError(
f"Worker {i} failed to set LoRA adapter to {lora_path}")
def collective_rpc(self,
method: str | Callable,
+2
View File
@@ -1,3 +1,4 @@
# test pr
[build-system]
requires = ["setuptools>=61.0"]
build-backend = "setuptools.build_meta"
@@ -78,6 +79,7 @@ exclude = ["assets*", "docker*", "docs", "scripts*"]
[tool.wheel]
exclude = ["assets*", "docker*", "docs", "scripts*"]
[tool.mypy]
warn_unused_configs = true
ignore_missing_imports = true
-62
View File
@@ -1,62 +0,0 @@
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 \
-67
View File
@@ -1,67 +0,0 @@
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
-66
View File
@@ -1,66 +0,0 @@
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 \
-66
View File
@@ -1,66 +0,0 @@
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 \
-66
View File
@@ -1,66 +0,0 @@
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 \
+2 -2
View File
@@ -10,8 +10,8 @@ fastvideo generate \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--height 768 \
--width 1280\
--height 448 \
--width 832 \
--num-frames 77 \
--num-inference-steps 50 \
--fps 16 \
-113
View File
@@ -1,113 +0,0 @@
import os
import unittest
import torch
from transformers import AutoTokenizer, T5EncoderModel
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
class TestAutoencoderKLCausal3D(unittest.TestCase):
@classmethod
def setUpClass(cls):
"""
setUpClass is called once, before any test is run.
We can set environment variables or load heavy resources here.
"""
os.environ["HF_HUB_DISABLE_PROGRESS_BARS"] = "1"
# Load tokenizer/model that can be reused across all tests
cls.tokenizer = AutoTokenizer.from_pretrained("hf-internal-testing/tiny-random-t5")
cls.text_encoder = T5EncoderModel.from_pretrained("hf-internal-testing/tiny-random-t5")
def setUp(self):
"""
setUp is called before each test method to prepare fresh state.
"""
self.batch_size = 1
self.init_time_len = 9
self.init_height = 16
self.init_width = 16
self.latent_channels = 4
self.spatial_compression_ratio = 8
self.time_compression_ratio = 4
# Model initialization config
self.init_dict = {
"in_channels":
3,
"out_channels":
3,
"latent_channels":
self.latent_channels,
"down_block_types": (
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
),
"up_block_types": (
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
),
"block_out_channels": (8, 8, 8, 8),
"layers_per_block":
1,
"act_fn":
"silu",
"norm_num_groups":
4,
"scaling_factor":
0.476986,
"spatial_compression_ratio":
self.spatial_compression_ratio,
"time_compression_ratio":
self.time_compression_ratio,
"mid_block_add_attention":
True,
}
# Instantiate the model
self.model = AutoencoderKLCausal3D(**self.init_dict)
# Create a random input tensor
self.input_tensor = torch.rand(self.batch_size, 3, self.init_time_len, self.init_height, self.init_width)
def test_encode_shape(self):
"""
Check that the shape of the encoded output matches expectations.
"""
vae_encoder_output = self.model.encode(self.input_tensor)
# The distribution from the VAE has a .sample() method
# so we verify the shape of that sample.
sample_shape = vae_encoder_output["latent_dist"].sample().shape
# We expect shape: [batch_size, latent_channels,
# (init_time_len // time_compression_ratio) + 1,
# init_height // spatial_compression_ratio,
# init_width // spatial_compression_ratio]
expected_shape = (
self.batch_size,
self.latent_channels,
(self.init_time_len // self.time_compression_ratio) + 1,
self.init_height // self.spatial_compression_ratio,
self.init_width // self.spatial_compression_ratio,
)
# (Optional) Print them if you like, or just rely on assertions:
print(f"sample_shape: {sample_shape}")
print(f"expected_shape: {expected_shape}")
self.assertEqual(
sample_shape,
expected_shape,
f"Encoded sample shape {sample_shape} does not match {expected_shape}.",
)
if __name__ == "__main__":
unittest.main()
-39
View File
@@ -1,39 +0,0 @@
import os
import shutil
import pytest
import torch
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
@pytest.fixture(scope="module", autouse=True)
def setup_distributed():
os.environ["RANK"] = "0"
os.environ["WORLD_SIZE"] = "1"
os.environ["LOCAL_RANK"] = "0"
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "12345"
dist.init_process_group("nccl")
yield
dist.destroy_process_group()
@pytest.mark.skipif(not torch.cuda.is_available(), reason="Requires at least 2 GPUs to run NCCL tests")
def test_save_and_remove_checkpoint():
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.utils.checkpoint import save_checkpoint
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs
transformer = MochiTransformer3DModel(num_layers=0)
fsdp_kwargs, _ = get_dit_fsdp_kwargs(transformer, "none")
transformer = FSDP(transformer, **fsdp_kwargs)
test_folder = "./test_checkpoint"
save_checkpoint(transformer, 0, test_folder, 0)
assert os.path.exists(test_folder), "Checkpoint folder was not created."
shutil.rmtree(test_folder)
assert not os.path.exists(test_folder), "Checkpoint folder still exists."
-111
View File
@@ -1,111 +0,0 @@
from functools import partial
from multiprocessing import Manager
import pytest
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
from fastvideo.utils.communications import nccl_info, prepare_sequence_parallel_data
def _init_distributed_test_gpu(rank, world_size, backend, port, data, results):
dist.init_process_group(
backend=backend,
init_method=f"tcp://127.0.0.1:{port}",
world_size=world_size,
rank=rank,
)
device = torch.device(f"cuda:{rank}")
nccl_info.sp_size = world_size
nccl_info.rank_within_group = rank
nccl_info.group_id = 0
seq_group = dist.new_group(ranks=list(range(world_size)))
nccl_info.group = seq_group
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask = data
hidden_states = hidden_states[rank].unsqueeze(dim=0).to(device)
encoder_hidden_states = encoder_hidden_states.to(device)
attention_mask = attention_mask.to(device)
encoder_attention_mask = encoder_attention_mask.to(device)
print(f"Rank {rank} input hidden_states:\n", hidden_states)
print(f"Rank {rank} input hidden_states shape:\n", hidden_states.shape)
out_hidden, out_encoder, out_attn_mask, out_encoder_mask = prepare_sequence_parallel_data(
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask)
print(f"Rank {rank} output out_hidden:\n", out_hidden)
shapes = (
out_hidden.shape,
out_encoder.shape,
out_attn_mask.shape,
out_encoder_mask.shape,
)
shape_tensor = torch.tensor([*shapes[0], *shapes[1], *shapes[2], *shapes[3]], dtype=torch.int32, device=device)
shape_list = [torch.zeros_like(shape_tensor) for _ in range(world_size)]
dist.all_gather(shape_list, shape_tensor, group=seq_group)
gathered_shapes = [tuple(s.tolist()) for s in shape_list]
out_hidden_cpu = out_hidden.to("cpu")
results[rank] = {
"shapes": gathered_shapes,
"out_hidden": out_hidden_cpu,
}
dist.barrier()
dist.destroy_process_group()
@pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.device_count() < 2,
reason="Requires at least 2 GPUs to run NCCL tests")
def test_prepare_sequence_parallel_data_gpu():
world_size = 2
backend = "nccl"
port = 12355 # or use a random free port if collisions occur
# Create test tensors on CPU; the dimension at index=2 should be divisible by world_size=2 (if applicable).
hidden_states = torch.randn(2, 1, 2, 1, 1)
encoder_hidden_states = torch.randn(2, 2)
attention_mask = torch.randn(2, 2)
encoder_attention_mask = torch.randn(2, 2)
print("init hidden states", hidden_states)
manager = Manager()
results_dict = manager.dict()
# Wrap our helper function with partial
mp_func = partial(_init_distributed_test_gpu,
world_size=world_size,
backend=backend,
port=port,
data=(hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask),
results=results_dict)
# Spawn two GPU processes (rank=0, rank=1)
mp.spawn(mp_func, nprocs=world_size)
first_rank_shapes = None
overall_hidden_out = []
for rank in sorted(results_dict.keys()):
rank_data = results_dict[rank]
rank_shapes = rank_data["shapes"]
if first_rank_shapes is None:
first_rank_shapes = rank_shapes
assert rank_shapes == first_rank_shapes, (
f"Mismatch in shapes across ranks: {rank_shapes} != {first_rank_shapes}")
overall_hidden_out.append(rank_data["out_hidden"])
overall_hidden_out = torch.cat(overall_hidden_out, dim=2)
print("overall_hidden_out", overall_hidden_out)
print("overall_hidden_out_shape", overall_hidden_out.shape)
assert torch.allclose(hidden_states, torch.tensor(overall_hidden_out), rtol=1e-7, atol=1e-6)
if __name__ == "__main__":
test_prepare_sequence_parallel_data_gpu()