Compare commits

...
3 Commits
Author SHA1 Message Date
“BrianChen1129” 586d144a0f update pre commit 2025-07-27 05:56:13 +00:00
“BrianChen1129” 3601e33c4a update pre commit 2025-07-27 05:54:53 +00:00
“BrianChen1129” 6345b6bff7 add DMD training pipeline 2025-07-27 05:08:10 +00:00
17 changed files with 1406 additions and 20 deletions
+30
View File
@@ -655,6 +655,12 @@ class TrainingArgs(FastVideoArgs):
lora_alpha: int | None = None
lora_training: bool = False
# distillation args
generator_update_interval: int = 5
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
@@ -954,4 +960,28 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
# Distillation arguments
parser.add_argument("--generator-update-interval",
type=int,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
help="Minimum step ratio")
parser.add_argument("--max-timestep-ratio",
type=float,
default=TrainingArgs.max_timestep_ratio,
help="Maximum step ratio")
parser.add_argument("--real-score-guidance-scale",
type=float,
default=TrainingArgs.real_score_guidance_scale,
help="Teacher guidance scale")
return parser
def parse_int_list(value: str) -> list[int]:
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"),
"real_score_transformer": (TransformerLoader, "diffusers"),
"fake_score_transformer": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
+11 -3
View File
@@ -242,8 +242,11 @@ class ComposedPipelineBase(ABC):
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(
"model_index.json does not contain a %s module, adding %s to model_index",
module_name, module_name)
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
@@ -259,7 +262,12 @@ 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,
@@ -148,6 +148,8 @@ class TrainingBatch:
# 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
@@ -155,6 +157,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
@@ -162,6 +165,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
@@ -173,3 +177,13 @@ 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
generator_loss: float = 0.0
fake_score_loss: float = 0.0
+4 -7
View File
@@ -119,13 +119,13 @@ 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
@@ -758,9 +758,7 @@ class DmdDenoisingStage(DenoisingStage):
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
pred_video.shape[:2],
dtype=torch.long,
device=pred_video.device)
[1], dtype=torch.long, device=pred_video.device)
noise = torch.randn(video_raw_latent_shape,
device=self.device,
dtype=pred_video.dtype)
@@ -771,8 +769,7 @@ class DmdDenoisingStage(DenoisingStage):
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])
next_timestep).unflatten(0, pred_video.shape[:2])
else:
latents = pred_video
@@ -0,0 +1,81 @@
# 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.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ImageEncodingStage, ConditioningStage,
DecodingStage, DmdDenoisingStage,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
logger = init_logger(__name__)
class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanImageToVideoDmdPipeline
+2 -1
View File
@@ -1,4 +1,5 @@
from .distillation_pipeline import DistillationPipeline
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
+801
View File
@@ -0,0 +1,801 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import os
import time
from abc import abstractmethod
from collections import deque
from collections.abc import Iterator
from typing import Any
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.envs as envs
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_local_torch_device, get_sp_group,
get_world_group)
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases, load_checkpoint,
pred_noise_to_pred_video, save_checkpoint, shift_timestep)
from fastvideo.utils import is_vsa_available, set_random_seed
import wandb # isort: skip
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class DistillationPipeline(TrainingPipeline):
"""
A distillation pipeline for training a 3 step model.
Inherits from TrainingPipeline to reuse training infrastructure.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
video_latent_shape: tuple[int, ...]
video_latent_shape_sp: tuple[int, ...]
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 pipeline...")
super().initialize_training_pipeline(training_args)
self.noise_scheduler = self.get_module("scheduler")
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
assert self.training_args is not None
self.timestep_shift = self.training_args.pipeline_config.flow_shift
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
# self.transformer is the generator model
self.real_score_transformer = self.get_module("real_score_transformer")
self.fake_score_transformer = self.get_module("fake_score_transformer")
self.real_score_transformer.requires_grad_(False)
self.real_score_transformer.eval()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.fake_score_transformer = apply_activation_checkpointing(
self.fake_score_transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
self.real_score_transformer = apply_activation_checkpointing(
self.real_score_transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
# Initialize optimizers
fake_score_params = list(
filter(lambda p: p.requires_grad,
self.fake_score_transformer.parameters()))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.fake_score_lr_scheduler = get_scheduler(
training_args.lr_scheduler,
optimizer=self.fake_score_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: generator and fake_score")
self.generator_update_interval = self.training_args.generator_update_interval
logger.info(
"Distillation pipeline initialized with generator_update_interval=%s",
self.generator_update_interval)
self.denoising_step_list = torch.tensor(
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps",
len(self.denoising_step_list))
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.min_timestep = int(self.training_args.min_timestep_ratio *
self.num_train_timestep)
self.max_timestep = int(self.training_args.max_timestep_ratio *
self.num_train_timestep)
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
@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.transformer.requires_grad_(True)
self.transformer.train()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
return training_batch
def _build_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
training_batch: TrainingBatch) -> TrainingBatch:
if text_dict is None:
raise ValueError(
"text_dict cannot be None for distillation pipeline")
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,
"return_dict": False,
}
training_batch.noise_latents = noise_input
# After setting noise_latents, it's guaranteed to be not None
return training_batch
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
assert training_batch.latents is not None
latents = training_batch.latents
dtype = latents.dtype
index = torch.randint(0,
len(self.denoising_step_list), [1],
device=self.device,
dtype=torch.long)
timestep = self.denoising_step_list[index]
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_latent = self.noise_scheduler.add_noise(latents.flatten(0, 1),
noise.flatten(0, 1),
timestep).unflatten(
0,
(1, latents.shape[1]))
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
assert training_batch.noise_latents is not None
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
return pred_video
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
with torch.no_grad():
timestep = torch.randint(0,
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
assert self.timestep_shift is not None
timestep = shift_timestep(
timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
timestep = timestep.clamp(self.min_timestep, self.max_timestep)
noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_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(
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
fake_score_pred_noise = self.fake_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
assert training_batch.noise_latents is not None
pred_fake_video = pred_noise_to_pred_video(
pred_noise=fake_score_pred_noise.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(
0, fake_score_pred_noise.shape[:2])
# real_score_transformer cond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
real_score_pred_noise_cond = self.real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
assert training_batch.noise_latents is not None
pred_real_video_cond = pred_noise_to_pred_video(
pred_noise=real_score_pred_noise_cond.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(
0, real_score_pred_noise_cond.shape[:2])
# real_score_transformer uncond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.unconditional_dict,
training_batch)
real_score_pred_noise_uncond = self.real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
assert training_batch.noise_latents is not None
pred_real_video_uncond = pred_noise_to_pred_video(
pred_noise=real_score_pred_noise_uncond.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(
0, real_score_pred_noise_uncond.shape[:2])
real_score_pred_video = pred_real_video_cond + (
pred_real_video_cond -
pred_real_video_uncond) * self.real_score_guidance_scale
grad = (pred_fake_video - real_score_pred_video) / torch.abs(
generator_pred_video - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
generator_pred_video.float(),
(generator_pred_video.float() - grad.float()).detach())
return dmd_loss
def faker_score_forward(
self, training_batch: TrainingBatch
) -> tuple[TrainingBatch, torch.Tensor]:
with torch.no_grad(), set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
generator_pred_video = self._generator_forward(training_batch)
fake_score_timestep = torch.randint(0,
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
assert self.timestep_shift is not None
fake_score_timestep = shift_timestep(
fake_score_timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
fake_score_timestep = fake_score_timestep.clamp(self.min_timestep,
self.max_timestep)
fake_score_noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_pred_video.dtype)
if self.sp_world_size > 1:
fake_score_noise = rearrange(fake_score_noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
fake_score_noise = fake_score_noise[:, self.
rank_in_sp_group, :, :, :, :]
noisy_generator_pred_video = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
fake_score_timestep).unflatten(0,
(1, generator_pred_video.shape[1]))
with set_forward_context(current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
training_batch = self._build_distill_input_kwargs(
noisy_generator_pred_video, fake_score_timestep,
training_batch.conditional_dict, training_batch)
fake_score_pred_noise = self.fake_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
target = fake_score_noise - generator_pred_video
denoising_loss = torch.mean((fake_score_pred_noise - target)**2)
return training_batch, denoising_loss
def _clip_model_grad_norm_(self, training_batch: TrainingBatch,
transformer) -> TrainingBatch:
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
if max_grad_norm is not None:
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
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:
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 # type: ignore
batches.append(batch)
self.optimizer.zero_grad()
total_dmd_loss = 0.0
if (self.current_trainstep % self.generator_update_interval == 0):
for batch in batches:
batch_stu = copy.deepcopy(batch)
with set_forward_context(
current_timestep=batch_stu.timesteps,
attn_metadata=batch_stu.attn_metadata_vsa):
generator_pred_video = self._generator_forward(batch_stu)
with set_forward_context(current_timestep=batch_stu.timesteps,
attn_metadata=batch_stu.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=batch_stu)
with set_forward_context(
current_timestep=batch_stu.timesteps,
attn_metadata=batch_stu.attn_metadata_vsa):
(dmd_loss / gradient_accumulation_steps).backward()
total_dmd_loss += dmd_loss.detach().item()
self._clip_model_grad_norm_(batch_stu, self.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.generator_loss = avg_dmd_loss.item()
else:
training_batch.generator_loss = 0.0
self.fake_score_optimizer.zero_grad()
total_fake_score_loss = 0.0
for batch in batches:
batch_critic = copy.deepcopy(batch)
batch_critic, fake_score_loss = self.faker_score_forward(
batch_critic)
with set_forward_context(current_timestep=batch_critic.timesteps,
attn_metadata=batch_critic.attn_metadata):
(fake_score_loss / gradient_accumulation_steps).backward()
total_fake_score_loss += fake_score_loss.detach().item()
self._clip_model_grad_norm_(batch_critic, self.fake_score_transformer)
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
self.fake_score_optimizer.zero_grad(set_to_none=True)
avg_fake_score_loss = torch.tensor(total_fake_score_loss /
gradient_accumulation_steps,
device=self.device)
world_group = get_world_group()
world_group.all_reduce(avg_fake_score_loss,
op=torch.distributed.ReduceOp.AVG)
training_batch.fake_score_loss = avg_fake_score_loss.item()
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_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.transformer, 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(" Generator update ratio: %s",
self.generator_update_interval)
assert isinstance(self.training_args, TrainingArgs)
logger.info(" Max gradient norm: %s", self.training_args.max_grad_norm)
assert self.real_score_transformer is not None
logger.info(
" Real score transformer parameters: %s B",
sum(p.numel()
for p in self.real_score_transformer.parameters()) / 1e9)
assert self.fake_score_transformer is not None
logger.info(
" Fake score transformer parameters: %s B",
sum(p.numel()
for p in self.fake_score_transformer.parameters()) / 1e9)
@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.dit_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
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( # type: ignore
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)
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("Rank %s: Using SP group seed %s", self.global_rank,
sp_group_seed)
else:
set_random_seed(seed + self.global_rank)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(seed)
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.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,
)
use_vsa = vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if use_vsa:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
if vsa_decay_interval_steps > 1:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = vsa_sparsity
else:
current_vsa_sparsity = 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
generator_loss = training_batch.generator_loss
fake_score_loss = training_batch.fake_score_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}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_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_fake_score_loss": fake_score_loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
}
# Only log generator loss when generator is actually trained
if (step % self.generator_update_interval == 0):
log_data["train_generator_loss"] = generator_loss
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
wandb.log(log_data, step=step)
if step % self.training_args.checkpointing_steps == 0 and step > 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._log_validation(self.transformer, self.training_args, step)
wandb.finish()
save_checkpoint(self.transformer, 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()
+7 -6
View File
@@ -262,23 +262,24 @@ class TrainingPipeline(LoRAPipeline, 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
assert training_batch.timesteps is not None
assert training_batch.raw_latent_shape is not None
latents_shape = training_batch.raw_latent_shape
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,
@@ -469,7 +470,7 @@ class TrainingPipeline(LoRAPipeline, 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(
+9
View File
@@ -576,3 +576,12 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
return timestep
t = timestep / num_train_timestep
denominator = 1 + (shift - 1) * t
return num_train_timestep * (shift * t / denominator)
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.wan.wan_dmd_pipeline import WanDMDPipeline
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.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", "real_score_transformer",
"fake_score_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.dit_cpu_offload = True
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
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 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.fastvideo_args import TrainingArgs
from fastvideo.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,245 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from einops import rearrange
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.dataloader.schema import (pyarrow_schema_i2v,
pyarrow_schema_i2v_validation)
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.pipelines.wan.wan_i2v_dmd_pipeline import (
WanImageToVideoDmdPipeline)
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.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", "real_score_transformer",
"fake_score_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.dit_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,
dit_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_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
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)
assert text_dict is not None
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,
"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.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
+1 -1
View File
@@ -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, # type: ignore
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
+58
View File
@@ -0,0 +1,58 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY=
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_raw/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
export TOKENIZERS_PARALLELISM=false
# make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/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 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 8 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 500 \
--validation_steps 100 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune" \
--tracker_project_name Wan_distillation \
--num_height 448 \
--num_width 832 \
--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 \
--generator_update_interval 5 \
--dmd_denoising_steps '1000,757,522' \
--min_timestep_ratio 0.02 \
--max_timestep_ratio 0.98 \
--real_score_guidance_scale 3.5 \
--seed 1024
+60
View File
@@ -0,0 +1,60 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY=
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_raw/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
export TOKENIZERS_PARALLELISM=false
# Train generator with VSA
# Make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/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 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 8 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 500 \
--validation_steps 100 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/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 \
--generator_update_interval 5 \
--dmd_denoising_steps '1000,757,522' \
--min_timestep_ratio 0.02 \
--max_timestep_ratio 0.98 \
--real_score_guidance_scale 3.5 \
--seed 1024 \
--VSA_sparsity 0.8
+1 -1
View File
@@ -9,7 +9,7 @@ NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# Make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/training/wan_training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
+1 -1
View File
@@ -14,7 +14,7 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_finetune/checkpoint-5"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# Make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \