Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1757d3dba0 | ||
|
|
a9d0c29ed9 | ||
|
|
a335811869 | ||
|
|
357b0533fe | ||
|
|
2ec3732758 | ||
|
|
a004408a93 | ||
|
|
007e237e69 | ||
|
|
8e18dc9f71 | ||
|
|
7ab32539af | ||
|
|
6ef8fcb61d | ||
|
|
016e24da63 | ||
|
|
85b8717545 | ||
|
|
657fd745e1 |
@@ -77,6 +77,8 @@ jobs:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
|
||||
@@ -10,7 +10,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
|
||||
@@ -33,7 +33,7 @@ repos:
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
rev: v0.11.12
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
@@ -48,7 +48,7 @@ repos:
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -72,6 +72,12 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
Then download STA mask strategy from Hugging Face
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/STA_Mask_Strategy --local_dir=assets/ --repo_type=dataset
|
||||
```
|
||||
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
|
||||
(optimizations-sage)=
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
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__)
|
||||
|
||||
def main(args):
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
# 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(args.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=args.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(args.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(
|
||||
"--samples_per_file",
|
||||
type=int,
|
||||
default=64
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flush_frequency",
|
||||
type=int,
|
||||
default=256,
|
||||
help="how often to save to parquet files"
|
||||
)
|
||||
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)
|
||||
@@ -63,7 +63,7 @@ class WanVAEArchConfig(VAEArchConfig):
|
||||
|
||||
@dataclass
|
||||
class WanVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
arch_config: WanVAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
use_feature_cache: bool = True
|
||||
|
||||
use_tiling: bool = False
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import os
|
||||
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
|
||||
|
||||
def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
|
||||
]
|
||||
resize = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width)),
|
||||
]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
start_idx=start_idx)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
@@ -0,0 +1,44 @@
|
||||
# schema.py
|
||||
"""
|
||||
Unified data schema and format for saving and loading image/video data after
|
||||
preprocessing.
|
||||
|
||||
It uses apache arrow in-memory format that can be consumed by modern data
|
||||
frameworks that can handle parquet or lance file.
|
||||
"""
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
pyarrow_schema = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
@@ -0,0 +1,129 @@
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
) -> None:
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
self.cfg_rate = cfg_rate
|
||||
self.datase_dir_path = os.path.dirname(json_path)
|
||||
self.video_dir = os.path.join(self.datase_dir_path, "video")
|
||||
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
|
||||
self.prompt_embed_dir = os.path.join(self.datase_dir_path,
|
||||
"prompt_embed")
|
||||
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
|
||||
"prompt_attention_mask")
|
||||
with open(self.json_path) as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
|
||||
self.num_latent_t = num_latent_t
|
||||
|
||||
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
|
||||
|
||||
self.uncond_prompt_mask = torch.zeros(256).bool()
|
||||
self.lengths = [
|
||||
data_item.get("length", 1) for data_item in self.data_anno
|
||||
]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
latent_file = self.data_anno[idx]["latent_path"]
|
||||
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
|
||||
prompt_attention_mask_file = self.data_anno[idx][
|
||||
"prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(
|
||||
os.path.join(self.latent_dir, latent_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
if random.random() < self.cfg_rate:
|
||||
prompt_embed = self.uncond_prompt_embed
|
||||
prompt_attention_mask = self.uncond_prompt_mask
|
||||
else:
|
||||
prompt_embed = torch.load(
|
||||
os.path.join(self.prompt_embed_dir, prompt_embed_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
prompt_attention_mask = torch.load(
|
||||
os.path.join(self.prompt_attention_mask_dir,
|
||||
prompt_attention_mask_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
return latent, prompt_embed, prompt_attention_mask
|
||||
|
||||
def __len__(self):
|
||||
return len(self.data_anno)
|
||||
|
||||
|
||||
def latent_collate_function(batch):
|
||||
# return latent, prompt, latent_attn_mask, text_attn_mask
|
||||
# latent_attn_mask: # b t h w
|
||||
# text_attn_mask: b 1 l
|
||||
# needs to check if the latent/prompt' size and apply padding & attn mask
|
||||
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
|
||||
# calculate max shape
|
||||
max_t = max([latent.shape[1] for latent in latents])
|
||||
max_h = max([latent.shape[2] for latent in latents])
|
||||
max_w = max([latent.shape[3] for latent in latents])
|
||||
|
||||
# padding
|
||||
latent_list: list[torch.Tensor] = [
|
||||
torch.nn.functional.pad(
|
||||
latent,
|
||||
(
|
||||
0,
|
||||
max_t - latent.shape[1],
|
||||
0,
|
||||
max_h - latent.shape[2],
|
||||
0,
|
||||
max_w - latent.shape[3],
|
||||
),
|
||||
) for latent in latents
|
||||
]
|
||||
# attn mask
|
||||
latent_attn_mask = torch.ones(len(latent_list), max_t, max_h, max_w)
|
||||
# set to 0 if padding
|
||||
for i, latent in enumerate(latent_list):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
latents = torch.stack(latent_list, dim=0)
|
||||
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
|
||||
num_latent_t=28,
|
||||
cfg_rate=0.0)
|
||||
dataloader = torch.utils.data.DataLoader(dataset,
|
||||
batch_size=2,
|
||||
shuffle=False,
|
||||
collate_fn=latent_collate_function)
|
||||
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
|
||||
print(
|
||||
latent.shape,
|
||||
prompt_embed.shape,
|
||||
latent_attn_mask.shape,
|
||||
prompt_attention_mask.shape,
|
||||
)
|
||||
import pdb
|
||||
|
||||
pdb.set_trace()
|
||||
@@ -0,0 +1,372 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from einops import rearrange
|
||||
from torch import distributed as dist
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sp_group)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ParquetVideoTextDataset(Dataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
def __init__(self,
|
||||
path: str,
|
||||
batch_size: int = 1024,
|
||||
rank: int = 0,
|
||||
world_size: int = 1,
|
||||
cfg_rate: float = 0.0,
|
||||
num_latent_t: int = 2,
|
||||
seed: int = 0):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.rank = rank
|
||||
self.local_rank = get_sequence_model_parallel_rank()
|
||||
self.sp_world_size = world_size
|
||||
self.world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
self.cfg_rate = cfg_rate
|
||||
self.num_latent_t = num_latent_t
|
||||
self.local_indices = None
|
||||
self.plan_output_dir = os.path.join(
|
||||
self.path, f"data_plan_{self.world_size}_{self.sp_world_size}.json")
|
||||
|
||||
ranks = get_sp_group().ranks
|
||||
group_ranks: List[List] = [[] for _ in range(self.world_size)]
|
||||
torch.distributed.all_gather_object(group_ranks, ranks)
|
||||
|
||||
if rank == 0:
|
||||
# If a plan already exists, then skip creating a new plan
|
||||
# This will be useful when resume training
|
||||
if os.path.exists(self.plan_output_dir):
|
||||
print(f"Using existing plan from {self.plan_output_dir}")
|
||||
dist.barrier()
|
||||
return
|
||||
|
||||
# Find all parquet files recursively, and record num_rows for each file
|
||||
print(f"Scanning for parquet files in {self.path}")
|
||||
metadatas = []
|
||||
for root, _, files in os.walk(self.path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
for row_idx in range(num_rows):
|
||||
metadatas.append((file_path, row_idx))
|
||||
|
||||
# Generate the plan that distribute rows among workers
|
||||
random.seed(seed)
|
||||
random.shuffle(metadatas)
|
||||
|
||||
# Get all sp groups
|
||||
# e.g. if num_gpus = 4, sp_size = 2
|
||||
# group_ranks = [(0, 1), (2, 3)]
|
||||
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
|
||||
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
|
||||
group_ranks_list: List[Any] = list(
|
||||
set(tuple(r) for r in group_ranks))
|
||||
num_sp_groups = len(group_ranks_list)
|
||||
plan = defaultdict(list)
|
||||
for idx, metadata in enumerate(metadatas):
|
||||
sp_group_idx = idx % num_sp_groups
|
||||
for global_rank in group_ranks_list[sp_group_idx]:
|
||||
plan[global_rank].append(metadata)
|
||||
|
||||
with open(self.plan_output_dir, "w") as f:
|
||||
json.dump(plan, f)
|
||||
dist.barrier()
|
||||
|
||||
def __len__(self):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[str(self.rank)]
|
||||
except Exception as err:
|
||||
raise Exception(
|
||||
"The data plan hasn't been created yet") from err
|
||||
assert self.local_indices is not None
|
||||
return len(self.local_indices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[self.rank]
|
||||
except Exception as err:
|
||||
raise Exception(
|
||||
"The data plan hasn't been created yet") from err
|
||||
assert self.local_indices is not None
|
||||
file_path, row_idx = self.local_indices[idx]
|
||||
parquet_file = pq.ParquetFile(file_path)
|
||||
|
||||
# Calculate the row group to read into memory and the local idx
|
||||
# This way we can avoid reading in the entire parquet file
|
||||
cumulative = 0
|
||||
for i in range(parquet_file.num_row_groups):
|
||||
num_rows = parquet_file.metadata.row_group(i).num_rows
|
||||
if cumulative + num_rows > idx:
|
||||
row_group_index = i
|
||||
local_index = idx - cumulative
|
||||
break
|
||||
cumulative += num_rows
|
||||
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
processed = self._process_row(row_dict)
|
||||
lat, emb, mask, info = processed["latents"], processed[
|
||||
"embeddings"], processed["masks"], processed["info"]
|
||||
if lat.numel() == 0: # Validation parquet
|
||||
return lat, emb, mask, info
|
||||
else:
|
||||
lat = lat[:, -self.num_latent_t:]
|
||||
if self.sp_world_size > 1:
|
||||
lat = rearrange(lat,
|
||||
"t (n s) h w -> t n s h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
lat = lat[:, self.local_rank, :, :, :]
|
||||
return lat, emb, mask, info
|
||||
|
||||
def _process_row(self, row) -> Dict[str, Any]:
|
||||
"""Process a PyArrow batch into tensors."""
|
||||
|
||||
vae_latent_bytes = row["vae_latent_bytes"]
|
||||
vae_latent_shape = row["vae_latent_shape"]
|
||||
text_embedding_bytes = row["text_embedding_bytes"]
|
||||
text_embedding_shape = row["text_embedding_shape"]
|
||||
text_attention_mask_bytes = row["text_attention_mask_bytes"]
|
||||
text_attention_mask_shape = row["text_attention_mask_shape"]
|
||||
|
||||
# Process latent
|
||||
if not vae_latent_shape: # No VAE latent is stored. Split is validation
|
||||
lat = np.array([])
|
||||
else:
|
||||
lat = np.frombuffer(vae_latent_bytes,
|
||||
dtype=np.float32).reshape(vae_latent_shape)
|
||||
# Make array writable
|
||||
lat = np.copy(lat)
|
||||
|
||||
if random.random() < self.cfg_rate:
|
||||
emb = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
emb = np.frombuffer(text_embedding_bytes,
|
||||
dtype=np.float32).reshape(text_embedding_shape)
|
||||
# Make array writable
|
||||
emb = np.copy(emb)
|
||||
if emb.shape[0] < 512:
|
||||
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
|
||||
padded_emb[:emb.shape[0], :] = emb
|
||||
emb = padded_emb
|
||||
elif emb.shape[0] > 512:
|
||||
emb = emb[:512, :]
|
||||
|
||||
# Process mask
|
||||
if len(text_attention_mask_bytes) > 0 and len(
|
||||
text_attention_mask_shape) > 0:
|
||||
msk = np.frombuffer(text_attention_mask_bytes,
|
||||
dtype=np.uint8).astype(np.bool_)
|
||||
msk = msk.reshape(1, -1)
|
||||
# Make array writable
|
||||
msk = np.copy(msk)
|
||||
if msk.shape[1] < 512:
|
||||
padded_msk = np.zeros((1, 512), dtype=np.bool_)
|
||||
padded_msk[:, :msk.shape[1]] = msk
|
||||
msk = padded_msk
|
||||
elif msk.shape[1] > 512:
|
||||
msk = msk[:, :512]
|
||||
else:
|
||||
msk = np.ones((1, 512), dtype=np.bool_)
|
||||
|
||||
# Collect metadata
|
||||
info = {
|
||||
"width": row["width"],
|
||||
"height": row["height"],
|
||||
"num_frames": row["num_frames"],
|
||||
"duration_sec": row["duration_sec"],
|
||||
"fps": row["fps"],
|
||||
"file_name": row["file_name"],
|
||||
"caption": row["caption"],
|
||||
}
|
||||
|
||||
return {
|
||||
"latents": torch.from_numpy(lat),
|
||||
"embeddings": torch.from_numpy(emb),
|
||||
"masks": torch.from_numpy(msk),
|
||||
"info": info
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Benchmark Parquet dataset loading speed')
|
||||
parser.add_argument('--path',
|
||||
type=str,
|
||||
default="your/dataset/path",
|
||||
help='Path to Parquet dataset')
|
||||
parser.add_argument('--batch_size',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Batch size for DataLoader')
|
||||
parser.add_argument('--num_batches',
|
||||
type=int,
|
||||
default=100,
|
||||
help='Number of batches to benchmark')
|
||||
parser.add_argument('--vae_debug', action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
# Initialize distributed training
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
# Initialize CUDA device first
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
|
||||
# Initialize distributed training
|
||||
if world_size > 1:
|
||||
dist.init_process_group(backend="nccl",
|
||||
init_method="env://",
|
||||
world_size=world_size,
|
||||
rank=rank)
|
||||
print(
|
||||
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
|
||||
)
|
||||
|
||||
# Create dataset
|
||||
dataset = ParquetVideoTextDataset(
|
||||
args.path,
|
||||
batch_size=args.batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataloader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_size=args.batch_size,
|
||||
num_workers=1, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
drop_last=True)
|
||||
|
||||
# Example of how to load dataloader state
|
||||
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
|
||||
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
|
||||
# dataloader.load_state_dict(dataloader_state[rank])
|
||||
|
||||
# Warm-up with synchronization
|
||||
if rank == 0:
|
||||
print("Warming up...")
|
||||
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
|
||||
# Example of how to save dataloader state
|
||||
# if i == 30:
|
||||
# dist.barrier()
|
||||
# local_data = {rank: dataloader.state_dict()}
|
||||
# gathered_data = [None] * world_size
|
||||
# dist.all_gather_object(gathered_data, local_data)
|
||||
# if rank == 0:
|
||||
# global_state_dict = {}
|
||||
# for d in gathered_data:
|
||||
# global_state_dict.update(d)
|
||||
# torch.save(global_state_dict, "dataloader_state.pt")
|
||||
assert torch.sum(masks[0]).item() == torch.count_nonzero(
|
||||
embeddings[0]).item() // 4096
|
||||
if args.vae_debug:
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=VAE_PATH,
|
||||
vae_config=WanVAEConfig(load_encoder=False),
|
||||
vae_precision="fp32")
|
||||
fastvideo_args.device = device
|
||||
vae_loader = VAELoader()
|
||||
vae = vae_loader.load(model_path=VAE_PATH,
|
||||
architecture="",
|
||||
fastvideo_args=fastvideo_args)
|
||||
|
||||
videoprocessor = VideoProcessor(vae_scale_factor=8)
|
||||
|
||||
with torch.inference_mode():
|
||||
video = vae.decode(latents[0].unsqueeze(0).to(device))
|
||||
video = videoprocessor.postprocess_video(video)
|
||||
video_path = os.path.join("/workspace/FastVideo/debug_videos",
|
||||
infos["caption"][0][:50] + ".mp4")
|
||||
export_to_video(video[0], video_path, fps=16)
|
||||
|
||||
# Move data to device
|
||||
# latents = latents.to(device)
|
||||
# embeddings = embeddings.to(device)
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
# Benchmark
|
||||
if rank == 0:
|
||||
print(f"Benchmarking with batch_size={args.batch_size}")
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
for i, (latents, embeddings, masks,
|
||||
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
|
||||
if i >= args.num_batches:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(device)
|
||||
embeddings = embeddings.to(device)
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
total_samples += batch_size
|
||||
|
||||
# Print progress only from rank 0
|
||||
if rank == 0 and (i + 1) % 10 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
print(
|
||||
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
|
||||
)
|
||||
|
||||
# Final statistics
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if rank == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
|
||||
print("\nBenchmark Results:")
|
||||
print(f"Total time: {elapsed:.2f} seconds")
|
||||
print(f"Total samples: {total_samples}")
|
||||
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
|
||||
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
|
||||
|
||||
if world_size > 1:
|
||||
dist.destroy_process_group()
|
||||
@@ -0,0 +1,349 @@
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances: dict[type, 'SingletonMeta'] = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cap_list: list[dict] = []
|
||||
self.elements: list[int] = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements: dict[int, list[int]] = {}
|
||||
self.n_used_elements: dict[int, int] = {}
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info) -> int:
|
||||
worker_id = 0 if work_info is None else work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] %
|
||||
len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h: int,
|
||||
w: int,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16) -> bool:
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0) -> None:
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx) -> dict:
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx) -> dict:
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx) -> dict:
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images"
|
||||
in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps: list[str] = (image_data["cap"] if isinstance(
|
||||
image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
single_text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
single_text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices) -> torch.Tensor:
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data) -> list[dict]:
|
||||
cap_lists = []
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self) -> list:
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -0,0 +1,153 @@
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip) -> bool:
|
||||
if not torch.is_tensor(clip):
|
||||
raise TypeError(f"clip should be Tensor. Got {type(clip)}")
|
||||
|
||||
if not clip.ndimension() == 4:
|
||||
raise ValueError(f"clip should be 4D. Got {clip.dim()}D")
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i:i + h, j:j + w]
|
||||
|
||||
|
||||
def resize(clip, target_size, interpolation_mode) -> torch.Tensor:
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(
|
||||
f"target size should be tuple (height, width), instead got {target_size}"
|
||||
)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
size=target_size,
|
||||
mode=interpolation_mode,
|
||||
align_corners=True,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
|
||||
def center_crop_th_tw(clip, th, tw, top_crop) -> torch.Tensor:
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
tr = th / tw
|
||||
if h / w > tr:
|
||||
new_h = int(w * tr)
|
||||
new_w = w
|
||||
else:
|
||||
new_h = h
|
||||
new_w = int(h / tr)
|
||||
|
||||
i = 0 if top_crop else int(round((h - new_h) / 2.0))
|
||||
j = int(round((w - new_w) / 2.0))
|
||||
return crop(clip, i, j, new_h, new_w)
|
||||
|
||||
|
||||
def normalize_video(clip) -> torch.Tensor:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError(
|
||||
f"clip tensor should have data type uint8. Got {clip.dtype}")
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
|
||||
class CenterCropResizeVideo:
|
||||
"""
|
||||
First use the short side for cropping length,
|
||||
center crop video, then resize to the specified size
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
) -> None:
|
||||
if len(size) != 2:
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
self.top_crop = top_crop
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_center_crop = center_crop_th_tw(clip,
|
||||
self.size[0],
|
||||
self.size[1],
|
||||
top_crop=self.top_crop)
|
||||
clip_center_crop_resize = resize(
|
||||
clip_center_crop,
|
||||
target_size=self.size,
|
||||
interpolation_mode=self.interpolation_mode,
|
||||
)
|
||||
return clip_center_crop_resize
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class Normalize255:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def __call__(self, clip) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
return normalize_video(clip)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return self.__class__.__name__
|
||||
|
||||
|
||||
class TemporalRandomCrop:
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, size) -> None:
|
||||
self.size = size
|
||||
|
||||
def __call__(self, total_frames) -> tuple[int, int]:
|
||||
rand_end = max(0, total_frames - self.size - 1)
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
@@ -0,0 +1,10 @@
|
||||
from huggingface_hub import HfApi, upload_folder
|
||||
|
||||
api = HfApi()
|
||||
repo_id = "weizhou03/HD-Mixkit-Finetune-Wan" # customize this
|
||||
api.create_repo(repo_id=repo_id, repo_type="dataset")
|
||||
|
||||
upload_folder(repo_id=repo_id,
|
||||
folder_path="/workspace/data/HD-Mixkit-Finetune-Wan",
|
||||
repo_type="dataset",
|
||||
path_in_repo="")
|
||||
@@ -5,7 +5,8 @@ from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size, get_world_group,
|
||||
init_distributed_environment, initialize_model_parallel)
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
|
||||
__all__ = [
|
||||
@@ -17,4 +18,5 @@ __all__ = [
|
||||
"get_tensor_model_parallel_world_size",
|
||||
"cleanup_dist_env_and_memory",
|
||||
"get_world_group",
|
||||
"model_parallel_is_initialized",
|
||||
]
|
||||
|
||||
@@ -1,16 +1,182 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
|
||||
|
||||
from typing import Optional
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
from torch import Tensor
|
||||
from torch.distributed import ProcessGroup, ReduceOp
|
||||
|
||||
|
||||
class DistributedAutograd:
|
||||
"""Collection of autograd functions for distributed operations.
|
||||
|
||||
This class provides custom autograd functions for distributed operations like all_reduce,
|
||||
all_gather, and all_to_all. Each operation is implemented as a static inner class with
|
||||
proper forward and backward implementations.
|
||||
"""
|
||||
|
||||
class AllReduce(torch.autograd.Function):
|
||||
"""Differentiable all_reduce operation.
|
||||
|
||||
The gradient of all_reduce is another all_reduce operation since the operation
|
||||
combines values from all ranks equally.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any,
|
||||
group: ProcessGroup,
|
||||
input_: Tensor,
|
||||
op: Optional[dist.ReduceOp] = None) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.op = op
|
||||
output = input_.clone()
|
||||
dist.all_reduce(output, group=group, op=op)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None]:
|
||||
grad_output = grad_output.clone()
|
||||
dist.all_reduce(grad_output, group=ctx.group, op=ctx.op)
|
||||
return None, grad_output, None
|
||||
|
||||
class AllGather(torch.autograd.Function):
|
||||
"""Differentiable all_gather operation.
|
||||
|
||||
The operation gathers tensors from all ranks and concatenates them along a specified dimension.
|
||||
The backward pass uses reduce_scatter to efficiently distribute gradients back to source ranks.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
|
||||
world_size: int, dim: int) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.world_size = world_size
|
||||
ctx.dim = dim
|
||||
ctx.input_shape = input_.shape
|
||||
|
||||
input_size = input_.size()
|
||||
output_size = (input_size[0] * world_size, ) + input_size[1:]
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
|
||||
dist.all_gather_into_tensor(output_tensor, input_, group=group)
|
||||
|
||||
output_tensor = output_tensor.reshape((world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
|
||||
# Split the gradient tensor along the gathered dimension
|
||||
dim_size = grad_output.size(ctx.dim) // ctx.world_size
|
||||
grad_chunks = grad_output.reshape(grad_output.shape[:ctx.dim] +
|
||||
(ctx.world_size, dim_size) +
|
||||
grad_output.shape[ctx.dim + 1:])
|
||||
grad_chunks = grad_chunks.movedim(ctx.dim, 0)
|
||||
|
||||
# Each rank only needs its corresponding gradient
|
||||
grad_input = torch.empty(ctx.input_shape,
|
||||
dtype=grad_output.dtype,
|
||||
device=grad_output.device)
|
||||
dist.reduce_scatter_tensor(grad_input,
|
||||
grad_chunks.contiguous(),
|
||||
group=ctx.group)
|
||||
|
||||
return None, grad_input, None, None
|
||||
|
||||
class AllToAll4D(torch.autograd.Function):
|
||||
"""Differentiable all_to_all operation specialized for 4D tensors.
|
||||
|
||||
This operation is particularly useful for attention operations where we need to
|
||||
redistribute data across ranks for efficient parallel processing.
|
||||
|
||||
The operation supports two modes:
|
||||
1. scatter_dim=2, gather_dim=1: Used for redistributing attention heads
|
||||
2. scatter_dim=1, gather_dim=2: Used for redistributing sequence dimensions
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
|
||||
world_size: int, scatter_dim: int,
|
||||
gather_dim: int) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.world_size = world_size
|
||||
ctx.scatter_dim = scatter_dim
|
||||
ctx.gather_dim = gather_dim
|
||||
|
||||
if world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * world_size
|
||||
shard_hc = hc // world_size
|
||||
|
||||
input_t = input_.reshape(bs, shard_seqlen, world_size, shard_hc,
|
||||
hs).transpose(0, 2).contiguous()
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous()
|
||||
output = output.reshape(bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * world_size
|
||||
shard_seqlen = seqlen // world_size
|
||||
|
||||
input_t = input_.reshape(bs, world_size, shard_seqlen, shard_hc,
|
||||
hs)
|
||||
input_t = input_t.transpose(0, 3).transpose(0, 1).contiguous()
|
||||
input_t = input_t.reshape(world_size, shard_hc, shard_seqlen,
|
||||
bs, hs)
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
|
||||
output = output.reshape(hc, shard_seqlen, bs, hs)
|
||||
output = output.transpose(0, 2).contiguous()
|
||||
output = output.reshape(bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Invalid scatter_dim={scatter_dim}, gather_dim={gather_dim}. "
|
||||
f"Only (scatter_dim=2, gather_dim=1) and (scatter_dim=1, gather_dim=2) are supported."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def backward(
|
||||
ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None, None, None]:
|
||||
if ctx.world_size == 1:
|
||||
return None, grad_output, None, None, None
|
||||
|
||||
# For backward pass, we swap scatter_dim and gather_dim
|
||||
output = DistributedAutograd.AllToAll4D.apply(
|
||||
ctx.group, grad_output, ctx.world_size, ctx.gather_dim,
|
||||
ctx.scatter_dim)
|
||||
return None, output, None, None, None
|
||||
|
||||
|
||||
class DeviceCommunicatorBase:
|
||||
"""
|
||||
Base class for device-specific communicator.
|
||||
Base class for device-specific communicator with autograd support.
|
||||
It can use the `cpu_group` to initialize the communicator.
|
||||
If the device has PyTorch integration (PyTorch can recognize its
|
||||
communication backend), the `device_group` will also be given.
|
||||
@@ -33,35 +199,28 @@ class DeviceCommunicatorBase:
|
||||
self.rank_in_group = dist.get_group_rank(self.cpu_group,
|
||||
self.global_rank)
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
dist.all_reduce(input_, group=self.device_group)
|
||||
return input_
|
||||
def all_reduce(self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[dist.ReduceOp] = ReduceOp.SUM) -> torch.Tensor:
|
||||
"""Performs an all_reduce operation with gradient support."""
|
||||
return DistributedAutograd.AllReduce.apply(self.device_group, input_,
|
||||
op)
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
"""Performs an all_gather operation with gradient support."""
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
input_size = input_.size()
|
||||
# NOTE: we have to use concat-style all-gather here,
|
||||
# stack-style all-gather has compatibility issues with
|
||||
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
|
||||
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
|
||||
# Allocate output tensor.
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
# All-gather.
|
||||
dist.all_gather_into_tensor(output_tensor,
|
||||
input_,
|
||||
group=self.device_group)
|
||||
# Reshape
|
||||
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(self.world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
return DistributedAutograd.AllGather.apply(self.device_group, input_,
|
||||
self.world_size, dim)
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Performs a 4D all-to-all operation with gradient support."""
|
||||
return DistributedAutograd.AllToAll4D.apply(self.device_group, input_,
|
||||
self.world_size,
|
||||
scatter_dim, gather_dim)
|
||||
|
||||
def gather(self,
|
||||
input_: torch.Tensor,
|
||||
@@ -95,81 +254,6 @@ class DeviceCommunicatorBase:
|
||||
output_tensor = None
|
||||
return output_tensor
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
|
||||
|
||||
Args:
|
||||
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
|
||||
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
|
||||
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after all-to-all operation.
|
||||
"""
|
||||
# Bypass the function if we are using only 1 GPU.
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * self.world_size
|
||||
shard_hc = hc // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, shard_seqlen, self.world_size,
|
||||
shard_hc, hs).transpose(0,
|
||||
2).contiguous())
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous().reshape(
|
||||
bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * self.world_size
|
||||
shard_seqlen = seqlen // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, self.world_size, shard_seqlen,
|
||||
shard_hc, hs).transpose(0, 3).transpose(
|
||||
0, 1).contiguous().reshape(
|
||||
self.world_size, shard_hc,
|
||||
shard_seqlen, bs, hs))
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(hc, shard_seqlen, bs,
|
||||
hs).transpose(0, 2).contiguous().reshape(
|
||||
bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
|
||||
@@ -29,17 +29,19 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def all_reduce(self, input_):
|
||||
def all_reduce(self,
|
||||
input_,
|
||||
op: Optional[torch.distributed.ReduceOp] = None):
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
out = pynccl_comm.all_reduce(input_)
|
||||
out = pynccl_comm.all_reduce(input_, op=op)
|
||||
if out is None:
|
||||
# fall back to the default all-reduce using PyTorch.
|
||||
# this usually happens during testing.
|
||||
# when we run the model, allreduce only happens for the TP
|
||||
# group, where we always have either custom allreduce or pynccl.
|
||||
out = input_.clone()
|
||||
torch.distributed.all_reduce(out, group=self.device_group)
|
||||
torch.distributed.all_reduce(out, group=self.device_group, op=op)
|
||||
return out
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
|
||||
@@ -35,7 +35,7 @@ from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
from torch.distributed import Backend, ProcessGroup
|
||||
from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
@@ -260,7 +260,11 @@ class GroupCoordinator:
|
||||
with torch.cuda.stream(stream):
|
||||
yield graph_capture_context
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
def all_reduce(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
User-facing all-reduce function before we actually call the
|
||||
all-reduce operation.
|
||||
@@ -283,10 +287,14 @@ class GroupCoordinator:
|
||||
return torch.ops.vllm.all_reduce(input_,
|
||||
group_name=self.unique_name)
|
||||
else:
|
||||
return self._all_reduce_out_place(input_)
|
||||
return self._all_reduce_out_place(input_, op=op)
|
||||
|
||||
def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
return self.device_communicator.all_reduce(input_)
|
||||
def _all_reduce_out_place(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
return self.device_communicator.all_reduce(input_, op=op)
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
world_size = self.world_size
|
||||
@@ -647,7 +655,7 @@ class GroupCoordinator:
|
||||
tensor_dict[key] = value
|
||||
return tensor_dict
|
||||
|
||||
def barrier(self):
|
||||
def barrier(self) -> None:
|
||||
"""Barrier synchronization among the group.
|
||||
NOTE: don't use `device_group` here! `barrier` in NCCL is
|
||||
terrible because it is internally a broadcast operation with
|
||||
|
||||
@@ -100,6 +100,10 @@ class FastVideoArgs:
|
||||
device_str: Optional[str] = None
|
||||
device = None
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
|
||||
def __post_init__(self):
|
||||
pass
|
||||
|
||||
@@ -132,6 +136,13 @@ class FastVideoArgs:
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--inference-mode",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.inference_mode,
|
||||
help="Whether to use inference mode",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
@@ -423,3 +434,334 @@ def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
# TODO(will): may need to handle this for CI.
|
||||
raise ValueError("Current fastvideo args is not set.")
|
||||
return _current_fastvideo_args
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TrainingArgs(FastVideoArgs):
|
||||
"""
|
||||
Training arguments. Inherits from FastVideoArgs and adds training-specific
|
||||
arguments. If there are any conflicts, the training arguments will take
|
||||
precedence.
|
||||
"""
|
||||
data_path: str = ""
|
||||
dataloader_num_workers: int = 0
|
||||
num_height: int = 0
|
||||
num_width: int = 0
|
||||
num_frames: int = 0
|
||||
|
||||
train_batch_size: int = 0
|
||||
num_latent_t: int = 0
|
||||
group_frame: bool = False
|
||||
group_resolution: bool = False
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
cache_dir: str = ""
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
cfg: float = 0.0
|
||||
precondition_outputs: bool = False
|
||||
|
||||
# validation & logs
|
||||
validation_prompt_dir: str = ""
|
||||
validation_sampling_steps: str = ""
|
||||
validation_guidance_scale: str = ""
|
||||
validation_steps: float = 0.0
|
||||
log_validation: bool = False
|
||||
tracker_project_name: str = ""
|
||||
# seed: int
|
||||
|
||||
# output
|
||||
output_dir: str = ""
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: bool = False
|
||||
logging_dir: str = ""
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
max_train_steps: int = 0
|
||||
gradient_accumulation_steps: int = 0
|
||||
learning_rate: float = 0.0
|
||||
scale_lr: bool = False
|
||||
lr_scheduler: str = ""
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
gradient_checkpointing: bool = False
|
||||
selective_checkpointing: float = 0.0
|
||||
allow_tf32: bool = False
|
||||
mixed_precision: str = ""
|
||||
train_sp_batch_size: int = 0
|
||||
fsdp_sharding_startegy: str = ""
|
||||
|
||||
weighting_scheme: str = ""
|
||||
logit_mean: float = 0.0
|
||||
logit_std: float = 1.0
|
||||
mode_scale: float = 0.0
|
||||
|
||||
num_euler_timesteps: int = 0
|
||||
lr_num_cycles: int = 0
|
||||
lr_power: float = 0.0
|
||||
not_apply_cfg_solver: bool = False
|
||||
distill_cfg: float = 0.0
|
||||
scheduler_type: str = ""
|
||||
linear_quadratic_threshold: float = 0.0
|
||||
linear_range: float = 0.0
|
||||
weight_decay: float = 0.0
|
||||
use_ema: bool = False
|
||||
multi_phased_distill_schedule: str = ""
|
||||
pred_decay_weight: float = 0.0
|
||||
pred_decay_type: str = ""
|
||||
hunyuan_teacher_disable_cfg: bool = False
|
||||
|
||||
# master_weight_type
|
||||
master_weight_type: str = ""
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
parser.add_argument("--data-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to parquet files")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of workers for dataloader")
|
||||
parser.add_argument("--num-height",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of heights")
|
||||
parser.add_argument("--num-width",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of widths")
|
||||
parser.add_argument("--num-frames",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of frames")
|
||||
|
||||
# Training batch and model configuration
|
||||
parser.add_argument("--train-batch-size",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Training batch size")
|
||||
parser.add_argument("--num-latent-t",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of latent time steps")
|
||||
parser.add_argument("--group-frame",
|
||||
action=StoreBoolean,
|
||||
help="Whether to group frames during training")
|
||||
parser.add_argument("--group-resolution",
|
||||
action=StoreBoolean,
|
||||
help="Whether to group resolutions during training")
|
||||
|
||||
# Model paths
|
||||
parser.add_argument("--pretrained-model-name-or-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to pretrained model or model name")
|
||||
parser.add_argument("--dit-model-name-or-path",
|
||||
type=str,
|
||||
required=False,
|
||||
help="Path to DiT model or model name")
|
||||
parser.add_argument("--cache-dir",
|
||||
type=str,
|
||||
help="Directory to cache models")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
default=0.999,
|
||||
help="EMA decay rate")
|
||||
parser.add_argument("--ema-start-step",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Step to start EMA")
|
||||
parser.add_argument("--cfg",
|
||||
type=float,
|
||||
help="Classifier-free guidance scale")
|
||||
parser.add_argument(
|
||||
"--precondition-outputs",
|
||||
action=StoreBoolean,
|
||||
help="Whether to precondition the outputs of the model")
|
||||
|
||||
# Validation and logging
|
||||
parser.add_argument("--validation-prompt-dir",
|
||||
type=str,
|
||||
help="Directory containing validation prompts")
|
||||
parser.add_argument("--validation-sampling-steps",
|
||||
type=str,
|
||||
help="Validation sampling steps")
|
||||
parser.add_argument("--validation-guidance-scale",
|
||||
type=str,
|
||||
help="Validation guidance scale")
|
||||
parser.add_argument("--validation-steps",
|
||||
type=float,
|
||||
help="Number of validation steps")
|
||||
parser.add_argument("--log-validation",
|
||||
action=StoreBoolean,
|
||||
help="Whether to log validation results")
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
|
||||
# Output configuration
|
||||
parser.add_argument("--output-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Output directory for checkpoints and logs")
|
||||
parser.add_argument("--checkpoints-total-limit",
|
||||
type=int,
|
||||
help="Maximum number of checkpoints to keep")
|
||||
parser.add_argument("--checkpointing-steps",
|
||||
type=int,
|
||||
help="Steps between checkpoints")
|
||||
parser.add_argument("--resume-from-checkpoint",
|
||||
type=str,
|
||||
help="Path to checkpoint to resume from")
|
||||
parser.add_argument("--logging-dir",
|
||||
type=str,
|
||||
help="Directory for logging")
|
||||
|
||||
# Training configuration
|
||||
parser.add_argument("--num-train-epochs",
|
||||
type=int,
|
||||
help="Number of training epochs")
|
||||
parser.add_argument("--max-train-steps",
|
||||
type=int,
|
||||
help="Maximum number of training steps")
|
||||
parser.add_argument("--gradient-accumulation-steps",
|
||||
type=int,
|
||||
help="Number of steps to accumulate gradients")
|
||||
parser.add_argument("--learning-rate",
|
||||
type=float,
|
||||
required=True,
|
||||
help="Learning rate")
|
||||
parser.add_argument("--scale-lr",
|
||||
action=StoreBoolean,
|
||||
help="Whether to scale learning rate")
|
||||
parser.add_argument("--lr-scheduler",
|
||||
type=str,
|
||||
default="constant",
|
||||
help="Learning rate scheduler type")
|
||||
parser.add_argument("--lr-warmup-steps",
|
||||
type=int,
|
||||
default=10,
|
||||
help="Number of warmup steps for learning rate")
|
||||
parser.add_argument("--max-grad-norm",
|
||||
type=float,
|
||||
help="Maximum gradient norm")
|
||||
parser.add_argument("--gradient-checkpointing",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use gradient checkpointing")
|
||||
parser.add_argument("--selective-checkpointing",
|
||||
type=float,
|
||||
help="Selective checkpointing threshold")
|
||||
parser.add_argument("--allow-tf32",
|
||||
action=StoreBoolean,
|
||||
help="Whether to allow TF32")
|
||||
parser.add_argument("--mixed-precision",
|
||||
type=str,
|
||||
help="Mixed precision training type")
|
||||
parser.add_argument("--train-sp-batch-size",
|
||||
type=int,
|
||||
help="Training spatial parallelism batch size")
|
||||
|
||||
parser.add_argument("--fsdp-sharding-strategy",
|
||||
type=str,
|
||||
help="FSDP sharding strategy")
|
||||
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
default="uniform",
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="mean to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_std",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="std to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode_scale",
|
||||
type=float,
|
||||
default=1.29,
|
||||
help=
|
||||
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
)
|
||||
|
||||
# Additional training parameters
|
||||
parser.add_argument("--num-euler-timesteps",
|
||||
type=int,
|
||||
help="Number of Euler timesteps")
|
||||
parser.add_argument("--lr-num-cycles",
|
||||
type=int,
|
||||
help="Number of learning rate cycles")
|
||||
parser.add_argument("--lr-power",
|
||||
type=float,
|
||||
help="Learning rate power")
|
||||
parser.add_argument("--not-apply-cfg-solver",
|
||||
action=StoreBoolean,
|
||||
help="Whether to not apply CFG solver")
|
||||
parser.add_argument("--distill-cfg",
|
||||
type=float,
|
||||
help="Distillation CFG scale")
|
||||
parser.add_argument("--scheduler-type", type=str, help="Scheduler type")
|
||||
parser.add_argument("--linear-quadratic-threshold",
|
||||
type=float,
|
||||
help="Linear quadratic threshold")
|
||||
parser.add_argument("--linear-range", type=float, help="Linear range")
|
||||
parser.add_argument("--weight-decay", type=float, help="Weight decay")
|
||||
parser.add_argument("--use-ema",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA")
|
||||
parser.add_argument("--multi-phased-distill-schedule",
|
||||
type=str,
|
||||
help="Multi-phased distillation schedule")
|
||||
parser.add_argument("--pred-decay-weight",
|
||||
type=float,
|
||||
help="Prediction decay weight")
|
||||
parser.add_argument("--pred-decay-type",
|
||||
type=str,
|
||||
help="Prediction decay type")
|
||||
parser.add_argument("--hunyuan-teacher-disable-cfg",
|
||||
action=StoreBoolean,
|
||||
help="Whether to disable CFG for Hunyuan teacher")
|
||||
parser.add_argument("--master-weight-type",
|
||||
type=str,
|
||||
help="Master weight type")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -33,9 +33,11 @@ class BaseDiT(nn.Module, ABC):
|
||||
f"Subclasses of BaseDiT must define '{attr}' class variable"
|
||||
)
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any],
|
||||
**kwargs) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.hf_config = hf_config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -442,8 +442,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
|
||||
|
||||
def __init__(self, config: HunyuanVideoConfig):
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
self.patch_size = [
|
||||
config.patch_size_t, config.patch_size, config.patch_size
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
# ==============================================================================
|
||||
from typing import Dict, Optional, Tuple
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
@@ -462,8 +462,9 @@ class StepVideoModel(BaseDiT):
|
||||
_supported_attention_backends = StepVideoConfig(
|
||||
)._supported_attention_backends
|
||||
|
||||
def __init__(self, config: StepVideoConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.attention_head_dim = config.attention_head_dim
|
||||
self.in_channels = config.in_channels
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -298,7 +298,7 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
assert orig_dtype != torch.float32
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
@@ -360,8 +360,9 @@ class WanTransformer3DModel(CachableDiT):
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = WanVideoConfig()._param_names_mapping
|
||||
|
||||
def __init__(self, config: WanVideoConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
|
||||
@@ -6,6 +6,7 @@ import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from copy import deepcopy
|
||||
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
@@ -366,6 +367,7 @@ class TransformerLoader(ComponentLoader):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
hf_config = deepcopy(config)
|
||||
cls_name = config.pop("_class_name")
|
||||
if cls_name is None:
|
||||
raise ValueError(
|
||||
@@ -392,13 +394,23 @@ 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)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params={"config": dit_config},
|
||||
weight_dir_list=safetensors_list,
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
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,
|
||||
# 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,14 +14,18 @@ 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._tensor import distribute_tensor
|
||||
from torch.distributed.fsdp import (CPUOffloadPolicy, MixedPrecisionPolicy,
|
||||
fully_shard)
|
||||
from torch.nn.modules.module import _IncompatibleKeys
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
@contextlib.contextmanager
|
||||
@@ -91,11 +95,21 @@ def load_fsdp_model(
|
||||
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"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(get_sequence_model_parallel_world_size(), ),
|
||||
@@ -104,6 +118,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 +144,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 +172,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)
|
||||
|
||||
@@ -5,19 +5,25 @@ Base class for composed pipelines.
|
||||
This module defines the base class for pipelines that are composed of multiple stages.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from copy import deepcopy
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.pipelines import (PipelineConfig,
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages import PipelineStage
|
||||
from fastvideo.v1.utils import (maybe_download_model,
|
||||
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -34,20 +40,35 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
is_video_pipeline: bool = False # To be overridden by video pipelines
|
||||
_required_config_modules: List[str] = []
|
||||
training_args: Optional[TrainingArgs] = None
|
||||
fastvideo_args: Optional[FastVideoArgs] = None
|
||||
|
||||
# TODO(will): args should support both inference args and training args
|
||||
def __init__(self,
|
||||
model_path: str,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
config: Optional[Dict[str, Any]] = None):
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
required_config_modules: Optional[List[str]] = None):
|
||||
"""
|
||||
Initialize the pipeline. After __init__, the pipeline should be ready to
|
||||
use. The pipeline should be stateless and not hold any batch state.
|
||||
"""
|
||||
|
||||
if fastvideo_args.training_mode:
|
||||
assert isinstance(fastvideo_args, TrainingArgs)
|
||||
self.training_args = fastvideo_args
|
||||
assert self.training_args is not None
|
||||
else:
|
||||
self.fastvideo_args = fastvideo_args
|
||||
assert self.fastvideo_args is not None
|
||||
|
||||
self.model_path = model_path
|
||||
self._stages: List[PipelineStage] = []
|
||||
self._stage_name_mapping: Dict[str, PipelineStage] = {}
|
||||
|
||||
if required_config_modules is not None:
|
||||
self._required_config_modules = required_config_modules
|
||||
|
||||
if self._required_config_modules is None:
|
||||
raise NotImplementedError(
|
||||
"Subclass must set _required_config_modules")
|
||||
@@ -59,16 +80,131 @@ class ComposedPipelineBase(ABC):
|
||||
else:
|
||||
self.config = config
|
||||
|
||||
self.maybe_init_distributed_environment(fastvideo_args)
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
self.modules = self.load_modules(fastvideo_args)
|
||||
|
||||
if fastvideo_args.training_mode:
|
||||
assert self.training_args is not None
|
||||
if self.training_args.log_validation:
|
||||
self.initialize_validation_pipeline(self.training_args)
|
||||
self.initialize_training_pipeline(self.training_args)
|
||||
|
||||
self.initialize_pipeline(fastvideo_args)
|
||||
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
if not fastvideo_args.training_mode:
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
|
||||
def get_module(self, module_name: str) -> Any:
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"if training_mode is True, the pipeline must implement this method")
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"if log_validation is True, the pipeline must implement this method"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls,
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
args: Optional[argparse.Namespace] = None,
|
||||
required_config_modules: Optional[List[str]] = None,
|
||||
**kwargs) -> "ComposedPipelineBase":
|
||||
config = None
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
if isinstance(pipeline_config, str):
|
||||
config.load_from_json(pipeline_config)
|
||||
|
||||
# 2. If users also provide some kwargs, it will override the pipeline config.
|
||||
# The user kwargs shouldn't contain model config parameters!
|
||||
if config is None:
|
||||
logger.warning("No config found for model %s, using default config",
|
||||
model_path)
|
||||
config_args = kwargs
|
||||
else:
|
||||
config_args = shallow_asdict(config)
|
||||
config_args.update(kwargs)
|
||||
|
||||
if args is None or args.inference_mode:
|
||||
fastvideo_args = FastVideoArgs(model_path=model_path,
|
||||
device_str=device or "cuda" if
|
||||
torch.cuda.is_available() else "cpu",
|
||||
**config_args)
|
||||
|
||||
fastvideo_args.model_path = model_path
|
||||
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
|
||||
) else "cpu"
|
||||
for key, value in config_args.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
else:
|
||||
assert args is not None, "args must be provided for training mode"
|
||||
fastvideo_args = TrainingArgs.from_cli_args(args)
|
||||
# TODO(will): fix this so that its not so ugly
|
||||
fastvideo_args.model_path = model_path
|
||||
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
|
||||
) else "cpu"
|
||||
for key, value in config_args.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
|
||||
fastvideo_args.use_cpu_offload = False
|
||||
# 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("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
return cls(model_path,
|
||||
fastvideo_args,
|
||||
required_config_modules=required_config_modules)
|
||||
|
||||
def maybe_init_distributed_environment(self, fastvideo_args: FastVideoArgs):
|
||||
if model_parallel_is_initialized():
|
||||
return
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", -1))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", -1))
|
||||
rank = int(os.environ.get("RANK", -1))
|
||||
|
||||
if local_rank == -1 or world_size == -1 or rank == -1:
|
||||
raise ValueError(
|
||||
"Local rank, world size, and rank must be set. Use torchrun to launch the script."
|
||||
)
|
||||
|
||||
torch.cuda.set_device(local_rank)
|
||||
init_distributed_environment(world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
assert fastvideo_args.tp_size is not None, "tp_size must be set"
|
||||
assert fastvideo_args.sp_size is not None, "sp_size must be set"
|
||||
initialize_model_parallel(
|
||||
tensor_model_parallel_size=fastvideo_args.tp_size,
|
||||
sequence_model_parallel_size=fastvideo_args.sp_size)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
fastvideo_args.device = device
|
||||
|
||||
def get_module(self, module_name: str, default_value: Any = None) -> Any:
|
||||
if module_name not in self.modules:
|
||||
return default_value
|
||||
return self.modules[module_name]
|
||||
|
||||
def add_module(self, module_name: str, module: Any):
|
||||
@@ -114,6 +250,12 @@ class ComposedPipelineBase(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
Create the training pipeline stages.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
@@ -136,19 +278,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 +308,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 +341,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,559 @@
|
||||
# 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
|
||||
from typing import Any, Dict
|
||||
|
||||
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: Dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: Dict[str, Any] = {} # 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
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=batch_captions,
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
assert hasattr(self, "prompt_encoding_stage")
|
||||
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)
|
||||
|
||||
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("Collected batch with %s samples", len(table))
|
||||
|
||||
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("Chunks per worker: %s", 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("Processed chunk with %s samples",
|
||||
written)
|
||||
except Exception as e:
|
||||
work_range = futures[future]
|
||||
failed_ranges.append(work_range)
|
||||
logger.error("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
# Retry failed ranges sequentially
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.process_chunk_range(
|
||||
work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
logger.info("Total samples written: %s", 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=[],
|
||||
)
|
||||
assert hasattr(self, "prompt_encoding_stage")
|
||||
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()
|
||||
|
||||
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(
|
||||
"Shape after removing padding - Embeddings: %s, Mask: %s",
|
||||
text_embedding.shape, 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("Saved validation sample: %s", 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("Total validation samples: %s", 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("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.process_chunk_range(work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
logger.info("Total validation samples written: %s", total_written)
|
||||
|
||||
# Clear memory
|
||||
del table
|
||||
gc.collect() # Force garbage collection
|
||||
|
||||
@staticmethod
|
||||
def process_chunk_range(args: Any) -> int:
|
||||
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("Error processing chunks %s-%s for worker %s: %s",
|
||||
start_idx, end_idx, worker_id, str(e))
|
||||
raise
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline
|
||||
@@ -23,7 +23,6 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
st_attn_available = False
|
||||
spec = importlib.util.find_spec("st_attn")
|
||||
@@ -74,7 +73,9 @@ class DenoisingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
# TODO(will): make the precision configurable for inference
|
||||
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -15,8 +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 +46,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(
|
||||
|
||||
@@ -0,0 +1,515 @@
|
||||
import gc
|
||||
import os
|
||||
import traceback
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from diffusers.optimization import get_scheduler
|
||||
from einops import rearrange
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Note: if checking with float32, cannot use flash-attn.
|
||||
GRADIENT_CHECK_DTYPE = torch.bfloat16
|
||||
|
||||
|
||||
class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
"""
|
||||
A pipeline for training a model. All training pipelines should inherit from this class.
|
||||
All reusable components and code should be implemented in this class.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer"]
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing training pipeline...")
|
||||
self.device = training_args.device
|
||||
self.sp_group = get_sp_group()
|
||||
self.world_size = self.sp_group.world_size
|
||||
self.rank = self.sp_group.rank
|
||||
self.local_rank = self.sp_group.local_rank
|
||||
self.transformer = self.get_module("transformer")
|
||||
assert self.transformer is not None
|
||||
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
params_to_optimize = self.transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
self.init_steps = 0
|
||||
logger.info("optimizer: %s", self.optimizer)
|
||||
|
||||
self.lr_scheduler = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer,
|
||||
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
|
||||
num_training_steps=training_args.max_train_steps * self.world_size,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset = ParquetVideoTextDataset(
|
||||
training_args.data_path,
|
||||
batch_size=training_args.train_batch_size,
|
||||
rank=self.rank,
|
||||
world_size=self.world_size,
|
||||
cfg_rate=training_args.cfg,
|
||||
num_latent_t=training_args.num_latent_t)
|
||||
|
||||
self.train_dataloader = StatefulDataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=training_args.train_batch_size,
|
||||
num_workers=training_args.
|
||||
dataloader_num_workers, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
|
||||
drop_last=True)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
|
||||
if self.rank <= 0:
|
||||
project = training_args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=training_args)
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"Training pipelines must implement this method")
|
||||
|
||||
@abstractmethod
|
||||
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
|
||||
loader, noise_scheduler, noise_random_generator,
|
||||
gradient_accumulation_steps, sp_size,
|
||||
precondition_outputs, max_grad_norm, weighting_scheme,
|
||||
logit_mean, logit_std, mode_scale):
|
||||
"""
|
||||
Train one step of the model.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Training pipeline must implement this method")
|
||||
|
||||
def log_validation(self, transformer, training_args, global_step) -> None:
|
||||
assert training_args is not None
|
||||
training_args.inference_mode = True
|
||||
training_args.use_cpu_offload = False
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
raise ValueError("Validation pipeline is not set")
|
||||
|
||||
# Create sampling parameters if not provided
|
||||
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
|
||||
|
||||
# Prepare validation prompts
|
||||
logger.info('fastvideo_args.validation_prompt_dir: %s',
|
||||
training_args.validation_prompt_dir)
|
||||
validation_dataset = ParquetVideoTextDataset(
|
||||
training_args.validation_prompt_dir,
|
||||
batch_size=1,
|
||||
rank=0,
|
||||
world_size=1,
|
||||
cfg_rate=0,
|
||||
num_latent_t=training_args.num_latent_t)
|
||||
|
||||
validation_dataloader = StatefulDataLoader(
|
||||
validation_dataset,
|
||||
batch_size=1,
|
||||
num_workers=1, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
drop_last=False)
|
||||
|
||||
transformer.requires_grad_(False)
|
||||
for p in transformer.parameters():
|
||||
p.requires_grad = False
|
||||
transformer.eval()
|
||||
|
||||
# Add the transformer to the validation pipeline
|
||||
self.validation_pipeline.add_module("transformer", transformer)
|
||||
self.validation_pipeline.latent_preparation_stage.transformer = transformer # type: ignore[attr-defined]
|
||||
self.validation_pipeline.denoising_stage.transformer = transformer # type: ignore[attr-defined]
|
||||
|
||||
# Process each validation prompt
|
||||
videos = []
|
||||
captions = []
|
||||
for _, embeddings, masks, infos in validation_dataloader:
|
||||
logger.info("infos: %s", infos)
|
||||
caption = infos['caption']
|
||||
captions.append(caption)
|
||||
prompt_embeds = embeddings.to(training_args.device)
|
||||
prompt_attention_mask = masks.to(training_args.device)
|
||||
|
||||
# Calculate sizes
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8,
|
||||
sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
temporal_compression_factor = training_args.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
logger.info(
|
||||
"validation num_frames: %s, temporal_compression_factor: %s, num_latent_t: %s",
|
||||
num_frames, temporal_compression_factor,
|
||||
training_args.num_latent_t)
|
||||
# Prepare batch for validation
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
latents=None,
|
||||
# seed=sampling_param.seed,
|
||||
prompt_embeds=[prompt_embeds],
|
||||
prompt_attention_mask=[prompt_attention_mask],
|
||||
# make sure we use the same height, width, and num_frames as the training pipeline
|
||||
height=training_args.num_height,
|
||||
width=training_args.num_width,
|
||||
num_frames=num_frames,
|
||||
# num_inference_steps=fastvideo_args.validation_sampling_steps,
|
||||
num_inference_steps=sampling_param.num_inference_steps,
|
||||
# guidance_scale=fastvideo_args.validation_guidance_scale,
|
||||
guidance_scale=1,
|
||||
n_tokens=n_tokens,
|
||||
do_classifier_free_guidance=False,
|
||||
eta=0.0,
|
||||
extra={},
|
||||
)
|
||||
|
||||
# Run validation inference
|
||||
with torch.inference_mode(), torch.autocast("cuda",
|
||||
dtype=torch.bfloat16):
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
|
||||
# Process outputs
|
||||
video = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in video:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
videos.append(frames)
|
||||
|
||||
# Log validation results
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
if rank == 0:
|
||||
video_filenames = []
|
||||
video_captions = []
|
||||
for i, video in enumerate(videos):
|
||||
caption = captions[i]
|
||||
os.makedirs(training_args.output_dir, exist_ok=True)
|
||||
filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_video_{i}.mp4")
|
||||
imageio.mimsave(filename, video, fps=sampling_param.fps)
|
||||
video_filenames.append(filename)
|
||||
video_captions.append(
|
||||
caption) # Store the caption for each video
|
||||
|
||||
logs = {
|
||||
"validation_videos": [
|
||||
wandb.Video(filename,
|
||||
caption=caption) for filename, caption in zip(
|
||||
video_filenames, video_captions)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
|
||||
# Re-enable gradients for training
|
||||
transformer.requires_grad_(True)
|
||||
transformer.train()
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def gradient_check_parameters(self,
|
||||
transformer,
|
||||
latents,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
timesteps,
|
||||
target,
|
||||
eps=5e-2,
|
||||
max_params_to_check=2000) -> float:
|
||||
"""
|
||||
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
|
||||
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
|
||||
"""
|
||||
assert self.training_args is not None
|
||||
# Move all inputs to CPU and clear GPU memory
|
||||
inputs_cpu = {
|
||||
'latents': latents.cpu(),
|
||||
'encoder_hidden_states': encoder_hidden_states.cpu(),
|
||||
'encoder_attention_mask': encoder_attention_mask.cpu(),
|
||||
'timesteps': timesteps.cpu(),
|
||||
'target': target.cpu()
|
||||
}
|
||||
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def compute_loss() -> torch.Tensor:
|
||||
assert self.training_args is not None
|
||||
# Move inputs to GPU, compute loss, cleanup
|
||||
inputs_gpu = {
|
||||
k:
|
||||
v.to(self.training_args.device,
|
||||
dtype=GRADIENT_CHECK_DTYPE
|
||||
if k != 'encoder_attention_mask' else None)
|
||||
for k, v in inputs_cpu.items()
|
||||
}
|
||||
|
||||
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
|
||||
# with torch.autocast(enabled=False, device_type="cuda"):
|
||||
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
|
||||
with set_forward_context(
|
||||
current_timestep=inputs_gpu['timesteps'],
|
||||
attn_metadata=None):
|
||||
model_pred = transformer(
|
||||
hidden_states=inputs_gpu['latents'],
|
||||
encoder_hidden_states=inputs_gpu[
|
||||
'encoder_hidden_states'],
|
||||
timestep=inputs_gpu['timesteps'],
|
||||
encoder_attention_mask=inputs_gpu[
|
||||
'encoder_attention_mask'],
|
||||
return_dict=False)[0]
|
||||
|
||||
if self.training_args.precondition_outputs:
|
||||
sigmas = get_sigmas(self.noise_scheduler,
|
||||
inputs_gpu['latents'].device,
|
||||
inputs_gpu['timesteps'],
|
||||
n_dim=inputs_gpu['latents'].ndim,
|
||||
dtype=inputs_gpu['latents'].dtype)
|
||||
model_pred = inputs_gpu['latents'] - model_pred * sigmas
|
||||
target_adjusted = inputs_gpu['target']
|
||||
else:
|
||||
target_adjusted = inputs_gpu['target']
|
||||
|
||||
loss = torch.mean((model_pred - target_adjusted)**2)
|
||||
|
||||
# Cleanup and return
|
||||
loss_cpu = loss.cpu()
|
||||
del inputs_gpu, model_pred, target_adjusted
|
||||
if 'sigmas' in locals():
|
||||
del sigmas
|
||||
torch.cuda.empty_cache()
|
||||
return loss_cpu.to(self.training_args.device)
|
||||
|
||||
try:
|
||||
# Get analytical gradients
|
||||
transformer.zero_grad()
|
||||
analytical_loss = compute_loss()
|
||||
analytical_loss.backward()
|
||||
|
||||
# Check gradients for selected parameters
|
||||
absolute_errors: list[float] = []
|
||||
param_count = 0
|
||||
|
||||
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):
|
||||
full_grad = param.grad.full_tensor()
|
||||
distributed = True
|
||||
else:
|
||||
full_grad = param.grad
|
||||
distributed = False
|
||||
continue
|
||||
if not (param.requires_grad and param.grad is not None
|
||||
and param_count < max_params_to_check
|
||||
and full_grad.abs().max() > 5e-4):
|
||||
continue
|
||||
if not distributed and 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
|
||||
|
||||
# Find first significant gradient element
|
||||
flat_param = local_param.data.view(-1)
|
||||
flat_grad = local_grad.view(-1)
|
||||
check_idx = next((i for i in range(min(10, flat_param.numel()))
|
||||
if abs(flat_grad[i]) > 1e-4), 0)
|
||||
|
||||
# Store original values
|
||||
orig_value = flat_param[check_idx].item()
|
||||
analytical_grad = flat_grad[check_idx].item()
|
||||
|
||||
# Compute numerical gradient
|
||||
for delta in [eps, -eps]:
|
||||
with torch.no_grad():
|
||||
# only have a single rank modify the parameter
|
||||
# because we are using FSDP
|
||||
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(
|
||||
"%s[%s]: analytical=%.5f, numerical=%.5f, abs_error=%.2e, rel_error=%.2f%%",
|
||||
name, check_idx, analytical_grad, numerical_grad,
|
||||
abs_error, rel_error * 100)
|
||||
|
||||
# param_count += 1
|
||||
|
||||
# Compute and log statistics
|
||||
if rank <= 0 and absolute_errors:
|
||||
min_err, max_err, mean_err = min(absolute_errors), max(
|
||||
absolute_errors
|
||||
), sum(absolute_errors) / len(absolute_errors)
|
||||
logger.info("Gradient check stats: min=%s, max=%s, mean=%s",
|
||||
min_err, max_err, mean_err)
|
||||
|
||||
wandb.log({
|
||||
"grad_check/min_abs_error": min_err,
|
||||
"grad_check/max_abs_error": max_err,
|
||||
"grad_check/mean_abs_error": mean_err,
|
||||
"grad_check/analytical_loss": analytical_loss.item(),
|
||||
})
|
||||
return max_err
|
||||
|
||||
return float('inf')
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Gradient check failed: %s", e)
|
||||
traceback.print_exc()
|
||||
return float('inf')
|
||||
|
||||
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
|
||||
noise_random_generator) -> float | None:
|
||||
"""
|
||||
Setup and perform gradient check on a fresh batch.
|
||||
Args:
|
||||
args: Training arguments
|
||||
loader_iter: Data loader iterator
|
||||
noise_scheduler: Noise scheduler for diffusion
|
||||
noise_random_generator: Random number generator for noise
|
||||
Returns:
|
||||
float or None: Maximum gradient error or None if check is disabled/fails
|
||||
"""
|
||||
assert self.training_args is not None
|
||||
|
||||
try:
|
||||
# Get a fresh batch and process it exactly like train_one_step
|
||||
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
|
||||
loader_iter)
|
||||
|
||||
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
|
||||
check_latents = check_latents.to(self.training_args.device,
|
||||
dtype=GRADIENT_CHECK_DTYPE)
|
||||
check_encoder_hidden_states = check_encoder_hidden_states.to(
|
||||
self.training_args.device, dtype=GRADIENT_CHECK_DTYPE)
|
||||
check_latents = normalize_dit_input("wan", check_latents)
|
||||
batch_size = check_latents.shape[0]
|
||||
check_noise = torch.randn_like(check_latents)
|
||||
|
||||
check_u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=noise_random_generator,
|
||||
logit_mean=args.logit_mean,
|
||||
logit_std=args.logit_std,
|
||||
mode_scale=args.mode_scale,
|
||||
)
|
||||
check_indices = (check_u *
|
||||
noise_scheduler.config.num_train_timesteps).long()
|
||||
check_timesteps = noise_scheduler.timesteps[check_indices].to(
|
||||
device=check_latents.device)
|
||||
|
||||
check_sigmas = get_sigmas(
|
||||
noise_scheduler,
|
||||
check_latents.device,
|
||||
check_timesteps,
|
||||
n_dim=check_latents.ndim,
|
||||
dtype=check_latents.dtype,
|
||||
)
|
||||
check_noisy_model_input = (
|
||||
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
|
||||
|
||||
# Compute target exactly like train_one_step
|
||||
if args.precondition_outputs:
|
||||
check_target = check_latents
|
||||
else:
|
||||
check_target = check_noise - check_latents
|
||||
|
||||
# Perform gradient check with the exact same inputs as training
|
||||
max_grad_error = self.gradient_check_parameters(
|
||||
transformer=self.transformer,
|
||||
latents=
|
||||
check_noisy_model_input, # Use noisy input like in training
|
||||
encoder_hidden_states=check_encoder_hidden_states,
|
||||
encoder_attention_mask=check_encoder_attention_mask,
|
||||
timesteps=check_timesteps,
|
||||
target=check_target,
|
||||
max_params_to_check=100 # Check more parameters
|
||||
)
|
||||
|
||||
if max_grad_error > 5e-2:
|
||||
logger.error("❌ Large gradient error detected: %s",
|
||||
max_grad_error)
|
||||
else:
|
||||
logger.info("✅ Gradient check passed: max error %s",
|
||||
max_grad_error)
|
||||
|
||||
return max_grad_error
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Gradient check setup failed: %s", e)
|
||||
traceback.print_exc()
|
||||
return None
|
||||
@@ -0,0 +1,325 @@
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.fsdp import FullStateDictConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import StateDictType
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
|
||||
|
||||
|
||||
def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str,
|
||||
batch_size: int,
|
||||
generator,
|
||||
logit_mean: Optional[float] = None,
|
||||
logit_std: Optional[float] = None,
|
||||
mode_scale: Optional[float] = None,
|
||||
):
|
||||
"""
|
||||
Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(
|
||||
mean=logit_mean,
|
||||
std=logit_std,
|
||||
size=(batch_size, ),
|
||||
device="cpu",
|
||||
generator=generator,
|
||||
)
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2)**2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
|
||||
return u
|
||||
|
||||
|
||||
def get_sigmas(noise_scheduler,
|
||||
device,
|
||||
timesteps,
|
||||
n_dim=4,
|
||||
dtype=torch.float32) -> torch.Tensor:
|
||||
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
|
||||
schedule_timesteps = noise_scheduler.timesteps.to(device)
|
||||
timesteps = timesteps.to(device)
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
return sigma
|
||||
|
||||
|
||||
def save_checkpoint(transformer, rank, output_dir, step) -> None:
|
||||
# Configure FSDP to save full state dict
|
||||
FSDP.set_state_dict_type(
|
||||
transformer,
|
||||
state_dict_type=StateDictType.FULL_STATE_DICT,
|
||||
state_dict_config=FullStateDictConfig(offload_to_cpu=True,
|
||||
rank0_only=True),
|
||||
)
|
||||
|
||||
# Now get the state dict
|
||||
cpu_state = transformer.state_dict()
|
||||
|
||||
# Save it (only on rank 0 since we used rank0_only=True)
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
|
||||
torch.save(cpu_state, weight_path)
|
||||
config_dict = transformer.hf_config
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"] # TODO
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents, args=None) -> torch.Tensor:
|
||||
if model_type == "hunyuan_hf" or model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
elif model_type == "wan":
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
|
||||
vae_config = WanVAEConfig()
|
||||
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
|
||||
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
|
||||
|
||||
latents_mean = latents_mean.view(1, -1, 1, 1,
|
||||
1).to(device=latents.device)
|
||||
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
|
||||
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
|
||||
return latents
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
|
||||
|
||||
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(
|
||||
"An error occurred while clipping gradients: %s. Gradient clipping will be skipped and gradient "
|
||||
"norm will not be logged.", e)
|
||||
_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:
|
||||
raise NotImplementedError("Pipeline parallel is not supported")
|
||||
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:
|
||||
tensors = [tensors] if isinstance(tensors, torch.Tensor) else 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( # type: ignore[no-any-return]
|
||||
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,317 @@
|
||||
import sys
|
||||
import time
|
||||
from collections import deque
|
||||
from copy import deepcopy
|
||||
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
|
||||
from fastvideo.v1.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input,
|
||||
save_checkpoint)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Manual gradient checking flag - set to True to enable gradient verification
|
||||
ENABLE_GRADIENT_CHECK = False
|
||||
|
||||
|
||||
class WanTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Wan.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer"]
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.vae_config.load_encoder = False
|
||||
validation_pipeline = WanValidationPipeline.from_pretrained(
|
||||
args.model_path, args=None, inference_mode=True)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
def train_one_step(
|
||||
self,
|
||||
transformer,
|
||||
model_type,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader_iter,
|
||||
noise_scheduler,
|
||||
noise_random_generator,
|
||||
gradient_accumulation_steps,
|
||||
sp_size,
|
||||
precondition_outputs,
|
||||
max_grad_norm,
|
||||
weighting_scheme,
|
||||
logit_mean,
|
||||
logit_std,
|
||||
mode_scale,
|
||||
) -> tuple[float, float]:
|
||||
assert self.training_args is not None
|
||||
self.modules["transformer"].requires_grad_(True)
|
||||
self.modules["transformer"].train()
|
||||
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
(
|
||||
latents,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
infos,
|
||||
) = next(loader_iter)
|
||||
latents = latents.to(self.training_args.device,
|
||||
dtype=torch.bfloat16)
|
||||
encoder_hidden_states = encoder_hidden_states.to(
|
||||
self.training_args.device, dtype=torch.bfloat16)
|
||||
latents = normalize_dit_input(model_type, latents)
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn_like(latents)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=noise_random_generator,
|
||||
logit_mean=logit_mean,
|
||||
logit_std=logit_std,
|
||||
mode_scale=mode_scale,
|
||||
)
|
||||
indices = (u * noise_scheduler.config.num_train_timesteps).long()
|
||||
timesteps = noise_scheduler.timesteps[indices].to(
|
||||
device=latents.device)
|
||||
if sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
sp_group = get_sp_group()
|
||||
sp_group.broadcast(timesteps, src=0)
|
||||
sigmas = get_sigmas(
|
||||
noise_scheduler,
|
||||
latents.device,
|
||||
timesteps,
|
||||
n_dim=latents.ndim,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
input_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if 'hunyuan' in model_type:
|
||||
input_kwargs["guidance"] = torch.tensor(
|
||||
[1000.0],
|
||||
device=noisy_model_input.device,
|
||||
dtype=torch.bfloat16)
|
||||
with set_forward_context(current_timestep=timesteps,
|
||||
attn_metadata=None):
|
||||
model_pred = transformer(**input_kwargs)[0]
|
||||
|
||||
if precondition_outputs:
|
||||
model_pred = noisy_model_input - model_pred * sigmas
|
||||
target = latents if precondition_outputs else noise - latents
|
||||
|
||||
loss = (torch.mean((model_pred.float() - target.float())**2) /
|
||||
gradient_accumulation_steps)
|
||||
|
||||
loss.backward()
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
sp_group = get_sp_group()
|
||||
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
# TODO(will): perhaps move this into transformer api so that we can do
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
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,
|
||||
)
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
return total_loss, grad_norm
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
):
|
||||
assert self.training_args is not None
|
||||
noise_random_generator = None
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
# Train!
|
||||
assert self.training_args.sp_size is not None
|
||||
assert self.training_args.gradient_accumulation_steps is not None
|
||||
total_batch_size = (self.world_size *
|
||||
self.training_args.gradient_accumulation_steps /
|
||||
self.training_args.sp_size *
|
||||
self.training_args.train_sp_batch_size)
|
||||
logger.info("***** Running training *****")
|
||||
# logger.info(f" Num examples = {len(train_dataset)}")
|
||||
# logger.info(f" Dataloader size = {len(train_dataloader)}")
|
||||
# logger.info(f" Num Epochs = {args.num_train_epochs}")
|
||||
logger.info(" Resume training from step %s", self.init_steps)
|
||||
logger.info(" Instantaneous batch size per device = %s",
|
||||
self.training_args.train_batch_size)
|
||||
logger.info(
|
||||
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
|
||||
total_batch_size)
|
||||
logger.info(" Gradient Accumulation steps = %s",
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
logger.info(" Total optimization steps = %s",
|
||||
self.training_args.max_train_steps)
|
||||
logger.info(
|
||||
" Total training parameters per FSDP shard = %s B",
|
||||
sum(p.numel()
|
||||
for p in self.transformer.parameters() if p.requires_grad) /
|
||||
1e9)
|
||||
# print dtype
|
||||
logger.info(" Master weight dtype: %s",
|
||||
self.transformer.parameters().__next__().dtype)
|
||||
|
||||
# Potentially load in the weights and states from a previous save
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
assert NotImplementedError(
|
||||
"resume_from_checkpoint is not supported now.")
|
||||
# TODO
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, self.training_args.max_train_steps),
|
||||
initial=self.init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
|
||||
loader_iter = iter(self.train_dataloader)
|
||||
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
# TODO(will): fix this
|
||||
# for i in range(self.init_steps):
|
||||
# next(loader_iter)
|
||||
# get gpu memory usage
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage before train_one_step: %s MB",
|
||||
gpu_memory_usage)
|
||||
|
||||
for step in range(self.init_steps + 1, args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
|
||||
loss, grad_norm = self.train_one_step(
|
||||
self.transformer,
|
||||
# args.model_type,
|
||||
"wan",
|
||||
self.optimizer,
|
||||
self.lr_scheduler,
|
||||
loader_iter,
|
||||
noise_scheduler,
|
||||
noise_random_generator,
|
||||
self.training_args.gradient_accumulation_steps,
|
||||
self.training_args.sp_size,
|
||||
self.training_args.precondition_outputs,
|
||||
self.training_args.max_grad_norm,
|
||||
self.training_args.weighting_scheme,
|
||||
self.training_args.logit_mean,
|
||||
self.training_args.logit_std,
|
||||
self.training_args.mode_scale,
|
||||
)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage after train_one_step: %s MB",
|
||||
gpu_memory_usage)
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
# Manual gradient checking - only at first step
|
||||
if step == 1 and ENABLE_GRADIENT_CHECK:
|
||||
logger.info("Performing gradient check at step %s", step)
|
||||
self.setup_gradient_check(args, loader_iter, noise_scheduler,
|
||||
noise_random_generator)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if self.rank <= 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
"learning_rate": self.lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
},
|
||||
step=step,
|
||||
)
|
||||
if step % self.training_args.checkpointing_steps == 0:
|
||||
# Your existing checkpoint saving code
|
||||
save_checkpoint(self.transformer, self.rank,
|
||||
self.training_args.output_dir, step)
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
self.log_validation(self.transformer, self.training_args, step)
|
||||
|
||||
save_checkpoint(self.transformer, self.rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = WanTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.forward(None, args)
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.v1.fastvideo_args import TrainingArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.use_cpu_offload = False
|
||||
main(args)
|
||||
+8
-2
@@ -19,7 +19,7 @@ dependencies = [
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
|
||||
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.0", "bitsandbytes",
|
||||
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
|
||||
"torch==2.6.0", "torchvision",
|
||||
|
||||
# Acceleration & Optimization
|
||||
@@ -47,6 +47,12 @@ dependencies = [
|
||||
|
||||
# flash-attn: pip install flash-attn==2.7.4.post1 --no-cache-dir --no-build-isolation
|
||||
|
||||
train = [
|
||||
"torchdata",
|
||||
"pyarrow",
|
||||
"datasets",
|
||||
]
|
||||
|
||||
lint = [
|
||||
"pre-commit==4.0.1",
|
||||
]
|
||||
@@ -57,7 +63,7 @@ test = [
|
||||
"pytest",
|
||||
]
|
||||
|
||||
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
|
||||
dev = [ "fastvideo[lint]", "fastvideo[test]", "fastvideo[train]", ]
|
||||
|
||||
[project.scripts]
|
||||
fastvideo = "fastvideo.v1.entrypoints.cli.main:main"
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
DATA_DIR=data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
|
||||
VALIDATION_DIR=data/HD-Mixkit-Finetune-Wan/validation_parquet_dataset
|
||||
NUM_GPUS=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
fastvideo/v1/training/wan_training_pipeline.py\
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_prompt_dir "$VALIDATION_DIR"\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 4 \
|
||||
--sp_size $NUM_GPUS \
|
||||
--tp_size $NUM_GPUS \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 5\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=120 \
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=50 \
|
||||
--validation_steps 20\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="$DATA_DIR/outputs/wan_finetune"\
|
||||
--tracker_project_name wan_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 81 \
|
||||
--shift 3 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
@@ -0,0 +1,23 @@
|
||||
# export WANDB_MODE="offline"
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="your/path/to/Mixkit-Src/merge.txt"
|
||||
OUTPUT_DIR="your/path"
|
||||
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 \
|
||||
--max_height=480 \
|
||||
--max_width=832 \
|
||||
--num_frames=81 \
|
||||
--dataloader_num_workers 1 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_prompt_txt $VALIDATION_PATH \
|
||||
--samples_per_file 108 \
|
||||
--flush_frequency 108
|
||||
Reference in New Issue
Block a user