Compare commits
73
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e1dca7aa71 | ||
|
|
759f29f64e | ||
|
|
899da621ff | ||
|
|
6eaebf0322 | ||
|
|
8997fd708a | ||
|
|
f615f56a7f | ||
|
|
2144ddf819 | ||
|
|
e6c4e53323 | ||
|
|
9edcc0adb4 | ||
|
|
f35ef32066 | ||
|
|
056f1774b8 | ||
|
|
c8f4a6378f | ||
|
|
a486a15bf4 | ||
|
|
207ce0e9b5 | ||
|
|
bc818a0928 | ||
|
|
8295fabda1 | ||
|
|
c81eabd50f | ||
|
|
81b404b17d | ||
|
|
693a0f361d | ||
|
|
c963128781 | ||
|
|
d7143bb00d | ||
|
|
a652992506 | ||
|
|
8955d15f4e | ||
|
|
411302db13 | ||
|
|
a3e53f0e97 | ||
|
|
7420c7d163 | ||
|
|
ca958fa4cc | ||
|
|
1c99f157ec | ||
|
|
3e3e0f6ccb | ||
|
|
493d8c659a | ||
|
|
b1f7acc4b5 | ||
|
|
3d6a69b154 | ||
|
|
a3a3eea735 | ||
|
|
fb9e18ed20 | ||
|
|
2bf1f5510b | ||
|
|
5a9f2d2f05 | ||
|
|
b5bdf51697 | ||
|
|
5bad51acf2 | ||
|
|
6cdd5dca30 | ||
|
|
0c0fb63c56 | ||
|
|
91ef7567ac | ||
|
|
09f2c64bbe | ||
|
|
fc5e823875 | ||
|
|
0c910981e5 | ||
|
|
5d809f19d1 | ||
|
|
8da919967f | ||
|
|
21ffa0d96a | ||
|
|
d6a1e875f0 | ||
|
|
c025ea80e5 | ||
|
|
98653a2261 | ||
|
|
9589938583 | ||
|
|
13fa3d83c7 | ||
|
|
db154cc313 | ||
|
|
8fb2708ff0 | ||
|
|
36dab19ff9 | ||
|
|
3c1179b4c2 | ||
|
|
8b60ac2964 | ||
|
|
0b3de5e5eb | ||
|
|
c627dc6421 | ||
|
|
117d23fcc0 | ||
|
|
4932793d00 | ||
|
|
152e6a6c77 | ||
|
|
3fd5531f4f | ||
|
|
26ab6c84c3 | ||
|
|
65c0733a48 | ||
|
|
4f40eef184 | ||
|
|
1d5bf58b34 | ||
|
|
4ad1281377 | ||
|
|
ee5da0838e | ||
|
|
4006b3a5ee | ||
|
|
6c731260f0 | ||
|
|
1de2eb2afd | ||
|
|
4b9782e3f4 |
@@ -59,5 +59,8 @@ docs/source/inference/examples/
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
|
||||
# Local scripts (keep local but don't track in git)
|
||||
local_scripts/
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
@@ -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)
|
||||
@@ -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,12 @@ 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
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
@@ -721,5 +735,29 @@ 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")
|
||||
|
||||
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"),
|
||||
@@ -554,6 +556,5 @@ class PipelineComponentLoader:
|
||||
# Get the appropriate loader for this module type
|
||||
loader = ComponentLoader.for_module_type(module_name,
|
||||
transformers_or_diffusers)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, fastvideo_args)
|
||||
|
||||
@@ -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,156 @@ 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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -51,6 +51,7 @@ class ComposedPipelineBase(ABC):
|
||||
Initialize the pipeline. After __init__, the pipeline should be ready to
|
||||
use. The pipeline should be stateless and not hold any batch state.
|
||||
"""
|
||||
|
||||
self.fastvideo_args = fastvideo_args
|
||||
|
||||
self.model_path: str = model_path
|
||||
@@ -122,7 +123,7 @@ class ComposedPipelineBase(ABC):
|
||||
for key, value in kwargs.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
|
||||
fastvideo_args.use_cpu_offload = False
|
||||
# fastvideo_args.use_cpu_offload = False
|
||||
# make sure we are in training mode
|
||||
fastvideo_args.inference_mode = False
|
||||
# we hijack the precision to be the master weight type so that the
|
||||
@@ -130,7 +131,7 @@ class ComposedPipelineBase(ABC):
|
||||
# use FSDP2's MixedPrecisionPolicy to set the precision for the
|
||||
# fwd, bwd, and other operations' precision.
|
||||
# fastvideo_args.precision = fastvideo_args.master_weight_type
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
# assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
@@ -222,12 +223,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 +244,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,12 @@ class ForwardBatch:
|
||||
class TrainingBatch:
|
||||
current_timestep: int = 0
|
||||
current_vsa_sparsity: float = 0.0
|
||||
|
||||
|
||||
# Dataloader batch outputs
|
||||
latents: torch.Tensor | None = None
|
||||
raw_latent_shape: 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 +165,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 +173,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 +185,20 @@ 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
|
||||
regression_loss: float = 0.0
|
||||
|
||||
# Training control
|
||||
dmd_log_dict: dict[str, Any] = field(default_factory=dict)
|
||||
critic_log_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DmdDenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.v1.pipelines.stages.image_encoding import ImageEncodingStage
|
||||
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
|
||||
@@ -28,6 +29,7 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
|
||||
@@ -103,6 +103,7 @@ class DecodingStage(PipelineStage):
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import inspect, copy
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
@@ -27,6 +27,7 @@ from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import (
|
||||
@@ -105,18 +106,19 @@ class DenoisingStage(PipelineStage):
|
||||
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",
|
||||
"b c (n t) h w -> b c n t 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",
|
||||
"b c (n t) h w -> b c n t 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")
|
||||
@@ -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,244 @@ 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
|
||||
|
||||
# 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))
|
||||
video_raw_latent_shape = latents.shape
|
||||
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())
|
||||
|
||||
# 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(latents,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=sp_world_size).contiguous()
|
||||
latents = latents[:, rank_in_sp_group, :, :, :, :]
|
||||
if batch.image_latent is not None:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b c (n t) h w -> b c n t h w",
|
||||
n=sp_world_size).contiguous()
|
||||
|
||||
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
|
||||
# timesteps = batch.timesteps
|
||||
# num_inference_steps = batch.num_inference_steps
|
||||
|
||||
# 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)
|
||||
noise = torch.randn(
|
||||
video_raw_latent_shape, device=self.device, dtype=pred_video.dtype)
|
||||
if sp_group:
|
||||
noise = rearrange(noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=sp_world_size).contiguous()
|
||||
noise = noise[:, rank_in_sp_group, :, :, :, :]
|
||||
latents = self.scheduler.add_noise(
|
||||
pred_video.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
next_timestep.flatten(0, 1)
|
||||
).unflatten(0, pred_video.shape[:2])
|
||||
else:
|
||||
latents = pred_video
|
||||
|
||||
# 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=1)
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
# 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
|
||||
@@ -0,0 +1,182 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from huggingface_hub import snapshot_download
|
||||
import subprocess
|
||||
import sys
|
||||
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
|
||||
|
||||
# Import the training pipeline
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
|
||||
NUM_NODES = "1"
|
||||
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
|
||||
# preprocessing
|
||||
DATA_DIR = "data"
|
||||
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol"))
|
||||
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
|
||||
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
|
||||
|
||||
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol_processed_t2v"))
|
||||
|
||||
|
||||
# training
|
||||
NUM_GPUS_PER_NODE_TRAINING = "4"
|
||||
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_distillation_pipeline.py"
|
||||
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
|
||||
LOCAL_VALIDATION_DATASET_FILE = "examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
|
||||
|
||||
def download_data():
|
||||
# create the data dir if it doesn't exist
|
||||
data_dir = Path(DATA_DIR)
|
||||
|
||||
print(f"Creating data directory at {data_dir}")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
|
||||
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
|
||||
try:
|
||||
result = snapshot_download(
|
||||
repo_id="wlsaidhi/crush-smol-merged",
|
||||
local_dir=str(LOCAL_RAW_DATA_DIR),
|
||||
repo_type="dataset",
|
||||
resume_download=True,
|
||||
token=os.environ.get("HF_TOKEN"), # In case authentication is needed
|
||||
)
|
||||
print(f"Download completed successfully. Files downloaded to: {result}")
|
||||
|
||||
# Verify the download
|
||||
if not LOCAL_RAW_DATA_DIR.exists():
|
||||
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
|
||||
|
||||
# List downloaded files
|
||||
print("Downloaded files:")
|
||||
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
|
||||
if file.is_file():
|
||||
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during download: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
|
||||
PREPROCESSING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--seed", "42",
|
||||
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge.txt"),
|
||||
"--preprocess_video_batch_size", "1",
|
||||
"--max_height", "480",
|
||||
"--max_width", "832",
|
||||
"--num_frames", "77",
|
||||
"--dataloader_num_workers", "0",
|
||||
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
|
||||
"--train_fps", "16",
|
||||
"--validation_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
|
||||
"--samples_per_file", "1",
|
||||
"--flush_frequency", "1",
|
||||
"--video_length_tolerance_range", "5",
|
||||
"--preprocess_task", "t2v",
|
||||
]
|
||||
|
||||
process = subprocess.run(cmd, check=True)
|
||||
|
||||
|
||||
def run_training():
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
|
||||
TRAINING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", LOCAL_TRAINING_DATA_DIR,
|
||||
"--validation_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
|
||||
"--train_batch_size", "2",
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", "2",
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "2",
|
||||
"--hsdp_shard_dim", "2",
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "10",
|
||||
"--gradient_accumulation_steps", "2",
|
||||
"--max_train_steps", "901",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--checkpointing_steps", "6000",
|
||||
"--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", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "ci_wan_t2v_dmd_overfit",
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--flow_shift", "8",
|
||||
"--validation_guidance_scale", "1.0",
|
||||
"--master_weight_type", "fp32",
|
||||
"--vae_precision", "bf16",
|
||||
"--num_euler_timesteps", "50",
|
||||
"--multi_phased_distill_schedule", "4000-1",
|
||||
"--weight_decay", "0.01",
|
||||
"--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",
|
||||
"--dit_precision", "fp32",
|
||||
"--max_grad_norm", "1.0",
|
||||
"--enable_gradient_checkpointing_type", "full",
|
||||
]
|
||||
|
||||
print(f"Running training with command: {cmd}")
|
||||
process = subprocess.run(cmd, check=True)
|
||||
|
||||
|
||||
def test_e2e_overfit_single_sample():
|
||||
os.environ["WANDB_MODE"] = "online"
|
||||
|
||||
# download_data()
|
||||
# run_preprocessing()
|
||||
run_training()
|
||||
|
||||
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
|
||||
print(f"reference_video_file: {reference_video_file}")
|
||||
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
|
||||
print(f"final_validation_video_file: {final_validation_video_file}")
|
||||
|
||||
|
||||
# Ensure both files exist
|
||||
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
|
||||
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
|
||||
|
||||
# Compute SSIM
|
||||
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
|
||||
reference_video_file,
|
||||
final_validation_video_file,
|
||||
use_ms_ssim=True # Using MS-SSIM for better quality assessment
|
||||
)
|
||||
|
||||
print("\n===== SSIM Results for Step 900 Validation =====")
|
||||
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
|
||||
print(f"Min MS-SSIM: {min_ssim:.4f}")
|
||||
print(f"Max MS-SSIM: {max_ssim:.4f}")
|
||||
|
||||
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_e2e_overfit_single_sample()
|
||||
@@ -1,4 +1,5 @@
|
||||
from .training_pipeline import TrainingPipeline
|
||||
from .wan_training_pipeline import WanTrainingPipeline
|
||||
from .distillation_pipeline import DistillationPipeline
|
||||
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
|
||||
@@ -0,0 +1,951 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import gc
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from abc import abstractmethod
|
||||
from collections import deque
|
||||
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
from diffusers.optimization import get_scheduler
|
||||
from einops import rearrange
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
|
||||
get_local_torch_device, get_world_group, get_sp_parallel_rank, get_sp_world_size)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs,TrainingArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.v1.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
|
||||
normalize_dit_input, save_checkpoint, prepare_for_saving)
|
||||
from fastvideo.v1.utils import set_random_seed, is_vsa_available
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.v1.training.activation_checkpoint import (
|
||||
apply_activation_checkpointing)
|
||||
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.v1.training.training_utils import DiffusionWrapper
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class DistillationPipeline(TrainingPipeline):
|
||||
"""
|
||||
A distillation pipeline for training a student model using teacher model guidance.
|
||||
Inherits from TrainingPipeline to reuse training infrastructure.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae", "teacher_transformer", "critic_transformer"]
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
|
||||
Dict[str, Any]]]
|
||||
current_epoch: int = 0
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize the distillation training pipeline with multiple models."""
|
||||
logger.info("Initializing distillation training pipeline...")
|
||||
|
||||
# 1. Call parent initialization first
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=self.timestep_shift)
|
||||
|
||||
# 2. Distillation-specific initialization
|
||||
# The parent class already sets self.transformer as the student model
|
||||
self.student_transformer = DiffusionWrapper(self.transformer, self.noise_scheduler)
|
||||
self.teacher_transformer = DiffusionWrapper(self.get_module("teacher_transformer"), self.noise_scheduler)
|
||||
self.critic_transformer = DiffusionWrapper(self.get_module("critic_transformer"), self.noise_scheduler)
|
||||
# self.student_transformer.to(torch.bfloat16)
|
||||
# self.teacher_transformer.to(torch.bfloat16)
|
||||
# self.critic_transformer.to(torch.bfloat16)
|
||||
# torch.distributed.breakpoint()
|
||||
self.teacher_transformer.requires_grad_(False)
|
||||
self.teacher_transformer.eval()
|
||||
self.critic_transformer.requires_grad_(True)
|
||||
self.critic_transformer.train()
|
||||
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
self.critic_transformer = apply_activation_checkpointing(
|
||||
self.critic_transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
self.teacher_transformer = apply_activation_checkpointing(
|
||||
self.teacher_transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
# Initialize optimizers
|
||||
critic_params = list(filter(lambda p: p.requires_grad, self.critic_transformer.parameters()))
|
||||
self.critic_transformer_optimizer = torch.optim.AdamW(
|
||||
critic_params,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
self.critic_lr_scheduler = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.critic_transformer_optimizer,
|
||||
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
|
||||
num_training_steps=training_args.max_train_steps * self.world_size,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
logger.info("Distillation optimizers initialized: student and critic")
|
||||
|
||||
self.student_critic_update_ratio = self.training_args.student_critic_update_ratio
|
||||
logger.info(f"Distillation pipeline initialized with student_critic_update_ratio={self.student_critic_update_ratio}")
|
||||
|
||||
self.denoising_step_list = torch.tensor(
|
||||
self.training_args.denoising_step_list, dtype=torch.long, device=get_local_torch_device())
|
||||
logger.info(f"Distillation student model to {len(self.denoising_step_list)} denoising steps")
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
# TODO(yongqi): hardcode for bidirectional distillation
|
||||
self.distill_task_type = "bidirectional_video"
|
||||
self.denoising_loss_type = 'flow'
|
||||
# TODO(yongqi): hardcode for causal distillation
|
||||
self.num_frame_per_block = 3
|
||||
|
||||
self.min_step = int(self.training_args.min_step_ratio * self.num_train_timestep)
|
||||
self.max_step = int(self.training_args.max_step_ratio * self.num_train_timestep)
|
||||
|
||||
self.teacher_guidance_scale = self.training_args.teacher_guidance_scale
|
||||
self.denoising_loss_func = FlowPredLoss()
|
||||
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
"""Initialize validation pipeline - must be implemented by subclasses."""
|
||||
raise NotImplementedError(
|
||||
"Distillation pipelines must implement this method")
|
||||
|
||||
def _prepare_distillation(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Prepare training environment for distillation."""
|
||||
self.student_transformer.requires_grad_(True)
|
||||
self.student_transformer.train()
|
||||
self.critic_transformer.requires_grad_(True)
|
||||
self.critic_transformer.train()
|
||||
|
||||
return training_batch
|
||||
|
||||
def _process_timestep(self, timestep: torch.Tensor, type: str) -> torch.Tensor:
|
||||
"""
|
||||
Pre-process the randomly generated timestep based on the generator's task type.
|
||||
Input:
|
||||
- timestep: [batch_size, num_frame] tensor containing the randomly generated timestep.
|
||||
- type: a string indicating the type of the current model (image, bidirectional_video, or causal_video).
|
||||
Output Behavior:
|
||||
- image: check that the second dimension (num_frame) is 1.
|
||||
- bidirectional_video: broadcast the timestep to be the same for all frames.
|
||||
- causal_video: broadcast the timestep to be the same for all frames **in a block**.
|
||||
"""
|
||||
if type == "image":
|
||||
assert timestep.shape[1] == 1
|
||||
return timestep
|
||||
elif type == "bidirectional_video":
|
||||
# Create a new tensor to avoid in-place operations
|
||||
new_timestep = timestep.clone()
|
||||
for index in range(timestep.shape[0]):
|
||||
new_timestep[index] = timestep[index, 0]
|
||||
return new_timestep
|
||||
elif type == "causal_video":
|
||||
# make the noise level the same within every motion block
|
||||
timestep = timestep.reshape(
|
||||
timestep.shape[0], -1, self.num_frame_per_block)
|
||||
timestep[:, :, 1:] = timestep[:, :, 0:1]
|
||||
timestep = timestep.reshape(timestep.shape[0], -1)
|
||||
return timestep
|
||||
else:
|
||||
raise NotImplementedError("Unsupported model type {}".format(type))
|
||||
|
||||
def _student_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
|
||||
"""Forward pass through student transformer and compute student losses."""
|
||||
latents = training_batch.latents
|
||||
dtype = latents.dtype
|
||||
simulated_noisy_input = []
|
||||
for timestep in self.denoising_step_list:
|
||||
# Use cross-codebase generator for reproducible noise generation
|
||||
noise = torch.randn(
|
||||
self.video_latent_shape, device=self.device, dtype=dtype)
|
||||
if self.sp_world_size > 1:
|
||||
noise = rearrange(noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
noisy_timestep = timestep * torch.ones(
|
||||
self.video_latent_shape_sp[:2], device=self.device, dtype=torch.long)
|
||||
|
||||
if timestep != 0:
|
||||
noisy_video = self.noise_scheduler.add_noise(
|
||||
latents.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
noisy_timestep.flatten(0, 1)
|
||||
).unflatten(0, self.video_latent_shape_sp[:2])
|
||||
else:
|
||||
noisy_video = latents
|
||||
|
||||
simulated_noisy_input.append(noisy_video)
|
||||
|
||||
simulated_noisy_input = torch.stack(simulated_noisy_input, dim=1)
|
||||
|
||||
# Step 2: Randomly sample a timestep and pick the corresponding input
|
||||
# Use cross-codebase generator for reproducible index generation
|
||||
index = torch.randint(0, len(self.denoising_step_list), [
|
||||
self.video_latent_shape_sp[0], self.video_latent_shape_sp[1]], device=self.device, dtype=torch.long)
|
||||
|
||||
index = self._process_timestep(index, type=self.distill_task_type)
|
||||
|
||||
# select the corresponding timestep's noisy input from the stacked tensor [B, T, F, C, H, W]
|
||||
|
||||
noisy_input = torch.gather(
|
||||
simulated_noisy_input, dim=1,
|
||||
index=index.reshape(index.shape[0], 1, index.shape[1], 1, 1, 1).expand(
|
||||
-1, -1, -1, *self.video_latent_shape_sp[2:])
|
||||
).squeeze(1)
|
||||
|
||||
timestep = self.denoising_step_list[index]
|
||||
|
||||
training_batch = self._build_input_kwargs(noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_video = self.student_transformer(training_batch, timestep)
|
||||
|
||||
pred_video = pred_video.type_as(noisy_input)
|
||||
return pred_video, timestep.float().detach()
|
||||
|
||||
def _compute_kl_grad(
|
||||
self, noisy_video: torch.Tensor,
|
||||
estimated_clean_video: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
training_batch: TrainingBatch,
|
||||
normalization: bool = True
|
||||
) -> Tuple[torch.Tensor, dict]:
|
||||
# critic_transformer forward
|
||||
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_fake_video = self.critic_transformer(training_batch, timestep)
|
||||
|
||||
# teacher_transformer cond forward
|
||||
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.conditional_dict, training_batch)
|
||||
pred_real_video_cond = self.teacher_transformer(training_batch, timestep)
|
||||
|
||||
# teacher_transformer uncond forward
|
||||
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.unconditional_dict, training_batch)
|
||||
pred_real_video_uncond = self.teacher_transformer(training_batch, timestep)
|
||||
|
||||
pred_real_video = pred_real_video_cond + (
|
||||
pred_real_video_cond - pred_real_video_uncond
|
||||
) * self.teacher_guidance_scale
|
||||
|
||||
grad = (pred_fake_video - pred_real_video)
|
||||
|
||||
if normalization:
|
||||
p_real = (estimated_clean_video - pred_real_video)
|
||||
normalizer = torch.abs(p_real).mean(dim=[1, 2, 3, 4], keepdim=True)
|
||||
grad = grad / normalizer
|
||||
grad = torch.nan_to_num(grad)
|
||||
|
||||
return grad, {
|
||||
"dmdtrain_latents": estimated_clean_video.detach(),
|
||||
"dmdtrain_noisy_latent": noisy_video.detach(),
|
||||
"dmdtrain_pred_real_video": pred_real_video.detach(),
|
||||
"dmdtrain_pred_fake_video": pred_fake_video.detach(),
|
||||
"dmdtrain_gradient_norm": torch.mean(torch.abs(grad)).detach(),
|
||||
"timestep": timestep.float().detach()
|
||||
}
|
||||
|
||||
def _compute_dmd_loss(self, pred_video: torch.Tensor, training_batch: TrainingBatch) -> Tuple[torch.Tensor, dict]:
|
||||
"""Compute DMD (Diffusion Model Distillation) loss."""
|
||||
|
||||
original_latent = pred_video
|
||||
batch_size, latent_t = self.video_latent_shape_sp[:2]
|
||||
with torch.no_grad():
|
||||
# Use cross-codebase generator for reproducible timestep generation
|
||||
timestep = torch.randint(
|
||||
0,
|
||||
self.num_train_timestep,
|
||||
[batch_size, latent_t],
|
||||
device=self.device,
|
||||
dtype=torch.long
|
||||
)
|
||||
|
||||
timestep = self._process_timestep(
|
||||
timestep, type=self.distill_task_type)
|
||||
|
||||
if self.timestep_shift > 1:
|
||||
timestep = self.timestep_shift * \
|
||||
(timestep / self.num_train_timestep) / \
|
||||
(1 + (self.timestep_shift - 1) * (timestep / self.num_train_timestep)) * self.num_train_timestep
|
||||
|
||||
timestep = timestep.clamp(self.min_step, self.max_step)
|
||||
|
||||
# Use cross-codebase generator for reproducible noise generation
|
||||
noise = torch.randn(
|
||||
self.video_latent_shape, device=self.device, dtype=pred_video.dtype)
|
||||
if self.sp_world_size > 1:
|
||||
noise = rearrange(noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
pred_video.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
timestep.flatten(0, 1)
|
||||
).detach().unflatten(0, (batch_size, latent_t))
|
||||
|
||||
grad, dmd_log_dict = self._compute_kl_grad(
|
||||
noisy_video=noisy_latent,
|
||||
estimated_clean_video=original_latent,
|
||||
timestep=timestep,
|
||||
training_batch=training_batch
|
||||
)
|
||||
|
||||
dmd_loss = 0.5 * F.mse_loss(original_latent.double(
|
||||
), (original_latent.double() - grad.double()).detach(), reduction="mean")
|
||||
|
||||
return dmd_loss, dmd_log_dict
|
||||
|
||||
def _student_forward_and_compute_dmd_loss(self, training_batch: TrainingBatch) -> Tuple[TrainingBatch, torch.Tensor, dict]:
|
||||
"""Forward pass through student transformer and compute student losses."""
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
|
||||
pred_video, timestep_dmd = self._student_forward(training_batch)
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata):
|
||||
dmd_loss, dmd_log_dict = self._compute_dmd_loss(
|
||||
pred_video=pred_video,
|
||||
training_batch=training_batch
|
||||
)
|
||||
|
||||
dmd_log_dict['dmd_timestep_stu'] = timestep_dmd
|
||||
|
||||
return training_batch, dmd_loss, dmd_log_dict
|
||||
|
||||
def _critic_forward_and_compute_loss(self, training_batch: TrainingBatch) -> Tuple[TrainingBatch, torch.Tensor, dict]:
|
||||
with torch.no_grad():
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
|
||||
generated_video, timestep_gen = self._student_forward(training_batch)
|
||||
|
||||
critic_timestep = torch.randint(
|
||||
0,
|
||||
self.num_train_timestep,
|
||||
self.video_latent_shape_sp[:2],
|
||||
device=self.device,
|
||||
dtype=torch.long
|
||||
)
|
||||
critic_timestep = self._process_timestep(
|
||||
critic_timestep, type=self.distill_task_type)
|
||||
|
||||
# TODO: Add timestep warping
|
||||
if self.timestep_shift > 1:
|
||||
critic_timestep = self.timestep_shift * \
|
||||
(critic_timestep / self.num_train_timestep) / (1 + (self.timestep_shift - 1) * (critic_timestep / self.num_train_timestep)) * self.num_train_timestep
|
||||
|
||||
critic_timestep = critic_timestep.clamp(self.min_step, self.max_step)
|
||||
|
||||
# Use cross-codebase generator for reproducible noise generation
|
||||
critic_noise = torch.randn(
|
||||
self.video_latent_shape, device=self.device, dtype=generated_video.dtype)
|
||||
if self.sp_world_size > 1:
|
||||
critic_noise = rearrange(critic_noise,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
critic_noise = critic_noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
noisy_generated_video = self.noise_scheduler.add_noise(
|
||||
generated_video.flatten(0, 1),
|
||||
critic_noise.flatten(0, 1),
|
||||
critic_timestep.flatten(0, 1)
|
||||
).unflatten(0, self.video_latent_shape_sp[:2])
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata):
|
||||
training_batch = self._build_input_kwargs(noisy_generated_video, critic_timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
pred_fake_video = self.critic_transformer(training_batch, critic_timestep)
|
||||
|
||||
# # Step 3: Compute the denoising loss for the fake critic
|
||||
pred_fake_video_noise = DiffusionWrapper._convert_x0_to_flow_pred(
|
||||
x0_pred=pred_fake_video.flatten(0, 1),
|
||||
xt=noisy_generated_video.flatten(0, 1),
|
||||
timestep=critic_timestep.flatten(0, 1),
|
||||
scheduler=self.noise_scheduler
|
||||
)
|
||||
|
||||
denoising_loss = self.denoising_loss_func(
|
||||
x=generated_video.flatten(0, 1),
|
||||
noise=critic_noise.flatten(0, 1),
|
||||
flow_pred=pred_fake_video_noise
|
||||
)
|
||||
|
||||
critic_log_dict = {
|
||||
"critictrain_latent": generated_video.detach(),
|
||||
"critictrain_noisy_latent": noisy_generated_video.detach(),
|
||||
"critictrain_pred_video": pred_fake_video.detach(),
|
||||
"critic_timestep": critic_timestep.float().detach(),
|
||||
"critic_timestep_stu": timestep_gen.float().detach(),
|
||||
}
|
||||
|
||||
return training_batch, denoising_loss, critic_log_dict
|
||||
|
||||
def _clip_grad_norm(self, training_batch: TrainingBatch, transformer) -> TrainingBatch:
|
||||
assert self.training_args is not None
|
||||
max_grad_norm = self.training_args.max_grad_norm
|
||||
|
||||
# TODO(will): perhaps move this into transformer api so that we can do
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
# Clip gradients for both student and critic models
|
||||
model_parts = [transformer]
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
foreach=None,
|
||||
)
|
||||
assert grad_norm is not float('nan') or grad_norm is not float(
|
||||
'inf')
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
training_batch.grad_norm = grad_norm
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
super()._prepare_dit_inputs(training_batch)
|
||||
conditional_dict = {
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
|
||||
training_batch.conditional_dict = conditional_dict
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
assert training_batch.latents is not None
|
||||
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
self.video_latent_shape = training_batch.latents.shape # [B, C, T, H, W]
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
|
||||
if self.sp_world_size > 1:
|
||||
training_batch.latents = rearrange(training_batch.latents,
|
||||
"b (n t) c h w -> b n t c h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
training_batch.latents = training_batch.latents[:, self.rank_in_sp_group, :, :, :, :]
|
||||
|
||||
self.video_latent_shape_sp = training_batch.latents.shape
|
||||
|
||||
return training_batch
|
||||
|
||||
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Train one step with alternating student and critic updates, supporting gradient accumulation."""
|
||||
import copy
|
||||
gradient_accumulation_steps = getattr(self.training_args, 'gradient_accumulation_steps', 1)
|
||||
batches = []
|
||||
# Collect N batches for gradient accumulation
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
batch = self._prepare_distillation(training_batch)
|
||||
batch = self._get_next_batch(batch)
|
||||
batch = self._normalize_dit_input(batch)
|
||||
batch = self._prepare_dit_inputs(batch)
|
||||
batch = self._build_attention_metadata(batch)
|
||||
batch.attn_metadata_vsa = copy.deepcopy(batch.attn_metadata)
|
||||
if batch.attn_metadata is not None:
|
||||
batch.attn_metadata.VSA_sparsity = 0.0
|
||||
batches.append(batch)
|
||||
|
||||
# Student accumulation
|
||||
self.optimizer.zero_grad()
|
||||
total_dmd_loss = 0.0
|
||||
total_dmd_log_dict = None
|
||||
if (self.current_trainstep % self.student_critic_update_ratio == 0):
|
||||
for batch in batches:
|
||||
batch_stu = copy.deepcopy(batch)
|
||||
batch_stu, dmd_loss, dmd_log_dict = self._student_forward_and_compute_dmd_loss(batch_stu)
|
||||
# Ensure backward is under the correct forward context
|
||||
with set_forward_context(
|
||||
current_timestep=batch_stu.timesteps, attn_metadata=batch_stu.attn_metadata):
|
||||
(dmd_loss / gradient_accumulation_steps).backward()
|
||||
total_dmd_loss += dmd_loss.detach().item()
|
||||
if total_dmd_log_dict is None:
|
||||
total_dmd_log_dict = dmd_log_dict
|
||||
# Only keep the first log dict, ignore subsequent ones
|
||||
self._clip_grad_norm(batch_stu, self.student_transformer)
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
avg_dmd_loss = torch.tensor(total_dmd_loss / gradient_accumulation_steps, device=self.device)
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_dmd_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
training_batch.student_loss = avg_dmd_loss.item()
|
||||
training_batch.dmd_log_dict = total_dmd_log_dict if total_dmd_log_dict is not None else {}
|
||||
else:
|
||||
training_batch.student_loss = 0.0
|
||||
training_batch.dmd_log_dict = {}
|
||||
|
||||
# Critic accumulation
|
||||
self.critic_transformer_optimizer.zero_grad()
|
||||
total_critic_loss = 0.0
|
||||
total_critic_log_dict = None
|
||||
for batch in batches:
|
||||
batch_critic = copy.deepcopy(batch)
|
||||
batch_critic, critic_loss, critic_log_dict = self._critic_forward_and_compute_loss(batch_critic)
|
||||
with set_forward_context(
|
||||
current_timestep=batch_critic.timesteps, attn_metadata=batch_critic.attn_metadata):
|
||||
(critic_loss / gradient_accumulation_steps).backward()
|
||||
total_critic_loss += critic_loss.detach().item()
|
||||
if total_critic_log_dict is None:
|
||||
total_critic_log_dict = critic_log_dict
|
||||
# Only keep the first log dict, ignore subsequent ones
|
||||
self._clip_grad_norm(batch_critic, self.critic_transformer)
|
||||
self.critic_transformer_optimizer.step()
|
||||
self.critic_lr_scheduler.step()
|
||||
self.critic_transformer_optimizer.zero_grad(set_to_none=True)
|
||||
avg_critic_loss = torch.tensor(total_critic_loss / gradient_accumulation_steps, device=self.device)
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_critic_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
training_batch.critic_loss = avg_critic_loss.item()
|
||||
training_batch.critic_log_dict = total_critic_log_dict if total_critic_log_dict is not None else {}
|
||||
|
||||
training_batch.total_loss = training_batch.student_loss + training_batch.critic_loss
|
||||
return training_batch
|
||||
|
||||
def _resume_from_checkpoint(self) -> None: #TODO(yongqi)
|
||||
"""Resume training from checkpoint with distillation models."""
|
||||
assert self.training_args is not None
|
||||
logger.info("Loading distillation checkpoint from %s",
|
||||
self.training_args.resume_from_checkpoint)
|
||||
|
||||
resumed_step = load_checkpoint(
|
||||
self.student_transformer.model, self.global_rank,
|
||||
self.training_args.resume_from_checkpoint, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
self.noise_random_generator)
|
||||
|
||||
# TODO: Add checkpoint loading for critic and teacher models
|
||||
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
logger.info("Successfully resumed from step %s", resumed_step)
|
||||
else:
|
||||
logger.warning("Failed to load checkpoint, starting from step 0")
|
||||
self.init_steps = -1
|
||||
|
||||
def _log_training_info(self) -> None:
|
||||
"""Log distillation-specific training information."""
|
||||
# First call parent class method to get basic training info
|
||||
super()._log_training_info()
|
||||
|
||||
# Then add distillation-specific information
|
||||
logger.info("Distillation-specific settings:")
|
||||
logger.info(" Student/Critic update ratio: %s", self.student_critic_update_ratio)
|
||||
assert isinstance(self.training_args, TrainingArgs)
|
||||
logger.info(" Max gradient norm: %s", self.training_args.max_grad_norm)
|
||||
assert self.teacher_transformer is not None
|
||||
logger.info(" Teacher transformer parameters: %s B",
|
||||
sum(p.numel() for p in self.teacher_transformer.parameters()) / 1e9)
|
||||
assert self.critic_transformer is not None
|
||||
logger.info(" Critic transformer parameters: %s B",
|
||||
sum(p.numel() for p in self.critic_transformer.parameters()) / 1e9)
|
||||
|
||||
def add_visualization(self, generator_log_dict: Dict[str, Any], critic_log_dict: Dict[str, Any], training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
wandb_loss_dict = {}
|
||||
|
||||
# Clear GPU cache before VAE decoding to prevent OOM
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# # Use consistent decoding approach - use decode_stage for all
|
||||
# decode_stage = self.validation_pipeline._stages[-1]
|
||||
|
||||
# Process critic training data
|
||||
critic_latents_name = ['critictrain_latent', 'critictrain_noisy_latent', 'critictrain_pred_video']
|
||||
# critic_latents_name = ['critictrain_pred_video']
|
||||
|
||||
for latent_key in critic_latents_name:
|
||||
latents = critic_log_dict[latent_key] # bs, t,c, h, w
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
# decoded_latent = decode_stage(ForwardBatch(data_type="video", latents=latents), training_args)
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
|
||||
wandb_loss_dict[latent_key] = prepare_for_saving(video)
|
||||
# Clean up references
|
||||
del video, latents
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Process DMD training data if available - use decode_stage instead of self.vae.decode
|
||||
if 'dmdtrain_pred_fake_video' in generator_log_dict:
|
||||
dmd_latents_name = ['dmdtrain_pred_fake_video', 'dmdtrain_pred_real_video', 'dmdtrain_latents', 'dmdtrain_noisy_latent']
|
||||
for latent_key in dmd_latents_name:
|
||||
latents = generator_log_dict[latent_key]
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
# decoded_latent = decode_stage(ForwardBatch(data_type="video", latents=latents), training_args)
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
video = self.vae.decode(latents)
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
video = video.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
|
||||
wandb_loss_dict[latent_key] = prepare_for_saving(video)
|
||||
# Clean up references
|
||||
del video, latents
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Log to wandb
|
||||
if self.global_rank == 0:
|
||||
wandb.log(wandb_loss_dict, step=step)
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
assert training_args is not None
|
||||
training_args.inference_mode = True
|
||||
training_args.use_cpu_offload = True
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
raise ValueError("Validation pipeline is not set")
|
||||
|
||||
logger.info("Starting validation")
|
||||
|
||||
# Create sampling parameters if not provided
|
||||
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
|
||||
|
||||
# Set deterministic seed for validation
|
||||
# set_random_seed(self.seed)
|
||||
logger.info("Using validation seed: %s", self.seed)
|
||||
|
||||
# Prepare validation prompts
|
||||
logger.info('rank: %s: fastvideo_args.validation_dataset_file: %s',
|
||||
self.global_rank,
|
||||
training_args.validation_dataset_file,
|
||||
local_main_process_only=False)
|
||||
validation_dataset = ValidationDataset(
|
||||
training_args.validation_dataset_file)
|
||||
validation_dataloader = DataLoader(validation_dataset,
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
transformer.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
validation_steps = [step for step in validation_steps if step > 0]
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
# Process each validation prompt for each validation step
|
||||
for num_inference_steps in validation_steps:
|
||||
logger.info("rank: %s: num_inference_steps: %s",
|
||||
self.global_rank,
|
||||
num_inference_steps,
|
||||
local_main_process_only=False)
|
||||
step_videos: list[np.ndarray] = []
|
||||
step_captions: list[str] = []
|
||||
|
||||
for validation_batch in validation_dataloader:
|
||||
batch = self._prepare_validation_batch(sampling_param,
|
||||
training_args,
|
||||
validation_batch,
|
||||
num_inference_steps)
|
||||
|
||||
negative_prompt = batch.negative_prompt
|
||||
batch_negative = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=negative_prompt,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
result_batch = self.validation_pipeline.prompt_encoding_stage(batch_negative, training_args)
|
||||
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
|
||||
# logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
|
||||
# self.global_rank,
|
||||
# self.rank_in_sp_group,
|
||||
# batch.prompt,
|
||||
# local_main_process_only=False)
|
||||
|
||||
assert batch.prompt is not None and isinstance(
|
||||
batch.prompt, str)
|
||||
step_captions.append(batch.prompt)
|
||||
|
||||
# # # Run validation inference
|
||||
with torch.no_grad():
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in video:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
step_videos.append(frames)
|
||||
|
||||
# Log validation results for this step
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
|
||||
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
|
||||
# results to global rank 0
|
||||
if self.rank_in_sp_group == 0:
|
||||
if self.global_rank == 0:
|
||||
# Global rank 0 collects results from all sp_group leaders
|
||||
all_videos = step_videos # Start with own results
|
||||
all_captions = step_captions
|
||||
|
||||
# Receive from other sp_group leaders
|
||||
for sp_group_idx in range(1, num_sp_groups):
|
||||
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
|
||||
recv_videos = world_group.recv_object(src=src_rank)
|
||||
recv_captions = world_group.recv_object(src=src_rank)
|
||||
all_videos.extend(recv_videos)
|
||||
all_captions.extend(recv_captions)
|
||||
|
||||
video_filenames = []
|
||||
for i, (video, caption) in enumerate(
|
||||
zip(all_videos, all_captions, strict=True)):
|
||||
os.makedirs(training_args.output_dir, exist_ok=True)
|
||||
filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
|
||||
)
|
||||
imageio.mimsave(filename, video, fps=sampling_param.fps)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
f"validation_videos_{num_inference_steps}_steps": [
|
||||
wandb.Video(filename, caption=caption)
|
||||
for filename, caption in zip(
|
||||
video_filenames, all_captions, strict=True)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
|
||||
# Save all prompts from all cards to txt file
|
||||
# prompt_filename = os.path.join(
|
||||
# training_args.output_dir,
|
||||
# f"validation_step_{global_step}_inference_steps_{num_inference_steps}_prompts.txt"
|
||||
# )
|
||||
# with open(prompt_filename, 'w', encoding='utf-8') as f:
|
||||
# for i, caption in enumerate(all_captions):
|
||||
# f.write(f"Video_{i}: {caption}\n")
|
||||
# logger.info(f"Saved {len(all_captions)} prompts to {prompt_filename}")
|
||||
|
||||
else:
|
||||
# Other sp_group leaders send their results to global rank 0
|
||||
world_group.send_object(step_videos, dst=0)
|
||||
world_group.send_object(step_captions, dst=0)
|
||||
|
||||
# Re-enable gradients for training
|
||||
transformer.train()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def train(self) -> None:
|
||||
"""Main training loop with distillation-specific logging."""
|
||||
assert self.training_args is not None
|
||||
|
||||
assert self.training_args.seed is not None, "seed must be set"
|
||||
seed = self.training_args.seed
|
||||
|
||||
# Set the same seed within each SP group to ensure reproducibility
|
||||
if self.sp_world_size > 1:
|
||||
# Use the same seed for all processes within the same SP group
|
||||
sp_group_seed = seed + (self.global_rank // self.sp_world_size)
|
||||
set_random_seed(sp_group_seed)
|
||||
logger.info(f"Rank {self.global_rank}: Using SP group seed {sp_group_seed}")
|
||||
else:
|
||||
set_random_seed(seed + self.global_rank)
|
||||
|
||||
self.noise_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(seed)
|
||||
|
||||
self.validation_generator = torch.Generator(device=get_local_torch_device()).manual_seed(42)
|
||||
|
||||
logger.info("Initialized random seeds with seed: %s", seed)
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
self._log_training_info()
|
||||
self._log_validation(self.student_transformer, self.training_args, 0)
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, self.training_args.max_train_steps),
|
||||
initial=self.init_steps,
|
||||
desc="Steps",
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
|
||||
for step in range(self.init_steps + 1,
|
||||
self.training_args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
current_vsa_sparsity = self.training_args.VSA_sparsity if vsa_available else 0.0
|
||||
|
||||
training_batch = TrainingBatch()
|
||||
self.current_trainstep = step
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
total_loss = training_batch.total_loss
|
||||
student_loss = training_batch.student_loss
|
||||
critic_loss = training_batch.critic_loss
|
||||
grad_norm = training_batch.grad_norm
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"total_loss": f"{total_loss:.4f}",
|
||||
"student_loss": f"{student_loss:.4f}",
|
||||
"critic_loss": f"{critic_loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
|
||||
if self.global_rank == 0:
|
||||
# Prepare logging data
|
||||
log_data = {
|
||||
"train_total_loss": total_loss,
|
||||
"train_student_loss": student_loss,
|
||||
"train_critic_loss": critic_loss,
|
||||
"learning_rate": self.lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
}
|
||||
|
||||
# Add DMD training metrics if available
|
||||
if hasattr(training_batch, 'dmd_log_dict') and training_batch.dmd_log_dict:
|
||||
dmd_metrics = {
|
||||
"dmd_gradient_norm": training_batch.dmd_log_dict.get("dmdtrain_gradient_norm", 0.0),
|
||||
"dmd_timestep": training_batch.dmd_log_dict.get("timestep", 0.0).mean().item(),
|
||||
"dmd_timestep_stu": training_batch.dmd_log_dict.get("dmd_timestep_stu", 0.0).mean().item()
|
||||
}
|
||||
log_data.update(dmd_metrics)
|
||||
|
||||
# Add critic training metrics if available
|
||||
if hasattr(training_batch, 'critic_log_dict') and training_batch.critic_log_dict:
|
||||
critic_metrics = {
|
||||
"critic_timestep": training_batch.critic_log_dict.get("critic_timestep", 0.0).mean().item(),
|
||||
"critic_timestep_stu": training_batch.critic_log_dict.get("critic_timestep_stu", 0.0).mean().item(),
|
||||
}
|
||||
log_data.update(critic_metrics)
|
||||
wandb.log(log_data, step=step)
|
||||
|
||||
if step % self.training_args.checkpointing_steps == 0:
|
||||
print("rank", self.global_rank, "save checkpoint at step", step)
|
||||
save_checkpoint(self.transformer, self.global_rank, #TODO(yongqi)
|
||||
self.training_args.output_dir, step,
|
||||
self.optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, self.noise_random_generator)
|
||||
if self.transformer:
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
|
||||
|
||||
self.add_visualization(training_batch.dmd_log_dict, training_batch.critic_log_dict, self.training_args, step)
|
||||
self._log_validation(self.student_transformer, self.training_args, step)
|
||||
|
||||
|
||||
|
||||
wandb.finish()
|
||||
# save_checkpoint(self.student_transformer.model, self.global_rank,
|
||||
# self.training_args.output_dir,
|
||||
# self.training_args.max_train_steps, self.optimizer,
|
||||
# self.train_dataloader, self.lr_scheduler,
|
||||
# self.noise_random_generator)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
class FlowPredLoss():
|
||||
def __call__(
|
||||
self, x: torch.Tensor,
|
||||
noise: torch.Tensor,
|
||||
flow_pred: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
return torch.mean((flow_pred - (noise - x)) ** 2)
|
||||
@@ -5,7 +5,7 @@ import time
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import deque
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
from typing import Any, Dict
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
@@ -39,7 +39,7 @@ from fastvideo.v1.training.activation_checkpoint import (
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
|
||||
normalize_dit_input, save_checkpoint, shard_latents_across_sp)
|
||||
normalize_dit_input, save_checkpoint, shard_latents_across_sp, prepare_for_saving)
|
||||
from fastvideo.v1.utils import is_vsa_available, set_random_seed, shallow_asdict
|
||||
|
||||
import wandb # isort: skip
|
||||
@@ -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:
|
||||
@@ -242,23 +243,22 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
training_batch.timesteps = timesteps
|
||||
training_batch.sigmas = sigmas
|
||||
training_batch.noise = noise
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_attention_metadata(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
assert self.training_args is not None
|
||||
latents = training_batch.latents
|
||||
assert latents is not None
|
||||
latents_shape = training_batch.raw_latent_shape
|
||||
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],
|
||||
latents.shape[3] // patch_size[1],
|
||||
latents.shape[4] // patch_size[2]
|
||||
latents_shape[2] // patch_size[0],
|
||||
latents_shape[3] // patch_size[1],
|
||||
latents_shape[4] // patch_size[2]
|
||||
]
|
||||
training_batch.attn_metadata = VideoSparseAttentionMetadata(
|
||||
current_timestep=training_batch.timesteps,
|
||||
@@ -319,12 +319,14 @@ 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
|
||||
|
||||
# make sure no implicit broadcasting happens
|
||||
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
|
||||
|
||||
loss = (torch.mean((model_pred.float() - target.float())**2) /
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
|
||||
@@ -439,7 +441,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
self._log_training_info()
|
||||
self._log_validation(self.transformer, self.training_args, 1)
|
||||
self._log_validation(self.transformer, self.training_args, 0)
|
||||
|
||||
# Train!
|
||||
progress_bar = tqdm(
|
||||
@@ -449,7 +451,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
# Only show the progress bar once on each machine.
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
for step in range(self.init_steps + 1,
|
||||
for step in range(self.init_steps,
|
||||
self.training_args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
if vsa_available:
|
||||
@@ -467,6 +469,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
training_batch.current_vsa_sparsity = current_vsa_sparsity
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
|
||||
loss = training_batch.total_loss
|
||||
grad_norm = training_batch.grad_norm
|
||||
|
||||
@@ -705,3 +708,4 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
transformer.train()
|
||||
|
||||
|
||||
@@ -3,13 +3,16 @@ import json
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dcp
|
||||
from torchvision.utils import make_grid
|
||||
from einops import rearrange
|
||||
from safetensors.torch import save_file
|
||||
import wandb
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
@@ -18,6 +21,8 @@ from fastvideo.v1.training.checkpointing_utils import (ModelWrapper,
|
||||
OptimizerWrapper,
|
||||
RandomStateWrapper,
|
||||
SchedulerWrapper)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from abc import ABC
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -550,3 +555,80 @@ def convert_custom_format_to_diffusers_format(state_dict: dict[str, Any],
|
||||
new_state_dict[training_key] = v
|
||||
|
||||
return new_state_dict
|
||||
|
||||
def prepare_for_saving(tensor: torch.Tensor, fps: int = 16, caption: str | None = None) -> wandb.Image | wandb.Video:
|
||||
if tensor.ndim == 4:
|
||||
# Assuming it's an image and has shape [batch_size, 3, height, width]
|
||||
tensor = make_grid(tensor, 4, padding=0, normalize=False)
|
||||
return wandb.Image((tensor * 255).numpy().astype(np.uint8), caption=caption)
|
||||
elif tensor.ndim == 5:
|
||||
# Assuming it's a video and has shape [batch_size, num_frames, 3, height, width]
|
||||
return wandb.Video((tensor * 255).numpy().astype(np.uint8), fps=fps, format="webm", caption=caption)
|
||||
else:
|
||||
raise ValueError("Unsupported tensor shape for saving. Expected 4D (image) or 5D (video) tensor.")
|
||||
|
||||
class DiffusionWrapper(torch.nn.Module, ABC):
|
||||
def __init__(self, transformer, scheduler):
|
||||
super().__init__()
|
||||
self.model = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
def forward(self, training_batch: TrainingBatch, timestep: torch.Tensor):
|
||||
pred_noise = self.model(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
pred_video = self._convert_flow_pred_to_x0(
|
||||
flow_pred=pred_noise.flatten(0, 1),
|
||||
xt=training_batch.noise_latents.flatten(0, 1),
|
||||
timestep=timestep.flatten(0, 1),
|
||||
scheduler=self.scheduler
|
||||
).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
return pred_video
|
||||
|
||||
@staticmethod
|
||||
def _convert_x0_to_flow_pred(x0_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor, scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert x0 prediction to flow matching's prediction.
|
||||
x0_pred: the x0 prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = (x_t - x_0) / sigma_t
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = x0_pred.dtype
|
||||
x0_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(x0_pred.device), [x0_pred, xt,
|
||||
scheduler.sigmas,
|
||||
scheduler.timesteps]
|
||||
)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
flow_pred = (xt - x0_pred) / sigma_t
|
||||
return flow_pred.to(original_dtype)
|
||||
|
||||
@staticmethod
|
||||
def _convert_flow_pred_to_x0(flow_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor, scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert flow matching's prediction to x0 prediction.
|
||||
flow_pred: the prediction with shape [B, C, H, W]
|
||||
xt: the input noisy data with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [B]
|
||||
|
||||
pred = noise - x0
|
||||
x_t = (1-sigma_t) * x0 + sigma_t * noise
|
||||
we have x0 = x_t - sigma_t * pred
|
||||
see derivations https://chatgpt.com/share/67bf8589-3d04-8008-bc6e-4cf1a24e2d0e
|
||||
"""
|
||||
# use higher precision for calculations
|
||||
original_dtype = flow_pred.dtype
|
||||
flow_pred, xt, sigmas, timesteps = map(
|
||||
lambda x: x.double().to(flow_pred.device), [flow_pred, xt,
|
||||
scheduler.sigmas,
|
||||
scheduler.timesteps]
|
||||
)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
x0_pred = xt - sigma_t * flow_pred
|
||||
return x0_pred.to(original_dtype)
|
||||
@@ -0,0 +1,93 @@
|
||||
# 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_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
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"] = FlowMatchEulerDiscreteScheduler(
|
||||
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 = True
|
||||
args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
validation_pipeline = WanDmdPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy,
|
||||
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()
|
||||
main(args)
|
||||
@@ -0,0 +1,238 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.distributed import get_local_torch_device, get_sp_parallel_rank, get_sp_world_size
|
||||
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
|
||||
|
||||
if self.sp_world_size > 1:
|
||||
image_latents = rearrange(image_latents,
|
||||
"b c (n t) h w -> b c n t h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
image_latents = image_latents[:, :, self.rank_in_sp_group, :, :, :]
|
||||
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), # bs, c, t, h, w
|
||||
"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)
|
||||
@@ -28,7 +28,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Wan.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer"]
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
|
||||
@@ -40,7 +40,7 @@ class WanTrainingPipeline(TrainingPipeline):
|
||||
args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
validation_pipeline = WanPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
args=args_copy,
|
||||
inference_mode=True,
|
||||
loaded_modules={"transformer": self.get_module("transformer")},
|
||||
tp_size=training_args.tp_size,
|
||||
|
||||
@@ -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 \
|
||||
@@ -0,0 +1,67 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/val/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
NUM_GPUS=8
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 8 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 29 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
|
||||
# validation_preprocessed_path
|
||||
@@ -0,0 +1,66 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/test_8/
|
||||
NUM_GPUS=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.1-T2V-14B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-14B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_preprocessed_path "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 20 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 768 \
|
||||
--num_width 1280 \
|
||||
--num_frames 77 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
--VSA_sparsity 0.0 \
|
||||
@@ -0,0 +1,66 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
NUM_GPUS=1
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_distillation_pipeline.py \
|
||||
--model_path data/Wan2.1-T2V-1.3B-Diffusers-VT \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path data/Wan2.1-T2V-1.3B-Diffusers-VT \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim $NUM_GPUS \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 0 \
|
||||
--gradient_accumulation_steps 16 \
|
||||
--max_train_steps 3000 \
|
||||
--learning_rate 4e-6 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 10 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 61 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 5 \
|
||||
--denoising_step_list '1000,757,522' \
|
||||
--min_step_ratio 0.02 \
|
||||
--max_step_ratio 0.98 \
|
||||
--teacher_guidance_scale 3.5 \
|
||||
--enable_gradient_checkpointing_type "full" \
|
||||
--seed 1000 \
|
||||
--VSA_sparsity 0.9 \
|
||||
@@ -0,0 +1,67 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=mini_i2v_dataset/crush-smol_preprocessed/combined_parquet_dataset
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/val/
|
||||
VALIDATION_DIR=mini_i2v_dataset/crush-smol_raw/validation.json
|
||||
NUM_GPUS=2
|
||||
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 1 \
|
||||
--max_train_steps 5 \
|
||||
--learning_rate 1e-5 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 10 \
|
||||
--validation_steps 3 \
|
||||
--validation_sampling_steps "3" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.0 \
|
||||
--output_dir "outputs_dmd/wan_finetune" \
|
||||
--tracker_project_name Wan_distillation \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 29 \
|
||||
--flow_shift 8 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--vae_precision "bf16" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0 \
|
||||
--student_critic_update_ratio 3 \
|
||||
--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 1024 \
|
||||
|
||||
# validation_preprocessed_path
|
||||
@@ -0,0 +1,66 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
Reference in New Issue
Block a user