Compare commits

...
19 Commits
Author SHA1 Message Date
Will Lin 6a7e4767b5 overfit SP 2025-05-29 00:34:00 -07:00
Will Lin 3642928abe mp training 2025-05-28 18:48:06 -07:00
Will Lin a94743819d check grad for SP 2025-05-28 18:29:13 -07:00
Will Lin ed7c50698f fix validation for sp 2025-05-28 16:14:58 -07:00
Will Lin 51e1d027fa update 2025-05-28 11:59:46 -07:00
Will Lin dbfccfef40 update 2025-05-28 11:58:03 -07:00
Will Lin 0ec4e8ecf0 update 2025-05-28 11:57:05 -07:00
Will Lin 5b72c993ed rebase fixes 2025-05-28 11:53:21 -07:00
Will Lin ec6ba01eb1 move utils into training_utils 2025-05-28 11:40:15 -07:00
Zihang-He 086c9600a8 added gradient clipping 2025-05-28 11:38:55 -07:00
Will Lin 5630eb5041 gradient checking 2025-05-28 11:37:46 -07:00
Will Lin 31f62ab249 cleanup 2025-05-28 11:36:36 -07:00
Will Lin 9db9c84cc9 add validation 2025-05-28 11:36:34 -07:00
JerryZhou54 ab9eac9ee8 Small fix 2025-05-28 11:35:34 -07:00
JerryZhou54 3cce02193d Small fix 2025-05-28 11:35:32 -07:00
JerryZhou54 a1a42f7afc Integrate the new parquet dataloader into training pipeline 2025-05-28 11:35:07 -07:00
d36c209191 Will/training (#425)
Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-05-28 11:34:17 -07:00
JerryZhou54 7f4379acc0 Finish data preprocessing and loading 2025-05-28 11:31:41 -07:00
JerryZhou54 7ba8cdf148 Add data preprocessing script for WAN 2025-05-28 11:27:39 -07:00
21 changed files with 2295 additions and 67 deletions
+111
View File
@@ -0,0 +1,111 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
def main(args):
# Assume using torchrun
local_rank = int(os.getenv("RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
num_gpus=world_size,
device_str="cuda",
**pipeline_config_args,
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--preprocess_video_batch_size",
type=int,
default=2,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--preprocess_text_batch_size",
type=int,
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
+2 -1
View File
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -38,6 +38,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
linear_range=0.5,
):
if linear_quadratic:
raise NotImplementedError("Linear quadratic schedule is not implemented")
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
mochi_scaling_factor = 1.0
def normalize_dit_input(model_type, latents):
def normalize_dit_input(model_type, latents, args=None):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
@@ -41,5 +41,16 @@ def normalize_dit_input(model_type, latents):
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
vae_config = WanVAEConfig()
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
@@ -0,0 +1,29 @@
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from fastvideo.v1.pipelines.wan.wan_latent_pipeline import WanLatentPipeline
def main():
print("Starting data preprocessor")
pipeline = WanLatentPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
for batch in train_dataloader:
pipeline(batch)
if __name__ == "__main__":
main()
+2 -1
View File
@@ -70,7 +70,7 @@ class FastVideoArgs:
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
"fp16",
# "fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
@@ -478,6 +478,7 @@ class TrainingArgs(FastVideoArgs):
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = ""
logging_dir: str = ""
# optimizer & scheduler
+16 -2
View File
@@ -394,7 +394,16 @@ class TransformerLoader(ComponentLoader):
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s", cls_name)
logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype)
# model = load_fsdp_model(model_cls=model_cls,
# init_params={
# "config": dit_config,
# "hf_config": hf_config
# },
# weight_dir_list=safetensors_list,
# device=fastvideo_args.device,
# cpu_offload=fastvideo_args.use_cpu_offload,
# default_dtype=default_dtype)
model = load_fsdp_model(model_cls=model_cls,
init_params={
"config": dit_config,
@@ -403,7 +412,12 @@ class TransformerLoader(ComponentLoader):
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype)
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
)
if fastvideo_args.enable_torch_compile:
logger.info("Torch Compile enabled for DiT")
for n, m in reversed(list(model.named_modules())):
+29 -4
View File
@@ -14,13 +14,16 @@ from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
from torch.distributed._tensor import distribute_tensor
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move this to utils elsewhere
@@ -86,16 +89,29 @@ def get_param_names_mapping(
# TODO(PY): add compile option
# param_dtype: torch.dtype,
# reduce_dtype: torch.dtype,
# output_dtype: torch.dtype,
# pp_enabled: bool = False,
# cpu_offload: bool = False,
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
default_dtype: Optional[torch.dtype] = torch.bfloat16,
output_dtype: Optional[torch.dtype] = None,
) -> torch.nn.Module:
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
# with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(get_sequence_model_parallel_world_size(), ),
@@ -104,6 +120,7 @@ def load_fsdp_model(
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
dp_mesh=device_mesh["dp"])
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
@@ -129,6 +146,7 @@ def shard_model(
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
) -> None:
"""
@@ -156,14 +174,17 @@ def shard_model(
"""
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": dp_mesh
"mesh": dp_mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
# Shard the model with FSDP, iterating in reverse to start with
# iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
# TODO(will): don't reshard after forward for the last layer to save on the
# all-gather that will immediately happen Shard the model with FSDP,
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
@@ -210,6 +231,10 @@ def load_fsdp_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
# s = fully_shard.state(model)
# logger.info(f"type(s): {type(s)}")
# logger.info(f"s: {s}")
# import pdb; pdb.set_trace()
sharded_sd = {}
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
+158 -14
View File
@@ -5,19 +5,25 @@ Base class for composed pipelines.
This module defines the base class for pipelines that are composed of multiple stages.
"""
import argparse
import os
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Dict, List, Optional, cast
from typing import Any, Dict, List, Optional, Union, cast
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.distributed import (init_distributed_environment,
initialize_model_parallel,
model_parallel_is_initialized)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.utils import (maybe_download_model,
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
verify_model_config_and_directory)
logger = init_logger(__name__)
@@ -39,15 +45,20 @@ class ComposedPipelineBase(ABC):
def __init__(self,
model_path: str,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None):
config: Optional[Dict[str, Any]] = None,
required_config_modules: Optional[List[str]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
self.fastvideo_args = fastvideo_args
self.model_path = model_path
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
if required_config_modules is not None:
self._required_config_modules = required_config_modules
if self._required_config_modules is None:
raise NotImplementedError(
"Subclass must set _required_config_modules")
@@ -59,16 +70,135 @@ class ComposedPipelineBase(ABC):
else:
self.config = config
self.maybe_init_distributed_environment(fastvideo_args)
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(fastvideo_args)
if fastvideo_args.training_mode:
if fastvideo_args.log_validation:
self.initialize_validation_pipeline(fastvideo_args)
self.initialize_training_pipeline(fastvideo_args)
self.initialize_pipeline(fastvideo_args)
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
# logger.info("Creating pipeline stages...")
# self.create_pipeline_stages(fastvideo_args)
def get_module(self, module_name: str) -> Any:
if fastvideo_args.training_mode:
logger.info("Creating training pipeline stages...")
self.create_training_stages(fastvideo_args)
else:
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
def initialize_training_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if training_mode is True, the pipeline must implement this method")
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if log_validation is True, the pipeline must implement this method"
)
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
args: Optional[argparse.Namespace] = None,
required_config_modules: Optional[List[str]] = None,
**kwargs) -> "ComposedPipelineBase":
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path,
device_str=device or "cuda" if
torch.cuda.is_available() else "cpu",
**config_args)
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
# we use cpu offload for training
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
# we hijack the precision to be the master weight type so that the
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
return cls(model_path,
fastvideo_args,
required_config_modules=required_config_modules)
def maybe_init_distributed_environment(self, fastvideo_args: FastVideoArgs):
if model_parallel_is_initialized():
return
local_rank = int(os.environ.get("LOCAL_RANK", -1))
world_size = int(os.environ.get("WORLD_SIZE", -1))
rank = int(os.environ.get("RANK", -1))
if local_rank == -1 or world_size == -1 or rank == -1:
raise ValueError(
"Local rank, world size, and rank must be set. Use torchrun to launch the script."
)
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(
tensor_model_parallel_size=fastvideo_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size)
device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device = device
def get_module(self, module_name: str, default_value: Any = None) -> Any:
if module_name not in self.modules:
return default_value
return self.modules[module_name]
def add_module(self, module_name: str, module: Any):
@@ -114,6 +244,19 @@ class ComposedPipelineBase(ABC):
"""
raise NotImplementedError
# @abstractmethod
# def create_validation_stages(self, fastvideo_args: FastVideoArgs):
# """
# Create the validation pipeline stages.
# """
# raise NotImplementedError
def create_training_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the training pipeline stages.
"""
raise NotImplementedError
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
@@ -136,19 +279,21 @@ class ComposedPipelineBase(ABC):
modules_config
) > 1, "model_index.json must contain at least one pipeline module"
required_modules = [
"vae", "text_encoder", "transformer", "scheduler", "tokenizer"
]
for module_name in required_modules:
for module_name in self.required_config_modules:
if module_name not in modules_config:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.info("Diffusers config passed sanity checks")
# all the component models used by the pipeline
required_modules = self.required_config_modules
logger.info("Loading required modules: %s", required_modules)
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
@@ -164,7 +309,6 @@ class ComposedPipelineBase(ABC):
logger.warning("Overwriting module %s", module_name)
modules[module_name] = module
required_modules = self.required_config_modules
# Check if all required modules were loaded
for module_name in required_modules:
if module_name not in modules or modules[module_name] is None:
@@ -198,7 +342,7 @@ class ComposedPipelineBase(ABC):
# Execute each stage
logger.info("Running pipeline stages: %s",
self._stage_name_mapping.keys())
logger.info("Batch: %s", batch)
# logger.info("Batch: %s", batch)
for stage in self.stages:
batch = stage(batch, fastvideo_args)
@@ -0,0 +1,563 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class PreprocessPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
args,
):
# Initialize class variables for data sharing
self.video_data = {} # Store video metadata and paths
self.latent_data = {} # Store latent tensors
self.preprocess_validation_text(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def preprocess_video_and_text(self, fastvideo_args: FastVideoArgs, args):
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
for root, _, files in os.walk(combined_parquet_dir):
for file in files:
if file.endswith('.parquet'):
table = pq.read_table(os.path.join(root, file))
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
num_processed_samples = 0
# Add progress bar for video preprocessing
pbar = tqdm(train_dataloader,
desc="Processing videos",
unit="batch",
disable=local_rank != 0)
for batch_idx, data in enumerate(pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
# Get corresponding captions for this batch
batch_captions = valid_data["text"]
batch = ForwardBatch(
data_type="video",
prompt=batch_captions,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Create a list to store non-padded embeddings and masks
non_padded_embeds = []
non_padded_masks = []
# Process each item in the batch
for i in range(prompt_embeds.size(0)):
seq_len = seq_lens[i].item()
# Slice the embeddings and masks to keep only non-padding parts
non_padded_embeds.append(prompt_embeds[i, :seq_len])
non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
prompt_embeds = non_padded_embeds
prompt_attention_mask = non_padded_masks
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
height, width = valid_data["pixel_values"][idx].shape[-2:]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# Create record for Parquet dataset
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape":
list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "video",
"width": width,
"height": height,
"num_frames": latents[idx].shape[1],
"duration_sec": float(valid_data["duration"][idx]),
"fps": float(valid_data["fps"][idx]),
}
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array(
[record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["vae_latent_dtype"] for record in batch_data]),
pa.array([
record["text_embedding_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_embedding_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_embedding_dtype"] for record in batch_data
]),
pa.array([
record["text_attention_mask_bytes"]
for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"]
for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"]
for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(
arrays, names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info(f"Collected batch with {len(table)} samples")
if num_processed_samples >= args.flush_frequency:
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(
num_processed_samples // args.samples_per_file, 1)
print(
f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(),
total_chunks))
chunks_per_worker = (total_chunks + num_workers -
1) // num_workers
print(
f"Using {num_workers} workers to process {total_chunks} chunks"
)
logger.info(f"Chunks per worker: {chunks_per_worker}")
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info(
f"Processed chunk with {written} samples")
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially"
)
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(
work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total samples written: {total_written}")
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
with torch.inference_mode():
# Text Encoder
batch = ForwardBatch(
data_type="video",
prompt=prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds = result_batch.prompt_embeds[0]
prompt_attention_mask = result_batch.prompt_attention_mask[0]
file_name = prompt.split(".")[0]
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
# Slice the embeddings to keep only the non-padding parts
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
# Log the shapes after removing padding
logger.info(
f"Shape after removing padding - Embeddings: {text_embedding.shape}, Mask: {text_attention_mask.shape}"
)
# Create record for Parquet dataset
record = {
"id": file_name,
"vae_latent_bytes": b"", # Not available for validation
"vae_latent_shape": [],
"vae_latent_dtype": "",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": file_name,
"caption": prompt,
"media_type": "video",
"width": 0, # Not available for validation
"height": 0, # Not available for validation
"num_frames": 0, # Not available for validation
"duration_sec": 0.0, # Not available for validation
"fps": 0.0, # Not available for validation
}
batch_data.append(record)
logger.info(f"Saved validation sample: {file_name}")
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array([record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array([record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array([record["vae_latent_dtype"] for record in batch_data]),
pa.array(
[record["text_embedding_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["text_embedding_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["text_embedding_dtype"] for record in batch_data]),
pa.array([
record["text_attention_mask_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"] for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(arrays,
names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
logger.info(f"Total validation samples: {len(table)}")
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=1) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
}
for future in tqdm(futures, desc="Processing chunks"):
try:
total_written += future.result()
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially")
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total validation samples written: {total_written}")
# Clear memory
del table
gc.collect() # Force garbage collection
@staticmethod
def process_chunk_range(args):
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
num_samples = len(table)
# Create worker-specific subdirectory
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Check how many files there are already in the dir, and update i accordingly
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_idx, end_idx):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
chunk = table.slice(start_sample, end_sample - start_sample)
# Create chunk file in worker's directory
chunk_path = os.path.join(
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
# Write to temporary file
pq.write_table(chunk, temp_path, compression='zstd')
# Rename temporary file to final file
if os.path.exists(chunk_path):
os.remove(
chunk_path) # Remove existing file if it exists
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception as e:
# Clean up temporary file if it exists
if os.path.exists(temp_path):
os.remove(temp_path)
raise e
return total_written
except Exception as e:
logger.error(
f"Error processing chunks {start_idx}-{end_idx} for worker {worker_id}: {str(e)}"
)
raise
EntryClass = PreprocessPipeline
+4 -2
View File
@@ -74,7 +74,8 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
@@ -83,6 +84,7 @@ class DenoisingStage(PipelineStage):
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
# b c t h w -> b t n s h w
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
@@ -188,7 +190,7 @@ class DenoisingStage(PipelineStage):
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
dtype=torch.bfloat16,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
+14 -4
View File
@@ -63,10 +63,15 @@ class TextEncodingStage(PipelineStage):
if fastvideo_args.use_cpu_offload:
text_encoder = text_encoder.to(fastvideo_args.device)
assert isinstance(batch.prompt, str)
text = preprocess_func(batch.prompt)
text_inputs = tokenizer(text, **encoder_config.tokenizer_kwargs).to(
fastvideo_args.device)
assert isinstance(batch.prompt, (str, list))
if isinstance(batch.prompt, str):
batch.prompt = [batch.prompt]
texts = []
for prompt_str in batch.prompt:
texts.append(preprocess_func(prompt_str))
text_inputs = tokenizer(texts,
**encoder_config.tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
with set_forward_context(current_timestep=0, attn_metadata=None):
@@ -78,6 +83,8 @@ class TextEncodingStage(PipelineStage):
prompt_embeds = postprocess_func(outputs)
batch.prompt_embeds.append(prompt_embeds)
if batch.prompt_attention_mask is not None:
batch.prompt_attention_mask.append(attention_mask)
if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str)
@@ -98,6 +105,9 @@ class TextEncodingStage(PipelineStage):
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_prompt_embeds)
if batch.negative_attention_mask is not None:
batch.negative_attention_mask.append(
negative_attention_mask)
if fastvideo_args.use_cpu_offload:
text_encoder.to('cpu')
+841
View File
@@ -0,0 +1,841 @@
import gc
import os
import sys
import time
import traceback
from abc import ABC, abstractmethod
from collections import deque
from copy import deepcopy
import imageio
import numpy as np
import torch
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.training_utils import (
_clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, save_checkpoint)
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
GRADIENT_CHECK_DTYPE = torch.bfloat16
class TrainingPipeline(ComposedPipelineBase, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
All reusable components and code should be implemented in this class.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_training_pipeline(self, fastvideo_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = fastvideo_args.device
self.sp_group = get_sp_group()
self.world_size = self.sp_group.world_size
self.rank = self.sp_group.rank
self.local_rank = self.sp_group.local_rank
self.transformer = self.get_module("transformer")
assert self.transformer is not None
self.transformer.requires_grad_(True)
self.transformer.train()
args = fastvideo_args
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * self.world_size,
num_training_steps=args.max_train_steps * self.world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = ParquetVideoTextDataset(
args.data_path,
batch_size=args.train_batch_size,
rank=self.rank,
world_size=self.world_size,
cfg_rate=args.cfg,
num_latent_t=args.num_latent_t)
train_dataloader = StatefulDataLoader(
train_dataset,
batch_size=args.train_batch_size,
num_workers=args.
dataloader_num_workers, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
self.lr_scheduler = lr_scheduler
self.train_dataset = train_dataset
self.train_dataloader = train_dataloader
self.init_steps = init_steps
self.optimizer = optimizer
self.noise_scheduler = noise_scheduler
# self.noise_random_generator = noise_random_generator
# num_update_steps_per_epoch = math.ceil(
# len(train_dataloader) / args.gradient_accumulation_steps *
# args.sp_size / args.train_sp_batch_size)
# args.num_train_epochs = math.ceil(args.max_train_steps /
# num_update_steps_per_epoch)
if self.rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@abstractmethod
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"Training pipelines must implement this method")
@abstractmethod
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
loader, noise_scheduler, noise_random_generator,
gradient_accumulation_steps, sp_size,
precondition_outputs, max_grad_norm, weighting_scheme,
logit_mean, logit_std, mode_scale):
"""
Train one step of the model.
"""
raise NotImplementedError(
"Training pipeline must implement this method")
def log_validation(self, transformer, fastvideo_args, global_step):
fastvideo_args.inference_mode = True
fastvideo_args.use_cpu_offload = False
if not fastvideo_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
# Prepare validation prompts
print('fastvideo_args.validation_prompt_dir',
fastvideo_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
fastvideo_args.validation_prompt_dir,
batch_size=1,
rank=0,
world_size=1,
cfg_rate=0,
num_latent_t=args.num_latent_t)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=False)
transformer.requires_grad_(False)
for p in transformer.parameters():
p.requires_grad = False
transformer.eval()
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
self.validation_pipeline.latent_preparation_stage.transformer = transformer
self.validation_pipeline.denoising_stage.transformer = transformer
# Process each validation prompt
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
logger.info(f"infos: {infos}")
caption = infos['caption']
captions.append(caption)
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
# Calculate sizes
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]
logger.info('embed dtype', prompt_embeds.dtype)
num_frames = (fastvideo_args.num_latent_t - 1) * 4 + 1
logger.info(f"validation num_frames: {num_frames}")
# Prepare batch for validation
# print('shape of embeddings', prompt_embeds.shape)
batch = ForwardBatch(
# **shallow_asdict(sampling_param),
data_type="video",
latents=None,
# seed=sampling_param.seed,
# data_type="video",
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=args.num_height,
width=args.num_width,
num_frames=num_frames,
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=50,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
extra={},
)
# Run validation inference
with torch.autocast("cuda", dtype=torch.bfloat16):
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
# 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))
videos.append(frames)
# Log validation results
rank = int(os.environ.get("RANK", 0))
if rank == 0:
video_filenames = []
video_captions = []
for i, video in enumerate(videos):
caption = captions[i]
os.makedirs(fastvideo_args.output_dir, exist_ok=True)
filename = os.path.join(
fastvideo_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_captions.append(
caption) # Store the caption for each video
logs = {
"validation_videos": [
wandb.Video(filename,
caption=caption) for filename, caption in zip(
video_filenames, video_captions)
]
}
wandb.log(logs, step=global_step)
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def gradient_check_parameters(self,
transformer,
latents,
encoder_hidden_states,
encoder_attention_mask,
timesteps,
target,
eps=5e-2,
max_params_to_check=2000):
"""
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
"""
# Move all inputs to CPU and clear GPU memory
inputs_cpu = {
'latents': latents.cpu(),
'encoder_hidden_states': encoder_hidden_states.cpu(),
'encoder_attention_mask': encoder_attention_mask.cpu(),
'timesteps': timesteps.cpu(),
'target': target.cpu()
}
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
torch.cuda.empty_cache()
def compute_loss():
# Move inputs to GPU, compute loss, cleanup
inputs_gpu = {
k:
v.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE
if k != 'encoder_attention_mask' else None)
for k, v in inputs_cpu.items()
}
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
# with torch.autocast(enabled=False, device_type="cuda"):
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
with set_forward_context(
current_timestep=inputs_gpu['timesteps'],
attn_metadata=None):
model_pred = transformer(
hidden_states=inputs_gpu['latents'],
encoder_hidden_states=inputs_gpu[
'encoder_hidden_states'],
timestep=inputs_gpu['timesteps'],
encoder_attention_mask=inputs_gpu[
'encoder_attention_mask'],
return_dict=False)[0]
if self.fastvideo_args.precondition_outputs:
sigmas = get_sigmas(self.noise_scheduler,
inputs_gpu['latents'].device,
inputs_gpu['timesteps'],
n_dim=inputs_gpu['latents'].ndim,
dtype=inputs_gpu['latents'].dtype)
model_pred = inputs_gpu['latents'] - model_pred * sigmas
target_adjusted = inputs_gpu['target']
else:
target_adjusted = inputs_gpu['target']
loss = torch.mean((model_pred - target_adjusted)**2)
# Cleanup and return
loss_cpu = loss.cpu()
del inputs_gpu, model_pred, target_adjusted
if 'sigmas' in locals(): del sigmas
torch.cuda.empty_cache()
return loss_cpu.to(self.fastvideo_args.device)
try:
# Get analytical gradients
transformer.zero_grad()
analytical_loss = compute_loss()
analytical_loss.backward()
# Check gradients for selected parameters
absolute_errors = []
param_count = 0
rank = int(os.environ.get("RANK", 0))
sp_group = get_sp_group()
for name, param in transformer.named_parameters():
sp_group.barrier()
# skip scale_shift_table because it is not sharded
if 'scale_shift_table' in name:
continue
if isinstance(param.grad, torch.distributed.tensor.DTensor):
l = param.grad.full_tensor()
distributed = True
else:
l = param.grad
distributed = False
continue
# logger.info(f"rank: {rank}, name: {name}, param: {param.shape}, grad: {param.grad.shape}, distributed: {distributed}", local_main_process_only=False)
# logger.info(f"rank: {rank}, name: {name}, type of param: {type(param)}, type of grad: {type(param.grad)}", local_main_process_only=False)
# logger.info(f"rank: {rank}, name: {name}, param: {param}, grad: {param.grad}", local_main_process_only=False)
if not (param.requires_grad and param.grad is not None
and param_count < max_params_to_check
and l.abs().max() > 5e-4):
continue
if not distributed:
if rank != 0:
continue
# Get local parameter and gradient tensors
local_param = param._local_tensor if hasattr(
param, '_local_tensor') else param
local_grad = param.grad._local_tensor if hasattr(
param.grad, '_local_tensor') else param.grad
# logger.info(f"rank: {rank}, local_param: {local_param.shape}, local_grad: {local_grad.shape}", local_main_process_only=False)
# Find first significant gradient element
flat_param = local_param.data.view(-1)
flat_grad = local_grad.view(-1)
# logger.info(f"rank: {rank}, flat_param: {flat_param.shape}, flat_grad: {flat_grad.shape}", local_main_process_only=False)
check_idx = next((i for i in range(min(10, flat_param.numel()))
if abs(flat_grad[i]) > 1e-4), 0)
# logger.info(f"rank: {rank}, check_idx: {check_idx}", local_main_process_only=False)
# Store original values
orig_value = flat_param[check_idx].item()
analytical_grad = flat_grad[check_idx].item()
# Compute numerical gradient
for delta in [eps, -eps]:
with torch.no_grad():
# only have a single rank modify the parameter
if rank == 0:
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0: loss_plus = loss.item()
else: loss_minus = loss.item()
# Restore parameter and compute error
with torch.no_grad():
flat_param[check_idx] = orig_value
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
abs_error = abs(analytical_grad - numerical_grad)
rel_error = abs_error / max(abs(analytical_grad),
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
if self.rank <= 0:
logger.info(
f"{name}[{check_idx}]: analytical={analytical_grad:.6f}, "
f"numerical={numerical_grad:.6f}, abs_error={abs_error:.2e}, rel_error={rel_error:.2%}"
)
# param_count += 1
# Compute and log statistics
if self.rank <= 0:
if absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
logger.info(
f"Gradient check stats: min={min_err:.2e}, max={max_err:.2e}, mean={mean_err:.2e}"
)
wandb.log({
"grad_check/min_abs_error":
min_err,
"grad_check/max_abs_error":
max_err,
"grad_check/mean_abs_error":
mean_err,
"grad_check/analytical_loss":
analytical_loss.item(),
})
return max_err
return float('inf')
except Exception as e:
logger.error(f"Gradient check failed: {e}")
traceback.print_exc()
return float('inf')
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
noise_random_generator):
"""
Setup and perform gradient check on a fresh batch.
Args:
args: Training arguments
loader_iter: Data loader iterator
noise_scheduler: Noise scheduler for diffusion
noise_random_generator: Random number generator for noise
Returns:
float or None: Maximum gradient error or None if check is disabled/fails
"""
if not ENABLE_GRADIENT_CHECK:
return None
try:
# Get a fresh batch and process it exactly like train_one_step
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
loader_iter)
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
check_latents = check_latents.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE)
check_encoder_hidden_states = check_encoder_hidden_states.to(
self.fastvideo_args.device, dtype=GRADIENT_CHECK_DTYPE)
check_latents = normalize_dit_input("wan", check_latents)
batch_size = check_latents.shape[0]
check_noise = torch.randn_like(check_latents)
check_u = compute_density_for_timestep_sampling(
weighting_scheme=args.weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=args.logit_mean,
logit_std=args.logit_std,
mode_scale=args.mode_scale,
)
check_indices = (check_u *
noise_scheduler.config.num_train_timesteps).long()
check_timesteps = noise_scheduler.timesteps[check_indices].to(
device=check_latents.device)
check_sigmas = get_sigmas(
noise_scheduler,
check_latents.device,
check_timesteps,
n_dim=check_latents.ndim,
dtype=check_latents.dtype,
)
check_noisy_model_input = (
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
# Compute target exactly like train_one_step
if args.precondition_outputs:
check_target = check_latents
else:
check_target = check_noise - check_latents
# Perform gradient check with the exact same inputs as training
max_grad_error = self.gradient_check_parameters(
transformer=self.transformer,
latents=
check_noisy_model_input, # Use noisy input like in training
encoder_hidden_states=check_encoder_hidden_states,
encoder_attention_mask=check_encoder_attention_mask,
timesteps=check_timesteps,
target=check_target,
max_params_to_check=100 # Check more parameters
)
if max_grad_error > 5e-2:
logger.error(
f"❌ Large gradient error detected: {max_grad_error:.2e}")
else:
logger.info(
f"✅ Gradient check passed: max error {max_grad_error:.2e}")
return max_grad_error
except Exception as e:
logger.error(f"Gradient check setup failed: {e}")
traceback.print_exc()
return None
class WanTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def create_training_stages(self, fastvideo_args: FastVideoArgs):
pass
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
self.latents = None
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(fastvideo_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
# TODO(will): clean this up
args_copy.precision = "bf16"
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=args_copy)
self.validation_pipeline = validation_pipeline
def train_one_step(
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
):
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
logger.info(f"Rank {self.rank}: Training step {_}", local_main_process_only=False)
if self.latents is None:
(
self.latents,
self.encoder_hidden_states,
self.encoder_attention_mask,
self.infos,
) = next(loader_iter)
latents = self.latents
encoder_hidden_states = self.encoder_hidden_states
encoder_attention_mask = self.encoder_attention_mask
infos = self.infos
logger.info(f"Rank {self.rank}: Training step {_} loaded data", local_main_process_only=False)
latents = latents.to(self.fastvideo_args.device,
dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
self.fastvideo_args.device, dtype=torch.bfloat16)
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)[0]
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
if precondition_outputs:
target = latents
else:
target = noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
sp_group = get_sp_group()
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
model_parts = [self.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,
)
optimizer.step()
print('device after optimizer step',
next(transformer.named_parameters())[1].device)
lr_scheduler.step()
print('device after scheduler step',
next(transformer.named_parameters())[1].device)
return total_loss, grad_norm.item()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
args = fastvideo_args
self.fastvideo_args = args
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
noise_random_generator = None
from diffusers import FlowMatchEulerDiscreteScheduler
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
total_batch_size = (self.world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
# logger.info(f" Num examples = {len(train_dataset)}")
# logger.info(f" Dataloader size = {len(train_dataloader)}")
# logger.info(f" Num Epochs = {args.num_train_epochs}")
logger.info(f" Resume training from step {init_steps}")
logger.info(
f" Instantaneous batch size per device = {args.train_batch_size}")
logger.info(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
logger.info(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}"
)
logger.info(f" Total optimization steps = {args.max_train_steps}")
logger.info(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in self.transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
logger.info(
f" Master weight dtype: {self.transformer.parameters().__next__().dtype}"
)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
loader_iter = iter(train_dataloader)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader_iter)
# get gpu memory usage
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage before train_one_step: {gpu_memory_usage} MB")
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
args.weighting_scheme,
args.logit_mean,
args.logit_std,
args.mode_scale,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage after train_one_step: {gpu_memory_usage} MB")
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info(f"Performing gradient check at step {step}")
self.setup_gradient_check(args, loader_iter, noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % args.checkpointing_steps == 0:
# Your existing checkpoint saving code
save_checkpoint(self.transformer, self.rank, args.output_dir,
step)
self.transformer.train()
self.sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
self.log_validation(self.transformer, args, step)
save_checkpoint(self.transformer, self.rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
print(args)
main(args)
+272 -3
View File
@@ -1,16 +1,72 @@
import json
import math
import os
from typing import List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.tensor
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
from fastvideo.v1.logger import init_logger
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
logger = init_logger(__name__)
def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
generator,
logit_mean: float = None,
logit_std: float = None,
mode_scale: float = None,
):
"""
Compute the density for sampling the timesteps when doing SD3 training.
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
"""
if weighting_scheme == "logit_normal":
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
u = torch.normal(
mean=logit_mean,
std=logit_std,
size=(batch_size, ),
device="cpu",
generator=generator,
)
u = torch.nn.functional.sigmoid(u)
elif weighting_scheme == "mode":
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2)**2 - 1 + u)
else:
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
return u
def get_sigmas(noise_scheduler,
device,
timesteps,
n_dim=4,
dtype=torch.float32):
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
schedule_timesteps = noise_scheduler.timesteps.to(device)
timesteps = timesteps.to(device)
step_indices = [(schedule_timesteps == t).nonzero().item()
for t in timesteps]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < n_dim:
sigma = sigma.unsqueeze(-1)
return sigma
def save_checkpoint(transformer, rank, output_dir, step):
# Configure FSDP to save full state dict
FSDP.set_state_dict_type(
@@ -36,6 +92,219 @@ def save_checkpoint(transformer, rank, output_dir, step):
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> checkpoint saved at step {step} to {weight_path}",
step=step,
weight_path=weight_path)
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
def _clip_grad_norm_while_handling_failing_dtensor_cases(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
) -> Optional[torch.Tensor]:
global _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES
if not _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES:
try:
return clip_grad_norm_(parameters, max_norm, norm_type,
error_if_nonfinite, foreach, pp_mesh)
except NotImplementedError as e:
if "DTensor does not support cross-mesh operation" in str(e):
# https://github.com/pytorch/pytorch/issues/134212
logger.warning(
"DTensor does not support cross-mesh operation. If you haven't fully tensor-parallelized your "
"model, while combining other parallelisms such as FSDP, it could be the reason for this error. "
"Gradient clipping will be skipped and gradient norm will not be logged."
)
except Exception as e:
logger.warning(
f"An error occurred while clipping gradients: {e}. Gradient clipping will be skipped and gradient "
f"norm will not be logged.")
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = True
return None
# Copied from https://github.com/pytorch/torchtitan/blob/4a169701555ab9bd6ca3769f9650ae3386b84c6e/torchtitan/utils.py#L362
@torch.no_grad()
def clip_grad_norm_(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
) -> torch.Tensor:
r"""
Clip the gradient norm of parameters.
Gradient norm clipping requires computing the gradient norm over the entire model.
`torch.nn.utils.clip_grad_norm_` only computes gradient norm along DP/FSDP/TP dimensions.
We need to manually reduce the gradient norm across PP stages.
See https://github.com/pytorch/torchtitan/issues/596 for details.
Args:
parameters (`torch.Tensor` or `List[torch.Tensor]`):
Tensors that will have gradients normalized.
max_norm (`float`):
Maximum norm of the gradients after clipping.
norm_type (`float`, defaults to `2.0`):
Type of p-norm to use. Can be `inf` for infinity norm.
error_if_nonfinite (`bool`, defaults to `False`):
If `True`, an error is thrown if the total norm of the gradients from `parameters` is `nan`, `inf`, or `-inf`.
foreach (`bool`, defaults to `None`):
Use the faster foreach-based implementation. If `None`, use the foreach implementation for CUDA and CPU native tensors
and silently fall back to the slow implementation for other device types.
pp_mesh (`torch.distributed.device_mesh.DeviceMesh`, defaults to `None`):
Pipeline parallel device mesh. If not `None`, will reduce gradient norm across PP stages.
Returns:
`torch.Tensor`:
Total norm of the gradients
"""
grads = [p.grad for p in parameters if p.grad is not None]
# TODO(aryan): Wait for next Pytorch release to use `torch.nn.utils.get_total_norm`
# total_norm = torch.nn.utils.get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
total_norm = _get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
# If total_norm is a DTensor, the placements must be `torch.distributed._tensor.ops.math_ops._NormPartial`.
# We can simply reduce the DTensor to get the total norm in this tensor's process group
# and then convert it to a local tensor.
# It has two purposes:
# 1. to make sure the total norm is computed correctly when PP is used (see below)
# 2. to return a reduced total_norm tensor whose .item() would return the correct value
if isinstance(total_norm, torch.distributed.tensor.DTensor):
# Will reach here if any non-PP parallelism is used.
# If only using PP, total_norm will be a local tensor.
total_norm = total_norm.full_tensor()
if pp_mesh is not None:
if math.isinf(norm_type):
dist.all_reduce(total_norm,
op=dist.ReduceOp.MAX,
group=pp_mesh.get_group())
else:
total_norm **= norm_type
dist.all_reduce(total_norm,
op=dist.ReduceOp.SUM,
group=pp_mesh.get_group())
total_norm **= 1.0 / norm_type
_clip_grads_with_norm_(parameters, max_norm, total_norm, foreach)
return total_norm
@torch.no_grad()
def _clip_grads_with_norm_(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
total_norm: torch.Tensor,
foreach: Optional[bool] = None,
) -> None:
if isinstance(parameters, torch.Tensor):
parameters = [parameters]
grads = [p.grad for p in parameters if p.grad is not None]
max_norm = float(max_norm)
if len(grads) == 0:
return
grouped_grads: dict[Tuple[torch.device, torch.dtype],
Tuple[List[List[torch.Tensor]],
List[int]]] = (_group_tensors_by_device_and_dtype(
[grads])) # type: ignore[assignment]
clip_coef = max_norm / (total_norm + 1e-6)
# Note: multiplying by the clamped coef is redundant when the coef is clamped to 1, but doing so
# avoids a `if clip_coef < 1:` conditional which can require a CPU <=> device synchronization
# when the gradients do not reside in CPU memory.
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
for (device, _), ([device_grads], _) in grouped_grads.items():
if (foreach is None and _has_foreach_support(device_grads, device)) or (
foreach and _device_has_foreach_support(device)):
torch._foreach_mul_(device_grads, clip_coef_clamped.to(device))
elif foreach:
raise RuntimeError(
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
)
else:
clip_coef_clamped_device = clip_coef_clamped.to(device)
for g in device_grads:
g.mul_(clip_coef_clamped_device)
def _get_total_norm(
tensors: Union[torch.Tensor, List[torch.Tensor]],
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
) -> torch.Tensor:
if isinstance(tensors, torch.Tensor):
tensors = [tensors]
else:
tensors = list(tensors)
norm_type = float(norm_type)
if len(tensors) == 0:
return torch.tensor(0.0)
first_device = tensors[0].device
grouped_tensors: dict[tuple[torch.device, torch.dtype],
tuple[list[list[torch.Tensor]], list[int]]] = (
_group_tensors_by_device_and_dtype(
[tensors] # type: ignore[list-item]
)) # type: ignore[assignment]
norms: List[torch.Tensor] = []
for (device, _), ([device_tensors], _) in grouped_tensors.items():
local_tensors = [
t.to_local()
if isinstance(t, torch.distributed.tensor.DTensor) else t
for t in device_tensors
]
if (foreach is None and _has_foreach_support(local_tensors, device)
) or (foreach and _device_has_foreach_support(device)):
norms.extend(torch._foreach_norm(local_tensors, norm_type))
elif foreach:
raise RuntimeError(
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
)
else:
norms.extend(
[torch.linalg.vector_norm(g, norm_type) for g in local_tensors])
total_norm = torch.linalg.vector_norm(
torch.stack([norm.to(first_device) for norm in norms]), norm_type)
if error_if_nonfinite and torch.logical_or(total_norm.isnan(),
total_norm.isinf()):
raise RuntimeError(
f"The total norm of order {norm_type} for gradients from "
"`parameters` is non-finite, so it cannot be clipped. To disable "
"this error and scale the gradients by the non-finite norm anyway, "
"set `error_if_nonfinite=False`")
return total_norm
def _get_foreach_kernels_supported_devices() -> list[str]:
r"""Return the device type list that supports foreach kernels."""
return ["cuda", "xpu", torch._C._get_privateuse1_backend_name()]
@torch.no_grad()
def _group_tensors_by_device_and_dtype(
tensorlistlist: List[List[Optional[torch.Tensor]]],
with_indices: bool = False,
) -> dict[tuple[torch.device, torch.dtype], tuple[
List[List[Optional[torch.Tensor]]], List[int]]]:
return torch._C._group_tensors_by_device_and_dtype(tensorlistlist,
with_indices)
def _device_has_foreach_support(device: torch.device) -> bool:
return device.type in (_get_foreach_kernels_supported_devices() +
["cpu"]) and not torch.jit.is_scripting()
def _has_foreach_support(tensors: List[torch.Tensor],
device: torch.device) -> bool:
return _device_has_foreach_support(device) and all(
t is None or type(t) in [torch.Tensor] for t in tensors)
@@ -0,0 +1,19 @@
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
class WanLatentPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs):
logger.info("WAN Latent Pipeline forward")
pass
+27 -2
View File
@@ -15,7 +15,6 @@ from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
TextEncodingStage,
TimestepPreparationStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
@@ -48,7 +47,33 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
class WanValidationPipeline(ComposedPipelineBase):
"""
Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
"""
_required_config_modules = ["vae", "scheduler"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
@@ -1,11 +1,11 @@
import json
from pathlib import Path
import csv
import cv2
def get_video_info(video_path, prompt_text):
"""Extract video information using OpenCV and corresponding prompt text"""
def get_video_info(video_path, metadata):
"""Extract video information using OpenCV and corresponding metadata"""
cap = cv2.VideoCapture(str(video_path))
if not cap.isOpened():
@@ -23,60 +23,66 @@ def get_video_info(video_path, prompt_text):
return {
"path": video_path.name,
"title": metadata.get("Video Title", ""),
"description": metadata.get("Video Description", ""),
"video_url": metadata.get("Video URL", ""),
"download_url": metadata.get("Download URL", ""),
"resolution": {
"width": width,
"height": height
},
"fps": fps,
"duration": duration,
"cap": [prompt_text]
"cap": [metadata.get("Video Description", "")]
}
def read_prompt_file(prompt_path):
"""Read and return the content of a prompt file"""
def read_csv_file(csv_path):
"""Read and return the content of a CSV file"""
try:
with open(prompt_path, 'r', encoding='utf-8') as f:
return f.read().strip()
with open(csv_path, 'r', encoding='utf-8') as f:
reader = csv.DictReader(f)
return list(reader)
except Exception as e:
print(f"Error reading prompt file {prompt_path}: {e}")
print(f"Error reading CSV file {csv_path}: {e}")
return None
def process_videos_and_prompts(video_dir_path, prompt_dir_path, verbose=False):
"""Process videos and their corresponding prompt files
def process_videos_from_csv(video_dir_path, csv_path, verbose=False):
"""Process videos using metadata from CSV file
Args:
video_dir_path (str): Path to directory containing video files
prompt_dir_path (str): Path to directory containing prompt files
csv_path (str): Path to CSV file containing video metadata
verbose (bool): Whether to print verbose processing information
"""
video_dir = Path(video_dir_path)
prompt_dir = Path(prompt_dir_path)
csv_data = read_csv_file(csv_path)
processed_data = []
# Ensure directories exist
if not video_dir.exists() or not prompt_dir.exists():
print(f"Error: One or both directories do not exist:\nVideos: {video_dir}\nPrompts: {prompt_dir}")
if not video_dir.exists():
print(f"Error: Video directory does not exist: {video_dir}")
return []
if csv_data is None:
return []
# Process each video file
for video_file in video_dir.glob('*.mp4'):
video_name = video_file.stem
prompt_file = prompt_dir / f"{video_name}.txt"
# Check if corresponding prompt file exists
if not prompt_file.exists():
print(f"Warning: No prompt file found for video {video_name}")
for row in csv_data:
video_filename = row.get("Filename")
if not video_filename:
continue
# Read prompt content
prompt_text = read_prompt_file(prompt_file)
if prompt_text is None:
video_file = video_dir / video_filename
# Check if video file exists
if not video_file.exists():
print(f"Warning: Video file not found: {video_filename}")
continue
# Process video and add to results
video_info = get_video_info(video_file, prompt_text)
video_info = get_video_info(video_file, row)
if video_info:
processed_data.append(video_info)
@@ -105,9 +111,9 @@ def parse_args():
"""Parse command line arguments"""
import argparse
parser = argparse.ArgumentParser(description='Process videos and their corresponding prompt files')
parser = argparse.ArgumentParser(description='Process videos using metadata from CSV file')
parser.add_argument('--video_dir', '-v', required=True, help='Directory containing video files')
parser.add_argument('--prompt_dir', '-p', required=True, help='Directory containing prompt text files')
parser.add_argument('--csv_path', '-c', required=True, help='Path to CSV file containing video metadata')
parser.add_argument('--output_path',
'-o',
required=True,
@@ -121,8 +127,8 @@ if __name__ == "__main__":
# Parse command line arguments
args = parse_args()
# Process videos and prompts
processed_videos = process_videos_and_prompts(args.video_dir, args.prompt_dir, args.verbose)
# Process videos from CSV
processed_videos = process_videos_from_csv(args.video_dir, args.csv_path, args.verbose)
if processed_videos:
# Save results
+11 -3
View File
@@ -24,9 +24,9 @@ def is_16_9_ratio(width: int, height: int, tolerance: float = 0.1) -> bool:
def resize_video(args_tuple):
"""
Resize a single video file.
args_tuple: (input_file, output_dir, width, height, fps)
args_tuple: (input_file, output_dir, width, height, fps, num_frames)
"""
input_file, output_dir, width, height, fps = args_tuple
input_file, output_dir, width, height, fps, num_frames = args_tuple
video = None
resized = None
output_file = output_dir / f"{input_file.name}"
@@ -39,6 +39,13 @@ def resize_video(args_tuple):
if not is_16_9_ratio(video.w, video.h):
return (input_file.name, "skipped", "Not 16:9")
# Calculate target duration based on num_frames and fps
target_duration = num_frames / fps
# Trim video if it's longer than target duration
if video.duration > target_duration:
video = video.subclip(0, target_duration)
def process_frame(frame):
frame_float = frame.astype(float) / 255.0
resized = resize(frame_float, (height, width, 3), mode='reflect', anti_aliasing=True, preserve_range=True)
@@ -75,7 +82,7 @@ def process_folder(args):
print(f"Target: {args.width}x{args.height} at {args.fps}fps")
# Prepare arguments for parallel processing
process_args = [(video_file, output_path, args.width, args.height, args.fps) for video_file in video_files]
process_args = [(video_file, output_path, args.width, args.height, args.fps, args.num_frames) for video_file in video_files]
successful = 0
skipped = 0
@@ -115,6 +122,7 @@ def parse_args():
parser.add_argument('--width', type=int, default=1280, help='Target width in pixels (default: 848)')
parser.add_argument('--height', type=int, default=720, help='Target height in pixels (default: 480)')
parser.add_argument('--fps', type=int, default=30, help='Target frames per second (default: 30)')
parser.add_argument('--num_frames', type=int, default=163, help='Target number of frames (default: 163)')
parser.add_argument('--max_workers',
type=int,
default=4,
+45
View File
@@ -0,0 +1,45 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=/workspace/data
num_gpus=1
IP=127.0.0.1
torchrun --nnodes 1 --nproc_per_node $num_gpus \
--node_rank=0 \
--rdzv_id=456 \
--rdzv_backend=c10d \
--rdzv_endpoint=$IP:29500 \
fastvideo/distill_wan.py\
--seed 42\
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--train_batch_size=1 \
--num_latent_t 1 \
--sp_size $num_gpus \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--master_weight_type="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
+27
View File
@@ -0,0 +1,27 @@
#!/bin/bash
num_gpus=2
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/entrypoints/data_preprocessor.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 480 \
--width 832 \
--num_frames 77 \
--num_inference_steps 50 \
--fps 16 \
--guidance_scale 3.0 \
--prompt_path ./assets/prompt.txt \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload
+24
View File
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
TEXT_ENCODER_PATH="/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tokenizer"
MODEL_TYPE="wan"
DATA_MERGE_PATH="/workspace/data/Mixkit-Src/merge.txt"
OUTPUT_DIR="/workspace/data/HD-Mixkit-Finetune-Wan"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size=4 \
--preprocess_text_batch_size=4 \
--max_height=480 \
--max_width=832 \
--num_frames=81 \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--text_encoder_name $TEXT_ENCODER_PATH \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH
+53
View File
@@ -0,0 +1,53 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/cats_480_single_latents_parq/combined_parquet_dataset
VALIDATION_DIR=data/cats_480_single_latents_parq/validation_parquet_dataset
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
# --gradient_checkpointing\
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/pipelines/training_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_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 20 \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 6\
--gradient_accumulation_steps=1\
--max_train_steps=3000 \
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=5000 \
--validation_steps 200\
--validation_sampling_steps "2,4,8" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_finetune"\
--tracker_project_name wan_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 81 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0