diff --git a/.github/workflows/pre-commit.yml b/.github/workflows/pre-commit.yml index 964cacbb..6db44912 100644 --- a/.github/workflows/pre-commit.yml +++ b/.github/workflows/pre-commit.yml @@ -10,7 +10,7 @@ jobs: - uses: actions/checkout@v4 - uses: actions/setup-python@v5 with: - python-version: "3.10" + python-version: "3.12" - run: echo "::add-matcher::.github/workflows/matchers/actionlint.json" - run: echo "::add-matcher::.github/workflows/matchers/mypy.json" - uses: pre-commit/action@v3.0.1 diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index ad2ba3a7..d412c2b1 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -33,7 +33,7 @@ repos: args: [--in-place, --verbose] additional_dependencies: [toml] # TODO: Remove when yapf is upgraded - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.11.4 + rev: v0.11.12 hooks: - id: ruff args: [--output-format, github, --fix] @@ -48,7 +48,7 @@ repos: hooks: - id: isort - repo: https://github.com/jackdewinter/pymarkdown - rev: v0.9.29 + rev: v0.9.30 hooks: - id: pymarkdown args: [fix] diff --git a/fastvideo/v1/configs/models/vaes/wanvae.py b/fastvideo/v1/configs/models/vaes/wanvae.py index 6f38ba3b..2598f75f 100644 --- a/fastvideo/v1/configs/models/vaes/wanvae.py +++ b/fastvideo/v1/configs/models/vaes/wanvae.py @@ -63,7 +63,7 @@ class WanVAEArchConfig(VAEArchConfig): @dataclass class WanVAEConfig(VAEConfig): - arch_config: VAEArchConfig = field(default_factory=WanVAEArchConfig) + arch_config: WanVAEArchConfig = field(default_factory=WanVAEArchConfig) use_feature_cache: bool = True use_tiling: bool = False diff --git a/fastvideo/v1/distributed/parallel_state.py b/fastvideo/v1/distributed/parallel_state.py index cc75708a..7576d36c 100644 --- a/fastvideo/v1/distributed/parallel_state.py +++ b/fastvideo/v1/distributed/parallel_state.py @@ -655,7 +655,7 @@ class GroupCoordinator: tensor_dict[key] = value return tensor_dict - def barrier(self): + def barrier(self) -> None: """Barrier synchronization among the group. NOTE: don't use `device_group` here! `barrier` in NCCL is terrible because it is internally a broadcast operation with diff --git a/fastvideo/v1/fastvideo_args.py b/fastvideo/v1/fastvideo_args.py index eee1d6d8..fcb9030f 100644 --- a/fastvideo/v1/fastvideo_args.py +++ b/fastvideo/v1/fastvideo_args.py @@ -478,6 +478,7 @@ class TrainingArgs(FastVideoArgs): output_dir: str = "" checkpoints_total_limit: int = 0 checkpointing_steps: int = 0 + resume_from_checkpoint: bool = False logging_dir: str = "" # optimizer & scheduler diff --git a/fastvideo/v1/pipelines/composed_pipeline_base.py b/fastvideo/v1/pipelines/composed_pipeline_base.py index a8b22b0a..51b0117a 100644 --- a/fastvideo/v1/pipelines/composed_pipeline_base.py +++ b/fastvideo/v1/pipelines/composed_pipeline_base.py @@ -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__) @@ -34,20 +40,35 @@ class ComposedPipelineBase(ABC): is_video_pipeline: bool = False # To be overridden by video pipelines _required_config_modules: List[str] = [] + training_args: Optional[TrainingArgs] = None + fastvideo_args: Optional[FastVideoArgs] = None # TODO(will): args should support both inference args and training args 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. """ + + if fastvideo_args.training_mode: + assert isinstance(fastvideo_args, TrainingArgs) + self.training_args = fastvideo_args + assert self.training_args is not None + else: + self.fastvideo_args = fastvideo_args + assert self.fastvideo_args is not None + 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 +80,124 @@ 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: + assert self.training_args is not None + if self.training_args.log_validation: + self.initialize_validation_pipeline(self.training_args) + self.initialize_training_pipeline(self.training_args) + self.initialize_pipeline(fastvideo_args) - logger.info("Creating pipeline stages...") - self.create_pipeline_stages(fastvideo_args) + if not fastvideo_args.training_mode: + logger.info("Creating pipeline stages...") + self.create_pipeline_stages(fastvideo_args) - def get_module(self, module_name: str) -> Any: + def initialize_training_pipeline(self, training_args: TrainingArgs): + raise NotImplementedError( + "if training_mode is True, the pipeline must implement this method") + + def initialize_validation_pipeline(self, training_args: TrainingArgs): + 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 is None or 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) + + fastvideo_args.use_cpu_offload = False + fastvideo_args.inference_mode = False + + logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args) + + fastvideo_args.check_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) + assert fastvideo_args.tp_size is not None, "tp_size must be set" + assert fastvideo_args.sp_size is not None, "sp_size must be set" + 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 +243,12 @@ class ComposedPipelineBase(ABC): """ raise NotImplementedError + def create_training_stages(self, training_args: TrainingArgs): + """ + Create the training pipeline stages. + """ + raise NotImplementedError + def initialize_pipeline(self, fastvideo_args: FastVideoArgs): """ Initialize the pipeline. @@ -136,19 +271,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 +301,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 +334,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) diff --git a/fastvideo/v1/pipelines/training_utils.py b/fastvideo/v1/pipelines/training_utils.py deleted file mode 100644 index e1a0dbe9..00000000 --- a/fastvideo/v1/pipelines/training_utils.py +++ /dev/null @@ -1,41 +0,0 @@ -import json -import os - -import torch -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 - -logger = init_logger(__name__) - - -def save_checkpoint(transformer, rank, output_dir, step): - # Configure FSDP to save full state dict - FSDP.set_state_dict_type( - transformer, - state_dict_type=StateDictType.FULL_STATE_DICT, - state_dict_config=FullStateDictConfig(offload_to_cpu=True, - rank0_only=True), - ) - - # Now get the state dict - cpu_state = transformer.state_dict() - - # Save it (only on rank 0 since we used rank0_only=True) - if rank <= 0: - save_dir = os.path.join(output_dir, f"checkpoint-{step}") - os.makedirs(save_dir, exist_ok=True) - weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt") - torch.save(cpu_state, weight_path) - config_dict = transformer.hf_config - if "dtype" in config_dict: - del config_dict["dtype"] # TODO - config_path = os.path.join(save_dir, "config.json") - # 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) diff --git a/fastvideo/v1/pipelines/wan/wan_pipeline.py b/fastvideo/v1/pipelines/wan/wan_pipeline.py index 8159d5cc..f07af632 100644 --- a/fastvideo/v1/pipelines/wan/wan_pipeline.py +++ b/fastvideo/v1/pipelines/wan/wan_pipeline.py @@ -48,7 +48,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( diff --git a/fastvideo/v1/training/__init__.py b/fastvideo/v1/training/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/fastvideo/v1/training/training_pipeline.py b/fastvideo/v1/training/training_pipeline.py new file mode 100644 index 00000000..34e43545 --- /dev/null +++ b/fastvideo/v1/training/training_pipeline.py @@ -0,0 +1,492 @@ +import gc +import os +import traceback +from abc import ABC, abstractmethod + +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 fastvideo.v1.configs.sample import SamplingParam +from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset +from fastvideo.v1.distributed import 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.training.training_utils import ( + compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input) + +import wandb # isort: skip + +logger = init_logger(__name__) + +# Note: if checking with float32, cannot use flash-attn. +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"] + validation_pipeline: ComposedPipelineBase + + 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): + logger.info("Initializing training pipeline...") + self.device = training_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() + + noise_scheduler = self.modules["scheduler"] + params_to_optimize = self.transformer.parameters() + params_to_optimize = list( + filter(lambda p: p.requires_grad, params_to_optimize)) + + self.optimizer = torch.optim.AdamW( + params_to_optimize, + lr=training_args.learning_rate, + betas=(0.9, 0.999), + weight_decay=training_args.weight_decay, + eps=1e-8, + ) + + self.init_steps = 0 + logger.info("optimizer: %s", self.optimizer) + + self.lr_scheduler = get_scheduler( + training_args.lr_scheduler, + optimizer=self.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, + ) + + self.train_dataset = ParquetVideoTextDataset( + training_args.data_path, + batch_size=training_args.train_batch_size, + rank=self.rank, + world_size=self.world_size, + cfg_rate=training_args.cfg, + num_latent_t=training_args.num_latent_t) + + self.train_dataloader = StatefulDataLoader( + self.train_dataset, + batch_size=training_args.train_batch_size, + num_workers=training_args. + dataloader_num_workers, # Reduce number of workers to avoid memory issues + prefetch_factor=2, + shuffle=False, + pin_memory=True, + drop_last=True) + + self.noise_scheduler = noise_scheduler + + if self.rank <= 0: + project = training_args.tracker_project_name or "fastvideo" + wandb.init(project=project, config=training_args) + + @abstractmethod + def initialize_validation_pipeline(self, training_args: TrainingArgs): + 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, training_args, global_step) -> None: + assert training_args is not None + training_args.inference_mode = True + training_args.use_cpu_offload = False + if not training_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(training_args.model_path) + + # Prepare validation prompts + logger.info('fastvideo_args.validation_prompt_dir: %s', + training_args.validation_prompt_dir) + validation_dataset = ParquetVideoTextDataset( + training_args.validation_prompt_dir, + batch_size=1, + rank=0, + world_size=1, + cfg_rate=0, + num_latent_t=training_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 # type: ignore[attr-defined] + self.validation_pipeline.denoising_stage.transformer = transformer # type: ignore[attr-defined] + + # Process each validation prompt + videos = [] + captions = [] + for _, embeddings, masks, infos in validation_dataloader: + logger.info("infos: %s", infos) + caption = infos['caption'] + captions.append(caption) + prompt_embeds = embeddings.to(training_args.device) + prompt_attention_mask = masks.to(training_args.device) + + # 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] + + # Prepare batch for validation + # print('shape of embeddings', prompt_embeds.shape) + batch = ForwardBatch( + data_type="video", + latents=None, + # seed=sampling_param.seed, + 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=training_args.num_height, + width=training_args.num_width, + num_frames=training_args.num_frames, + # num_inference_steps=fastvideo_args.validation_sampling_steps, + num_inference_steps=10, + # 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.inference_mode(): + output_batch = self.validation_pipeline.forward( + batch, training_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] + filename = os.path.join( + training_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) -> float: + """ + Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE. + Uses standard tolerances for GRADIENT_CHECK_DTYPE precision. + """ + assert self.training_args is not None + # 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() -> torch.Tensor: + assert self.training_args is not None + # Move inputs to GPU, compute loss, cleanup + inputs_gpu = { + k: + v.to(self.training_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.training_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.training_args.device) + + try: + # Get analytical gradients + transformer.zero_grad() + analytical_loss = compute_loss() + analytical_loss.backward() + + # Check gradients for selected parameters + absolute_errors: list[float] = [] + param_count = 0 + + for name, param in transformer.named_parameters(): + if not (param.requires_grad and param.grad is not None + and param_count < max_params_to_check + and param.grad.abs().max() > 5e-4): + 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 + + # Find first significant gradient element + flat_param = local_param.data.view(-1) + flat_grad = local_grad.view(-1) + check_idx = next((i for i in range(min(10, flat_param.numel())) + if abs(flat_grad[i]) > 1e-4), 0) + + # 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(): + 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) + + logger.info( + "%s[%s]: analytical=%s, numerical=%s, abs_error=%s, rel_error=%s", + name, check_idx, analytical_grad, numerical_grad, abs_error, + rel_error) + + # param_count += 1 + + # Compute and log statistics + if absolute_errors: + min_err, max_err, mean_err = min(absolute_errors), max( + absolute_errors + ), sum(absolute_errors) / len(absolute_errors) + logger.info("Gradient check stats: min=%s, max=%s, mean=%s", + min_err, max_err, mean_err) + + if self.rank <= 0: + 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("Gradient check failed: %s", e) + traceback.print_exc() + return float('inf') + + def setup_gradient_check(self, args, loader_iter, noise_scheduler, + noise_random_generator) -> float | None: + """ + 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 + """ + assert self.training_args is not 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.training_args.device, + dtype=GRADIENT_CHECK_DTYPE) + check_encoder_hidden_states = check_encoder_hidden_states.to( + self.training_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("❌ Large gradient error detected: %s", + max_grad_error) + else: + logger.info("✅ Gradient check passed: max error %s", + max_grad_error) + + return max_grad_error + + except Exception as e: + logger.error("Gradient check setup failed: %s", e) + traceback.print_exc() + return None diff --git a/fastvideo/v1/training/training_utils.py b/fastvideo/v1/training/training_utils.py new file mode 100644 index 00000000..19d206e0 --- /dev/null +++ b/fastvideo/v1/training/training_utils.py @@ -0,0 +1,109 @@ +import json +import math +import os +from typing import Optional + +import torch +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 + +logger = init_logger(__name__) + + +def compute_density_for_timestep_sampling( + weighting_scheme: str, + batch_size: int, + generator, + logit_mean: Optional[float] = None, + logit_std: Optional[float] = None, + mode_scale: Optional[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) -> torch.Tensor: + 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) -> None: + # Configure FSDP to save full state dict + FSDP.set_state_dict_type( + transformer, + state_dict_type=StateDictType.FULL_STATE_DICT, + state_dict_config=FullStateDictConfig(offload_to_cpu=True, + rank0_only=True), + ) + + # Now get the state dict + cpu_state = transformer.state_dict() + + # Save it (only on rank 0 since we used rank0_only=True) + if rank <= 0: + save_dir = os.path.join(output_dir, f"checkpoint-{step}") + os.makedirs(save_dir, exist_ok=True) + weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt") + torch.save(cpu_state, weight_path) + config_dict = transformer.hf_config + if "dtype" in config_dict: + del config_dict["dtype"] # TODO + config_path = os.path.join(save_dir, "config.json") + # save dict as json + with open(config_path, "w") as f: + json.dump(config_dict, f, indent=4) + logger.info("--> checkpoint saved at step %s to %s", step, weight_path) + + +def normalize_dit_input(model_type, latents, args=None) -> torch.Tensor: + if model_type == "hunyuan_hf" or 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") diff --git a/fastvideo/v1/training/wan_training_pipeline.py b/fastvideo/v1/training/wan_training_pipeline.py new file mode 100644 index 00000000..19812912 --- /dev/null +++ b/fastvideo/v1/training/wan_training_pipeline.py @@ -0,0 +1,305 @@ +import sys +import time +from collections import deque +from copy import deepcopy + +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from tqdm.auto import tqdm + +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.pipeline_batch_info import ForwardBatch +from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline +from fastvideo.v1.training.training_pipeline import TrainingPipeline +from fastvideo.v1.training.training_utils import ( + compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input, + save_checkpoint) + +import wandb # isort: skip + +logger = init_logger(__name__) + +# Manual gradient checking flag - set to True to enable gradient verification +ENABLE_GRADIENT_CHECK = False + + +class WanTrainingPipeline(TrainingPipeline): + """ + A training pipeline for Wan. + """ + _required_config_modules = ["scheduler", "transformer"] + + 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.vae_config.load_encoder = False + validation_pipeline = WanValidationPipeline.from_pretrained( + args.model_path, args=None, inference_mode=True) + + 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, + ) -> tuple[float, float]: + assert self.training_args is not None + self.modules["transformer"].requires_grad_(True) + self.modules["transformer"].train() + + total_loss = 0.0 + optimizer.zero_grad() + for _ in range(gradient_accumulation_steps): + ( + latents, + encoder_hidden_states, + encoder_attention_mask, + infos, + ) = next(loader_iter) + latents = latents.to(self.training_args.device, + dtype=torch.bfloat16) + encoder_hidden_states = encoder_hidden_states.to( + self.training_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 + target = latents if precondition_outputs else 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() + + # TODO(will): clip grad norm + # grad_norm = transformer.clip_grad_norm_(max_grad_norm) + optimizer.step() + lr_scheduler.step() + return total_loss, 0.0 + # return total_loss, grad_norm.item() + + def forward( + self, + batch: ForwardBatch, + fastvideo_args: FastVideoArgs, + ): + assert self.training_args is not None + noise_random_generator = None + + noise_scheduler = FlowMatchEulerDiscreteScheduler() + + # Train! + assert self.training_args.sp_size is not None + assert self.training_args.gradient_accumulation_steps is not None + total_batch_size = (self.world_size * + self.training_args.gradient_accumulation_steps / + self.training_args.sp_size * + self.training_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(" Resume training from step %s", self.init_steps) + logger.info(" Instantaneous batch size per device = %s", + self.training_args.train_batch_size) + logger.info( + " Total train batch size (w. data & sequence parallel, accumulation) = %s", + total_batch_size) + logger.info(" Gradient Accumulation steps = %s", + self.training_args.gradient_accumulation_steps) + logger.info(" Total optimization steps = %s", + self.training_args.max_train_steps) + logger.info( + " Total training parameters per FSDP shard = %s B", + sum(p.numel() + for p in self.transformer.parameters() if p.requires_grad) / + 1e9) + # print dtype + logger.info(" Master weight dtype: %s", + self.transformer.parameters().__next__().dtype) + + # Potentially load in the weights and states from a previous save + if self.training_args.resume_from_checkpoint: + assert NotImplementedError( + "resume_from_checkpoint is not supported now.") + # TODO + + progress_bar = tqdm( + range(0, self.training_args.max_train_steps), + initial=self.init_steps, + desc="Steps", + # Only show the progress bar once on each machine. + disable=self.local_rank > 0, + ) + + loader_iter = iter(self.train_dataloader) + + step_times: deque[float] = deque(maxlen=100) + + # TODO(will): fix this + # for i in range(self.init_steps): + # next(loader_iter) + # get gpu memory usage + gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2 + logger.info("GPU memory usage before train_one_step: %s MB", + gpu_memory_usage) + + for step in range(self.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", + self.optimizer, + self.lr_scheduler, + loader_iter, + noise_scheduler, + noise_random_generator, + self.training_args.gradient_accumulation_steps, + self.training_args.sp_size, + self.training_args.precondition_outputs, + self.training_args.max_grad_norm, + self.training_args.weighting_scheme, + self.training_args.logit_mean, + self.training_args.logit_std, + self.training_args.mode_scale, + ) + gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2 + logger.info("GPU memory usage after train_one_step: %s MB", + gpu_memory_usage) + + 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("Performing gradient check at step %s", 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": self.lr_scheduler.get_last_lr()[0], + "step_time": step_time, + "avg_step_time": avg_step_time, + "grad_norm": grad_norm, + }, + step=step, + ) + if step % self.training_args.checkpointing_steps == 0: + # Your existing checkpoint saving code + save_checkpoint(self.transformer, self.rank, + self.training_args.output_dir, step) + 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) + + save_checkpoint(self.transformer, self.rank, + self.training_args.output_dir, + self.training_args.max_train_steps) + + if get_sp_group(): + cleanup_dist_env_and_memory() + + +def main(args) -> None: + logger.info("Starting training pipeline...") + + pipeline = WanTrainingPipeline.from_pretrained( + args.pretrained_model_name_or_path, args=args) + args = pipeline.training_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 + main(args) diff --git a/scripts/finetune/finetune_v1.sh b/scripts/finetune/finetune_v1.sh new file mode 100644 index 00000000..203fdc45 --- /dev/null +++ b/scripts/finetune/finetune_v1.sh @@ -0,0 +1,49 @@ +export WANDB_BASE_URL="https://api.wandb.ai" +export WANDB_MODE=online +# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA + +DATA_DIR=data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset +VALIDATION_DIR=data/HD-Mixkit-Finetune-Wan/validation_parquet_dataset +NUM_GPUS=1 +# 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 +torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\ + fastvideo/v1/training/wan_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 4 \ + --sp_size $NUM_GPUS \ + --tp_size $NUM_GPUS \ + --train_sp_batch_size 1\ + --dataloader_num_workers 5\ + --gradient_accumulation_steps=1\ + --max_train_steps=120 \ + --learning_rate=1e-6\ + --mixed_precision="bf16"\ + --checkpointing_steps=10 \ + --validation_steps 20\ + --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 "bf16"