Compare commits

..
Author SHA1 Message Date
JerryZhou54 1757d3dba0 Pass pre-commit tests 2025-05-30 17:59:34 +00:00
JerryZhou54 a9d0c29ed9 fix distributed datasets issue 2025-05-30 17:50:10 +00:00
William Lin a335811869 [Training] [8/n] SP Training (#450) 2025-05-29 17:02:26 -07:00
William LinandZihang-He 357b0533fe [Training] [7/n] gradient clipping (#449)
Co-authored-by: Zihang-He <z6he@ucsd.edu>
2025-05-29 15:07:29 -07:00
William Lin 2ec3732758 [Training] [6/n]Mixed precision training (#448) 2025-05-29 14:34:28 -07:00
Wei Zhouand“BrianChen1129” a004408a93 [Training] [0/n] Add preprocessing pipeline (#442)
Co-authored-by: “BrianChen1129” <yongqich@umich.edu>
2025-05-29 14:30:09 -07:00
007e237e69 [Training] [5/n] Add single gpu training pipeline (#447)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
Co-authored-by: Wei Zhou <69577934+JerryZhou54@users.noreply.github.com>
Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com>
Co-authored-by: “BrianChen1129” <yongqich@umich.edu>
2025-05-29 11:49:46 -07:00
Yongqi Chen 8e18dc9f71 Update STA mask strategy downloading (#445) 2025-05-28 12:49:22 -07:00
7ab32539af [Training] [1/n] Add latent datasets (#438)
Co-authored-by: Wei Zhou <wzhou322@gatech.edu>
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
Co-authored-by: “BrianChen1129” <yongqich@umich.edu>
2025-05-28 11:01:52 -07:00
William Lin 6ef8fcb61d [Training] [4/n] add training save checkpoint (#441) 2025-05-27 17:53:53 -07:00
William Lin 016e24da63 [Training] [3/n] Add training args and dependencies (#440) 2025-05-27 17:53:39 -07:00
William Lin 85b8717545 [Training] [2/n] add bwd for all2all and all_gather (#439) 2025-05-27 14:27:54 -07:00
Wenxuan Tan 657fd745e1 misc: Trigger transformers CI for layers and attention code change (#434) 2025-05-27 11:43:23 -07:00
applesaucethebunandBrayden Zhong 12647457a7 [Misc] Small fixes to Torch code (#395)
Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
2025-05-23 14:40:24 -07:00
Kevin Lin 298f74f956 Set device for encode (#420) 2025-05-23 14:19:45 -07:00
Wenxuan Tan ee8babb298 Unify env report script in issue template (#423) 2025-05-23 14:19:18 -07:00
Wenxuan Tan 60295cc03f Use version.py (#424) 2025-05-23 14:17:33 -07:00
William Lin 1572e13b6e [Tests] don't run 3.10 and 3.11 for SSIM (#427) 2025-05-23 12:59:50 -07:00
Wenxuan Tan a157275b4c Fix version number (#422) 2025-05-22 12:37:34 -07:00
William Lin c4dbe7dac3 [bug] fix bs > 1 (#418) 2025-05-21 21:07:01 -07:00
Kevin Lin d39591108e Fulfill worker response on interrupt (#417) 2025-05-21 20:59:11 -07:00
64 changed files with 2447 additions and 1340771 deletions
+1 -1
View File
@@ -8,7 +8,7 @@ body:
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
+4 -2
View File
@@ -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
@@ -141,8 +143,8 @@ jobs:
fail-fast: false
matrix:
python-version: [
{version: "3.10", tag: "latest"},
{version: "3.11", tag: "py3.11-latest"},
# {version: "3.10", tag: "latest"},
# {version: "3.11", tag: "py3.11-latest"},
{version: "3.12", tag: "py3.12-latest"}
]
uses: ./.github/workflows/runpod-test.yml
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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
+1 -2
View File
@@ -2,7 +2,6 @@ import torch
from flex_sta_ref import get_sliding_tile_attention_mask
from st_attn import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
flex_attention = torch.compile(flex_attention, dynamic=False)
@@ -23,7 +22,7 @@ def h100_fwd_kernel_test(Q, K, V, kernel_size):
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
-46
View File
@@ -1,46 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=./data
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# --gradient_checkpointing\
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/v1/pipelines/training_pipeline.py\
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--train_batch_size=1\
--num_latent_t 1 \
--sp_size 4 \
--tp_size 4 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16"
+6
View File
@@ -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)=
+2 -1
View File
@@ -1,5 +1,6 @@
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
from fastvideo.version import __version__
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
+5 -9
View File
@@ -15,12 +15,8 @@ from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
def main(args):
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))
@@ -31,7 +27,7 @@ def main(args):
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
@@ -39,7 +35,7 @@ def main(args):
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
fastvideo_args = FastVideoArgs(model_path=args.model_path,
num_gpus=world_size,
device_str="cuda",
**pipeline_config_args,
@@ -47,7 +43,7 @@ def main(args):
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -119,4 +115,4 @@ if __name__ == "__main__":
)
args = parser.parse_args()
main(args)
main(args)
@@ -1,199 +0,0 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from accelerate.logging import get_logger
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
# from fastvideo.utils.load import load_text_encoder, load_vae
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
logger = get_logger(__name__)
class T5dataset(Dataset):
def __init__(
self,
json_path,
vae_debug,
):
self.json_path = json_path
self.vae_debug = vae_debug
with open(self.json_path, "r") as f:
train_dataset = json.load(f)
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
def __getitem__(self, idx):
caption = self.train_dataset[idx]["caption"]
filename = self.train_dataset[idx]["latent_path"].split(".")[0]
length = self.train_dataset[idx]["length"]
if self.vae_debug:
latents = torch.load(
os.path.join(args.output_dir, "latent", self.train_dataset[idx]["latent_path"]),
map_location="cpu",
)
else:
latents = []
return dict(caption=caption, latents=latents, filename=filename, length=length)
def __len__(self):
return len(self.train_dataset)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
# device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 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)
videoprocessor = VideoProcessor(vae_scale_factor=8)
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
vae_precision = "fp16"
text_encoder_precision = "fp32"
fastvideo_args = FastVideoArgs(model_path=args.model_path,
use_cpu_offload=False,
vae_precision=vae_precision,
text_encoder_precisions=(text_encoder_precision,))
fastvideo_args.device = device
fastvideo_args.device_str = f"cuda:{local_rank}"
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
fastvideo_args.text_encoder_configs = (T5Config(),)
# vae_loader = VAELoader()
# vae = vae_loader.load_vae()
text_encoder_loader = TextEncoderLoader()
tokenizer_loader = TokenizerLoader()
model_path = args.model_path
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
ENCODER_PATH = os.path.join(path, "text_encoder")
TOKENIZER_PATH = os.path.join(path, "tokenizer")
print(ENCODER_PATH)
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
train_dataset = T5dataset(latents_json_path, args.vae_debug)
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
# vae.enable_tiling()
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
# with torch.autocast("cuda", dtype=torch.float32):
print(data["caption"])
text_inputs = tokenizer(data["caption"], **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
post_process_func = t5_postprocess_text
prompt_embeds = post_process_func(outputs)
prompt_attention_mask = attention_mask
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.postprocess_video(video)
for idx, video_name in enumerate(data["filename"]):
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
video_name + ".pt")
# save latent
torch.save(prompt_embeds[idx], prompt_embed_path)
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
print(f"sample {video_name} saved")
if args.vae_debug:
export_to_video(video[idx], video_path, fps=16)
item = {}
item["length"] = int(data["length"][idx])
item["latent_path"] = video_name + ".pt"
item["prompt_embed_path"] = video_name + ".pt"
item["prompt_attention_mask"] = video_name + ".pt"
item["caption"] = data["caption"][idx]
json_data.append(item)
dist.barrier()
local_data = json_data
gathered_data = [None] * world_size
dist.all_gather_object(gathered_data, local_data)
if local_rank == 0:
# os.remove(latents_json_path)
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
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")
# text encoder & vae & diffusion model
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(
"--train_batch_size",
type=int,
default=1,
help="Batch size (per device) for the training dataloader.",
)
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(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument("--vae_debug", action="store_true")
args = parser.parse_args()
main(args)
@@ -1,151 +0,0 @@
import argparse
import json
import os
import torch
# import torch.distributed as dist
# from accelerate.logging import get_logger
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.dataset import getdataset
# from fastvideo.utils.load import load_vae
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
logger = init_logger(__name__)
model_path = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
VAE_PATH = os.path.join(path, "vae")
print(VAE_PATH)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
vae_precision = "fp16"
fastvideo_args = FastVideoArgs(model_path=VAE_PATH,
use_cpu_offload=False,
vae_precision=vae_precision)
fastvideo_args.device = device
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
# encoder_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 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)
vae_loader = VAELoader()
vae = vae_loader.load(VAE_PATH, "vae", fastvideo_args)
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
# vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=torch.float16):
latents = vae.encode(data["pixel_values"].to(device)).sample()
for idx, video_path in enumerate(data["path"]):
video_name = os.path.basename(video_path).split(".")[0]
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
torch.save(latents[idx].to(torch.bfloat16), latent_path)
item = {}
item["length"] = latents[idx].shape[1]
item["latent_path"] = video_name + ".pt"
item["caption"] = data["text"][idx]
json_data.append(item)
print(f"{video_name} processed")
world_group.barrier()
local_data = json_data
gathered_data = [None] * world_size
for i in range(world_size):
if local_rank == i:
world_group.broadcast_object(local_data, src=i)
else:
gathered_data[i] = world_group.broadcast_object(None, src=i)
gathered_data[local_rank] = json_data
print(gathered_data)
if local_rank == 0:
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
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("--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(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
@@ -1,115 +0,0 @@
import argparse
import os
import torch
# import torch.distributed as dist
from accelerate.logging import get_logger
# from fastvideo.utils.load import load_text_encoder
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
logger = get_logger(__name__)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
vae_precision = "fp16"
text_encoder_precision = "fp32"
fastvideo_args = FastVideoArgs(model_path=args.model_path,
use_cpu_offload=False,
vae_precision=vae_precision,
text_encoder_precisions=(text_encoder_precision,))
fastvideo_args.device = device
fastvideo_args.device_str = f"cuda:{local_rank}"
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
fastvideo_args.text_encoder_configs = (T5Config(),)
# vae_loader = VAELoader()
# vae = vae_loader.load_vae()
text_encoder_loader = TextEncoderLoader()
tokenizer_loader = TokenizerLoader()
model_path = args.model_path
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
ENCODER_PATH = os.path.join(path, "text_encoder")
TOKENIZER_PATH = os.path.join(path, "tokenizer")
print(ENCODER_PATH)
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
# autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
# output_dir/validation/prompt_attention_mask
# output_dir/validation/prompt_embed
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
os.makedirs(
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
exist_ok=True,
)
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True)
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
for prompt in prompts:
with torch.inference_mode():
# with torch.autocast("cuda", dtype=autocast_type):
text_inputs = tokenizer(prompt, **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
post_process_func = t5_postprocess_text
prompt_embeds = post_process_func(outputs)
prompt_attention_mask = attention_mask
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
prompt_attention_mask_path = os.path.join(
args.output_dir,
"validation",
"prompt_attention_mask",
f"{file_name}.pt",
)
torch.save(prompt_embeds[0], prompt_embed_path)
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
print(f"sample {file_name} saved")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
+97
View File
@@ -0,0 +1,97 @@
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.dataset.t2v_datasets import T2V_dataset
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
def getdataset(args):
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 = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, 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,
)
raise NotImplementedError(args.dataset)
if __name__ == "__main__":
import random
from accelerate import Accelerator
from tqdm import tqdm
from fastvideo.dataset.t2v_datasets import dataset_prog
args = type(
"args",
(),
{
"ae": "CausalVAEModel_4x8x8",
"dataset": "t2v",
"attention_mode": "xformers",
"use_rope": True,
"text_max_length": 300,
"max_height": 320,
"max_width": 240,
"num_frames": 1,
"use_image_num": 0,
"interpolation_scale_t": 1,
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
"use_img_from_vid": False,
"speed_factor": 1.0,
"cfg": 0.1,
"text_encoder_name": "google/mt5-xxl",
"dataloader_num_workers": 10,
},
)
accelerator = Accelerator()
dataset = getdataset(args)
num = len(dataset_prog.img_cap_list)
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
print(e)
# import ipdb;ipdb.set_trace()
print(image_data)
zero += 1
continue
assert caps[0] is not None and len(caps[0]) > 0
print(num, zero)
import ipdb
ipdb.set_trace()
print("end")
+118
View File
@@ -0,0 +1,118 @@
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,
):
# 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, "r") 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
# just zero embeddings [256, 4096]
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [data_item["length"] if "length" in data_item else 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
latents = [
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(latents), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latents):
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(latents, 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)
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()
+324
View File
@@ -0,0 +1,324 @@
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 = {}
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):
self.cap_list = []
self.elements = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements = dict()
self.n_used_elements = dict()
def set_cap_list(self, num_workers, cap_list, n_elements):
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):
if work_info is None:
worker_id = 0
else:
worker_id = 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, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
class T2V_dataset(Dataset):
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
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):
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):
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,
)
def get_image(self, idx):
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 = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
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"] # 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):
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):
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):
cap_lists = []
with open(data, "r") 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, "r") 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):
cap_lists = self.read_jsons(self.data)
return cap_lists
+608
View File
@@ -0,0 +1,608 @@
import numbers
import random
import torch
from PIL import Image
def _is_tensor_video_clip(clip):
if not torch.is_tensor(clip):
raise TypeError("clip should be Tensor. Got %s" % type(clip))
if not clip.ndimension() == 4:
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
return True
def center_crop_arr(pil_image, image_size):
"""
Center cropping implementation from ADM.
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
"""
while min(*pil_image.size) >= 2 * image_size:
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
scale = image_size / min(*pil_image.size)
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
arr = np.array(pil_image)
crop_y = (arr.shape[0] - image_size) // 2
crop_x = (arr.shape[1] - image_size) // 2
return Image.fromarray(arr[crop_y:crop_y + image_size, crop_x:crop_x + image_size])
def crop(clip, i, j, h, w):
"""
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):
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 resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
clip,
scale_factor=scale_,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
"""
Do spatial cropping and resizing to the video clip
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
i (int): i in (i,j) i.e coordinates of the upper left corner.
j (int): j in (i,j) i.e coordinates of the upper left corner.
h (int): Height of the cropped region.
w (int): Width of the cropped region.
size (tuple(int, int)): height and width of resized clip
Returns:
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
clip = crop(clip, i, j, h, w)
clip = resize(clip, size, interpolation_mode)
return clip
def center_crop(clip, crop_size):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
th, tw = crop_size
if h < th or w < tw:
raise ValueError("height and width must be no smaller than crop_size")
i = int(round((h - th) / 2.0))
j = int(round((w - tw) / 2.0))
return crop(clip, i, j, th, tw)
def center_crop_using_short_edge(clip):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h < w:
th, tw = h, h
i = 0
j = int(round((w - tw) / 2.0))
else:
th, tw = w, w
i = int(round((h - th) / 2.0))
j = 0
return crop(clip, i, j, th, tw)
def center_crop_th_tw(clip, th, tw, top_crop):
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 random_shift_crop(clip):
"""
Slide along the long edge, with the short edge as crop size
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h <= w:
short_edge = h
else:
short_edge = w
th, tw = short_edge, short_edge
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return crop(clip, i, j, th, tw)
def normalize_video(clip):
"""
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("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
def normalize(clip, mean, std, inplace=False):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
mean (tuple): pixel RGB mean. Size is (3)
std (tuple): pixel standard deviation. Size is (3)
Returns:
normalized clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
if not inplace:
clip = clip.clone()
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
# print(mean)
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
return clip
def hflip(clip):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
Returns:
flipped clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
return clip.flip(-1)
class RandomCropVideo:
def __init__(self, size):
if isinstance(size, numbers.Number):
self.size = (int(size), int(size))
else:
self.size = size
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: randomly cropped video clip.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = self.size
if h < th or w < tw:
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
if w == tw and h == th:
return 0, 0, h, w
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return i, j, th, tw
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: cropped video clip by stride.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
return 0, 0, th, tw # from top-left
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class LongSideResizeVideo:
"""
First use the long side,
then resize to the specified size
"""
def __init__(
self,
size,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized video clip.
size is (T, C, 512, *) or (T, C, *, 512)
"""
_, _, h, w = clip.shape
if self.skip_low_resolution and max(h, w) <= self.size:
return clip
if h > w:
w = int(w * self.size / h)
h = self.size
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
return resize_clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
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",
):
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):
"""
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_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
# import ipdb;ipdb.set_trace()
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 UCFCenterCropVideo:
"""
First scale to the specified size in equal proportion to the short edge,
then center cropping
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
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_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class KineticsRandomCropResizeVideo:
"""
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
return clip_resize
class CenterCropVideo:
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop(clip, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize:
"""
Normalize the video clip by mean subtraction and division by standard deviation
Args:
mean (3-tuple): pixel RGB mean
std (3-tuple): pixel RGB standard deviation
inplace (boolean): whether do in-place normalization
"""
def __init__(self, mean, std, inplace=False):
self.mean = mean
self.std = std
self.inplace = inplace
def __call__(self, clip):
"""
Args:
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
"""
return normalize(clip, self.mean, self.std, self.inplace)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self):
pass
def __call__(self, clip):
"""
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 RandomHorizontalFlipVideo:
"""
Flip the video clip along the horizontal direction with a given probability
Args:
p (float): probability of the clip being flipped. Default value is 0.5
"""
def __init__(self, p=0.5):
self.p = p
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Size is (T, C, H, W)
Return:
clip (torch.tensor): Size is (T, C, H, W)
"""
if random.random() < self.p:
clip = hflip(clip)
return clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(p={self.p})"
# ------------------------------------------------------------
# --------------------- Sampling ---------------------------
# ------------------------------------------------------------
class TemporalRandomCrop(object):
"""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):
self.size = size
def __call__(self, total_frames):
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
class DynamicSampleDuration(object):
"""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, t_stride, extra_1):
self.t_stride = t_stride
self.extra_1 = extra_1
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
return 0, truncate_t
if __name__ == "__main__":
import os
import numpy as np
import torchvision.io as io
from torchvision import transforms
from torchvision.utils import save_image
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW")
trans = transforms.Compose([
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
target_video_len = 32
frame_interval = 1
total_frames = len(vframes)
print(total_frames)
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
# Sampling video frames
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
print(select_vframes.shape)
print(select_vframes.dtype)
select_vframes_trans = trans(select_vframes)
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
for i in range(target_video_len):
save_image(
select_vframes_trans[i],
os.path.join("./test000", "%04d.png" % i),
normalize=True,
value_range=(-1, 1),
)
+1 -2
View File
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -38,7 +38,6 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
linear_range=0.5,
):
if linear_quadratic:
raise NotImplementedError("Linear quadratic schedule is not implemented")
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
-870
View File
@@ -1,870 +0,0 @@
# !/bin/python3
# isort: skip_file
import argparse
import math
import os
import time
from collections import deque
import torch
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.utils.checkpoint import (save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
get_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.models.loader.component_loader import TransformerLoader, SchedulerLoader
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
SCHEDULER_PATH = os.path.join(MODEL_PATH, "scheduler")
def reshard_fsdp(model):
for m in FSDP.fsdp_modules(model):
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
gradient_accumulation_steps)
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
gradient_accumulation_steps)
absolute_mean = torch.mean(
torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(
torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
def distill_one_step(
transformer,
model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_teacher_disable_cfg,
):
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0, # codespell:ignore
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
}
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
# model_input = normalize_dit_input(model_type, latents)
model_input = latents
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
timesteps = (sigmas *
noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
noisy_model_input = noisy_model_input.to(torch.bfloat16)
forward_batch = ForwardBatch(data_type="video", enable_teacache=False)
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
teacher_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_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
with torch.autograd.graph.save_on_cpu(pin_memory=True):
model_pred = transformer(**teacher_kwargs)
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
).float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(
bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
).float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output,
index).to(torch.bfloat16)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
with torch.autograd.graph.save_on_cpu(pin_memory=True):
target_pred = transformer(
x_prev,
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
huber_c = 0.001
# loss = loss.mean()
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float()**2)) *
pred_decay_weight / gradient_accumulation_steps)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (torch.mean(model_pred.float()**2) *
pred_decay_weight /
gradient_accumulation_steps)
loss += pred_decay_loss
else:
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(model_pred.detach().float(), model_pred_norm,
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(ema_transformer.parameters(),
transformer.parameters()):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(),
1 - ema_decay))
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
grad_norm = torch.nn.utils.clip_grad_norm_(transformer.parameters(),
max_norm=max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item(), model_pred_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
torch.cuda.set_device(rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=args.sp_size,
sequence_model_parallel_size=args.sp_size)
fastvideo_args = FastVideoArgs(
model_path=MODEL_PATH,
num_gpus=world_size,
use_cpu_offload=False,
precision=args.master_weight_type,
dit_config=WanVideoConfig(),
device_str="cuda",
)
fastvideo_args.check_fastvideo_args()
device_str = f"cuda:{rank}"
device = torch.device(device_str)
fastvideo_args.device = device
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
logger.info("--> loading model from %s", TRANSFORMER_PATH)
fastvideo_args.device = device
transformer_loader = TransformerLoader()
transformer = transformer_loader.load(TRANSFORMER_PATH, "", fastvideo_args)
transformer = transformer.train()
transformer.requires_grad_(True)
teacher_loader = TransformerLoader()
teacher_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
fastvideo_args)
if args.use_ema:
ema_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
fastvideo_args)
else:
ema_transformer = None
logger.info(
" Total training parameters = %s M",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e6)
logger.info("--> model loaded")
teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
# scheduler
noise_scheduler_loader = SchedulerLoader()
noise_scheduler = noise_scheduler_loader.load(SCHEDULER_PATH, "",
fastvideo_args)
solver = EulerSolver(
noise_scheduler.sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(train_dataset))
logger.info(" Dataloader size = %s", len(train_dataloader))
logger.info(" Num Epochs = %s", args.num_train_epochs)
logger.info(" Resume training from step %s", init_steps)
logger.info(" Instantaneous batch size per device = %s",
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",
args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s", args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e9)
# print dtype
logger.info(" Master weight dtype: %s",
transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
loss, grad_norm, pred_norm = distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.95)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=str, default="64")
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
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(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
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***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="Threshold for linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--use_ema",
action="store_true",
help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+2 -2
View File
@@ -237,7 +237,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=str,
default="540p",
choices=["540p", "720p"],
help="Root path of all the models, including t2v models and extra models.",
help="The resolution of the model.",
)
group.add_argument(
"--load-key",
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
"--ring-degree",
type=int,
default=1,
help="Ulysses degree.",
help="Ring degree.",
)
return parser
+1 -1
View File
@@ -17,7 +17,7 @@ from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.utils.parallel_states import nccl_info
class Inference(object):
class Inference:
def __init__(
self,
+1 -1
View File
@@ -41,7 +41,7 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
elif mode == "Master":
prompt = master_mode_prompt.format(input=ori_prompt)
else:
raise Exception("Only supports Normal and Normal", mode)
raise Exception("Only supports Normal and Master mode, but got {}".format(mode))
return prompt
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
mochi_scaling_factor = 1.0
def normalize_dit_input(model_type, latents, args=None):
def normalize_dit_input(model_type, latents):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
@@ -41,16 +41,5 @@ def normalize_dit_input(model_type, latents, args=None):
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
vae_config = WanVAEConfig()
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
class STEP1TextEncoder(torch.nn.Module):
def __init__(self, model_dir, max_length=320):
super(STEP1TextEncoder, self).__init__()
super()
self.max_length = max_length
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
text_encoder = Step1Model.from_pretrained(model_dir)
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
@torch.no_grad
@torch.autocast(device_type='cuda', dtype=torch.bfloat16)
def forward(self, prompts, with_mask=True, max_length=None):
self.device = next(self.text_encoder.parameters()).device
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
if type(prompts) is str:
prompts = [prompts]
if type(prompts) is str:
prompts = [prompts]
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
+1 -39
View File
@@ -11,7 +11,6 @@ from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_
from torch.distributed.fsdp import FullOptimStateDictConfig, FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
import dataclasses
from fastvideo.utils.logging_ import main_print
@@ -45,50 +44,13 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
optimizer_path = os.path.join(save_dir, "optimizer.pt")
torch.save(optim_state, optimizer_path)
else:
weight_path = os.path.join(save_dstate_dictir, "discriminator_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
torch.save(optim_state, optimizer_path)
main_print(f"--> checkpoint saved at step {step}")
def save_checkpoint_v1(transformer, rank, output_dir, step):
# from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
# from torch.distributed.fsdp import StateDictType, FullStateDictConfig
# 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 torch.distributed.get_rank() == 0:
# torch.save(state_dict, "model_checkpoint.pt")
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
# weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
print(weight_path)
# save_file(cpu_state, weight_path)
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)
main_print(f"--> checkpoint saved at step {step}")
def save_checkpoint(transformer, rank, output_dir, step):
main_print(f"--> saving checkpoint at step {step}")
with FSDP.state_dict_type(
@@ -70,6 +70,8 @@ DEFAULT_CONDA_PATTERNS = {
"optree",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
@@ -85,6 +87,8 @@ DEFAULT_PIP_PATTERNS = {
"onnx",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
-38
View File
@@ -1,38 +0,0 @@
import platform
import accelerate
import peft
import torch
import transformers
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
VERSION = "1.2.0"
if __name__ == "__main__":
info = {
"FastVideo version": VERSION,
"Platform": platform.platform(),
"Python version": platform.python_version(),
"PyTorch version": torch.__version__,
"Transformers version": transformers.__version__,
"Accelerate version": accelerate.__version__,
"PEFT version": peft.__version__,
}
if is_torch_cuda_available():
info["PyTorch version"] += " (GPU)"
info["GPU type"] = torch.cuda.get_device_name()
if is_torch_npu_available():
info["PyTorch version"] += " (NPU)"
info["NPU type"] = torch.npu.get_device_name()
info["CANN version"] = torch.version.cann # codespell:ignore
try:
import bitsandbytes
info["Bitsandbytes version"] = bitsandbytes.__version__
except Exception:
pass
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
+1 -1
View File
@@ -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
+5 -66
View File
@@ -1,3 +1,5 @@
import os
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
@@ -7,7 +9,7 @@ from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
def getdataset(args, start_idx=0):
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 = [
@@ -25,8 +27,8 @@ def getdataset(args, start_idx=0):
*resize_topcrop,
norm_fun,
])
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name,
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,
@@ -37,66 +39,3 @@ def getdataset(args, start_idx=0):
start_idx=start_idx)
raise NotImplementedError(args.dataset)
if __name__ == "__main__":
import random
from accelerate import Accelerator
from tqdm import tqdm
from fastvideo.v1.dataset.t2v_datasets import dataset_prog
args = type(
"args",
(),
{
"ae": "CausalVAEModel_4x8x8",
"dataset": "t2v",
"attention_mode": "xformers",
"use_rope": True,
"text_max_length": 300,
"max_height": 320,
"max_width": 240,
"num_frames": 1,
"use_image_num": 0,
"interpolation_scale_t": 1,
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data":
"/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
"use_img_from_vid": False,
"speed_factor": 1.0,
"cfg": 0.1,
"text_encoder_name": "google/mt5-xxl",
"dataloader_num_workers": 10,
},
)
accelerator = Accelerator()
dataset = getdataset(args)
num = len(dataset_prog.img_cap_list)
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [
i["cap"] if isinstance(i["cap"], list) else [i["cap"]]
for i in image_data
]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
print(e)
# import ipdb;ipdb.set_trace()
print(image_data)
zero += 1
continue
assert caps[0] is not None and len(caps[0]) > 0
print(num, zero)
import ipdb
ipdb.set_trace()
print("end")
+10 -10
View File
@@ -13,7 +13,7 @@ class LatentDataset(Dataset):
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
@@ -29,13 +29,12 @@ class LatentDataset(Dataset):
# 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
# just zero embeddings [256, 4096]
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item["length"] if "length" in data_item else 1
for data_item in self.data_anno
data_item.get("length", 1) for data_item in self.data_anno
]
def __getitem__(self, idx):
@@ -83,7 +82,7 @@ def latent_collate_function(batch):
max_w = max([latent.shape[3] for latent in latents])
# padding
latents = [
latent_list: list[torch.Tensor] = [
torch.nn.functional.pad(
latent,
(
@@ -97,22 +96,23 @@ def latent_collate_function(batch):
) for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
latent_attn_mask = torch.ones(len(latent_list), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latents):
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(latents, 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)
num_latent_t=28,
cfg_rate=0.0)
dataloader = torch.utils.data.DataLoader(dataset,
batch_size=2,
shuffle=False,
+166 -212
View File
@@ -1,21 +1,28 @@
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 IterableDataset, get_worker_info
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
# Path to your dataset
dataset_path = "/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/train/"
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(IterableDataset):
class ParquetVideoTextDataset(Dataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
def __init__(self,
@@ -24,237 +31,187 @@ class ParquetVideoTextDataset(IterableDataset):
rank: int = 0,
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2):
num_latent_t: int = 2,
seed: int = 0):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.world_size = world_size
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")
# Find all parquet files recursively
print(f"Scanning for parquet files in {self.path}")
self.parquet_files = []
for root, _, files in os.walk(self.path):
for file in files:
if file.endswith('.parquet'):
self.parquet_files.append(os.path.join(root, file))
# Sort files for consistent ordering
self.parquet_files.sort()
ranks = get_sp_group().ranks
group_ranks: List[List] = [[] for _ in range(self.world_size)]
torch.distributed.all_gather_object(group_ranks, ranks)
# Distribute files among workers
# drop last unenven files
print(f"Total files: {len(self.parquet_files)}")
total_files = len(self.parquet_files)
base_count = total_files // world_size
extra_files = total_files % world_size
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
if rank < extra_files:
start_idx = rank * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = rank * base_count + extra_files
end_idx = start_idx + base_count
# 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))
self.parquet_files = self.parquet_files[start_idx:end_idx]
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
print(f"Files assigned to rank {rank}: {len(self.parquet_files)}")
if len(self.parquet_files) > 0:
print(f"First file: {self.parquet_files[0]}")
print(f"Last file: {self.parquet_files[-1]}")
# 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)
# Initialize current file index
self.current_file_idx = 0
self.current_reader = None
self.current_batches = None
self.total_samples = 0
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
dist.barrier()
def _open_next_file(self):
"""Open the next parquet file for reading."""
num_workers = get_worker_info().num_workers
worker_id = get_worker_info().id
total_files = len(self.parquet_files)
base_count = total_files // num_workers
extra_files = total_files % num_workers
if worker_id < extra_files:
start_idx = worker_id * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = worker_id * base_count + extra_files
end_idx = start_idx + base_count
worker_parquet_files = self.parquet_files[start_idx:end_idx]
if self.current_file_idx >= len(worker_parquet_files):
print(
f"Rank {self.rank}, Worker {worker_id}: No more files to open (current_idx={self.current_file_idx}, total_files={len(worker_parquet_files)})"
)
return False
if self.current_reader is not None:
self.current_reader.close()
file_path = worker_parquet_files[self.current_file_idx]
print(
f"Rank {self.rank}, Worker {worker_id}: Opening file {self.current_file_idx + 1}/{len(worker_parquet_files)}: {file_path}"
)
try:
self.current_reader = pq.ParquetFile(file_path)
self.current_batches = self.current_reader.iter_batches(
batch_size=self.batch_size)
self.current_file_idx += 1
return True
except Exception as e:
print(f"Error opening file {file_path}: {str(e)}")
return False
def __iter__(self):
"""Iterate over the dataset in a streaming fashion."""
print(f"Rank {self.rank}: Starting iteration")
# First try to open a file
if not self._open_next_file():
print(f"Rank {self.rank}: Failed to open first file")
return
while True:
def __len__(self):
if self.local_indices is None:
try:
# Get next batch from current file
batch = next(self.current_batches)
batch_dict = batch.to_pydict()
processed = self._process_batch(batch_dict)
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)
# Update sample count
batch_size = len(processed["latents"])
self.total_samples += batch_size
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)
# Print progress
if self.total_samples % 1000 == 0:
print(
f"Rank {self.rank}: Processed {self.total_samples} samples"
)
# 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
# Yield each item in the batch
for lat, emb, mask, info in zip(processed["latents"],
processed["embeddings"],
processed["masks"],
processed["info"]):
if lat.numel() == 0: # Split is validation
yield lat, emb, mask, info
else:
yield lat[:, -self.num_latent_t:], emb, mask, info
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
except StopIteration:
# Current file is exhausted, try next file
print(
f"Rank {self.rank}: Current file exhausted, trying next file"
)
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: No more files to process. Total samples: {self.total_samples}"
)
break
except Exception as e:
print(f"Error processing batch: {str(e)}")
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: Failed to open next file after error"
)
break
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
# Clean up
if self.current_reader is not None:
self.current_reader.close()
def _process_batch(self, batch):
def _process_row(self, row) -> Dict[str, Any]:
"""Process a PyArrow batch into tensors."""
out = {"lat": [], "emb": [], "msk": [], "info": []}
for i in range(len(batch["vae_latent_bytes"])):
vae_latent_bytes = batch["vae_latent_bytes"][i]
vae_latent_shape = batch["vae_latent_shape"][i]
text_embedding_bytes = batch["text_embedding_bytes"][i]
text_embedding_shape = batch["text_embedding_shape"][i]
text_attention_mask_bytes = batch["text_attention_mask_bytes"][i]
text_attention_mask_shape = batch["text_attention_mask_shape"][i]
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)
# 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, :]
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_)
# to string
file_name = str(batch["file_name"][i])
# Collect metadata
info = {
"width": batch["width"][i],
"height": batch["height"][i],
"num_frames": batch["num_frames"][i],
"duration_sec": batch["duration_sec"][i],
"fps": batch["fps"][i],
"file_name": batch["file_name"][i],
"caption": batch["caption"][i],
}
# 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_)
out["lat"].append(torch.from_numpy(lat))
out["emb"].append(torch.from_numpy(emb))
out["msk"].append(torch.from_numpy(msk))
out["info"].append(info)
return {
"latents": torch.stack(out["lat"]) if out["lat"] else None,
"embeddings": torch.stack(out["emb"]) if out["emb"] else None,
"masks": torch.stack(out["msk"]) if out["msk"] else None,
"info": out["info"]
# 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"],
}
def bind_cpu_cores(local_rank, cpu_per_process=16):
"""根据local_rank绑定固定cpu核。"""
start = local_rank * cpu_per_process
end = start + cpu_per_process
cores = list(range(start, end))
print(f"[Rank {local_rank}] Binding to CPU cores: {cores}")
os.sched_setaffinity(0, cores)
return {
"latents": torch.from_numpy(lat),
"embeddings": torch.from_numpy(emb),
"masks": torch.from_numpy(msk),
"info": info
}
if __name__ == "__main__":
@@ -262,7 +219,7 @@ if __name__ == "__main__":
description='Benchmark Parquet dataset loading speed')
parser.add_argument('--path',
type=str,
default=dataset_path,
default="your/dataset/path",
help='Path to Parquet dataset')
parser.add_argument('--batch_size',
type=int,
@@ -297,9 +254,6 @@ if __name__ == "__main__":
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
)
# Bind CPU cores after distributed initialization
# bind_cpu_cores(local_rank, cpu_per_process=16)
# Create dataset
dataset = ParquetVideoTextDataset(
args.path,
+26 -31
View File
@@ -17,7 +17,7 @@ from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances = {}
_instances: dict[type, 'SingletonMeta'] = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
@@ -28,15 +28,15 @@ class SingletonMeta(type):
class DataSetProg(metaclass=SingletonMeta):
def __init__(self):
self.cap_list = []
self.elements = []
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()
self.n_used_elements = dict()
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):
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
@@ -52,11 +52,8 @@ class DataSetProg(metaclass=SingletonMeta):
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info):
if work_info is None:
worker_id = 0
else:
worker_id = work_info.id
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] %
@@ -68,13 +65,11 @@ class DataSetProg(metaclass=SingletonMeta):
dataset_prog = DataSetProg()
def filter_resolution(h,
w,
max_h_div_w_ratio=17 / 16,
min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
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):
@@ -85,7 +80,7 @@ class T2V_dataset(Dataset):
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0):
start_idx=0) -> None:
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
@@ -132,14 +127,14 @@ class T2V_dataset(Dataset):
data = self.get_data(idx)
return data
def get_data(self, idx):
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):
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"]
@@ -183,7 +178,7 @@ class T2V_dataset(Dataset):
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx):
def get_image(self, idx) -> dict:
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
@@ -201,14 +196,14 @@ class T2V_dataset(Dataset):
image = image.float() / 127.5 - 1.0
caps = (image_data["cap"]
if isinstance(image_data["cap"], list) else [image_data["cap"]])
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 = [], []
text = text[0] if random.random() > self.cfg else ""
single_text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
single_text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
@@ -226,7 +221,7 @@ class T2V_dataset(Dataset):
path=image_data["path"],
)
def define_frame_index(self, cap_list):
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
@@ -326,14 +321,14 @@ class T2V_dataset(Dataset):
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices):
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):
def read_jsons(self, data) -> list[dict]:
cap_lists = []
with open(data) as f:
folder_anno = [
@@ -349,6 +344,6 @@ class T2V_dataset(Dataset):
cap_lists += sub_list
return cap_lists
def get_cap_list(self):
def get_cap_list(self) -> list:
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
+15 -504
View File
@@ -1,42 +1,19 @@
import numbers
import random
import torch
from PIL import Image
def _is_tensor_video_clip(clip):
def _is_tensor_video_clip(clip) -> bool:
if not torch.is_tensor(clip):
raise TypeError("clip should be Tensor. Got %s" % type(clip))
raise TypeError(f"clip should be Tensor. Got {type(clip)}")
if not clip.ndimension() == 4:
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
raise ValueError(f"clip should be 4D. Got {clip.dim()}D")
return True
def center_crop_arr(pil_image, image_size):
"""
Center cropping implementation from ADM.
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
"""
while min(*pil_image.size) >= 2 * image_size:
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size),
resample=Image.BOX)
scale = image_size / min(*pil_image.size)
pil_image = pil_image.resize(tuple(
round(x * scale) for x in pil_image.size),
resample=Image.BICUBIC)
arr = np.array(pil_image)
crop_y = (arr.shape[0] - image_size) // 2
crop_x = (arr.shape[1] - image_size) // 2
return Image.fromarray(arr[crop_y:crop_y + image_size,
crop_x:crop_x + image_size])
def crop(clip, i, j, h, w):
def crop(clip, i, j, h, w) -> torch.Tensor:
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
@@ -46,7 +23,7 @@ def crop(clip, i, j, h, w):
return clip[..., i:i + h, j:j + w]
def resize(clip, target_size, interpolation_mode):
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}"
@@ -60,71 +37,7 @@ def resize(clip, target_size, interpolation_mode):
)
def resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
clip,
scale_factor=scale_,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
"""
Do spatial cropping and resizing to the video clip
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
i (int): i in (i,j) i.e coordinates of the upper left corner.
j (int): j in (i,j) i.e coordinates of the upper left corner.
h (int): Height of the cropped region.
w (int): Width of the cropped region.
size (tuple(int, int)): height and width of resized clip
Returns:
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
clip = crop(clip, i, j, h, w)
clip = resize(clip, size, interpolation_mode)
return clip
def center_crop(clip, crop_size):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
th, tw = crop_size
if h < th or w < tw:
raise ValueError("height and width must be no smaller than crop_size")
i = int(round((h - th) / 2.0))
j = int(round((w - tw) / 2.0))
return crop(clip, i, j, th, tw)
def center_crop_using_short_edge(clip):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h < w:
th, tw = h, h
i = 0
j = int(round((w - tw) / 2.0))
else:
th, tw = w, w
i = int(round((h - th) / 2.0))
j = 0
return crop(clip, i, j, th, tw)
def center_crop_th_tw(clip, th, tw, top_crop):
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")
@@ -143,27 +56,7 @@ def center_crop_th_tw(clip, th, tw, top_crop):
return crop(clip, i, j, new_h, new_w)
def random_shift_crop(clip):
"""
Slide along the long edge, with the short edge as crop size
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h <= w:
short_edge = h
else:
short_edge = w
th, tw = short_edge, short_edge
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return crop(clip, i, j, th, tw)
def normalize_video(clip):
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
@@ -174,153 +67,12 @@ def normalize_video(clip):
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError("clip tensor should have data type uint8. Got %s" %
str(clip.dtype))
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
def normalize(clip, mean, std, inplace=False):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
mean (tuple): pixel RGB mean. Size is (3)
std (tuple): pixel standard deviation. Size is (3)
Returns:
normalized clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
if not inplace:
clip = clip.clone()
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
# print(mean)
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
return clip
def hflip(clip):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
Returns:
flipped clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
return clip.flip(-1)
class RandomCropVideo:
def __init__(self, size):
if isinstance(size, numbers.Number):
self.size = (int(size), int(size))
else:
self.size = size
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: randomly cropped video clip.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = self.size
if h < th or w < tw:
raise ValueError(
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
)
if w == tw and h == th:
return 0, 0, h, w
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return i, j, th, tw
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: cropped video clip by stride.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
return 0, 0, th, tw # from top-left
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class LongSideResizeVideo:
"""
First use the long side,
then resize to the specified size
"""
def __init__(
self,
size,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized video clip.
size is (T, C, 512, *) or (T, C, *, 512)
"""
_, _, h, w = clip.shape
if self.skip_low_resolution and max(h, w) <= self.size:
return clip
if h > w:
w = int(w * self.size / h)
h = self.size
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(clip,
target_size=(h, w),
interpolation_mode=self.interpolation_mode)
return resize_clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class CenterCropResizeVideo:
"""
First use the short side for cropping length,
@@ -332,7 +84,7 @@ class CenterCropResizeVideo:
size,
top_crop=False,
interpolation_mode="bilinear",
):
) -> None:
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
@@ -340,7 +92,7 @@ class CenterCropResizeVideo:
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
def __call__(self, clip) -> torch.Tensor:
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
@@ -348,12 +100,10 @@ class CenterCropResizeVideo:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
# clip_center_crop = center_crop_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(clip,
self.size[0],
self.size[1],
top_crop=self.top_crop)
# import ipdb;ipdb.set_trace()
clip_center_crop_resize = resize(
clip_center_crop,
target_size=self.size,
@@ -365,138 +115,15 @@ class CenterCropResizeVideo:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class UCFCenterCropVideo:
"""
First scale to the specified size in equal proportion to the short edge,
then center cropping
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
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_resize = resize_scale(clip=clip,
target_size=self.size,
interpolation_mode=self.interpolation_mode)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class KineticsRandomCropResizeVideo:
"""
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size,
self.interpolation_mode)
return clip_resize
class CenterCropVideo:
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop(clip, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize:
"""
Normalize the video clip by mean subtraction and division by standard deviation
Args:
mean (3-tuple): pixel RGB mean
std (3-tuple): pixel RGB standard deviation
inplace (boolean): whether do in-place normalization
"""
def __init__(self, mean, std, inplace=False):
self.mean = mean
self.std = std
self.inplace = inplace
def __call__(self, clip):
"""
Args:
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
"""
return normalize(clip, self.mean, self.std, self.inplace)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self):
def __init__(self) -> None:
pass
def __call__(self, clip):
def __call__(self, clip) -> torch.Tensor:
"""
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
@@ -509,34 +136,6 @@ class Normalize255:
return self.__class__.__name__
class RandomHorizontalFlipVideo:
"""
Flip the video clip along the horizontal direction with a given probability
Args:
p (float): probability of the clip being flipped. Default value is 0.5
"""
def __init__(self, p=0.5):
self.p = p
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Size is (T, C, H, W)
Return:
clip (torch.tensor): Size is (T, C, H, W)
"""
if random.random() < self.p:
clip = hflip(clip)
return clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(p={self.p})"
# ------------------------------------------------------------
# --------------------- Sampling ---------------------------
# ------------------------------------------------------------
class TemporalRandomCrop:
"""Temporally crop the given frame indices at a random location.
@@ -544,99 +143,11 @@ class TemporalRandomCrop:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, size):
def __init__(self, size) -> None:
self.size = size
def __call__(self, total_frames):
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
class DynamicSampleDuration:
"""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, t_stride, extra_1):
self.t_stride = t_stride
self.extra_1 = extra_1
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(
range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
return 0, truncate_t
if __name__ == "__main__":
import os
import numpy as np
import torchvision.io as io
from torchvision import transforms
from torchvision.utils import save_image
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi",
pts_unit="sec",
output_format="TCHW")
trans = transforms.Compose([
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5],
std=[0.5, 0.5, 0.5],
inplace=True),
])
target_video_len = 32
frame_interval = 1
total_frames = len(vframes)
print(total_frames)
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
# Sampling video frames
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(start_frame_ind,
end_frame_ind - 1,
target_video_len,
dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
print(select_vframes.shape)
print(select_vframes.dtype)
select_vframes_trans = trans(select_vframes)
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) *
255).to(dtype=torch.uint8)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi",
select_vframes_trans_int.permute(0, 2, 3, 1),
fps=8)
for i in range(target_video_len):
save_image(
select_vframes_trans[i],
os.path.join("./test000", "%04d.png" % i),
normalize=True,
value_range=(-1, 1),
)
+1 -1
View File
@@ -655,7 +655,7 @@ class GroupCoordinator:
tensor_dict[key] = value
return tensor_dict
def barrier(self):
def barrier(self) -> None:
"""Barrier synchronization among the group.
NOTE: don't use `device_group` here! `barrier` in NCCL is
terrible because it is internally a broadcast operation with
@@ -1,29 +0,0 @@
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from fastvideo.v1.pipelines.wan.wan_latent_pipeline import WanLatentPipeline
def main():
print("Starting data preprocessor")
pipeline = WanLatentPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
for batch in train_dataloader:
pipeline(batch)
if __name__ == "__main__":
main()
+7 -22
View File
@@ -70,7 +70,7 @@ class FastVideoArgs:
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
# "fp16",
"fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
@@ -438,6 +438,11 @@ def get_current_fastvideo_args() -> FastVideoArgs:
@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
@@ -473,8 +478,7 @@ class TrainingArgs(FastVideoArgs):
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = ""
resume_from_lora_checkpoint: str = ""
resume_from_checkpoint: bool = False
logging_dir: str = ""
# optimizer & scheduler
@@ -490,13 +494,7 @@ class TrainingArgs(FastVideoArgs):
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
use_cpu_offload: bool = False
# fp16_full_eval: bool
# fp16_backend: str
train_sp_batch_size: int = 0
use_lora: bool = False
lora_alpha: int = 0
lora_rank: int = 0
fsdp_sharding_startegy: str = ""
weighting_scheme: str = ""
@@ -504,8 +502,6 @@ class TrainingArgs(FastVideoArgs):
logit_std: float = 1.0
mode_scale: float = 0.0
# lr_scheduler
lr_scheduler: str = ""
num_euler_timesteps: int = 0
lr_num_cycles: int = 0
lr_power: float = 0.0
@@ -649,9 +645,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument("--resume-from-lora-checkpoint",
type=str,
help="Path to LoRA checkpoint to resume from")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
@@ -700,14 +693,6 @@ class TrainingArgs(FastVideoArgs):
type=int,
help="Training spatial parallelism batch size")
# LoRA configuration
parser.add_argument("--use-lora",
action=StoreBoolean,
help="Whether to use LoRA")
parser.add_argument("--lora-alpha",
type=int,
help="LoRA alpha parameter")
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--fsdp-sharding-strategy",
type=str,
help="FSDP sharding strategy")
+17 -24
View File
@@ -394,30 +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, default_dtype: %s", cls_name, default_dtype)
# model = load_fsdp_model(model_cls=model_cls,
# init_params={
# "config": dit_config,
# "hf_config": hf_config
# },
# weight_dir_list=safetensors_list,
# device=fastvideo_args.device,
# cpu_offload=fastvideo_args.use_cpu_offload,
# default_dtype=default_dtype)
model = load_fsdp_model(model_cls=model_cls,
init_params={
"config": dit_config,
"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,
)
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())):
+7 -14
View File
@@ -14,14 +14,15 @@ 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.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
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.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
logger = init_logger(__name__)
@@ -89,11 +90,6 @@ def get_param_names_mapping(
# TODO(PY): add compile option
# param_dtype: torch.dtype,
# reduce_dtype: torch.dtype,
# output_dtype: torch.dtype,
# pp_enabled: bool = False,
# cpu_offload: bool = False,
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
@@ -106,9 +102,11 @@ def load_fsdp_model(
output_dtype: Optional[torch.dtype] = None,
) -> torch.nn.Module:
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
mp_policy = MixedPrecisionPolicy(param_dtype,
reduce_dtype,
output_dtype,
cast_forward_inputs=True)
# with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
@@ -138,7 +136,6 @@ def load_fsdp_model(
f"Unexpected param or buffer {n} on meta device.")
for p in model.parameters():
p.requires_grad = False
# set_state_dict(model, StateDictType.LOCAL_STATE_DICT)
return model
@@ -232,10 +229,6 @@ def load_fsdp_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
# s = fully_shard.state(model)
# logger.info(f"type(s): {type(s)}")
# logger.info(f"s: {s}")
# import pdb; pdb.set_trace()
sharded_sd = {}
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
+3
View File
@@ -39,6 +39,9 @@ class ParallelTiledVAE(ABC):
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
def to(self, device) -> 'ParallelTiledVAE':
return self
@property
def temporal_compression_ratio(self) -> int:
return cast(int, self.config.temporal_compression_ratio)
@@ -40,6 +40,8 @@ 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,
@@ -51,7 +53,15 @@ class ComposedPipelineBase(ABC):
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
self.fastvideo_args = fastvideo_args
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] = {}
@@ -77,27 +87,22 @@ class ComposedPipelineBase(ABC):
self.modules = self.load_modules(fastvideo_args)
if fastvideo_args.training_mode:
if fastvideo_args.log_validation:
self.initialize_validation_pipeline(fastvideo_args)
self.initialize_training_pipeline(fastvideo_args)
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 fastvideo_args.training_mode:
logger.info("Creating training pipeline stages...")
self.create_training_stages(fastvideo_args)
else:
if not fastvideo_args.training_mode:
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
def initialize_training_pipeline(self, fastvideo_args: FastVideoArgs):
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, fastvideo_args: FastVideoArgs):
def initialize_validation_pipeline(self, training_args: TrainingArgs):
raise NotImplementedError(
"if log_validation is True, the pipeline must implement this method"
)
@@ -134,7 +139,7 @@ class ComposedPipelineBase(ABC):
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args.inference_mode:
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",
@@ -155,7 +160,6 @@ class ComposedPipelineBase(ABC):
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
# we use cpu offload for training
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
@@ -168,7 +172,7 @@ class ComposedPipelineBase(ABC):
fastvideo_args.check_fastvideo_args()
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
return cls(model_path,
fastvideo_args,
@@ -190,6 +194,8 @@ class ComposedPipelineBase(ABC):
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)
@@ -244,14 +250,7 @@ class ComposedPipelineBase(ABC):
"""
raise NotImplementedError
# @abstractmethod
# def create_validation_stages(self, fastvideo_args: FastVideoArgs):
# """
# Create the validation pipeline stages.
# """
# raise NotImplementedError
def create_training_stages(self, fastvideo_args: FastVideoArgs):
def create_training_stages(self, training_args: TrainingArgs):
"""
Create the training pipeline stages.
"""
-303
View File
@@ -1,303 +0,0 @@
import os
import sys
import time
from collections import deque
import torch
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.utils.checkpoint import save_checkpoint
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper
from fastvideo.utils.validation import log_validation
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class WanTrainingPipeline(ComposedPipelineBase): # == distill_one_step
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="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")))
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
device = fastvideo_args.device
local_rank = int(os.environ.get("LOCAL_RANK", -1))
rank = int(os.environ.get("RANK", -1))
assert rank != -1
assert local_rank != -1
sp_group = get_sp_group()
world_size = sp_group.world_size
rank = sp_group.rank
args = fastvideo_args
transformer = self.get_module("transformer")
teacher_transformer = self.get_module("teacher_transformer")
ema_transformer = None
assert not fastvideo_args.use_ema, "ema is not supported now"
assert teacher_transformer is not None
assert transformer is not None
train_dataset = self.train_dataset
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
solver = self.solver
noise_random_generator = None
uncond_prompt_embed = self.uncond_prompt_embed
uncond_prompt_mask = self.uncond_prompt_mask
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(train_dataset))
logger.info(" Dataloader size = %s", len(train_dataloader))
logger.info(" Num Epochs = %s", args.num_train_epochs)
logger.info(" Resume training from step %s", init_steps)
logger.info(" Instantaneous batch size per device = %s",
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",
args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s", args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e9)
# print dtype
logger.info(" Master weight dtype: %s",
transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
# loader = self.get_module("train_dataloader")
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule,
step)
loss, grad_norm, pred_norm = self.distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
raise NotImplementedError("lora is not supported now")
# save_lora_checkpoint(transformer, optimizer, rank,
# args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
raise NotImplementedError("ema is not supported now")
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir,
step)
sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.
linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
raise NotImplementedError("lora is not supported now")
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers", args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
print(args)
main(args)
+33 -41
View File
@@ -9,6 +9,7 @@ import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any, Dict
import numpy as np
import pyarrow as pa
@@ -52,8 +53,8 @@ class PreprocessPipeline(ComposedPipelineBase):
args,
):
# Initialize class variables for data sharing
self.video_data = {} # Store video metadata and paths
self.latent_data = {} # Store latent tensors
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)
@@ -126,7 +127,6 @@ class PreprocessPipeline(ComposedPipelineBase):
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
# Get corresponding captions for this batch
batch_captions = valid_data["text"]
batch = ForwardBatch(
@@ -135,15 +135,15 @@ class PreprocessPipeline(ComposedPipelineBase):
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]
# Remove padding from prompt_embeds using attention mask for all batches
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Create a list to store non-padded embeddings and masks
non_padded_embeds = []
non_padded_masks = []
@@ -267,7 +267,7 @@ class PreprocessPipeline(ComposedPipelineBase):
self.all_tables = []
self.all_tables.append(table)
logger.info(f"Collected batch with {len(table)} samples")
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
@@ -296,7 +296,7 @@ class PreprocessPipeline(ComposedPipelineBase):
print(
f"Using {num_workers} workers to process {total_chunks} chunks"
)
logger.info(f"Chunks per worker: {chunks_per_worker}")
logger.info("Chunks per worker: %s", chunks_per_worker)
# Prepare work ranges
work_ranges = []
@@ -320,30 +320,28 @@ class PreprocessPipeline(ComposedPipelineBase):
try:
written = future.result()
total_written += written
logger.info(
f"Processed chunk with {written} samples")
logger.info("Processed chunk with %s samples",
written)
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
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(
f"Retrying {len(failed_ranges)} failed ranges sequentially"
)
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(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info(f"Total samples written: {total_written}")
logger.info("Total samples written: %s", total_written)
num_processed_samples = 0
self.all_tables = []
@@ -354,10 +352,6 @@ class PreprocessPipeline(ComposedPipelineBase):
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
# Initialize Parquet dataset
validation_parquet_path = os.path.join(validation_parquet_dir,
"data.parquet")
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
@@ -378,6 +372,7 @@ class PreprocessPipeline(ComposedPipelineBase):
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]
@@ -386,15 +381,15 @@ class PreprocessPipeline(ComposedPipelineBase):
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
# Slice the embeddings to keep only the non-padding parts
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
# Log the shapes after removing padding
logger.info(
f"Shape after removing padding - Embeddings: {text_embedding.shape}, Mask: {text_attention_mask.shape}"
)
"Shape after removing padding - Embeddings: %s, Mask: %s",
text_embedding.shape, text_attention_mask.shape)
# Create record for Parquet dataset
record = {
@@ -419,7 +414,7 @@ class PreprocessPipeline(ComposedPipelineBase):
}
batch_data.append(record)
logger.info(f"Saved validation sample: {file_name}")
logger.info("Saved validation sample: %s", file_name)
if batch_data:
# Add progress bar for writing to Parquet dataset
@@ -472,7 +467,7 @@ class PreprocessPipeline(ComposedPipelineBase):
write_pbar.update(1)
write_pbar.close()
logger.info(f"Total validation samples: {len(table)}")
logger.info("Total validation samples: %s", len(table))
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
@@ -489,30 +484,28 @@ class PreprocessPipeline(ComposedPipelineBase):
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
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(
f"Retrying {len(failed_ranges)} failed ranges sequentially")
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(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info(f"Total validation samples written: {total_written}")
logger.info("Total validation samples written: %s", total_written)
# Clear memory
del table
gc.collect() # Force garbage collection
@staticmethod
def process_chunk_range(args):
def process_chunk_range(args: Any) -> int:
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
@@ -558,10 +551,9 @@ class PreprocessPipeline(ComposedPipelineBase):
return total_written
except Exception as e:
logger.error(
f"Error processing chunks {start_idx}-{end_idx} for worker {worker_id}: {str(e)}"
)
logger.error("Error processing chunks %s-%s for worker %s: %s",
start_idx, end_idx, worker_id, str(e))
raise
EntryClass = PreprocessPipeline
EntryClass = PreprocessPipeline
+2 -1
View File
@@ -7,6 +7,7 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
@@ -23,7 +24,7 @@ class DecodingStage(PipelineStage):
"""
def __init__(self, vae) -> None:
self.vae = vae
self.vae: ParallelTiledVAE = vae
def forward(
self,
+2 -3
View File
@@ -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,6 +73,7 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
# 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
@@ -84,7 +84,6 @@ class DenoisingStage(PipelineStage):
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
# b c t h w -> b t n s h w
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
@@ -190,7 +189,7 @@ class DenoisingStage(PipelineStage):
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=torch.bfloat16,
dtype=target_dtype,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
@@ -46,6 +46,8 @@ class EncodingStage(PipelineStage):
Returns:
The batch with encoded outputs.
"""
self.vae = self.vae.to(fastvideo_args.device)
image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
-824
View File
@@ -1,824 +0,0 @@
import gc
import os
import sys
import time
import traceback
from abc import ABC, abstractmethod
from collections import deque
from copy import deepcopy
import imageio
import numpy as np
import torch
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.checkpoint import save_checkpoint_v1
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.training_utils import (
_clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas)
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
GRADIENT_CHECK_DTYPE = torch.bfloat16
class TrainingPipeline(ComposedPipelineBase, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
All reusable components and code should be implemented in this class.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_training_pipeline(self, fastvideo_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = fastvideo_args.device
self.sp_group = get_sp_group()
self.world_size = self.sp_group.world_size
self.rank = self.sp_group.rank
self.local_rank = self.sp_group.local_rank
self.transformer = self.get_module("transformer")
assert self.transformer is not None
self.transformer.requires_grad_(True)
self.transformer.train()
args = fastvideo_args
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * self.world_size,
num_training_steps=args.max_train_steps * self.world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = ParquetVideoTextDataset(
args.data_path,
batch_size=args.train_batch_size,
rank=self.rank,
world_size=self.world_size,
cfg_rate=args.cfg,
num_latent_t=args.num_latent_t)
train_dataloader = StatefulDataLoader(
train_dataset,
batch_size=args.train_batch_size,
num_workers=args.
dataloader_num_workers, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
self.lr_scheduler = lr_scheduler
self.train_dataset = train_dataset
self.train_dataloader = train_dataloader
self.init_steps = init_steps
self.optimizer = optimizer
self.noise_scheduler = noise_scheduler
# self.noise_random_generator = noise_random_generator
# num_update_steps_per_epoch = math.ceil(
# len(train_dataloader) / args.gradient_accumulation_steps *
# args.sp_size / args.train_sp_batch_size)
# args.num_train_epochs = math.ceil(args.max_train_steps /
# num_update_steps_per_epoch)
if self.rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@abstractmethod
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"Training pipelines must implement this method")
@abstractmethod
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
loader, noise_scheduler, noise_random_generator,
gradient_accumulation_steps, sp_size,
precondition_outputs, max_grad_norm, weighting_scheme,
logit_mean, logit_std, mode_scale):
"""
Train one step of the model.
"""
raise NotImplementedError(
"Training pipeline must implement this method")
def log_validation(self, transformer, fastvideo_args, global_step):
fastvideo_args.inference_mode = True
fastvideo_args.use_cpu_offload = False
if not fastvideo_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
# Prepare validation prompts
print('fastvideo_args.validation_prompt_dir',
fastvideo_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
fastvideo_args.validation_prompt_dir,
batch_size=1,
rank=0,
world_size=1,
cfg_rate=0,
num_latent_t=args.num_latent_t)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=False)
transformer.requires_grad_(False)
for p in transformer.parameters():
p.requires_grad = False
transformer.eval()
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
self.validation_pipeline.latent_preparation_stage.transformer = transformer
self.validation_pipeline.denoising_stage.transformer = transformer
# Process each validation prompt
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
logger.info(f"infos: {infos}")
caption = infos['caption']
captions.append(caption)
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
logger.info('embed dtype', prompt_embeds.dtype)
# Prepare batch for validation
# print('shape of embeddings', prompt_embeds.shape)
batch = ForwardBatch(
# **shallow_asdict(sampling_param),
data_type="video",
latents=None,
# seed=sampling_param.seed,
# data_type="video",
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=args.num_height,
width=args.num_width,
num_frames=args.num_frames,
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=50,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
extra={},
)
# Run validation inference
with torch.autocast("cuda", dtype=torch.bfloat16):
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
# Log validation results
rank = int(os.environ.get("RANK", 0))
if rank == 0:
video_filenames = []
video_captions = []
for i, video in enumerate(videos):
caption = captions[i]
filename = os.path.join(
fastvideo_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_captions.append(
caption) # Store the caption for each video
logs = {
"validation_videos": [
wandb.Video(filename,
caption=caption) for filename, caption in zip(
video_filenames, video_captions)
]
}
wandb.log(logs, step=global_step)
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def gradient_check_parameters(self,
transformer,
latents,
encoder_hidden_states,
encoder_attention_mask,
timesteps,
target,
eps=5e-2,
max_params_to_check=2000):
"""
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
"""
# Move all inputs to CPU and clear GPU memory
inputs_cpu = {
'latents': latents.cpu(),
'encoder_hidden_states': encoder_hidden_states.cpu(),
'encoder_attention_mask': encoder_attention_mask.cpu(),
'timesteps': timesteps.cpu(),
'target': target.cpu()
}
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
torch.cuda.empty_cache()
def compute_loss():
# Move inputs to GPU, compute loss, cleanup
inputs_gpu = {
k:
v.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE
if k != 'encoder_attention_mask' else None)
for k, v in inputs_cpu.items()
}
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
# with torch.autocast(enabled=False, device_type="cuda"):
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
with set_forward_context(
current_timestep=inputs_gpu['timesteps'],
attn_metadata=None):
model_pred = transformer(
hidden_states=inputs_gpu['latents'],
encoder_hidden_states=inputs_gpu[
'encoder_hidden_states'],
timestep=inputs_gpu['timesteps'],
encoder_attention_mask=inputs_gpu[
'encoder_attention_mask'],
return_dict=False)[0]
if self.fastvideo_args.precondition_outputs:
sigmas = get_sigmas(self.noise_scheduler,
inputs_gpu['latents'].device,
inputs_gpu['timesteps'],
n_dim=inputs_gpu['latents'].ndim,
dtype=inputs_gpu['latents'].dtype)
model_pred = inputs_gpu['latents'] - model_pred * sigmas
target_adjusted = inputs_gpu['target']
else:
target_adjusted = inputs_gpu['target']
loss = torch.mean((model_pred - target_adjusted)**2)
# Cleanup and return
loss_cpu = loss.cpu()
del inputs_gpu, model_pred, target_adjusted
if 'sigmas' in locals(): del sigmas
torch.cuda.empty_cache()
return loss_cpu.to(self.fastvideo_args.device)
try:
# Get analytical gradients
transformer.zero_grad()
analytical_loss = compute_loss()
analytical_loss.backward()
# Check gradients for selected parameters
absolute_errors = []
param_count = 0
for name, param in transformer.named_parameters():
if not (param.requires_grad and param.grad is not None
and param_count < max_params_to_check
and param.grad.abs().max() > 5e-4):
continue
# Get local parameter and gradient tensors
local_param = param._local_tensor if hasattr(
param, '_local_tensor') else param
local_grad = param.grad._local_tensor if hasattr(
param.grad, '_local_tensor') else param.grad
# Find first significant gradient element
flat_param = local_param.data.view(-1)
flat_grad = local_grad.view(-1)
check_idx = next((i for i in range(min(10, flat_param.numel()))
if abs(flat_grad[i]) > 1e-4), 0)
# Store original values
orig_value = flat_param[check_idx].item()
analytical_grad = flat_grad[check_idx].item()
# Compute numerical gradient
for delta in [eps, -eps]:
with torch.no_grad():
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0: loss_plus = loss.item()
else: loss_minus = loss.item()
# Restore parameter and compute error
with torch.no_grad():
flat_param[check_idx] = orig_value
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
abs_error = abs(analytical_grad - numerical_grad)
rel_error = abs_error / max(abs(analytical_grad),
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
logger.info(
f"{name}[{check_idx}]: analytical={analytical_grad:.6f}, "
f"numerical={numerical_grad:.6f}, abs_error={abs_error:.2e}, rel_error={rel_error:.2%}"
)
# param_count += 1
# Compute and log statistics
if absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
logger.info(
f"Gradient check stats: min={min_err:.2e}, max={max_err:.2e}, mean={mean_err:.2e}"
)
if self.rank <= 0:
wandb.log({
"grad_check/min_abs_error":
min_err,
"grad_check/max_abs_error":
max_err,
"grad_check/mean_abs_error":
mean_err,
"grad_check/analytical_loss":
analytical_loss.item(),
})
return max_err
return float('inf')
except Exception as e:
logger.error(f"Gradient check failed: {e}")
traceback.print_exc()
return float('inf')
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
noise_random_generator):
"""
Setup and perform gradient check on a fresh batch.
Args:
args: Training arguments
loader_iter: Data loader iterator
noise_scheduler: Noise scheduler for diffusion
noise_random_generator: Random number generator for noise
Returns:
float or None: Maximum gradient error or None if check is disabled/fails
"""
if not ENABLE_GRADIENT_CHECK:
return None
try:
# Get a fresh batch and process it exactly like train_one_step
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
loader_iter)
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
check_latents = check_latents.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE)
check_encoder_hidden_states = check_encoder_hidden_states.to(
self.fastvideo_args.device, dtype=GRADIENT_CHECK_DTYPE)
check_latents = normalize_dit_input("wan", check_latents)
batch_size = check_latents.shape[0]
check_noise = torch.randn_like(check_latents)
check_u = compute_density_for_timestep_sampling(
weighting_scheme=args.weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=args.logit_mean,
logit_std=args.logit_std,
mode_scale=args.mode_scale,
)
check_indices = (check_u *
noise_scheduler.config.num_train_timesteps).long()
check_timesteps = noise_scheduler.timesteps[check_indices].to(
device=check_latents.device)
check_sigmas = get_sigmas(
noise_scheduler,
check_latents.device,
check_timesteps,
n_dim=check_latents.ndim,
dtype=check_latents.dtype,
)
check_noisy_model_input = (
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
# Compute target exactly like train_one_step
if args.precondition_outputs:
check_target = check_latents
else:
check_target = check_noise - check_latents
# Perform gradient check with the exact same inputs as training
max_grad_error = self.gradient_check_parameters(
transformer=self.transformer,
latents=
check_noisy_model_input, # Use noisy input like in training
encoder_hidden_states=check_encoder_hidden_states,
encoder_attention_mask=check_encoder_attention_mask,
timesteps=check_timesteps,
target=check_target,
max_params_to_check=100 # Check more parameters
)
if max_grad_error > 5e-2:
logger.error(
f"❌ Large gradient error detected: {max_grad_error:.2e}")
else:
logger.info(
f"✅ Gradient check passed: max error {max_grad_error:.2e}")
return max_grad_error
except Exception as e:
logger.error(f"Gradient check setup failed: {e}")
traceback.print_exc()
return None
class WanTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def create_training_stages(self, fastvideo_args: FastVideoArgs):
pass
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(fastvideo_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
# TODO(will): clean this up
args_copy.precision = "bf16"
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=args_copy)
self.validation_pipeline = validation_pipeline
def train_one_step(
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
):
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
encoder_attention_mask,
infos,
) = next(loader_iter)
latents = latents.to(self.fastvideo_args.device,
dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
self.fastvideo_args.device, dtype=torch.bfloat16)
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
print('device before forward ',
next(transformer.named_parameters())[1].device)
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)[0]
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
if precondition_outputs:
target = latents
else:
target = noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
print('device before backwardin context',
next(transformer.named_parameters())[1].device)
print('device before backward out context',
next(transformer.named_parameters())[1].device)
loss.backward()
print('device after backward out context',
next(transformer.named_parameters())[1].device)
avg_loss = loss.detach().clone()
sp_group = get_sp_group()
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
model_parts = [self.transformer]
grad_norm = _clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
optimizer.step()
print('device after optimizer step',
next(transformer.named_parameters())[1].device)
lr_scheduler.step()
print('device after scheduler step',
next(transformer.named_parameters())[1].device)
return total_loss, grad_norm.item()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
args = fastvideo_args
self.fastvideo_args = args
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
noise_random_generator = None
from diffusers import FlowMatchEulerDiscreteScheduler
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
total_batch_size = (self.world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
# logger.info(f" Num examples = {len(train_dataset)}")
# logger.info(f" Dataloader size = {len(train_dataloader)}")
# logger.info(f" Num Epochs = {args.num_train_epochs}")
logger.info(f" Resume training from step {init_steps}")
logger.info(
f" Instantaneous batch size per device = {args.train_batch_size}")
logger.info(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
logger.info(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}"
)
logger.info(f" Total optimization steps = {args.max_train_steps}")
logger.info(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in self.transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
logger.info(
f" Master weight dtype: {self.transformer.parameters().__next__().dtype}"
)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
loader_iter = iter(train_dataloader)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader_iter)
# get gpu memory usage
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage before train_one_step: {gpu_memory_usage} MB")
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
args.weighting_scheme,
args.logit_mean,
args.logit_std,
args.mode_scale,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage after train_one_step: {gpu_memory_usage} MB")
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info(f"Performing gradient check at step {step}")
self.setup_gradient_check(args, loader_iter, noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
raise NotImplementedError("LoRA is not supported now")
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step, pipe)
else:
# Your existing checkpoint saving code
save_checkpoint_v1(self.transformer, self.rank,
args.output_dir, step)
self.transformer.train()
self.sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
self.log_validation(self.transformer, args, step)
if args.use_lora:
raise NotImplementedError("LoRA is not supported now")
# save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps, pipe)
else:
save_checkpoint_v1(self.transformer, self.rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
print(args)
main(args)
@@ -1,19 +0,0 @@
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
class WanLatentPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs):
logger.info("WAN Latent Pipeline forward")
pass
@@ -15,7 +15,6 @@ from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
View File
+515
View File
@@ -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
@@ -1,12 +1,18 @@
import json
import math
import os
from typing import List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.tensor
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
@@ -14,9 +20,9 @@ def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
generator,
logit_mean: float = None,
logit_std: float = None,
mode_scale: float = None,
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.
@@ -47,7 +53,7 @@ def get_sigmas(noise_scheduler,
device,
timesteps,
n_dim=4,
dtype=torch.float32):
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)
@@ -60,10 +66,53 @@ def get_sigmas(noise_scheduler,
return sigma
logger = init_logger(__name__)
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 _clip_grad_norm_while_handling_failing_dtensor_cases(
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,
@@ -87,8 +136,8 @@ def _clip_grad_norm_while_handling_failing_dtensor_cases(
)
except Exception as e:
logger.warning(
f"An error occurred while clipping gradients: {e}. Gradient clipping will be skipped and gradient "
f"norm will not be logged.")
"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
@@ -148,6 +197,7 @@ def clip_grad_norm_(
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,
@@ -207,10 +257,7 @@ def _get_total_norm(
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
) -> torch.Tensor:
if isinstance(tensors, torch.Tensor):
tensors = [tensors]
else:
tensors = list(tensors)
tensors = [tensors] if isinstance(tensors, torch.Tensor) else list(tensors)
norm_type = float(norm_type)
if len(tensors) == 0:
return torch.tensor(0.0)
@@ -263,8 +310,8 @@ def _group_tensors_by_device_and_dtype(
with_indices: bool = False,
) -> dict[tuple[torch.device, torch.dtype], tuple[
List[List[Optional[torch.Tensor]]], List[int]]]:
return torch._C._group_tensors_by_device_and_dtype(tensorlistlist,
with_indices)
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:
@@ -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)
-953
View File
@@ -1,953 +0,0 @@
# !/bin/python3
# isort: skip_file
import argparse
import math
import os
import time
from collections import deque
from copy import deepcopy
import torch
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from peft import LoraConfig
# from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
# from torch.distributed.fsdp import ShardingStrategy
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint,
save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing)
# from fastvideo.utils.load import load_transformer
# from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
# initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
from fastvideo.v1.models.loader.component_loader import TransformerLoader
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.forward_context import set_forward_context
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
def main_print(content):
if int(os.environ["LOCAL_RANK"]) <= 0:
print(content)
# def reshard_fsdp(model):
# for m in FSDP.fsdp_modules(model):
# if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
# torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
gradient_accumulation_steps)
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
gradient_accumulation_steps)
absolute_mean = torch.mean(
torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(
torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
def distill_one_step(
transformer,
model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_teacher_disable_cfg,
):
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0, # codespell:ignore
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
}
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
model_input = normalize_dit_input(model_type, latents)
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
timesteps = (sigmas *
noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
print(f"--> noisy_model_input.dtype: {noisy_model_input.dtype}")
print(f"--> noisy_model_input.shape: {noisy_model_input.shape}")
print(
f"--> encoder_hidden_states.dtype: {encoder_hidden_states.dtype}"
)
print(
f"--> encoder_hidden_states.shape: {encoder_hidden_states.shape}"
)
print(f"--> timesteps.dtype: {timesteps.dtype}")
print(f"--> timesteps.shape: {timesteps.shape}")
print(
f"--> encoder_attention_mask.dtype: {encoder_attention_mask.dtype}"
)
print(
f"--> encoder_attention_mask.shape: {encoder_attention_mask.shape}"
)
noisy_model_input = noisy_model_input.to(dtype=torch.bfloat16)
teacher_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_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
# batch = ForwardBatch(
# enable_teacache=False,
# )
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=None,
fastvideo_args=None):
model_pred = transformer(**teacher_kwargs)[0]
print(f"--> model_pred shape: {model_pred.shape}")
huber_c = 0.001
target = torch.randn_like(model_pred)
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
loss.backward()
print(f"--> loss: {loss.item()}")
assert False, "stop here"
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
huber_c = 0.001
# loss = loss.mean()
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float()**2)) *
pred_decay_weight / gradient_accumulation_steps)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (torch.mean(model_pred.float()**2) *
pred_decay_weight /
gradient_accumulation_steps)
loss += pred_decay_loss
else:
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(model_pred.detach().float(), model_pred_norm,
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(ema_transformer.parameters(),
transformer.parameters()):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(),
1 - ema_decay))
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item(), model_pred_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=args.sp_size,
sequence_model_parallel_size=args.sp_size)
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
print(
f"--> local_rank: {local_rank}, rank: {rank}, world_size: {world_size}")
main_print(f"--> using model pipeline {args.pretrained_model_name_or_path}")
model_path = maybe_download_model(args.pretrained_model_name_or_path)
main_print(f"--> loading model from {model_path}")
transformer_path = os.path.join(model_path, "transformer")
main_print(f"--> loading transformer from {transformer_path}")
precision = torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16
precision_str = "fp32" if precision == torch.float32 else "bf16"
print(f"--> precision: {precision_str}")
print(f"--> precision: {precision}")
# transformer_path = os.path.join(args.pretrained_model_name_or_path, "transformer")
fastvideo_args = FastVideoArgs(model_path=transformer_path,
use_cpu_offload=False,
precision=precision_str)
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.dit_config = WanVideoConfig()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device_str = f"cuda:{local_rank}"
device = fastvideo_args.device
torch.cuda.set_device(device)
loader = TransformerLoader()
print(f"--> loading transformer to device {device} on rank {rank}")
transformer = loader.load(transformer_path, "",
fastvideo_args).to(device, dtype=precision)
# teacher_transformer = deepcopy(transformer)
if args.use_ema:
raise NotImplementedError("EMA is not supported for v1 distillation.")
ema_transformer = deepcopy(transformer)
else:
ema_transformer = None
if args.use_lora:
raise NotImplementedError("LoRA is not supported for v1 distillation.")
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer.add_adapter(transformer_lora_config)
main_print(
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
if args.use_lora:
raise NotImplementedError("LoRA is not supported for v1 distillation.")
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = [
"to_k", "to_q", "to_v", "to_out.0"
]
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
transformer)
main_print("--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules,
args.selective_checkpointing)
if args.use_ema:
apply_fsdp_checkpointing(ema_transformer, no_split_modules,
args.selective_checkpointing)
# Set model as trainable.
transformer.train()
transformer.requires_grad_(True)
# teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
linear_steps = int(noise_scheduler.config.num_train_timesteps *
args.linear_range)
sigmas = linear_quadratic_schedule(
noise_scheduler.config.num_train_timesteps,
args.linear_quadratic_threshold,
linear_steps,
)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
sigmas = noise_scheduler.sigmas
solver = EulerSolver(
sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
# l = list(params_to_optimize)
# for p in params_to_optimize:
# main_print(type(p))
# main_print(f"--> p: {p.shape}")
# main_print(f"--> p: {p.dtype}")
# main_print(f"--> p: {p.device}")
# main_print(f"--> p: {p.requires_grad}")
# main_print('=------------------------')
# break
# print(f"--> params_to_optimize: {list(params_to_optimize)}")
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
# print(f"--> params_to_optimize2: {params_to_optimize}")
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
# optimizer = None
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer)
main_print(f"optimizer: {optimizer}")
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
# if rank <= 0:
# project = args.tracker_project_name or "fastvideo"
# wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(
f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
loss, grad_norm, pred_norm = distill_one_step(
transformer,
args.model_type,
None, # teacher_transformer
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
# parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.95)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=str, default="64")
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
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(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
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***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="Threshold for linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--use_ema",
action="store_true",
help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+1
View File
@@ -0,0 +1 @@
__version__ = "0.1.0"
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.1.1"
version = "0.1.0"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
@@ -1,11 +1,11 @@
import json
from pathlib import Path
import csv
import cv2
def get_video_info(video_path, metadata):
"""Extract video information using OpenCV and corresponding metadata"""
def get_video_info(video_path, prompt_text):
"""Extract video information using OpenCV and corresponding prompt text"""
cap = cv2.VideoCapture(str(video_path))
if not cap.isOpened():
@@ -23,66 +23,60 @@ def get_video_info(video_path, metadata):
return {
"path": video_path.name,
"title": metadata.get("Video Title", ""),
"description": metadata.get("Video Description", ""),
"video_url": metadata.get("Video URL", ""),
"download_url": metadata.get("Download URL", ""),
"resolution": {
"width": width,
"height": height
},
"fps": fps,
"duration": duration,
"cap": [metadata.get("Video Description", "")]
"cap": [prompt_text]
}
def read_csv_file(csv_path):
"""Read and return the content of a CSV file"""
def read_prompt_file(prompt_path):
"""Read and return the content of a prompt file"""
try:
with open(csv_path, 'r', encoding='utf-8') as f:
reader = csv.DictReader(f)
return list(reader)
with open(prompt_path, 'r', encoding='utf-8') as f:
return f.read().strip()
except Exception as e:
print(f"Error reading CSV file {csv_path}: {e}")
print(f"Error reading prompt file {prompt_path}: {e}")
return None
def process_videos_from_csv(video_dir_path, csv_path, verbose=False):
"""Process videos using metadata from CSV file
def process_videos_and_prompts(video_dir_path, prompt_dir_path, verbose=False):
"""Process videos and their corresponding prompt files
Args:
video_dir_path (str): Path to directory containing video files
csv_path (str): Path to CSV file containing video metadata
prompt_dir_path (str): Path to directory containing prompt files
verbose (bool): Whether to print verbose processing information
"""
video_dir = Path(video_dir_path)
csv_data = read_csv_file(csv_path)
prompt_dir = Path(prompt_dir_path)
processed_data = []
# Ensure directories exist
if not video_dir.exists():
print(f"Error: Video directory does not exist: {video_dir}")
return []
if csv_data is None:
if not video_dir.exists() or not prompt_dir.exists():
print(f"Error: One or both directories do not exist:\nVideos: {video_dir}\nPrompts: {prompt_dir}")
return []
# Process each video file
for row in csv_data:
video_filename = row.get("Filename")
if not video_filename:
for video_file in video_dir.glob('*.mp4'):
video_name = video_file.stem
prompt_file = prompt_dir / f"{video_name}.txt"
# Check if corresponding prompt file exists
if not prompt_file.exists():
print(f"Warning: No prompt file found for video {video_name}")
continue
video_file = video_dir / video_filename
# Check if video file exists
if not video_file.exists():
print(f"Warning: Video file not found: {video_filename}")
# Read prompt content
prompt_text = read_prompt_file(prompt_file)
if prompt_text is None:
continue
# Process video and add to results
video_info = get_video_info(video_file, row)
video_info = get_video_info(video_file, prompt_text)
if video_info:
processed_data.append(video_info)
@@ -111,9 +105,9 @@ def parse_args():
"""Parse command line arguments"""
import argparse
parser = argparse.ArgumentParser(description='Process videos using metadata from CSV file')
parser = argparse.ArgumentParser(description='Process videos and their corresponding prompt files')
parser.add_argument('--video_dir', '-v', required=True, help='Directory containing video files')
parser.add_argument('--csv_path', '-c', required=True, help='Path to CSV file containing video metadata')
parser.add_argument('--prompt_dir', '-p', required=True, help='Directory containing prompt text files')
parser.add_argument('--output_path',
'-o',
required=True,
@@ -127,8 +121,8 @@ if __name__ == "__main__":
# Parse command line arguments
args = parse_args()
# Process videos from CSV
processed_videos = process_videos_from_csv(args.video_dir, args.csv_path, args.verbose)
# Process videos and prompts
processed_videos = process_videos_and_prompts(args.video_dir, args.prompt_dir, args.verbose)
if processed_videos:
# Save results
+3 -11
View File
@@ -24,9 +24,9 @@ def is_16_9_ratio(width: int, height: int, tolerance: float = 0.1) -> bool:
def resize_video(args_tuple):
"""
Resize a single video file.
args_tuple: (input_file, output_dir, width, height, fps, num_frames)
args_tuple: (input_file, output_dir, width, height, fps)
"""
input_file, output_dir, width, height, fps, num_frames = args_tuple
input_file, output_dir, width, height, fps = args_tuple
video = None
resized = None
output_file = output_dir / f"{input_file.name}"
@@ -39,13 +39,6 @@ def resize_video(args_tuple):
if not is_16_9_ratio(video.w, video.h):
return (input_file.name, "skipped", "Not 16:9")
# Calculate target duration based on num_frames and fps
target_duration = num_frames / fps
# Trim video if it's longer than target duration
if video.duration > target_duration:
video = video.subclip(0, target_duration)
def process_frame(frame):
frame_float = frame.astype(float) / 255.0
resized = resize(frame_float, (height, width, 3), mode='reflect', anti_aliasing=True, preserve_range=True)
@@ -82,7 +75,7 @@ def process_folder(args):
print(f"Target: {args.width}x{args.height} at {args.fps}fps")
# Prepare arguments for parallel processing
process_args = [(video_file, output_path, args.width, args.height, args.fps, args.num_frames) for video_file in video_files]
process_args = [(video_file, output_path, args.width, args.height, args.fps) for video_file in video_files]
successful = 0
skipped = 0
@@ -122,7 +115,6 @@ def parse_args():
parser.add_argument('--width', type=int, default=1280, help='Target width in pixels (default: 848)')
parser.add_argument('--height', type=int, default=720, help='Target height in pixels (default: 480)')
parser.add_argument('--fps', type=int, default=30, help='Target frames per second (default: 30)')
parser.add_argument('--num_frames', type=int, default=163, help='Target number of frames (default: 163)')
parser.add_argument('--max_workers',
type=int,
default=4,
-45
View File
@@ -1,45 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=/workspace/data
num_gpus=1
IP=127.0.0.1
torchrun --nnodes 1 --nproc_per_node $num_gpus \
--node_rank=0 \
--rdzv_id=456 \
--rdzv_backend=c10d \
--rdzv_endpoint=$IP:29500 \
fastvideo/distill_wan.py\
--seed 42\
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--train_batch_size=1 \
--num_latent_t 1 \
--sp_size $num_gpus \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--master_weight_type="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
@@ -9,11 +9,8 @@ NUM_GPUS=1
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# --gradient_checkpointing\
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/pipelines/training_pipeline.py\
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 \
@@ -25,12 +22,12 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--dataloader_num_workers 5\
--gradient_accumulation_steps=1\
--max_train_steps=120 \
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=60 \
--checkpointing_steps=50 \
--validation_steps 20\
--validation_sampling_steps "2,4,8" \
--log_validation \
@@ -50,4 +47,4 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0
--max_grad_norm 1.0
-27
View File
@@ -1,27 +0,0 @@
#!/bin/bash
num_gpus=2
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/entrypoints/data_preprocessor.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 480 \
--width 832 \
--num_frames 77 \
--num_inference_steps 50 \
--fps 16 \
--guidance_scale 3.0 \
--prompt_path ./assets/prompt.txt \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload
+3 -5
View File
@@ -1,10 +1,9 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
TEXT_ENCODER_PATH="/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tokenizer"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="/workspace/data/Mixkit-Src/merge.txt"
OUTPUT_DIR="/workspace/data/HD-Mixkit-Finetune-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 \
@@ -18,7 +17,6 @@ torchrun --nproc_per_node=$GPU_NUM \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--text_encoder_name $TEXT_ENCODER_PATH \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 108 \
@@ -1,30 +0,0 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
# MODEL_PATH="/home/ray/.cache/huggingface/hub/models--Wan-AI--Wan2.1-T2V-1.3B-Diffusers/snapshots/0fad780a534b6463e45facd96134c9f345acfa5b"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_MERGE_PATH="data/cats_480/merge.txt"
OUTPUT_DIR="data/cats_480_latents/"
VALIDATION_PATH="assets/prompt.txt"
# torchrun --nproc_per_node=$GPU_NUM \
# fastvideo/data_preprocess/preprocess_vae_latents_v1.py \
# --model_path $MODEL_PATH \
# --data_merge_path $DATA_MERGE_PATH \
# --train_batch_size=1 \
# --max_height=480 \
# --max_width=832 \
# --num_frames=81 \
# --dataloader_num_workers 1 \
# --output_dir=$OUTPUT_DIR \
# --train_fps 16
# torchrun --nproc_per_node=$GPU_NUM \
# fastvideo/data_preprocess/preprocess_text_embeddings_v1.py \
# --model_path $MODEL_PATH \
# --output_dir=$OUTPUT_DIR
torchrun --nproc_per_node=1 \
fastvideo/data_preprocess/preprocess_validation_text_embeddings_v1.py \
--model_path $MODEL_PATH \
--output_dir=$OUTPUT_DIR \
--validation_prompt_txt $VALIDATION_PATH