Compare commits

...
Author SHA1 Message Date
SolitaryThinker 5dbc110cd5 simulation 2025-07-18 05:06:43 +00:00
SolitaryThinker 60df5db74c what 2025-07-17 06:23:20 +00:00
SolitaryThinker a63a4419a9 add strats 2025-07-17 02:19:01 +00:00
SolitaryThinker 863c452cd7 debug 2025-07-15 17:57:34 -07:00
“BrianChen1129” c8f4a6378f update 2025-07-15 20:07:58 +00:00
“BrianChen1129” a486a15bf4 update 2025-07-15 05:24:37 +00:00
“BrianChen1129” 207ce0e9b5 fix wandb step issue 2025-07-15 00:28:48 +00:00
“BrianChen1129” bc818a0928 update i2v t2v distill scripts 2025-07-14 23:11:23 +00:00
“BrianChen1129” 8295fabda1 update 2025-07-14 23:09:47 +00:00
“BrianChen1129” c81eabd50f update new t2v dmd 2025-07-14 09:53:59 +00:00
“BrianChen1129” 81b404b17d update 2025-07-13 22:03:51 +00:00
“BrianChen1129” 693a0f361d update 2025-07-12 04:13:02 +00:00
“BrianChen1129” c963128781 update 2025-07-12 02:30:29 +00:00
“BrianChen1129” d7143bb00d Remove local_scripts from git tracking and add to .gitignore 2025-07-11 20:48:28 +00:00
“BrianChen1129” a652992506 update 2025-07-11 20:44:38 +00:00
“BrianChen1129” 8955d15f4e bkp 2025-07-11 09:49:09 +00:00
“BrianChen1129” 411302db13 update 2025-07-10 08:06:37 +00:00
“BrianChen1129” a3e53f0e97 update latents shape 2025-07-10 06:37:22 +00:00
“BrianChen1129” 7420c7d163 update 2025-07-10 02:43:45 +00:00
“BrianChen1129” ca958fa4cc fix seed diff 2025-07-09 22:45:58 +00:00
“BrianChen1129” 1c99f157ec add readme 2025-07-09 09:58:42 +00:00
“BrianChen1129” 3e3e0f6ccb update 2025-07-09 09:34:09 +00:00
“BrianChen1129” 493d8c659a training 2025-07-09 06:51:28 +00:00
“BrianChen1129” b1f7acc4b5 update 2025-07-08 20:47:05 +00:00
“BrianChen1129” 3d6a69b154 update 2025-07-08 04:12:16 +00:00
“BrianChen1129” a3a3eea735 update 2025-07-08 00:19:51 +00:00
“BrianChen1129” fb9e18ed20 update random seed and inference latent shape 2025-07-08 00:14:54 +00:00
“BrianChen1129” 2bf1f5510b update inference 2025-07-07 06:09:04 +00:00
“BrianChen1129” 5a9f2d2f05 update 2025-07-07 03:07:40 +00:00
“BrianChen1129” b5bdf51697 update 2025-07-06 05:24:22 +00:00
“BrianChen1129” 5bad51acf2 update 2025-07-06 05:24:08 +00:00
“BrianChen1129” 6cdd5dca30 add wrapper 2025-07-05 09:59:55 +00:00
“BrianChen1129” 0c0fb63c56 update 2025-07-05 09:10:10 +00:00
“BrianChen1129” 91ef7567ac update training 2025-07-05 04:13:22 +00:00
“BrianChen1129” 09f2c64bbe training val issue 2025-07-04 07:17:19 +00:00
“BrianChen1129” fc5e823875 update 2025-07-03 22:35:34 +00:00
“BrianChen1129” 0c910981e5 update 2025-07-03 21:37:45 +00:00
“BrianChen1129” 5d809f19d1 update 2025-07-03 20:37:51 +00:00
“BrianChen1129” 8da919967f update“ 2025-07-03 04:16:08 +00:00
“BrianChen1129” 21ffa0d96a update 2025-07-03 04:08:36 +00:00
“BrianChen1129” d6a1e875f0 merge main 2025-07-03 03:43:32 +00:00
“BrianChen1129” c025ea80e5 update 2025-07-03 03:40:28 +00:00
“BrianChen1129” 98653a2261 update validation 2025-07-03 02:28:29 +00:00
“BrianChen1129” 9589938583 update 2025-07-02 22:56:37 +00:00
“BrianChen1129” 13fa3d83c7 Merge branch 'main' of github.com:hao-ai-lab/FastVideo into yq/distill-dmd 2025-06-30 03:30:39 +00:00
“BrianChen1129” db154cc313 update 2025-06-30 03:30:26 +00:00
“BrianChen1129” 8fb2708ff0 update 2025-06-29 23:46:31 +00:00
“BrianChen1129” 36dab19ff9 update 2025-06-29 08:45:56 +00:00
“BrianChen1129” 3c1179b4c2 update 2025-06-29 06:44:06 +00:00
“BrianChen1129” 8b60ac2964 Merge branch 'main' of github.com:hao-ai-lab/FastVideo into yq/distill-dmd 2025-06-29 01:57:31 +00:00
“BrianChen1129” 0b3de5e5eb update 2025-06-29 00:38:52 +00:00
“BrianChen1129” c627dc6421 Merge branch 'main' of github.com:hao-ai-lab/FastVideo into yq/distill-dmd 2025-06-29 00:03:26 +00:00
“BrianChen1129” 117d23fcc0 Merge branch 'main' of github.com:hao-ai-lab/FastVideo into yq/distill-dmd 2025-06-28 23:49:01 +00:00
“BrianChen1129” 4932793d00 update 2025-06-28 04:20:47 +00:00
“BrianChen1129” 152e6a6c77 update 2025-06-27 21:25:22 +00:00
“BrianChen1129” 3fd5531f4f debugging 2025-06-27 04:30:35 +00:00
“BrianChen1129” 26ab6c84c3 update 2025-06-26 20:35:49 +00:00
“BrianChen1129” 65c0733a48 update 2025-06-25 23:36:05 +00:00
“BrianChen1129” 4f40eef184 e2e ready; fixing bugs 2025-06-25 07:07:56 +00:00
“BrianChen1129” 1d5bf58b34 debuging loss 2025-06-25 00:56:17 +00:00
“BrianChen1129” 4ad1281377 update 2025-06-24 23:13:05 +00:00
“BrianChen1129” ee5da0838e update 2025-06-24 03:29:09 +00:00
“BrianChen1129” 4006b3a5ee update framework 2025-06-24 00:33:32 +00:00
“BrianChen1129” 6c731260f0 update 2025-06-23 22:42:23 +00:00
“BrianChen1129” 1de2eb2afd update 2025-06-23 21:07:36 +00:00
“BrianChen1129” 4b9782e3f4 update 2025-06-23 07:10:57 +00:00
32 changed files with 3219 additions and 166 deletions
+3
View File
@@ -59,5 +59,8 @@ docs/source/inference/examples/
# Static images
!docs/source/_static/images/**/*.png
# Local scripts (keep local but don't track in git)
local_scripts/
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
+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=4, help='Batch size')
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
parser.add_argument('--batch_size', type=int, default=2, help='Batch size')
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
parser.add_argument('--topk', type=int, default=32, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
args = parser.parse_args()
main(args)
@@ -15,9 +15,10 @@ NUM_GPUS=8
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--wandb_run_name "wan_i2v_finetune"
--output_dir "outputs/wan_i2v_finetune"
--max_train_steps 2000
--train_batch_size 4
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
@@ -29,10 +30,10 @@ training_args=(
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 4
--tp_size 4
--hsdp_replicate_dim 2
--hsdp_shard_dim 4
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 8
--hsdp_shard_dim 1
)
# Model arguments
@@ -52,7 +53,7 @@ validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
@@ -61,7 +62,7 @@ optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--checkpointing_steps 2000
--weight_decay 1e-4
--weight_decay 0.1
--max_grad_norm 1.0
)
@@ -70,14 +71,14 @@ miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--training_cfg_rate 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
--enable_gradient_checkpointing_type "full"
)
# --enable_gradient_checkpointing_type "full"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
@@ -26,6 +26,24 @@
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
"image_path": null,
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 98
},
{
"caption": "The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.",
"image_path": null,
"video_path": "validation_dataset/EJqsC21GSBY-Scene-013.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 148
}
]
}
+3 -1
View File
@@ -9,7 +9,8 @@ from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam)
WanT2V_14B_SamplingParam,
Wan2_1_Fun_1_3B_InP_SamplingParam)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -23,6 +24,7 @@ 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,3 +94,20 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
-13.02252404
]))
# =============================================
# ============= Wan2.1 Fun Models =============
# =============================================
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
+64 -1
View File
@@ -6,7 +6,7 @@ import argparse
import dataclasses
from contextlib import contextmanager
from dataclasses import field
from typing import Any
from typing import Any, List
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.v1.logger import init_logger
@@ -78,6 +78,8 @@ class FastVideoArgs:
# Stage verification
enable_stage_verification: bool = True
denoising_step_list: List[int] | None = field(default=None)
@property
def training_mode(self) -> bool:
@@ -254,6 +256,12 @@ class FastVideoArgs:
help="Enable input/output verification for pipeline stages",
)
parser.add_argument("--denoising-step-list",
type=parse_int_list,
default=FastVideoArgs.denoising_step_list,
help="Comma-separated list of denoising steps (e.g., '1000,757,522')",
)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -455,6 +463,18 @@ 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
min_step_ratio: float = 0.2
max_step_ratio: float = 0.98
teacher_guidance_scale: float = 3.5
# I2V frame weighting for DMD loss
i2v_frame_weighting: bool = False # Enable frame-wise weighting for I2V
i2v_first_frame_weight: float = 0.1 # Weight for first frame (0.1 = 10% weight)
i2v_weighting_scheme: str = "temporal_consistency" # Weighting scheme: "first_frame_only", "linear_increase", "exponential_increase", "motion_aware", "temporal_consistency", "first_frame_masking"
i2v_temporal_scale_factor: float = 2.0 # Scale factor for temporal weighting schemes
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -721,5 +741,48 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.VSA_decay_interval_steps,
help="VSA decay interval steps")
# Distillation arguments
parser.add_argument("--student-critic-update-ratio",
type=int,
default=TrainingArgs.student_critic_update_ratio,
help="Ratio of student updates to critic updates.")
parser.add_argument("--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")
# I2V frame weighting for DMD loss
parser.add_argument("--i2v-frame-weighting",
action=StoreBoolean,
default=TrainingArgs.i2v_frame_weighting,
help="Enable frame-wise weighting for I2V distillation to discount first frame")
parser.add_argument("--i2v-first-frame-weight",
type=float,
default=TrainingArgs.i2v_first_frame_weight,
help="Weight for first frame in I2V distillation (0.1 = 10% weight)")
parser.add_argument("--i2v-weighting-scheme",
type=str,
default=TrainingArgs.i2v_weighting_scheme,
choices=["first_frame_only", "linear_increase", "exponential_increase", "motion_aware", "temporal_consistency", "first_frame_masking"],
help="Weighting scheme for I2V frame weighting")
parser.add_argument("--i2v-temporal-scale-factor",
type=float,
default=TrainingArgs.i2v_temporal_scale_factor,
help="Scale factor for temporal weighting schemes in I2V")
return parser
def parse_int_list(value: str) -> List[int]:
"""Parse a comma-separated string of integers into a list."""
if not value:
return []
return [int(x.strip()) for x in value.split(",")]
@@ -72,6 +72,8 @@ class ComponentLoader(ABC):
module_loaders = {
"scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"),
"teacher_transformer": (TransformerLoader, "diffusers"),
"critic_transformer": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
@@ -18,15 +18,14 @@
# Modified from diffusers==0.29.2
#
# ==============================================================================
import math
from dataclasses import dataclass
from typing import Any
from typing import Any, Optional, Tuple, Union, List
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput
from diffusers.utils import BaseOutput, is_scipy_available, logging
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.base import BaseScheduler
@@ -34,7 +33,7 @@ logger = init_logger(__name__)
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
@@ -46,8 +45,7 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
Euler scheduler.
@@ -57,16 +55,37 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
use_dynamic_shifting (`bool`, defaults to False):
Whether to apply timestep shifting on-the-fly based on the image resolution.
base_shift (`float`, defaults to 0.5):
Value to stabilize image generation. Increasing `base_shift` reduces variation and image is more consistent
with desired output.
max_shift (`float`, defaults to 1.15):
Value change allowed to latent vectors. Increasing `max_shift` encourages more variation and image may be
more exaggerated or stylized.
base_image_seq_len (`int`, defaults to 256):
The base image sequence length.
max_image_seq_len (`int`, defaults to 4096):
The maximum image sequence length.
invert_sigmas (`bool`, defaults to False):
Whether to invert the sigmas.
shift_terminal (`float`, defaults to None):
The end value of the shifted timestep schedule.
use_karras_sigmas (`bool`, defaults to False):
Whether to use Karras sigmas for step sizes in the noise schedule during sampling.
use_exponential_sigmas (`bool`, defaults to False):
Whether to use exponential sigmas for step sizes in the noise schedule during sampling.
use_beta_sigmas (`bool`, defaults to False):
Whether to use beta sigmas for step sizes in the noise schedule during sampling.
time_shift_type (`str`, defaults to "exponential"):
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
stochastic_sampling (`bool`, defaults to False):
Whether to use stochastic sampling.
"""
_compatibles: list[Any] = []
_compatibles = []
order = 1
@register_to_config
@@ -74,31 +93,53 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: int | None = None,
**kwargs,
use_dynamic_shifting: bool = False,
base_shift: Optional[float] = 0.5,
max_shift: Optional[float] = 1.15,
base_image_seq_len: Optional[int] = 256,
max_image_seq_len: Optional[int] = 4096,
invert_sigmas: bool = False,
shift_terminal: Optional[float] = None,
use_karras_sigmas: Optional[bool] = False,
use_exponential_sigmas: Optional[bool] = False,
use_beta_sigmas: Optional[bool] = False,
time_shift_type: str = "exponential",
stochastic_sampling: bool = False,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] *
num_train_timesteps).to(dtype=torch.float32)
self._step_index: int | None = None
self._begin_index = 0
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
if self.config.use_beta_sigmas and not is_scipy_available():
raise ImportError("Make sure to install scipy if you want to use beta sigmas.")
if sum([self.config.use_beta_sigmas, self.config.use_exponential_sigmas, self.config.use_karras_sigmas]) > 1:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
)
if time_shift_type not in {"exponential", "linear"}:
raise ValueError("`time_shift_type` must either be 'exponential' or 'linear'.")
BaseScheduler.__init__(self)
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.timesteps = sigmas * num_train_timesteps
self._step_index = None
self._begin_index = None
self._shift = shift
self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def shift(self):
"""
The value used for shifting.
"""
return self._shift
@property
def step_index(self):
@@ -125,44 +166,190 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
self._begin_index = begin_index
def set_shift(self, shift: float):
self._shift = shift
def scale_noise(
self,
sample: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
noise: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
"""
Forward process in flow-matching
Args:
sample (`torch.FloatTensor`):
The input sample.
timestep (`int`, *optional*):
The current timestep in the diffusion chain.
Returns:
`torch.FloatTensor`:
A scaled input sample.
"""
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
if sample.device.type == "mps" and torch.is_floating_point(timestep):
# mps does not support float64
schedule_timesteps = self.timesteps.to(sample.device, dtype=torch.float32)
timestep = timestep.to(sample.device, dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(sample.device)
timestep = timestep.to(sample.device)
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timestep.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timestep.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(sample.shape):
sigma = sigma.unsqueeze(-1)
sample = sigma * noise + (1.0 - sigma) * sample
return sample
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
if self.config.time_shift_type == "exponential":
return self._time_shift_exponential(mu, sigma, t)
elif self.config.time_shift_type == "linear":
return self._time_shift_linear(mu, sigma, t)
def stretch_shift_to_terminal(self, t: torch.Tensor) -> torch.Tensor:
r"""
Stretches and shifts the timestep schedule to ensure it terminates at the configured `shift_terminal` config
value.
Reference:
https://github.com/Lightricks/LTX-Video/blob/a01a171f8fe3d99dce2728d60a73fecf4d4238ae/ltx_video/schedulers/rf.py#L51
Args:
t (`torch.Tensor`):
A tensor of timesteps to be stretched and shifted.
Returns:
`torch.Tensor`:
A tensor of adjusted timesteps such that the final value equals `self.config.shift_terminal`.
"""
one_minus_z = 1 - t
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
stretched_t = 1 - (one_minus_z / scale_factor)
return stretched_t
def set_timesteps(
self,
num_inference_steps: int,
device: str | torch.device = None,
n_tokens: int = 0,
num_inference_steps: Optional[int] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[float] = None,
timesteps: Optional[List[float]] = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
num_inference_steps (`int`, *optional*):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
sigmas (`List[float]`, *optional*):
Custom values for sigmas to be used for each diffusion step. If `None`, the sigmas are computed
automatically.
mu (`float`, *optional*):
Determines the amount of shifting applied to sigmas when performing resolution-dependent timestep
shifting.
timesteps (`List[float]`, *optional*):
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
automatically.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError("`mu` must be passed when `use_dynamic_shifting` is set to be `True`")
if sigmas is not None and timesteps is not None:
if len(sigmas) != len(timesteps):
raise ValueError("`sigmas` and `timesteps` should have the same length")
if num_inference_steps is not None:
if (sigmas is not None and len(sigmas) != num_inference_steps) or (
timesteps is not None and len(timesteps) != num_inference_steps
):
raise ValueError(
"`sigmas` and `timesteps` should have the same length as num_inference_steps, if `num_inference_steps` is provided"
)
else:
num_inference_steps = len(sigmas) if sigmas is not None else len(timesteps)
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
# 1. Prepare default sigmas
is_timesteps_provided = timesteps is not None
if not self.config.reverse:
sigmas = 1 - sigmas
if is_timesteps_provided:
timesteps = np.array(timesteps).astype(np.float32)
self.sigmas = sigmas
if not getattr(self.config, "timesteps_scale", True):
self.timesteps = sigmas[:-1] # for stepvideo
if sigmas is None:
if timesteps is None:
timesteps = np.linspace(
self._sigma_to_t(self.sigma_max), self._sigma_to_t(self.sigma_min), num_inference_steps
)
sigmas = timesteps / self.config.num_train_timesteps
else:
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
dtype=torch.float32, device=device)
# Reset step index
self._step_index = None
sigmas = np.array(sigmas).astype(np.float32)
num_inference_steps = len(sigmas)
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
# 2. Perform timestep shifting. Either no shifting is applied, or resolution-dependent shifting of
# "exponential" or "linear" type is applied
if self.config.use_dynamic_shifting:
sigmas = self.time_shift(mu, 1.0, sigmas)
else:
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
# 3. If required, stretch the sigmas schedule to terminate at the configured `shift_terminal` value
if self.config.shift_terminal:
sigmas = self.stretch_shift_to_terminal(sigmas)
# 4. If required, convert sigmas to one of karras, exponential, or beta sigma schedules
if self.config.use_karras_sigmas:
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_exponential_sigmas:
sigmas = self._convert_to_exponential(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_beta_sigmas:
sigmas = self._convert_to_beta(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
# 5. Convert sigmas and timesteps to tensors and move to specified device
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
if not is_timesteps_provided:
timesteps = sigmas * self.config.num_train_timesteps
else:
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32, device=device)
# 6. Append the terminal sigma value.
# If a model requires inverted sigma schedule for denoising but timesteps without inversion, the
# `invert_sigmas` flag can be set to `True`. This case is only required in Mochi
if self.config.invert_sigmas:
sigmas = 1.0 - sigmas
timesteps = sigmas * self.config.num_train_timesteps
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
else:
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
self.timesteps = timesteps
self.sigmas = sigmas
self._step_index = None
self._begin_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
@@ -174,17 +361,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
idx: int = indices[pos].item()
return indices[pos].item()
return idx
def set_shift(self, shift: float) -> None:
self.config.shift = shift
def set_timesteps_scale(self, timesteps_scale: bool) -> None:
self.config.timesteps_scale = timesteps_scale
def _init_step_index(self, timestep) -> None:
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
@@ -192,22 +371,19 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
else:
self._step_index = self._begin_index
def scale_model_input(self,
sample: torch.Tensor,
timestep: int | None = None) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: float | torch.FloatTensor,
sample: torch.FloatTensor,
s_churn: float = 0.0,
s_tmin: float = 0.0,
s_tmax: float = float("inf"),
s_noise: float = 1.0,
generator: Optional[torch.Generator] = None,
per_token_timesteps: Optional[torch.Tensor] = None,
return_dict: bool = True,
**kwargs,
) -> FlowMatchDiscreteSchedulerOutput | tuple:
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
@@ -219,25 +395,38 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
s_churn (`float`):
s_tmin (`float`):
s_tmax (`float`):
s_noise (`float`, defaults to 1.0):
Scaling factor for noise added to the sample.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
per_token_timesteps (`torch.Tensor`, *optional*):
The timesteps for each token in the sample.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Whether or not to return a
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`,
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] is returned,
otherwise a tuple is returned where the first element is the sample tensor.
"""
if isinstance(timestep, (int | torch.IntTensor | torch.LongTensor)):
raise ValueError((
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `FlowMatchEulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if self.step_index is None:
self._init_step_index(timestep)
@@ -245,24 +434,454 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
assert self.step_index is not None
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if per_token_timesteps is not None:
per_token_sigmas = per_token_timesteps / self.config.num_train_timesteps
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
sigmas = self.sigmas[:, None, None]
lower_mask = sigmas < per_token_sigmas[None] - 1e-6
lower_sigmas = lower_mask * sigmas
lower_sigmas, _ = lower_sigmas.max(dim=0)
current_sigma = per_token_sigmas[..., None]
next_sigma = lower_sigmas[..., None]
dt = current_sigma - next_sigma
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
sigma_idx = self.step_index
sigma = self.sigmas[sigma_idx]
sigma_next = self.sigmas[sigma_idx + 1]
current_sigma = sigma
next_sigma = sigma_next
dt = sigma_next - sigma
if self.config.stochastic_sampling:
x0 = sample - current_sigma * model_output
noise = torch.randn_like(sample)
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
else:
prev_sample = sample + dt * model_output
# upon completion increase step index by one
assert self._step_index is not None
self._step_index += 1
if per_token_timesteps is None:
# Cast sample back to model compatible dtype
prev_sample = prev_sample.to(model_output.dtype)
if not return_dict:
return (prev_sample, )
return (prev_sample,)
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample)
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
def _convert_to_karras(self, in_sigmas: torch.Tensor, num_inference_steps) -> torch.Tensor:
"""Constructs the noise schedule of Karras et al. (2022)."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
rho = 7.0 # 7.0 is the value used in the paper
ramp = np.linspace(0, 1, num_inference_steps)
min_inv_rho = sigma_min ** (1 / rho)
max_inv_rho = sigma_max ** (1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
def _convert_to_exponential(self, in_sigmas: torch.Tensor, num_inference_steps: int) -> torch.Tensor:
"""Constructs an exponential noise schedule."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.exp(np.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps))
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
def _convert_to_beta(
self, in_sigmas: torch.Tensor, num_inference_steps: int, alpha: float = 0.6, beta: float = 0.6
) -> torch.Tensor:
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.array(
[
sigma_min + (ppf * (sigma_max - sigma_min))
for ppf in [
scipy.stats.beta.ppf(timestep, alpha, beta)
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
]
]
)
return sigmas
def _time_shift_exponential(self, mu, sigma, t):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
def _time_shift_linear(self, mu, sigma, t):
return mu / (mu + (1 / t - 1) ** sigma)
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B, C, H, W]
- noise: the noise with shape [B, C, H, W]
- timestep: the timestep with shape [B]
Output: the corrupted latent with shape [B, C, H, W]
"""
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timesteps.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
return sample
def __len__(self):
return self.config.num_train_timesteps
class FlowMatchScheduler():
order = 1
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.inverse_timesteps = inverse_timesteps
self.extra_one_step = extra_one_step
self.reverse_sigmas = reverse_sigmas
self.set_timesteps(num_inference_steps)
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False):
sigma_start = self.sigma_min + \
(self.sigma_max - self.sigma_min) * denoising_strength
if self.extra_one_step:
self.sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
else:
self.sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps)
if self.inverse_timesteps:
self.sigmas = torch.flip(self.sigmas, dims=[0])
self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas)
if self.reverse_sigmas:
self.sigmas = 1 - self.sigmas
self.timesteps = self.sigmas * self.num_train_timesteps
if training:
x = self.timesteps
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
num_inference_steps) ** 2)
y_shifted = y - y.min()
bsmntw_weighing = y_shifted * \
(num_inference_steps / y_shifted.sum())
self.linear_timesteps_weights = bsmntw_weighing
def step(self, model_output, timestep, sample, to_final=False):
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
sigma_ = 1 if (
self.inverse_timesteps or self.reverse_sigmas) else 0
else:
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
prev_sample = sample + model_output * (sigma_ - sigma)
return prev_sample
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B, C, H, W]
- noise: the noise with shape [B, C, H, W]
- timestep: the timestep with shape [B]
Output: the corrupted latent with shape [B, C, H, W]
"""
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) # [21, 1, 1, 1]
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
def training_weight(self, timestep):
timestep_id = torch.argmin(
(self.timesteps - timestep.to(self.timesteps.device)).abs())
weights = self.linear_timesteps_weights[timestep_id]
return weights
# class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# """
# Euler scheduler.
# This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
# methods the library implements for all schedulers such as loading and saving.
# Args:
# num_train_timesteps (`int`, defaults to 1000):
# The number of diffusion steps to train the model.
# timestep_spacing (`str`, defaults to `"linspace"`):
# The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
# Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
# shift (`float`, defaults to 1.0):
# The shift value for the timestep schedule.
# reverse (`bool`, defaults to `True`):
# Whether to reverse the timestep schedule.
# """
# _compatibles: list[Any] = []
# order = 1
# @register_to_config
# def __init__(
# self,
# num_train_timesteps: int = 1000,
# shift: float = 1.0,
# reverse: bool = True,
# solver: str = "euler",
# n_tokens: Optional[int] = None,
# **kwargs,
# ):
# sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
# if not reverse:
# sigmas = sigmas.flip(0)
# self.sigmas = sigmas
# # the value fed to model
# self.timesteps = (sigmas[:-1] *
# num_train_timesteps).to(dtype=torch.float32)
# self._step_index: int | None = None
# self._begin_index = 0
# self.supported_solver = ["euler"]
# if solver not in self.supported_solver:
# raise ValueError(
# f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
# )
# BaseScheduler.__init__(self)
# @property
# def step_index(self):
# """
# The index counter for current timestep. It will increase 1 after each scheduler step.
# """
# return self._step_index
# @property
# def begin_index(self):
# """
# The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
# """
# return self._begin_index
# # Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
# def set_begin_index(self, begin_index: int = 0):
# """
# Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
# Args:
# begin_index (`int`):
# The begin index for the scheduler.
# """
# self._begin_index = begin_index
# def _sigma_to_t(self, sigma):
# return sigma * self.config.num_train_timesteps
# def set_timesteps(
# self,
# num_inference_steps: int,
# device: Union[str, torch.device] = None,
# n_tokens: int = 0,
# ):
# """
# Sets the discrete timesteps used for the diffusion chain (to be run before inference).
# Args:
# num_inference_steps (`int`):
# The number of diffusion steps used when generating samples with a pre-trained model.
# device (`str` or `torch.device`, *optional*):
# The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
# n_tokens (`int`, *optional*):
# Number of tokens in the input sequence.
# """
# self.num_inference_steps = num_inference_steps
# sigmas = torch.linspace(1, 0, num_inference_steps + 1)
# sigmas = self.sd3_time_shift(sigmas)
# if not self.config.reverse:
# sigmas = 1 - sigmas
# self.sigmas = sigmas
# if not getattr(self.config, "timesteps_scale", True):
# self.timesteps = sigmas[:-1] # for stepvideo
# else:
# self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
# dtype=torch.float32, device=device)
# # Reset step index
# self._step_index = None
# def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
# if schedule_timesteps is None:
# schedule_timesteps = self.timesteps
# indices = (schedule_timesteps == timestep).nonzero()
# # The sigma index that is taken for the **very** first `step`
# # is always the second index (or the last index if there is only 1)
# # This way we can ensure we don't accidentally skip a sigma in
# # case we start in the middle of the denoising schedule (e.g. for image-to-image)
# pos = 1 if len(indices) > 1 else 0
# idx: int = indices[pos].item()
# return idx
# def set_shift(self, shift: float) -> None:
# self.config.shift = shift
# def set_timesteps_scale(self, timesteps_scale: bool) -> None:
# self.config.timesteps_scale = timesteps_scale
# def _init_step_index(self, timestep) -> None:
# if self.begin_index is None:
# if isinstance(timestep, torch.Tensor):
# timestep = timestep.to(self.timesteps.device)
# self._step_index = self.index_for_timestep(timestep)
# else:
# self._step_index = self._begin_index
# def scale_model_input(self,
# sample: torch.Tensor,
# timestep: Optional[int] = None) -> torch.Tensor:
# return sample
# def sd3_time_shift(self, t: torch.Tensor):
# return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
# def step(
# self,
# model_output: torch.FloatTensor,
# timestep: Union[float, torch.FloatTensor],
# sample: torch.FloatTensor,
# return_dict: bool = True,
# **kwargs,
# ) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
# """
# Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
# process from the learned model outputs (most often the predicted noise).
# Args:
# model_output (`torch.FloatTensor`):
# The direct output from learned diffusion model.
# timestep (`float`):
# The current discrete timestep in the diffusion chain.
# sample (`torch.FloatTensor`):
# A current instance of a sample created by the diffusion process.
# generator (`torch.Generator`, *optional*):
# A random number generator.
# n_tokens (`int`, *optional*):
# Number of tokens in the input sequence.
# return_dict (`bool`):
# Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
# tuple.
# Returns:
# [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
# If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
# returned, otherwise a tuple is returned where the first element is the sample tensor.
# """
# if isinstance(timestep, (int, torch.IntTensor, torch.LongTensor)):
# raise ValueError((
# "Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
# " `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
# " one of the `scheduler.timesteps` as a timestep."), )
# if self.step_index is None:
# self._init_step_index(timestep)
# # Upcast to avoid precision issues when computing prev_sample
# sample = sample.to(torch.float32)
# assert self.step_index is not None
# dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
# if self.config.solver == "euler":
# prev_sample = sample + model_output.to(torch.float32) * dt
# else:
# raise ValueError(
# f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
# )
# # upon completion increase step index by one
# assert self._step_index is not None
# self._step_index += 1
# if not return_dict:
# return (prev_sample, )
# return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
# def __len__(self):
# return self.config.num_train_timesteps
@@ -772,47 +772,5 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
return sample
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.add_noise
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=original_samples.device,
dtype=original_samples.dtype)
if original_samples.device.type == "mps" and torch.is_floating_point(
timesteps):
# mps does not support float64
schedule_timesteps = self.timesteps.to(original_samples.device,
dtype=torch.float32)
timesteps = timesteps.to(original_samples.device,
dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(original_samples.device)
timesteps = timesteps.to(original_samples.device)
# begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [
self.index_for_timestep(t, schedule_timesteps)
for t in timesteps
]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timesteps.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timesteps.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(original_samples.shape):
sigma = sigma.unsqueeze(-1)
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
noisy_samples = alpha_t * original_samples + sigma_t * noise
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
@@ -222,12 +222,13 @@ class ComposedPipelineBase(ABC):
assert len(
model_index
) > 1, "model_index.json must contain at least one pipeline module"
for module_name in self.required_config_modules:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.warning(
f"model_index.json does not contain a {module_name} module, adding {module_name} to model_index")
if 'transformer' in module_name:
model_index[module_name] = model_index['transformer']
# all the component models used by the pipeline
required_modules = self.required_config_modules
logger.info("Loading required modules: %s", required_modules)
@@ -242,7 +243,11 @@ class ComposedPipelineBase(ABC):
logger.info("Using module %s already provided", module_name)
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
if 'transformer' in module_name:
loading_module_name = module_name.split("_")[-1]
else:
loading_module_name = module_name
component_model_path = os.path.join(self.model_path, loading_module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
@@ -18,6 +18,12 @@ from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
WanTeaCacheParams)
__all__ = [
"ForwardBatch",
"TrainingBatch",
"AttentionMetadata",
"VideoSparseAttentionMetadata",
]
@dataclass
class ForwardBatch:
@@ -146,9 +152,11 @@ class ForwardBatch:
class TrainingBatch:
current_timestep: int = 0
current_vsa_sparsity: float = 0.0
# Dataloader batch outputs
latents: torch.Tensor | None = None
noise_latents: torch.Tensor | None = None
encoder_hidden_states: torch.Tensor | None = None
encoder_attention_mask: torch.Tensor | None = None
# i2v
@@ -156,6 +164,7 @@ class TrainingBatch:
image_embeds: torch.Tensor | None = None
image_latents: torch.Tensor | None = None
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
@@ -163,6 +172,7 @@ class TrainingBatch:
sigmas: torch.Tensor | None = None
noise: torch.Tensor | None = None
attn_metadata_vsa: AttentionMetadata | None = None
attn_metadata: AttentionMetadata | None = None
# input kwargs
@@ -174,3 +184,19 @@ class TrainingBatch:
# Training outputs
total_loss: float | None = None
grad_norm: float | None = None
# Distillation-specific attributes
encoder_hidden_states_neg: torch.Tensor | None = None
encoder_attention_mask_neg: torch.Tensor | None = None
conditional_dict: dict[str, Any] | None = None
unconditional_dict: dict[str, Any] | None = None
# Distillation losses
student_loss: float = 0.0
critic_loss: float = 0.0
# Training control
dmd_log_dict: dict[str, Any] = field(default_factory=dict)
critic_log_dict: dict[str, Any] = field(default_factory=dict)
@@ -10,6 +10,7 @@ from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
from fastvideo.v1.pipelines.stages.denoising import DmdDenoisingStage
from fastvideo.v1.pipelines.stages.encoding import EncodingStage
from fastvideo.v1.pipelines.stages.image_encoding import ImageEncodingStage
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
@@ -28,6 +29,7 @@ __all__ = [
"LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
@@ -103,6 +103,7 @@ class DecodingStage(PipelineStage):
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
+234 -2
View File
@@ -3,7 +3,7 @@
Denoising stage for diffusion pipelines.
"""
import inspect
import inspect, copy
from collections.abc import Iterable
from typing import Any
@@ -27,6 +27,7 @@ from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
from fastvideo.v1.platforms import AttentionBackendEnum
from fastvideo.v1.utils import dict_to_3d_list
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
try:
from fastvideo.v1.attention.backends.sliding_tile_attn import (
@@ -117,6 +118,7 @@ class DenoisingStage(PipelineStage):
batch.image_latent = image_latent
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
if timesteps is None:
raise ValueError("Timesteps must be provided")
@@ -176,7 +178,7 @@ class DenoisingStage(PipelineStage):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
# Expand latents for I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
@@ -541,3 +543,233 @@ class DenoisingStage(PipelineStage):
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
class DmdDenoisingStage(DenoisingStage):
"""
Denoising stage for DMD.
"""
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
self.scheduler = FlowMatchEulerDiscreteScheduler(
shift=8.0)
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run the denoising loop.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta
},
)
# Setup precision and autocast settings
# TODO(will): make the precision configurable for inference
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Handle sequence parallelism if enabled
sp_world_size, rank_in_sp_group = get_sp_world_size(
), get_sp_parallel_rank()
sp_group = sp_world_size > 1
if sp_group:
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=sp_world_size).contiguous()
latents = latents[:, :, rank_in_sp_group, :, :, :]
batch.latents = latents
if batch.image_latent is not None:
image_latent = rearrange(batch.image_latent,
"b t (n s) h w -> b t n s h w",
n=sp_world_size).contiguous()
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
batch.image_latent = image_latent
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_image": image_embeds,
"mask_strategy": dict_to_3d_list(
None, t_max=50, l_max=60, h_max=24)
},
)
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
neg_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_neg,
"encoder_attention_mask": batch.negative_attention_mask,
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Get latents and embeddings
latents = batch.latents
# TODO(yongqi) hard code prepare latents
latents = torch.randn(latents.permute(0, 2, 1, 3, 4).shape, dtype=torch.bfloat16, device="cuda", generator=torch.Generator(device="cuda").manual_seed(42))
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
timesteps = torch.tensor(
fastvideo_args.denoising_step_list, dtype=torch.long, device=get_local_torch_device())
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
# Expand latents for I2V
noise_latents = copy.deepcopy(latents)
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent.permute(0, 2, 1, 3, 4)],
dim=2).to(target_dtype)
assert torch.isnan(latent_model_input).sum() == 0
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (
torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_local_torch_device(),
).to(target_dtype) *
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
if (st_attn_available
and self.attn_backend == SlidingTileAttentionBackend
) or (vsa_available and self.attn_backend
== VideoSparseAttentionBackend):
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
# TODO(will): clean this up
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
forward_batch=batch,
fastvideo_args=fastvideo_args,
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
# TODO(will): finalize the interface. vLLM uses this to
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# fastvideo_args or training_args, and attn_metadata.
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch,
# fastvideo_args=fastvideo_args
):
# Run transformer
pred_noise = self.transformer(
latent_model_input.permute(0, 2, 1, 3, 4),
prompt_embeds,
t_expand,
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
t_shape = pred_noise.shape[1]
timestep = t_expand.expand(1, t_shape)
from fastvideo.v1.training.training_utils import DiffusionWrapper
pred_video = DiffusionWrapper._convert_flow_pred_to_x0(
flow_pred=pred_noise.flatten(0, 1),
xt=noise_latents.flatten(0, 1),
timestep=timestep.flatten(0, 1),
scheduler=self.scheduler
).unflatten(0, pred_noise.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
pred_video.shape[:2], dtype=torch.long, device=pred_video.device)
latents = self.scheduler.add_noise(
pred_video.flatten(0, 1),
torch.randn_like(pred_video.flatten(0, 1)),
next_timestep.flatten(0, 1)
).unflatten(0, pred_video.shape[:2])
else:
latents = pred_video.permute(0, 2, 1, 3, 4)
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
# Update batch with final latents
batch.latents = latents
# Save STA mask search results if needed
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
self.save_sta_search_results(batch)
return batch
@@ -14,6 +14,7 @@ from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
from fastvideo.v1.pipelines.stages.validators import VerificationResult
import numpy as np
logger = init_logger(__name__)
@@ -0,0 +1,70 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.v1.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
DmdDenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class WanDmdPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Wan video diffusion pipeline with LoRA support.
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanDmdPipeline
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.v1.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DmdDenoisingStage,
EncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
logger = init_logger(__name__)
class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanImageToVideoDmdPipeline
+2 -1
View File
@@ -1,4 +1,5 @@
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
from .distillation_pipeline import DistillationPipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
File diff suppressed because it is too large Load Diff
+5 -2
View File
@@ -221,7 +221,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
# indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
indices = (u * self.noise_scheduler.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=training_batch.latents.device)
if self.training_args.sp_size > 1:
@@ -253,7 +254,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],
@@ -319,6 +320,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
@@ -705,3 +707,4 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
+93
View File
@@ -3,17 +3,22 @@ import json
import math
import os
import time
from abc import ABC
from typing import Any
import numpy as np
import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from einops import rearrange
from safetensors.torch import save_file
from torchvision.utils import make_grid
import wandb
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.v1.training.checkpointing_utils import (ModelWrapper,
OptimizerWrapper,
RandomStateWrapper,
@@ -550,3 +555,91 @@ def convert_custom_format_to_diffusers_format(state_dict: dict[str, Any],
new_state_dict[training_key] = v
return new_state_dict
def prepare_for_saving(tensor: torch.Tensor,
fps: int = 16,
caption: str | None = None) -> wandb.Image | wandb.Video:
if tensor.ndim == 4:
# Assuming it's an image and has shape [batch_size, 3, height, width]
tensor = make_grid(tensor, 4, padding=0, normalize=False)
return wandb.Image((tensor * 255).numpy().astype(np.uint8),
caption=caption)
elif tensor.ndim == 5:
# Assuming it's a video and has shape [batch_size, num_frames, 3, height, width]
return wandb.Video((tensor * 255).numpy().astype(np.uint8),
fps=fps,
format="webm",
caption=caption)
else:
raise ValueError(
"Unsupported tensor shape for saving. Expected 4D (image) or 5D (video) tensor."
)
class DiffusionWrapper(torch.nn.Module, ABC):
def __init__(self, transformer, scheduler):
super().__init__()
self.model = transformer
self.scheduler = scheduler
def forward(self, training_batch: TrainingBatch, timestep: torch.Tensor):
pred_noise = self.model(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = self._convert_flow_pred_to_x0(
flow_pred=pred_noise.flatten(0, 1),
xt=training_batch.noise_latents.flatten(0, 1),
timestep=timestep.flatten(0, 1),
scheduler=self.scheduler).unflatten(0, pred_noise.shape[:2])
return pred_video
@staticmethod
def _convert_x0_to_flow_pred(x0_pred: torch.Tensor, xt: torch.Tensor,
timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert x0 prediction to flow matching's prediction.
x0_pred: the x0 prediction with shape [B, C, H, W]
xt: the input noisy data with shape [B, C, H, W]
timestep: the timestep with shape [B]
pred = (x_t - x_0) / sigma_t
"""
# use higher precision for calculations
original_dtype = x0_pred.dtype
x0_pred, xt, sigmas, timesteps = map(
lambda x: x.double().to(x0_pred.device),
[x0_pred, xt, scheduler.sigmas, scheduler.timesteps])
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
flow_pred = (xt - x0_pred) / sigma_t
return flow_pred.to(original_dtype)
@staticmethod
def _convert_flow_pred_to_x0(flow_pred: torch.Tensor, xt: torch.Tensor,
timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert flow matching's prediction to x0 prediction.
flow_pred: the prediction with shape [B, C, H, W]
xt: the input noisy data with shape [B, C, H, W]
timestep: the timestep with shape [B]
pred = noise - x0
x_t = (1-sigma_t) * x0 + sigma_t * noise
we have x0 = x_t - sigma_t * pred
see derivations https://chatgpt.com/share/67bf8589-3d04-8008-bc6e-4cf1a24e2d0e
"""
# use higher precision for calculations
original_dtype = flow_pred.dtype
flow_pred, xt, sigmas, timesteps = map(
lambda x: x.double().to(flow_pred.device),
[flow_pred, xt, scheduler.sigmas, scheduler.timesteps])
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
x0_pred = xt - sigma_t * flow_pred
return x0_pred.to(original_dtype)
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
import torch
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.wan.wan_dmd_pipeline import WanDmdPipeline
from fastvideo.v1.training.distillation_pipeline import DistillationPipeline
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
TrainingBatch)
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanDistillationPipeline(DistillationPipeline):
"""
A distillation pipeline for Wan that uses a single transformer model.
The main transformer serves as the student model, and copies are made for teacher and critic.
"""
_required_config_modules = ["scheduler", "transformer", "vae", "teacher_transformer", "critic_transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.use_cpu_offload = False
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanDmdPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus)
self.validation_pipeline = validation_pipeline
def _build_input_kwargs(self, noise_input: torch.Tensor, timestep: torch.Tensor, text_dict: dict[str, torch.Tensor],
training_batch: TrainingBatch) -> TrainingBatch:
training_batch.input_kwargs = {
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep[0][:1],
"return_dict":
False,
}
training_batch.noise_latents = noise_input
return training_batch
def main(args) -> None:
logger.info("Starting Wan distillation pipeline...")
# Create pipeline with original args
pipeline = WanDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
# Start training
pipeline.train()
logger.info("Wan distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)
@@ -0,0 +1,231 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.distributed import get_local_torch_device
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
TrainingBatch)
from fastvideo.v1.pipelines.wan.wan_i2v_dmd_pipeline import WanImageToVideoDmdPipeline
from fastvideo.v1.training.distillation_pipeline import DistillationPipeline
from fastvideo.v1.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanI2VDistillationPipeline(DistillationPipeline):
"""
A distillation pipeline for Wan that uses a single transformer model.
The main transformer serves as the student model, and copies are made for teacher and critic.
"""
_required_config_modules = ["scheduler", "transformer", "vae", "teacher_transformer", "critic_transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_i2v
self.validation_dataset_schema = pyarrow_schema_i2v_validation
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.use_cpu_offload = False
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
validation_pipeline = WanImageToVideoDmdPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
use_cpu_offload=True)
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.preprocessed_image = pil_image.to(
get_local_torch_device())
training_batch.image_embeds = clip_features.to(get_local_torch_device())
training_batch.image_latents = image_latents.to(
get_local_torch_device())
training_batch.infos = infos
return training_batch
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.image_path = validation_batch['video_path']
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert self.noise_random_generator is not None
assert training_batch.image_latents is not None
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(
get_local_torch_device(), dtype=torch.bfloat16)
temporal_compression_ratio = 4
num_frames = (self.training_args.num_latent_t -
1) * temporal_compression_ratio + 1
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, 1:] = 0
first_frame_mask = mask_lat_size[:, :, :1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
dim=2)
mask_lat_size = mask_lat_size.view(batch_size, -1,
temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(
image_latents.device).to(dtype=torch.bfloat16)
image_latents = torch.cat(
[mask_lat_size, image_latents],
dim=1)
training_batch.image_latents = image_latents
return training_batch
def _build_input_kwargs(self, noise_input: torch.Tensor, timestep: torch.Tensor, text_dict: dict[str, torch.Tensor],
training_batch: TrainingBatch) -> TrainingBatch:
assert training_batch.image_embeds is not None
assert training_batch.image_latents is not None
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_local_torch_device(),
dtype=torch.bfloat16)
noisy_model_input = torch.cat(
[noise_input, training_batch.image_latents.permute(0, 2, 1, 3, 4)], dim=2)
training_batch.input_kwargs = {
"hidden_states": noisy_model_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep[0][:1],
"encoder_hidden_states_image": image_embeds,
"return_dict":
False,
}
training_batch.noise_latents = noise_input
return training_batch
def main(args) -> None:
logger.info("Starting Wan distillation pipeline...")
# Create pipeline with original args
pipeline = WanI2VDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
# Start training
pipeline.train()
logger.info("Wan distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
main(args)
+78
View File
@@ -0,0 +1,78 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
VALIDATION_DIR=examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.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 NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
# export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
# export NODE_RANK=$SLURM_PROCID
# nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
# export MASTER_ADDR=${nodes[0]}
# export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
# echo "MASTER_ADDR: $MASTER_ADDR"
# echo "NODE_RANK: $NODE_RANK"
# 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 8 \
--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 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 500 \
--validation_steps 100 \
--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_train/wan_finetune_1e5" \
--tracker_project_name Wan_distillation \
--wandb_run_name "crush_smol_dmd_test" \
--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 \
--seed 1000 \
--teacher_guidance_scale 3.5 \
--enable_gradient_checkpointing_type "full"
+79
View File
@@ -0,0 +1,79 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# DATA_DIR=data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
# VALIDATION_DIR=examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json
# VALIDATION_DIR=/mnt/weka/home/hao.zhang/wl/FastVideo/data/mixkit/validation.json
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/wl/mixkit/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6
# IP=[MASTER NODE IP]
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export MASTER_PORT=29501
export TOKENIZERS_PARALLELISM=false
# 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 weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path weizhou03/Wan2.1-Fun-1.3B-InP-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 8 \
--hsdp_replicate_dim 8 \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 1 \
--max_train_steps 6000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 1000 \
--validation_steps 50 \
--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_train_i2v/wan_i2v_finetune_9e6" \
--tracker_project_name Wan_distillation \
--wandb_run_name "backward_sim_synth" \
--num_height 480 \
--num_width 832 \
--num_frames 61 \
--flow_shift 3 \
--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 \
--seed 1000 \
--teacher_guidance_scale 3.5
# --i2v_frame_weighting \
#--denoising_step_list '1000,757,522' \
# --i2v_weighting_scheme "first_frame_only" \
# --i2v_temporal_scale_factor 1.0 \
# --i2v_first_frame_weight 0.01 \
+62
View File
@@ -0,0 +1,62 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY=
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=mini_i2v_dataset/crush-smol_preprocessed/combined_parquet_dataset
VALIDATION_DIR=mini_i2v_dataset/crush-smol_preprocessed/validation_parquet_dataset
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_distillation_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_preprocessed_path "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim 8 \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 1 \
--max_train_steps 30000 \
--learning_rate 2e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 10 \
--validation_steps 10 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune" \
--tracker_project_name Wan_distillation \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 8 \
--validation_guidance_scale "1.0" \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--vae_precision "bf16" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--student_critic_update_ratio 5 \
--denoising_step_list '999,757,522' \
--min_step_ratio 0.02 \
--max_step_ratio 0.98 \
--teacher_guidance_scale 3.5 \
+70
View File
@@ -0,0 +1,70 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# 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=data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
VALIDATION_DIR=examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.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 \
--wandb_run_name "crush_smol_dmd_test" \
--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
@@ -0,0 +1,66 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/test_8/
NUM_GPUS=1
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_distillation_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-14B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-14B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_preprocessed_path "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 20 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 16 \
--max_train_steps 3000 \
--learning_rate 4e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 10 \
--validation_steps 10 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune" \
--tracker_project_name Wan_distillation \
--num_height 768 \
--num_width 1280 \
--num_frames 77 \
--flow_shift 8 \
--validation_guidance_scale "1.0" \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--vae_precision "bf16" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--student_critic_update_ratio 5 \
--denoising_step_list '1000,757,522' \
--min_step_ratio 0.02 \
--max_step_ratio 0.98 \
--teacher_guidance_scale 3.5 \
--enable_gradient_checkpointing_type "full" \
--seed 1000 \
--VSA_sparsity 0.0 \
+66
View File
@@ -0,0 +1,66 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
NUM_GPUS=1
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_distillation_pipeline.py \
--model_path data/Wan2.1-T2V-1.3B-Diffusers-VT \
--inference_mode False\
--pretrained_model_name_or_path data/Wan2.1-T2V-1.3B-Diffusers-VT \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 16 \
--max_train_steps 3000 \
--learning_rate 4e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 10 \
--validation_steps 10 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune" \
--tracker_project_name Wan_distillation \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 8 \
--validation_guidance_scale "1.0" \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--vae_precision "bf16" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--student_critic_update_ratio 5 \
--denoising_step_list '1000,757,522' \
--min_step_ratio 0.02 \
--max_step_ratio 0.98 \
--teacher_guidance_scale 3.5 \
--enable_gradient_checkpointing_type "full" \
--seed 1000 \
--VSA_sparsity 0.9 \
+66
View File
@@ -0,0 +1,66 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/val/
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
export TOKENIZERS_PARALLELISM=false
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_i2v_distillation_pipeline.py \
--model_path Wan2.1-I2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan2.1-I2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 8 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 16 \
--max_train_steps 3000 \
--learning_rate 4e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 10 \
--validation_steps 10 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune_i2v" \
--tracker_project_name Wan_distillation \
--num_height 448 \
--num_width 832 \
--num_frames 29 \
--flow_shift 8 \
--validation_guidance_scale "6.0" \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--vae_precision "bf16" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--student_critic_update_ratio 5 \
--denoising_step_list '1000,757,522' \
--min_step_ratio 0.02 \
--max_step_ratio 0.98 \
--teacher_guidance_scale 3.5 \
--enable_gradient_checkpointing_type "full" \
--seed 1000 \
+2 -2
View File
@@ -10,8 +10,8 @@ fastvideo generate \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--height 448 \
--width 832 \
--height 768 \
--width 1280\
--num-frames 77 \
--num-inference-steps 50 \
--fps 16 \