Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
586d144a0f | ||
|
|
3601e33c4a | ||
|
|
6345b6bff7 |
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
|
||||
@@ -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()
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
Reference in New Issue
Block a user