Compare commits
21
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1757d3dba0 | ||
|
|
a9d0c29ed9 | ||
|
|
a335811869 | ||
|
|
357b0533fe | ||
|
|
2ec3732758 | ||
|
|
a004408a93 | ||
|
|
007e237e69 | ||
|
|
8e18dc9f71 | ||
|
|
7ab32539af | ||
|
|
6ef8fcb61d | ||
|
|
016e24da63 | ||
|
|
85b8717545 | ||
|
|
657fd745e1 | ||
|
|
12647457a7 | ||
|
|
298f74f956 | ||
|
|
ee8babb298 | ||
|
|
60295cc03f | ||
|
|
1572e13b6e | ||
|
|
a157275b4c | ||
|
|
c4dbe7dac3 | ||
|
|
d39591108e |
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -10,7 +10,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
|
||||
@@ -33,7 +33,7 @@ repos:
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
rev: v0.11.12
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
@@ -48,7 +48,7 @@ repos:
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
@@ -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)=
|
||||
|
||||
@@ -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__"]
|
||||
|
||||
@@ -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)
|
||||
@@ -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")
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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),
|
||||
)
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
@@ -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")
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
|
||||
@@ -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())):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
+62
-15
@@ -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)
|
||||
@@ -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)
|
||||
@@ -0,0 +1 @@
|
||||
__version__ = "0.1.0"
|
||||
+1
-1
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Executable → Regular
+3
-5
@@ -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
|
||||
Reference in New Issue
Block a user