Compare commits
19
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6a7e4767b5 | ||
|
|
3642928abe | ||
|
|
a94743819d | ||
|
|
ed7c50698f | ||
|
|
51e1d027fa | ||
|
|
dbfccfef40 | ||
|
|
0ec4e8ecf0 | ||
|
|
5b72c993ed | ||
|
|
ec6ba01eb1 | ||
|
|
086c9600a8 | ||
|
|
5630eb5041 | ||
|
|
31f62ab249 | ||
|
|
9db9c84cc9 | ||
|
|
ab9eac9ee8 | ||
|
|
3cce02193d | ||
|
|
a1a42f7afc | ||
|
|
d36c209191 | ||
|
|
7f4379acc0 | ||
|
|
7ba8cdf148 |
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
@@ -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())):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Executable
+24
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user