Compare commits
17
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bf5726cd06 | ||
|
|
57cfd16136 | ||
|
|
91b7cc1be8 | ||
|
|
d6365373b4 | ||
|
|
50fb94b902 | ||
|
|
2caa0d4d0b | ||
|
|
24db823998 | ||
|
|
8c4704edf5 | ||
|
|
dfba7ec833 | ||
|
|
7a2e171f1b | ||
|
|
a8aac6090a | ||
|
|
191d1be3b4 | ||
|
|
338ea1e5f2 | ||
|
|
42a2f272d5 | ||
|
|
982bfcfdc8 | ||
|
|
2ac06379a7 | ||
|
|
baaa1673f7 |
@@ -8,7 +8,7 @@ body:
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
|
||||
@@ -141,8 +141,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
|
||||
|
||||
@@ -70,8 +70,6 @@ DEFAULT_CONDA_PATTERNS = {
|
||||
"optree",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
@@ -87,8 +85,6 @@ DEFAULT_PIP_PATTERNS = {
|
||||
"onnx",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
@@ -2,6 +2,7 @@ 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)
|
||||
@@ -22,7 +23,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.linalg.norm(tensor, dim=-1, keepdim=True)
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
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"
|
||||
+3
-1
@@ -18,6 +18,7 @@ import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
@@ -167,7 +168,8 @@ _cached_base: str = ""
|
||||
_cached_branch: str = ""
|
||||
|
||||
|
||||
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
|
||||
def get_repo_base_and_branch(
|
||||
pr_number: str) -> tuple[Optional[str], Optional[str]]:
|
||||
global _cached_base, _cached_branch
|
||||
if _cached_base and _cached_branch:
|
||||
return _cached_base, _cached_branch
|
||||
|
||||
@@ -5,6 +5,7 @@ import itertools
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../../../..'
|
||||
@@ -88,7 +89,7 @@ class Example:
|
||||
generate() -> str: Generates the documentation content.
|
||||
""" # noqa: E501
|
||||
path: Path
|
||||
category: str | None = None
|
||||
category: Optional[str] = None
|
||||
main_file: Path = field(init=False)
|
||||
other_files: list[Path] = field(init=False)
|
||||
title: str = field(init=False)
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
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", "__version__"]
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
|
||||
def main(args):
|
||||
# Assume using torchrun
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
|
||||
kwargs = {
|
||||
"use_cpu_offload": False,
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
pipeline_config_args = shallow_asdict(pipeline_config)
|
||||
pipeline_config_args.update(kwargs)
|
||||
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
|
||||
num_gpus=world_size,
|
||||
device_str="cuda",
|
||||
**pipeline_config_args,
|
||||
)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
|
||||
|
||||
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_video_batch_size",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_text_batch_size",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--samples_per_file",
|
||||
type=int,
|
||||
default=64
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flush_frequency",
|
||||
type=int,
|
||||
default=256,
|
||||
help="how often to save to parquet files"
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,199 @@
|
||||
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)
|
||||
@@ -0,0 +1,151 @@
|
||||
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)
|
||||
@@ -0,0 +1,115 @@
|
||||
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)
|
||||
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -38,6 +38,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
linear_range=0.5,
|
||||
):
|
||||
if linear_quadratic:
|
||||
raise NotImplementedError("Linear quadratic schedule is not implemented")
|
||||
linear_steps = int(num_train_timesteps * linear_range)
|
||||
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
|
||||
@@ -0,0 +1,870 @@
|
||||
# !/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="The resolution of the model.",
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--load-key",
|
||||
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
|
||||
"--ring-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ring degree.",
|
||||
help="Ulysses 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:
|
||||
class Inference(object):
|
||||
|
||||
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 Master mode, but got {}".format(mode))
|
||||
raise Exception("Only supports Normal and Normal", mode)
|
||||
return prompt
|
||||
|
||||
|
||||
|
||||
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents):
|
||||
def normalize_dit_input(model_type, latents, args=None):
|
||||
if model_type == "mochi":
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
@@ -41,5 +41,16 @@ def normalize_dit_input(model_type, latents):
|
||||
return latents * 0.476986
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
elif model_type == "wan":
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
|
||||
vae_config = WanVAEConfig()
|
||||
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
|
||||
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
|
||||
|
||||
|
||||
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
|
||||
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
|
||||
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
|
||||
return latents
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
|
||||
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
|
||||
class STEP1TextEncoder(torch.nn.Module):
|
||||
|
||||
def __init__(self, model_dir, max_length=320):
|
||||
super()
|
||||
super(STEP1TextEncoder, self).__init__()
|
||||
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
|
||||
if type(prompts) is str:
|
||||
prompts = [prompts]
|
||||
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
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,6 +11,7 @@ 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
|
||||
|
||||
@@ -44,13 +45,50 @@ 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_dir, "discriminator_pytorch_model.safetensors")
|
||||
weight_path = os.path.join(save_dstate_dictir, "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(
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
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")
|
||||
@@ -3,7 +3,8 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, fields
|
||||
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
|
||||
from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
|
||||
Type, TypeVar)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
@@ -26,12 +27,12 @@ class AttentionBackend(ABC):
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_impl_cls() -> type["AttentionImpl"]:
|
||||
def get_impl_cls() -> Type["AttentionImpl"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
# @staticmethod
|
||||
@@ -45,7 +46,7 @@ class AttentionBackend(ABC):
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -56,7 +57,8 @@ class AttentionMetadata:
|
||||
current_timestep: int
|
||||
|
||||
def asdict_zerocopy(self,
|
||||
skip_fields: set[str] | None = None) -> dict[str, Any]:
|
||||
skip_fields: Optional[Set[str]] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""Similar to dataclasses.asdict, but avoids deepcopying."""
|
||||
if skip_fields is None:
|
||||
skip_fields = set()
|
||||
@@ -122,7 +124,7 @@ class AttentionImpl(ABC, Generic[T]):
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
causal: bool = False,
|
||||
num_kv_heads: int | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
|
||||
@@ -26,7 +28,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
@@ -34,15 +36,15 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
return "FLASH_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["FlashAttentionImpl"]:
|
||||
def get_impl_cls() -> Type["FlashAttentionImpl"]:
|
||||
return FlashAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -54,7 +56,7 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from sageattention import sageattn
|
||||
|
||||
@@ -15,7 +17,7 @@ class SageAttentionBackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
@@ -23,7 +25,7 @@ class SageAttentionBackend(AttentionBackend):
|
||||
return "SAGE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SageAttentionImpl"]:
|
||||
def get_impl_cls() -> Type["SageAttentionImpl"]:
|
||||
return SageAttentionImpl
|
||||
|
||||
# @staticmethod
|
||||
@@ -39,7 +41,7 @@ class SageAttentionImpl(AttentionImpl):
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
@@ -14,7 +16,7 @@ class SDPABackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
@@ -22,7 +24,7 @@ class SDPABackend(AttentionBackend):
|
||||
return "SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SDPAImpl"]:
|
||||
def get_impl_cls() -> Type["SDPAImpl"]:
|
||||
return SDPAImpl
|
||||
|
||||
# @staticmethod
|
||||
@@ -38,7 +40,7 @@ class SDPAImpl(AttentionImpl):
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -19,7 +20,7 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(mask_strategy) -> list[list[list[torch.Tensor | None]]]:
|
||||
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
@@ -57,7 +58,7 @@ class SlidingTileAttentionBackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
# TODO(will-refactor): check this
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@@ -66,15 +67,15 @@ class SlidingTileAttentionBackend(AttentionBackend):
|
||||
return "SLIDING_TILE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SlidingTileAttentionImpl"]:
|
||||
def get_impl_cls() -> Type["SlidingTileAttentionImpl"]:
|
||||
return SlidingTileAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["SlidingTileAttentionMetadata"]:
|
||||
def get_metadata_cls() -> Type["SlidingTileAttentionMetadata"]:
|
||||
return SlidingTileAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["SlidingTileAttentionMetadataBuilder"]:
|
||||
def get_builder_cls() -> Type["SlidingTileAttentionMetadataBuilder"]:
|
||||
return SlidingTileAttentionMetadataBuilder
|
||||
|
||||
|
||||
@@ -109,7 +110,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -20,11 +22,11 @@ class DistributedAttention(nn.Module):
|
||||
def __init__(self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
num_kv_heads: int | None = None,
|
||||
softmax_scale: float | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[_Backend, ...]
|
||||
| None = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
@@ -60,10 +62,10 @@ class DistributedAttention(nn.Module):
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
replicated_q: torch.Tensor | None = None,
|
||||
replicated_k: torch.Tensor | None = None,
|
||||
replicated_v: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
replicated_q: Optional[torch.Tensor] = None,
|
||||
replicated_k: Optional[torch.Tensor] = None,
|
||||
replicated_v: Optional[torch.Tensor] = None,
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
Args:
|
||||
@@ -139,11 +141,11 @@ class LocalAttention(nn.Module):
|
||||
def __init__(self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
num_kv_heads: int | None = None,
|
||||
softmax_scale: float | None = None,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[_Backend, ...]
|
||||
| None = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
|
||||
@@ -2,10 +2,9 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py
|
||||
|
||||
import os
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from functools import cache
|
||||
from typing import cast
|
||||
from typing import Generator, Optional, Tuple, Type, cast
|
||||
|
||||
import torch
|
||||
|
||||
@@ -18,7 +17,7 @@ from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> _Backend | None:
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
@@ -32,7 +31,7 @@ def backend_name_to_enum(backend_name: str) -> _Backend | None:
|
||||
None
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> _Backend | None:
|
||||
def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
backend environment variable, if one is specified.
|
||||
@@ -54,10 +53,10 @@ def get_env_variable_attn_backend() -> _Backend | None:
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: _Backend | None = None
|
||||
forced_attn_backend: Optional[_Backend] = None
|
||||
|
||||
|
||||
def global_force_attn_backend(attn_backend: _Backend | None) -> None:
|
||||
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
@@ -72,7 +71,7 @@ def global_force_attn_backend(attn_backend: _Backend | None) -> None:
|
||||
forced_attn_backend = attn_backend
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> _Backend | None:
|
||||
def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
'''
|
||||
Get the currently-forced choice of attention backend,
|
||||
or None if auto-selection is currently enabled.
|
||||
@@ -83,8 +82,8 @@ def get_global_forced_attn_backend() -> _Backend | None:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[_Backend, ...] | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
|
||||
@@ -93,8 +92,8 @@ def get_attn_backend(
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[_Backend, ...] | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
#
|
||||
@@ -103,13 +102,13 @@ def _cached_get_attn_backend(
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: _Backend | None = (
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
else:
|
||||
# Check the environment variable and override if specified
|
||||
backend_by_env_var: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
|
||||
backend_by_env_var: Optional[str] = envs.FASTVIDEO_ATTENTION_BACKEND
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
@@ -121,7 +120,7 @@ def _cached_get_attn_backend(
|
||||
if not attention_cls:
|
||||
raise ValueError(
|
||||
f"Invalid attention backend for {current_platform.device_name}")
|
||||
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
|
||||
return cast(Type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
|
||||
|
||||
|
||||
@contextmanager
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any
|
||||
from typing import Any, Dict
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -41,7 +41,7 @@ class ModelConfig:
|
||||
self.__dict__.update(state)
|
||||
|
||||
# This should be used only when loading from transformers/diffusers
|
||||
def update_model_arch(self, source_model_dict: dict[str, Any]) -> None:
|
||||
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None:
|
||||
arch_config = self.arch_config
|
||||
valid_fields = {f.name for f in fields(arch_config)}
|
||||
|
||||
@@ -55,7 +55,7 @@ class ModelConfig:
|
||||
if hasattr(arch_config, "__post_init__"):
|
||||
arch_config.__post_init__()
|
||||
|
||||
def update_model_config(self, source_model_dict: dict[str, Any]) -> None:
|
||||
def update_model_config(self, source_model_dict: Dict[str, Any]) -> None:
|
||||
assert "arch_config" not in source_model_dict, "Source model config shouldn't contain arch_config."
|
||||
|
||||
valid_fields = {f.name for f in fields(self)}
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
@@ -11,7 +11,7 @@ class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: tuple[_Backend,
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
@@ -32,7 +32,7 @@ class DiTConfig(ModelConfig):
|
||||
|
||||
# FastVideoDiT-specific parameters
|
||||
prefix: str = ""
|
||||
quant_config: QuantizationConfig | None = None
|
||||
quant_config: Optional[QuantizationConfig] = None
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -155,9 +156,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
num_layers: int = 20
|
||||
num_single_layers: int = 40
|
||||
num_refiner_layers: int = 2
|
||||
rope_axes_dim: tuple[int, int, int] = (16, 56, 56)
|
||||
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56)
|
||||
guidance_embeds: bool = False
|
||||
dtype: torch.dtype | None = None
|
||||
dtype: Optional[torch.dtype] = None
|
||||
text_embed_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
rope_theta: int = 256
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -39,17 +40,17 @@ class StepVideoArchConfig(DiTArchConfig):
|
||||
num_attention_heads: int = 48
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = 64
|
||||
out_channels: Optional[int] = 64
|
||||
num_layers: int = 48
|
||||
dropout: float = 0.0
|
||||
patch_size: int = 1
|
||||
norm_type: str = "ada_norm_single"
|
||||
norm_elementwise_affine: bool = False
|
||||
norm_eps: float = 1e-6
|
||||
caption_channels: int | list[int] | tuple[int, ...] | None = field(
|
||||
caption_channels: Optional[Union[int, List[int], Tuple[int, ...]]] = field(
|
||||
default_factory=lambda: [6144, 1024])
|
||||
attention_type: str | None = "torch"
|
||||
use_additional_conditions: bool | None = False
|
||||
attention_type: Optional[str] = "torch"
|
||||
use_additional_conditions: Optional[bool] = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -51,7 +52,7 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2)
|
||||
text_len = 512
|
||||
num_attention_heads: int = 40
|
||||
attention_head_dim: int = 128
|
||||
@@ -64,8 +65,8 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
cross_attn_norm: bool = True
|
||||
qk_norm: str = "rms_norm_across_heads"
|
||||
eps: float = 1e-6
|
||||
image_dim: int | None = None
|
||||
added_kv_proj_dim: int | None = None
|
||||
image_dim: Optional[int] = None
|
||||
added_kv_proj_dim: Optional[int] = None
|
||||
rope_max_seq_len: int = 1024
|
||||
|
||||
def __post_init__(self):
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -10,8 +10,8 @@ from fastvideo.v1.platforms import _Backend
|
||||
|
||||
@dataclass
|
||||
class EncoderArchConfig(ArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
architectures: List[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
output_hidden_states: bool = False
|
||||
use_return_dict: bool = True
|
||||
@@ -32,7 +32,7 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
scalable_attention: bool = True
|
||||
tie_word_embeddings: bool = False
|
||||
|
||||
tokenizer_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
@@ -49,11 +49,11 @@ class ImageEncoderArchConfig(EncoderArchConfig):
|
||||
|
||||
@dataclass
|
||||
class BaseEncoderOutput:
|
||||
last_hidden_state: torch.FloatTensor | None = None
|
||||
pooler_output: torch.FloatTensor | None = None
|
||||
hidden_states: tuple[torch.FloatTensor, ...] | None = None
|
||||
attentions: tuple[torch.FloatTensor, ...] | None = None
|
||||
attention_mask: torch.Tensor | None = None
|
||||
last_hidden_state: Optional[torch.FloatTensor] = None
|
||||
pooler_output: Optional[torch.FloatTensor] = None
|
||||
hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
|
||||
attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
|
||||
attention_mask: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -61,8 +61,8 @@ class EncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
|
||||
|
||||
prefix: str = ""
|
||||
quant_config: QuantizationConfig | None = None
|
||||
lora_config: Any | None = None
|
||||
quant_config: Optional[QuantizationConfig] = None
|
||||
lora_config: Optional[Any] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
@@ -50,8 +51,8 @@ class CLIPTextConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=CLIPTextArchConfig)
|
||||
|
||||
num_hidden_layers_override: int | None = None
|
||||
require_post_norm: bool | None = None
|
||||
num_hidden_layers_override: Optional[int] = None
|
||||
require_post_norm: Optional[bool] = None
|
||||
prefix: str = "clip"
|
||||
|
||||
|
||||
@@ -60,6 +61,6 @@ class CLIPVisionConfig(ImageEncoderConfig):
|
||||
arch_config: ImageEncoderArchConfig = field(
|
||||
default_factory=CLIPVisionArchConfig)
|
||||
|
||||
num_hidden_layers_override: int | None = None
|
||||
require_post_norm: bool | None = None
|
||||
num_hidden_layers_override: Optional[int] = None
|
||||
require_post_norm: Optional[bool] = None
|
||||
prefix: str = "clip"
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
@@ -11,7 +12,7 @@ class LlamaArchConfig(TextEncoderArchConfig):
|
||||
intermediate_size: int = 11008
|
||||
num_hidden_layers: int = 32
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int | None = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
hidden_act: str = "silu"
|
||||
max_position_embeddings: int = 2048
|
||||
initializer_range: float = 0.02
|
||||
@@ -23,11 +24,11 @@ class LlamaArchConfig(TextEncoderArchConfig):
|
||||
pretraining_tp: int = 1
|
||||
tie_word_embeddings: bool = False
|
||||
rope_theta: float = 10000.0
|
||||
rope_scaling: float | None = None
|
||||
rope_scaling: Optional[float] = None
|
||||
attention_bias: bool = False
|
||||
attention_dropout: float = 0.0
|
||||
mlp_bias: bool = False
|
||||
head_dim: int | None = None
|
||||
head_dim: Optional[int] = None
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
@@ -11,7 +12,7 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
d_kv: int = 64
|
||||
d_ff: int = 2048
|
||||
num_layers: int = 6
|
||||
num_decoder_layers: int | None = None
|
||||
num_decoder_layers: Optional[int] = None
|
||||
num_heads: int = 8
|
||||
relative_attention_num_buckets: int = 32
|
||||
relative_attention_max_distance: int = 128
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import Any, Union
|
||||
|
||||
import torch
|
||||
|
||||
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import StoreBoolean
|
||||
|
||||
@dataclass
|
||||
class VAEArchConfig(ArchConfig):
|
||||
scaling_factor: float | torch.Tensor = 0
|
||||
scaling_factor: Union[float, torch.tensor] = 0
|
||||
|
||||
temporal_compression_ratio: int = 4
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
@@ -8,19 +9,19 @@ class HunyuanVAEArchConfig(VAEArchConfig):
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 16
|
||||
down_block_types: tuple[str, ...] = (
|
||||
down_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
)
|
||||
up_block_types: tuple[str, ...] = (
|
||||
up_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
)
|
||||
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
|
||||
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
|
||||
layers_per_block: int = 2
|
||||
act_fn: str = "silu"
|
||||
norm_num_groups: int = 32
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -9,12 +10,12 @@ from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
class WanVAEArchConfig(VAEArchConfig):
|
||||
base_dim: int = 96
|
||||
z_dim: int = 16
|
||||
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
|
||||
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: tuple[float, ...] = ()
|
||||
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||
attn_scales: Tuple[float, ...] = ()
|
||||
temperal_downsample: Tuple[bool, ...] = (False, True, True)
|
||||
dropout: float = 0.0
|
||||
latents_mean: tuple[float, ...] = (
|
||||
latents_mean: Tuple[float, ...] = (
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
@@ -32,7 +33,7 @@ class WanVAEArchConfig(VAEArchConfig):
|
||||
0.2503,
|
||||
-0.2921,
|
||||
)
|
||||
latents_std: tuple[float, ...] = (
|
||||
latents_std: Tuple[float, ...] = (
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
@@ -54,9 +55,9 @@ class WanVAEArchConfig(VAEArchConfig):
|
||||
spatial_compression_ratio = 8
|
||||
|
||||
def __post_init__(self):
|
||||
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(
|
||||
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(
|
||||
self.latents_std).view(1, self.z_dim, 1, 1, 1)
|
||||
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(
|
||||
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from typing import Any, cast
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
|
||||
@@ -18,7 +17,7 @@ def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def postprocess_text(output: BaseEncoderOutput) -> torch.Tensor:
|
||||
def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -27,7 +26,7 @@ class PipelineConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
@@ -44,18 +43,18 @@ class PipelineConfig:
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(postprocess_text, ))
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: str | None = None
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
@@ -108,7 +107,7 @@ class PipelineConfig:
|
||||
input_pipeline_dict = json.load(f)
|
||||
self.update_pipeline_config(input_pipeline_dict)
|
||||
|
||||
def update_pipeline_config(self, source_pipeline_dict: dict[str,
|
||||
def update_pipeline_config(self, source_pipeline_dict: Dict[str,
|
||||
Any]) -> None:
|
||||
for f in fields(self):
|
||||
key = f.name
|
||||
@@ -124,9 +123,8 @@ class PipelineConfig:
|
||||
assert len(current_value) == len(
|
||||
new_value
|
||||
), "Users shouldn't delete or add text encoder config objects in your json"
|
||||
for target_config, source_config in zip(current_value,
|
||||
new_value,
|
||||
strict=False):
|
||||
for target_config, source_config in zip(
|
||||
current_value, new_value):
|
||||
target_config.update_model_config(source_config)
|
||||
else:
|
||||
setattr(self, key, new_value)
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TypedDict
|
||||
from typing import Callable, Tuple, TypedDict
|
||||
|
||||
import torch
|
||||
|
||||
@@ -36,11 +35,11 @@ def llama_preprocess_text(prompt: str) -> str:
|
||||
return prompt_template_video["template"].format(prompt)
|
||||
|
||||
|
||||
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
hidden_state_skip_layer = 2
|
||||
assert outputs.hidden_states is not None
|
||||
hidden_states: tuple[torch.Tensor, ...] = outputs.hidden_states
|
||||
last_hidden_state: torch.Tensor = hidden_states[-(hidden_state_skip_layer +
|
||||
hidden_states: Tuple[torch.Tensor, ...] = outputs.hidden_states
|
||||
last_hidden_state: torch.tensor = hidden_states[-(hidden_state_skip_layer +
|
||||
1)]
|
||||
crop_start = prompt_template_video.get("crop_start", -1)
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
@@ -51,8 +50,8 @@ def clip_preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
pooler_output: torch.Tensor = outputs.pooler_output
|
||||
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
pooler_output: torch.tensor = outputs.pooler_output
|
||||
return pooler_output
|
||||
|
||||
|
||||
@@ -73,19 +72,19 @@ class HunyuanConfig(PipelineConfig):
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
|
||||
postprocess_text_funcs: tuple[
|
||||
Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
postprocess_text_funcs: Tuple[
|
||||
Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(llama_postprocess_text, clip_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", "fp16"))
|
||||
|
||||
def __post_init__(self):
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Registry for pipeline weight-specific configurations."""
|
||||
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from typing import Callable, Dict, Optional, Type
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
@@ -18,7 +18,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: dict[str, type[PipelineConfig]] = {
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
@@ -30,7 +30,7 @@ WEIGHT_CONFIG_REGISTRY: dict[str, type[PipelineConfig]] = {
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
@@ -39,7 +39,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
@@ -51,7 +51,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
|
||||
|
||||
def get_pipeline_config_cls_for_name(
|
||||
pipeline_name_or_path: str) -> type[PipelineConfig] | None:
|
||||
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -11,15 +11,13 @@ from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
mask: torch.tensor = outputs.attention_mask
|
||||
hidden_state: torch.tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
assert torch.isnan(hidden_state).sum() == 0
|
||||
prompt_embeds = [
|
||||
u[:v] for u, v in zip(hidden_state, seq_lens, strict=False)
|
||||
]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack([
|
||||
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens)]
|
||||
prompt_embeds_tensor: torch.tensor = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
@@ -46,16 +44,16 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
flow_shift: int = 3
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (T5Config(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(t5_postprocess_text, ))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp32", ))
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -15,12 +15,12 @@ class SamplingParam:
|
||||
data_type: str = "video"
|
||||
|
||||
# Image inputs
|
||||
image_path: str | None = None
|
||||
image_path: Optional[str] = None
|
||||
|
||||
# Text inputs
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str | None = None
|
||||
prompt_path: str | None = None
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
negative_prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
|
||||
# Batch info
|
||||
@@ -53,7 +53,7 @@ class SamplingParam:
|
||||
if self.prompt_path and not self.prompt_path.endswith(".txt"):
|
||||
raise ValueError("prompt_path must be a txt file")
|
||||
|
||||
def update(self, source_dict: dict[str, Any]) -> None:
|
||||
def update(self, source_dict: Dict[str, Any]) -> None:
|
||||
for key, value in source_dict.items():
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
@@ -15,7 +14,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
|
||||
@@ -27,7 +26,7 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
SAMPLING_PARAM_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
@@ -36,7 +35,7 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
|
||||
"hunyuan":
|
||||
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
@@ -47,7 +46,8 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
}
|
||||
|
||||
|
||||
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
|
||||
def get_sampling_param_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[Any]:
|
||||
"""Get the appropriate sampling param for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
|
||||
@@ -2,11 +2,12 @@ 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
|
||||
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
|
||||
|
||||
def getdataset(args):
|
||||
def getdataset(args, start_idx=0):
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
@@ -25,15 +26,15 @@ def getdataset(args):
|
||||
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)
|
||||
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,
|
||||
)
|
||||
return T2V_dataset(args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
start_idx=start_idx)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
@@ -44,7 +45,7 @@ if __name__ == "__main__":
|
||||
from accelerate import Accelerator
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import dataset_prog
|
||||
from fastvideo.v1.dataset.t2v_datasets import dataset_prog
|
||||
|
||||
args = type(
|
||||
"args",
|
||||
@@ -63,7 +64,8 @@ if __name__ == "__main__":
|
||||
"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",
|
||||
"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,
|
||||
@@ -80,7 +82,10 @@ if __name__ == "__main__":
|
||||
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]
|
||||
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:
|
||||
@@ -0,0 +1,44 @@
|
||||
# schema.py
|
||||
"""
|
||||
Unified data schema and format for saving and loading image/video data after
|
||||
preprocessing.
|
||||
|
||||
It uses apache arrow in-memory format that can be consumed by modern data
|
||||
frameworks that can handle parquet or lance file.
|
||||
"""
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
pyarrow_schema = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
@@ -20,9 +20,11 @@ class LatentDataset(Dataset):
|
||||
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.prompt_embed_dir = os.path.join(self.datase_dir_path,
|
||||
"prompt_embed")
|
||||
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
|
||||
"prompt_attention_mask")
|
||||
with open(self.json_path) as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
|
||||
@@ -31,12 +33,16 @@ class LatentDataset(Dataset):
|
||||
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]
|
||||
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"]
|
||||
prompt_attention_mask_file = self.data_anno[idx][
|
||||
"prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(
|
||||
os.path.join(self.latent_dir, latent_file),
|
||||
@@ -54,7 +60,8 @@ class LatentDataset(Dataset):
|
||||
weights_only=True,
|
||||
)
|
||||
prompt_attention_mask = torch.load(
|
||||
os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file),
|
||||
os.path.join(self.prompt_attention_mask_dir,
|
||||
prompt_attention_mask_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
@@ -104,8 +111,12 @@ def latent_collate_function(batch):
|
||||
|
||||
|
||||
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)
|
||||
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,
|
||||
@@ -0,0 +1,371 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections import defaultdict
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from einops import rearrange
|
||||
from torch import distributed as dist
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sp_group)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
# Path to your dataset
|
||||
dataset_path = "/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/train/"
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ParquetVideoTextDataset(Dataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
def __init__(self,
|
||||
path: str,
|
||||
batch_size: int = 1024,
|
||||
rank: int = 0,
|
||||
world_size: int = 1,
|
||||
cfg_rate: float = 0.0,
|
||||
num_latent_t: int = 2,
|
||||
seed: int = 0):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.rank = rank
|
||||
self.local_rank = get_sequence_model_parallel_rank()
|
||||
self.sp_world_size = world_size
|
||||
self.world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
self.cfg_rate = cfg_rate
|
||||
self.num_latent_t = num_latent_t
|
||||
self.local_indices = None
|
||||
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
|
||||
|
||||
ranks = get_sp_group().ranks
|
||||
group_ranks = [None for _ in range(self.world_size)]
|
||||
torch.distributed.all_gather_object(group_ranks, ranks)
|
||||
|
||||
if rank == 0:
|
||||
# If a plan already exists, then skip creating a new plan
|
||||
# This will be useful when resume training
|
||||
if os.path.exists(self.plan_output_dir):
|
||||
print(f"Using existing plan from {self.plan_output_dir}")
|
||||
return
|
||||
|
||||
# Find all parquet files recursively, and record num_rows for each file
|
||||
print(f"Scanning for parquet files in {self.path}")
|
||||
metadatas = []
|
||||
for root, _, files in os.walk(self.path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
for row_idx in range(num_rows):
|
||||
metadatas.append((file_path, row_idx))
|
||||
|
||||
# Generate the plan that distribute rows among workers
|
||||
random.seed(seed)
|
||||
random.shuffle(metadatas)
|
||||
|
||||
# Get all sp groups
|
||||
# e.g. if num_gpus = 4, sp_size = 2
|
||||
# group_ranks = [(0, 1), (2, 3)]
|
||||
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
|
||||
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
|
||||
group_ranks = list(set(tuple(r) for r in group_ranks))
|
||||
num_sp_groups = len(group_ranks)
|
||||
plan = defaultdict(list)
|
||||
for idx, metadata in enumerate(metadatas):
|
||||
sp_group_idx = idx % num_sp_groups
|
||||
for global_rank in group_ranks[sp_group_idx]:
|
||||
plan[global_rank].append(metadata)
|
||||
|
||||
with open(self.plan_output_dir, "w") as f:
|
||||
json.dump(plan, f)
|
||||
|
||||
def __len__(self):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[str(self.rank)]
|
||||
except:
|
||||
raise Exception("The data plan hasn't been created yet")
|
||||
return len(self.local_indices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[self.rank]
|
||||
except:
|
||||
raise Exception("The data plan hasn't been created yet")
|
||||
file_path, row_idx = self.local_indices[idx]
|
||||
parquet_file = pq.ParquetFile(file_path)
|
||||
|
||||
# Calculate the row group to read into memory and the local idx
|
||||
# This way we can avoid reading in the entire parquet file
|
||||
cumulative = 0
|
||||
for i in range(parquet_file.num_row_groups):
|
||||
num_rows = parquet_file.metadata.row_group(i).num_rows
|
||||
if cumulative + num_rows > idx:
|
||||
row_group_index = i
|
||||
local_index = idx - cumulative
|
||||
break
|
||||
cumulative += num_rows
|
||||
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
processed = self._process_row(row_dict)
|
||||
lat, emb, mask, info = processed["latents"], processed[
|
||||
"embeddings"], processed["masks"], processed["info"]
|
||||
if lat.numel() == 0: # Validation parquet
|
||||
return lat, emb, mask, info
|
||||
else:
|
||||
lat = lat[:, -self.num_latent_t:]
|
||||
if self.sp_world_size > 1:
|
||||
lat = rearrange(lat,
|
||||
"t (n s) h w -> t n s h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
lat = lat[:, self.local_rank, :, :, :]
|
||||
return lat, emb, mask, info
|
||||
|
||||
def _process_row(self, row):
|
||||
"""Process a PyArrow batch into tensors."""
|
||||
out = {"lat": None, "emb": None, "msk": None, "info": None}
|
||||
|
||||
vae_latent_bytes = row["vae_latent_bytes"]
|
||||
vae_latent_shape = row["vae_latent_shape"]
|
||||
text_embedding_bytes = row["text_embedding_bytes"]
|
||||
text_embedding_shape = row["text_embedding_shape"]
|
||||
text_attention_mask_bytes = row["text_attention_mask_bytes"]
|
||||
text_attention_mask_shape = row["text_attention_mask_shape"]
|
||||
|
||||
# Process latent
|
||||
if not vae_latent_shape: # No VAE latent is stored. Split is validation
|
||||
lat = np.array([])
|
||||
else:
|
||||
lat = np.frombuffer(vae_latent_bytes,
|
||||
dtype=np.float32).reshape(vae_latent_shape)
|
||||
# Make array writable
|
||||
lat = np.copy(lat)
|
||||
|
||||
if random.random() < self.cfg_rate:
|
||||
emb = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
emb = np.frombuffer(text_embedding_bytes,
|
||||
dtype=np.float32).reshape(text_embedding_shape)
|
||||
# Make array writable
|
||||
emb = np.copy(emb)
|
||||
if emb.shape[0] < 512:
|
||||
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
|
||||
padded_emb[:emb.shape[0], :] = emb
|
||||
emb = padded_emb
|
||||
elif emb.shape[0] > 512:
|
||||
emb = emb[:512, :]
|
||||
|
||||
# Process mask
|
||||
if len(text_attention_mask_bytes) > 0 and len(
|
||||
text_attention_mask_shape) > 0:
|
||||
msk = np.frombuffer(text_attention_mask_bytes,
|
||||
dtype=np.uint8).astype(np.bool_)
|
||||
msk = msk.reshape(1, -1)
|
||||
# Make array writable
|
||||
msk = np.copy(msk)
|
||||
if msk.shape[1] < 512:
|
||||
padded_msk = np.zeros((1, 512), dtype=np.bool_)
|
||||
padded_msk[:, :msk.shape[1]] = msk
|
||||
msk = padded_msk
|
||||
elif msk.shape[1] > 512:
|
||||
msk = msk[:, :512]
|
||||
else:
|
||||
msk = np.ones((1, 512), dtype=np.bool_)
|
||||
|
||||
# Collect metadata
|
||||
info = {
|
||||
"width": row["width"],
|
||||
"height": row["height"],
|
||||
"num_frames": row["num_frames"],
|
||||
"duration_sec": row["duration_sec"],
|
||||
"fps": row["fps"],
|
||||
"file_name": row["file_name"],
|
||||
"caption": row["caption"],
|
||||
}
|
||||
|
||||
out["lat"] = torch.from_numpy(lat)
|
||||
out["emb"] = torch.from_numpy(emb)
|
||||
out["msk"] = torch.from_numpy(msk)
|
||||
out["info"] = info
|
||||
|
||||
return {
|
||||
"latents": out["lat"],
|
||||
"embeddings": out["emb"],
|
||||
"masks": out["msk"],
|
||||
"info": out["info"]
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Benchmark Parquet dataset loading speed')
|
||||
parser.add_argument('--path',
|
||||
type=str,
|
||||
default=dataset_path,
|
||||
help='Path to Parquet dataset')
|
||||
parser.add_argument('--batch_size',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Batch size for DataLoader')
|
||||
parser.add_argument('--num_batches',
|
||||
type=int,
|
||||
default=100,
|
||||
help='Number of batches to benchmark')
|
||||
parser.add_argument('--vae_debug', action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
# Initialize distributed training
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
# Initialize CUDA device first
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
|
||||
# Initialize distributed training
|
||||
if world_size > 1:
|
||||
dist.init_process_group(backend="nccl",
|
||||
init_method="env://",
|
||||
world_size=world_size,
|
||||
rank=rank)
|
||||
print(
|
||||
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
|
||||
)
|
||||
|
||||
# Create dataset
|
||||
dataset = ParquetVideoTextDataset(
|
||||
args.path,
|
||||
batch_size=args.batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataloader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_size=args.batch_size,
|
||||
num_workers=1, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
drop_last=True)
|
||||
|
||||
# Example of how to load dataloader state
|
||||
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
|
||||
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
|
||||
# dataloader.load_state_dict(dataloader_state[rank])
|
||||
|
||||
# Warm-up with synchronization
|
||||
if rank == 0:
|
||||
print("Warming up...")
|
||||
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
|
||||
# Example of how to save dataloader state
|
||||
# if i == 30:
|
||||
# dist.barrier()
|
||||
# local_data = {rank: dataloader.state_dict()}
|
||||
# gathered_data = [None] * world_size
|
||||
# dist.all_gather_object(gathered_data, local_data)
|
||||
# if rank == 0:
|
||||
# global_state_dict = {}
|
||||
# for d in gathered_data:
|
||||
# global_state_dict.update(d)
|
||||
# torch.save(global_state_dict, "dataloader_state.pt")
|
||||
assert torch.sum(masks[0]).item() == torch.count_nonzero(
|
||||
embeddings[0]).item() // 4096
|
||||
if args.vae_debug:
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=VAE_PATH,
|
||||
vae_config=WanVAEConfig(load_encoder=False),
|
||||
vae_precision="fp32")
|
||||
fastvideo_args.device = device
|
||||
vae_loader = VAELoader()
|
||||
vae = vae_loader.load(model_path=VAE_PATH,
|
||||
architecture="",
|
||||
fastvideo_args=fastvideo_args)
|
||||
|
||||
videoprocessor = VideoProcessor(vae_scale_factor=8)
|
||||
|
||||
with torch.inference_mode():
|
||||
video = vae.decode(latents[0].unsqueeze(0).to(device))
|
||||
video = videoprocessor.postprocess_video(video)
|
||||
video_path = os.path.join("/workspace/FastVideo/debug_videos",
|
||||
infos["caption"][0][:50] + ".mp4")
|
||||
export_to_video(video[0], video_path, fps=16)
|
||||
|
||||
# Move data to device
|
||||
# latents = latents.to(device)
|
||||
# embeddings = embeddings.to(device)
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
# Benchmark
|
||||
if rank == 0:
|
||||
print(f"Benchmarking with batch_size={args.batch_size}")
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
for i, (latents, embeddings, masks,
|
||||
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
|
||||
if i >= args.num_batches:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(device)
|
||||
embeddings = embeddings.to(device)
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
total_samples += batch_size
|
||||
|
||||
# Print progress only from rank 0
|
||||
if rank == 0 and (i + 1) % 10 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
print(
|
||||
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
|
||||
)
|
||||
|
||||
# Final statistics
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if rank == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
|
||||
print("\nBenchmark Results:")
|
||||
print(f"Total time: {elapsed:.2f} seconds")
|
||||
print(f"Total samples: {total_samples}")
|
||||
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
|
||||
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
|
||||
|
||||
if world_size > 1:
|
||||
dist.destroy_process_group()
|
||||
@@ -46,7 +46,8 @@ class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
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)))
|
||||
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]
|
||||
@@ -57,7 +58,9 @@ class DataSetProg(metaclass=SingletonMeta):
|
||||
else:
|
||||
worker_id = work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_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
|
||||
|
||||
@@ -65,7 +68,10 @@ 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):
|
||||
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
|
||||
@@ -73,7 +79,14 @@ def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16)
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0):
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
@@ -102,7 +115,8 @@ class T2V_dataset(Dataset):
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
@@ -129,7 +143,8 @@ class T2V_dataset(Dataset):
|
||||
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")
|
||||
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")
|
||||
@@ -160,16 +175,17 @@ class T2V_dataset(Dataset):
|
||||
)
|
||||
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,
|
||||
)
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx):
|
||||
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
|
||||
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]
|
||||
@@ -178,13 +194,15 @@ class T2V_dataset(Dataset):
|
||||
# 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)
|
||||
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 = (image_data["cap"]
|
||||
if isinstance(image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
@@ -238,10 +256,12 @@ class T2V_dataset(Dataset):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None or resolution.get("width", None) is None):
|
||||
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"]
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
@@ -259,29 +279,34 @@ class T2V_dataset(Dataset):
|
||||
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)
|
||||
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)
|
||||
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):
|
||||
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))
|
||||
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)
|
||||
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
|
||||
@@ -290,13 +315,15 @@ class T2V_dataset(Dataset):
|
||||
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")
|
||||
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)}")
|
||||
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):
|
||||
@@ -308,11 +335,14 @@ class T2V_dataset(Dataset):
|
||||
|
||||
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]
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno, "r") as f:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
@@ -320,5 +350,5 @@ class T2V_dataset(Dataset):
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self):
|
||||
cap_lists = self.read_jsons(self.data)
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -21,15 +21,19 @@ def center_crop_arr(pil_image, image_size):
|
||||
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)
|
||||
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)
|
||||
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])
|
||||
return Image.fromarray(arr[crop_y:crop_y + image_size,
|
||||
crop_x:crop_x + image_size])
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w):
|
||||
@@ -44,7 +48,9 @@ def crop(clip, i, j, h, 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}")
|
||||
raise ValueError(
|
||||
f"target size should be tuple (height, width), instead got {target_size}"
|
||||
)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
size=target_size,
|
||||
@@ -56,7 +62,9 @@ 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}")
|
||||
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(
|
||||
@@ -166,7 +174,8 @@ 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("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
|
||||
|
||||
@@ -227,7 +236,9 @@ class RandomCropVideo:
|
||||
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)}")
|
||||
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
|
||||
@@ -301,7 +312,9 @@ class LongSideResizeVideo:
|
||||
else:
|
||||
h = int(h * self.size / w)
|
||||
w = self.size
|
||||
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
|
||||
resize_clip = resize(clip,
|
||||
target_size=(h, w),
|
||||
interpolation_mode=self.interpolation_mode)
|
||||
return resize_clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
@@ -321,7 +334,8 @@ class CenterCropResizeVideo:
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
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
|
||||
@@ -335,7 +349,10 @@ class CenterCropResizeVideo:
|
||||
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)
|
||||
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,
|
||||
@@ -361,7 +378,8 @@ class UCFCenterCropVideo:
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
@@ -376,7 +394,9 @@ class UCFCenterCropVideo:
|
||||
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_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
|
||||
|
||||
@@ -396,7 +416,8 @@ class KineticsRandomCropResizeVideo:
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
@@ -405,7 +426,8 @@ class KineticsRandomCropResizeVideo:
|
||||
|
||||
def __call__(self, clip):
|
||||
clip_random_crop = random_shift_crop(clip)
|
||||
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
|
||||
clip_resize = resize(clip_random_crop, self.size,
|
||||
self.interpolation_mode)
|
||||
return clip_resize
|
||||
|
||||
|
||||
@@ -418,7 +440,8 @@ class CenterCropVideo:
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
@@ -514,7 +537,7 @@ class RandomHorizontalFlipVideo:
|
||||
# ------------------------------------------------------------
|
||||
# --------------------- Sampling ---------------------------
|
||||
# ------------------------------------------------------------
|
||||
class TemporalRandomCrop(object):
|
||||
class TemporalRandomCrop:
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
@@ -531,7 +554,7 @@ class TemporalRandomCrop(object):
|
||||
return begin_index, end_index
|
||||
|
||||
|
||||
class DynamicSampleDuration(object):
|
||||
class DynamicSampleDuration:
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
@@ -545,7 +568,8 @@ class DynamicSampleDuration(object):
|
||||
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_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
|
||||
@@ -560,14 +584,18 @@ if __name__ == "__main__":
|
||||
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")
|
||||
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),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5],
|
||||
std=[0.5, 0.5, 0.5],
|
||||
inplace=True),
|
||||
])
|
||||
|
||||
target_video_len = 32
|
||||
@@ -582,7 +610,10 @@ if __name__ == "__main__":
|
||||
# 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)
|
||||
frame_indice = np.linspace(start_frame_ind,
|
||||
end_frame_ind - 1,
|
||||
target_video_len,
|
||||
dtype=int)
|
||||
print(frame_indice)
|
||||
|
||||
select_vframes = vframes[frame_indice]
|
||||
@@ -593,11 +624,14 @@ if __name__ == "__main__":
|
||||
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)
|
||||
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)
|
||||
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(
|
||||
@@ -0,0 +1,10 @@
|
||||
from huggingface_hub import HfApi, upload_folder
|
||||
|
||||
api = HfApi()
|
||||
repo_id = "weizhou03/HD-Mixkit-Finetune-Wan" # customize this
|
||||
api.create_repo(repo_id=repo_id, repo_type="dataset")
|
||||
|
||||
upload_folder(repo_id=repo_id,
|
||||
folder_path="/workspace/data/HD-Mixkit-Finetune-Wan",
|
||||
repo_type="dataset",
|
||||
path_in_repo="")
|
||||
@@ -5,7 +5,8 @@ from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size, get_world_group,
|
||||
init_distributed_environment, initialize_model_parallel)
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
|
||||
__all__ = [
|
||||
@@ -17,4 +18,5 @@ __all__ = [
|
||||
"get_tensor_model_parallel_world_size",
|
||||
"cleanup_dist_env_and_memory",
|
||||
"get_world_group",
|
||||
"model_parallel_is_initialized",
|
||||
]
|
||||
|
||||
@@ -1,14 +1,182 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
|
||||
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
from torch import Tensor
|
||||
from torch.distributed import ProcessGroup, ReduceOp
|
||||
|
||||
|
||||
class DistributedAutograd:
|
||||
"""Collection of autograd functions for distributed operations.
|
||||
|
||||
This class provides custom autograd functions for distributed operations like all_reduce,
|
||||
all_gather, and all_to_all. Each operation is implemented as a static inner class with
|
||||
proper forward and backward implementations.
|
||||
"""
|
||||
|
||||
class AllReduce(torch.autograd.Function):
|
||||
"""Differentiable all_reduce operation.
|
||||
|
||||
The gradient of all_reduce is another all_reduce operation since the operation
|
||||
combines values from all ranks equally.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any,
|
||||
group: ProcessGroup,
|
||||
input_: Tensor,
|
||||
op: Optional[dist.ReduceOp] = None) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.op = op
|
||||
output = input_.clone()
|
||||
dist.all_reduce(output, group=group, op=op)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None]:
|
||||
grad_output = grad_output.clone()
|
||||
dist.all_reduce(grad_output, group=ctx.group, op=ctx.op)
|
||||
return None, grad_output, None
|
||||
|
||||
class AllGather(torch.autograd.Function):
|
||||
"""Differentiable all_gather operation.
|
||||
|
||||
The operation gathers tensors from all ranks and concatenates them along a specified dimension.
|
||||
The backward pass uses reduce_scatter to efficiently distribute gradients back to source ranks.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
|
||||
world_size: int, dim: int) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.world_size = world_size
|
||||
ctx.dim = dim
|
||||
ctx.input_shape = input_.shape
|
||||
|
||||
input_size = input_.size()
|
||||
output_size = (input_size[0] * world_size, ) + input_size[1:]
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
|
||||
dist.all_gather_into_tensor(output_tensor, input_, group=group)
|
||||
|
||||
output_tensor = output_tensor.reshape((world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
|
||||
# Split the gradient tensor along the gathered dimension
|
||||
dim_size = grad_output.size(ctx.dim) // ctx.world_size
|
||||
grad_chunks = grad_output.reshape(grad_output.shape[:ctx.dim] +
|
||||
(ctx.world_size, dim_size) +
|
||||
grad_output.shape[ctx.dim + 1:])
|
||||
grad_chunks = grad_chunks.movedim(ctx.dim, 0)
|
||||
|
||||
# Each rank only needs its corresponding gradient
|
||||
grad_input = torch.empty(ctx.input_shape,
|
||||
dtype=grad_output.dtype,
|
||||
device=grad_output.device)
|
||||
dist.reduce_scatter_tensor(grad_input,
|
||||
grad_chunks.contiguous(),
|
||||
group=ctx.group)
|
||||
|
||||
return None, grad_input, None, None
|
||||
|
||||
class AllToAll4D(torch.autograd.Function):
|
||||
"""Differentiable all_to_all operation specialized for 4D tensors.
|
||||
|
||||
This operation is particularly useful for attention operations where we need to
|
||||
redistribute data across ranks for efficient parallel processing.
|
||||
|
||||
The operation supports two modes:
|
||||
1. scatter_dim=2, gather_dim=1: Used for redistributing attention heads
|
||||
2. scatter_dim=1, gather_dim=2: Used for redistributing sequence dimensions
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
|
||||
world_size: int, scatter_dim: int,
|
||||
gather_dim: int) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.world_size = world_size
|
||||
ctx.scatter_dim = scatter_dim
|
||||
ctx.gather_dim = gather_dim
|
||||
|
||||
if world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * world_size
|
||||
shard_hc = hc // world_size
|
||||
|
||||
input_t = input_.reshape(bs, shard_seqlen, world_size, shard_hc,
|
||||
hs).transpose(0, 2).contiguous()
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous()
|
||||
output = output.reshape(bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * world_size
|
||||
shard_seqlen = seqlen // world_size
|
||||
|
||||
input_t = input_.reshape(bs, world_size, shard_seqlen, shard_hc,
|
||||
hs)
|
||||
input_t = input_t.transpose(0, 3).transpose(0, 1).contiguous()
|
||||
input_t = input_t.reshape(world_size, shard_hc, shard_seqlen,
|
||||
bs, hs)
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
|
||||
output = output.reshape(hc, shard_seqlen, bs, hs)
|
||||
output = output.transpose(0, 2).contiguous()
|
||||
output = output.reshape(bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Invalid scatter_dim={scatter_dim}, gather_dim={gather_dim}. "
|
||||
f"Only (scatter_dim=2, gather_dim=1) and (scatter_dim=1, gather_dim=2) are supported."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def backward(
|
||||
ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None, None, None]:
|
||||
if ctx.world_size == 1:
|
||||
return None, grad_output, None, None, None
|
||||
|
||||
# For backward pass, we swap scatter_dim and gather_dim
|
||||
output = DistributedAutograd.AllToAll4D.apply(
|
||||
ctx.group, grad_output, ctx.world_size, ctx.gather_dim,
|
||||
ctx.scatter_dim)
|
||||
return None, output, None, None, None
|
||||
|
||||
|
||||
class DeviceCommunicatorBase:
|
||||
"""
|
||||
Base class for device-specific communicator.
|
||||
Base class for device-specific communicator with autograd support.
|
||||
It can use the `cpu_group` to initialize the communicator.
|
||||
If the device has PyTorch integration (PyTorch can recognize its
|
||||
communication backend), the `device_group` will also be given.
|
||||
@@ -16,8 +184,8 @@ class DeviceCommunicatorBase:
|
||||
|
||||
def __init__(self,
|
||||
cpu_group: ProcessGroup,
|
||||
device: torch.device | None = None,
|
||||
device_group: ProcessGroup | None = None,
|
||||
device: Optional[torch.device] = None,
|
||||
device_group: Optional[ProcessGroup] = None,
|
||||
unique_name: str = ""):
|
||||
self.device = device or torch.device("cpu")
|
||||
self.cpu_group = cpu_group
|
||||
@@ -31,40 +199,33 @@ class DeviceCommunicatorBase:
|
||||
self.rank_in_group = dist.get_group_rank(self.cpu_group,
|
||||
self.global_rank)
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
dist.all_reduce(input_, group=self.device_group)
|
||||
return input_
|
||||
def all_reduce(self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[dist.ReduceOp] = ReduceOp.SUM) -> torch.Tensor:
|
||||
"""Performs an all_reduce operation with gradient support."""
|
||||
return DistributedAutograd.AllReduce.apply(self.device_group, input_,
|
||||
op)
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
"""Performs an all_gather operation with gradient support."""
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
input_size = input_.size()
|
||||
# NOTE: we have to use concat-style all-gather here,
|
||||
# stack-style all-gather has compatibility issues with
|
||||
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
|
||||
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
|
||||
# Allocate output tensor.
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
# All-gather.
|
||||
dist.all_gather_into_tensor(output_tensor,
|
||||
input_,
|
||||
group=self.device_group)
|
||||
# Reshape
|
||||
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(self.world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
return DistributedAutograd.AllGather.apply(self.device_group, input_,
|
||||
self.world_size, dim)
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Performs a 4D all-to-all operation with gradient support."""
|
||||
return DistributedAutograd.AllToAll4D.apply(self.device_group, input_,
|
||||
self.world_size,
|
||||
scatter_dim, gather_dim)
|
||||
|
||||
def gather(self,
|
||||
input_: torch.Tensor,
|
||||
dst: int = 0,
|
||||
dim: int = -1) -> torch.Tensor | None:
|
||||
dim: int = -1) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
NOTE: We assume that the input tensor is on the same device across
|
||||
all the ranks.
|
||||
@@ -93,82 +254,7 @@ class DeviceCommunicatorBase:
|
||||
output_tensor = None
|
||||
return output_tensor
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
|
||||
|
||||
Args:
|
||||
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
|
||||
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
|
||||
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after all-to-all operation.
|
||||
"""
|
||||
# Bypass the function if we are using only 1 GPU.
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * self.world_size
|
||||
shard_hc = hc // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, shard_seqlen, self.world_size,
|
||||
shard_hc, hs).transpose(0,
|
||||
2).contiguous())
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous().reshape(
|
||||
bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * self.world_size
|
||||
shard_seqlen = seqlen // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, self.world_size, shard_seqlen,
|
||||
shard_hc, hs).transpose(0, 3).transpose(
|
||||
0, 1).contiguous().reshape(
|
||||
self.world_size, shard_hc,
|
||||
shard_seqlen, bs, hs))
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(hc, shard_seqlen, bs,
|
||||
hs).transpose(0, 2).contiguous().reshape(
|
||||
bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
if dst is None:
|
||||
@@ -178,7 +264,7 @@ class DeviceCommunicatorBase:
|
||||
def recv(self,
|
||||
size: torch.Size,
|
||||
dtype: torch.dtype,
|
||||
src: int | None = None) -> torch.Tensor:
|
||||
src: Optional[int] = None) -> torch.Tensor:
|
||||
"""Receives a tensor from the source rank."""
|
||||
"""NOTE: `src` is the local rank of the source rank."""
|
||||
if src is None:
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/cuda_communicator.py
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
@@ -12,35 +14,37 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
|
||||
def __init__(self,
|
||||
cpu_group: ProcessGroup,
|
||||
device: torch.device | None = None,
|
||||
device_group: ProcessGroup | None = None,
|
||||
device: Optional[torch.device] = None,
|
||||
device_group: Optional[ProcessGroup] = None,
|
||||
unique_name: str = ""):
|
||||
super().__init__(cpu_group, device, device_group, unique_name)
|
||||
|
||||
from fastvideo.v1.distributed.device_communicators.pynccl import (
|
||||
PyNcclCommunicator)
|
||||
|
||||
self.pynccl_comm: PyNcclCommunicator | None = None
|
||||
self.pynccl_comm: Optional[PyNcclCommunicator] = None
|
||||
if self.world_size > 1:
|
||||
self.pynccl_comm = PyNcclCommunicator(
|
||||
group=self.cpu_group,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def all_reduce(self, input_):
|
||||
def all_reduce(self,
|
||||
input_,
|
||||
op: Optional[torch.distributed.ReduceOp] = None):
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
out = pynccl_comm.all_reduce(input_)
|
||||
out = pynccl_comm.all_reduce(input_, op=op)
|
||||
if out is None:
|
||||
# fall back to the default all-reduce using PyTorch.
|
||||
# this usually happens during testing.
|
||||
# when we run the model, allreduce only happens for the TP
|
||||
# group, where we always have either custom allreduce or pynccl.
|
||||
out = input_.clone()
|
||||
torch.distributed.all_reduce(out, group=self.device_group)
|
||||
torch.distributed.all_reduce(out, group=self.device_group, op=op)
|
||||
return out
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
if dst is None:
|
||||
@@ -55,7 +59,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
def recv(self,
|
||||
size: torch.Size,
|
||||
dtype: torch.dtype,
|
||||
src: int | None = None) -> torch.Tensor:
|
||||
src: Optional[int] = None) -> torch.Tensor:
|
||||
"""Receives a tensor from the source rank."""
|
||||
"""NOTE: `src` is the local rank of the source rank."""
|
||||
if src is None:
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl.py
|
||||
|
||||
from typing import Optional, Union
|
||||
|
||||
# ===================== import region =====================
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
@@ -20,9 +22,9 @@ class PyNcclCommunicator:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
group: ProcessGroup | StatelessProcessGroup,
|
||||
device: int | str | torch.device,
|
||||
library_path: str | None = None,
|
||||
group: Union[ProcessGroup, StatelessProcessGroup],
|
||||
device: Union[int, str, torch.device],
|
||||
library_path: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
|
||||
@@ -27,7 +27,7 @@
|
||||
import ctypes
|
||||
import platform
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch.distributed import ReduceOp
|
||||
@@ -124,7 +124,7 @@ class ncclRedOpTypeEnum:
|
||||
class Function:
|
||||
name: str
|
||||
restype: Any
|
||||
argtypes: list[Any]
|
||||
argtypes: List[Any]
|
||||
|
||||
|
||||
class NCCLLibrary:
|
||||
@@ -212,13 +212,13 @@ class NCCLLibrary:
|
||||
|
||||
# class attribute to store the mapping from the path to the library
|
||||
# to avoid loading the same library multiple times
|
||||
path_to_library_cache: dict[str, Any] = {}
|
||||
path_to_library_cache: Dict[str, Any] = {}
|
||||
|
||||
# class attribute to store the mapping from library path
|
||||
# to the corresponding dictionary
|
||||
path_to_dict_mapping: dict[str, dict[str, Any]] = {}
|
||||
path_to_dict_mapping: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
def __init__(self, so_file: str | None = None):
|
||||
def __init__(self, so_file: Optional[str] = None):
|
||||
|
||||
so_file = so_file or find_nccl_library()
|
||||
|
||||
@@ -240,7 +240,7 @@ class NCCLLibrary:
|
||||
raise e
|
||||
|
||||
if so_file not in NCCLLibrary.path_to_dict_mapping:
|
||||
_funcs: dict[str, Any] = {}
|
||||
_funcs: Dict[str, Any] = {}
|
||||
for func in NCCLLibrary.exported_functions:
|
||||
f = getattr(self.lib, func.name)
|
||||
f.restype = func.restype
|
||||
|
||||
@@ -27,16 +27,15 @@ import gc
|
||||
import pickle
|
||||
import weakref
|
||||
from collections import namedtuple
|
||||
from collections.abc import Callable
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from multiprocessing import shared_memory
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
from torch.distributed import Backend, ProcessGroup
|
||||
from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
@@ -58,15 +57,15 @@ TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
|
||||
|
||||
|
||||
def _split_tensor_dict(
|
||||
tensor_dict: dict[str, torch.Tensor | Any]
|
||||
) -> tuple[list[tuple[str, Any]], list[torch.Tensor]]:
|
||||
tensor_dict: Dict[str, Union[torch.Tensor, Any]]
|
||||
) -> Tuple[List[Tuple[str, Any]], List[torch.Tensor]]:
|
||||
"""Split the tensor dictionary into two parts:
|
||||
1. A list of (key, value) pairs. If the value is a tensor, it is replaced
|
||||
by its metadata.
|
||||
2. A list of tensors.
|
||||
"""
|
||||
metadata_list: list[tuple[str, Any]] = []
|
||||
tensor_list: list[torch.Tensor] = []
|
||||
metadata_list: List[Tuple[str, Any]] = []
|
||||
tensor_list: List[torch.Tensor] = []
|
||||
for key, value in tensor_dict.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
# Note: we cannot use `value.device` here,
|
||||
@@ -82,7 +81,7 @@ def _split_tensor_dict(
|
||||
return metadata_list, tensor_list
|
||||
|
||||
|
||||
_group_name_counter: dict[str, int] = {}
|
||||
_group_name_counter: Dict[str, int] = {}
|
||||
|
||||
|
||||
def _get_unique_name(name: str) -> str:
|
||||
@@ -98,7 +97,7 @@ def _get_unique_name(name: str) -> str:
|
||||
return newname
|
||||
|
||||
|
||||
_groups: dict[str, Callable[[], Optional["GroupCoordinator"]]] = {}
|
||||
_groups: Dict[str, Callable[[], Optional["GroupCoordinator"]]] = {}
|
||||
|
||||
|
||||
def _register_group(group: "GroupCoordinator") -> None:
|
||||
@@ -129,7 +128,7 @@ class GroupCoordinator:
|
||||
|
||||
# available attributes:
|
||||
rank: int # global rank
|
||||
ranks: list[int] # global ranks in the group
|
||||
ranks: List[int] # global ranks in the group
|
||||
world_size: int # size of the group
|
||||
# difference between `local_rank` and `rank_in_group`:
|
||||
# if we have a group of size 4 across two nodes:
|
||||
@@ -144,16 +143,16 @@ class GroupCoordinator:
|
||||
device_group: ProcessGroup # group for device communication
|
||||
use_device_communicator: bool # whether to use device communicator
|
||||
device_communicator: DeviceCommunicatorBase # device communicator
|
||||
mq_broadcaster: Any | None # shared memory broadcaster
|
||||
mq_broadcaster: Optional[Any] # shared memory broadcaster
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
group_ranks: list[list[int]],
|
||||
group_ranks: List[List[int]],
|
||||
local_rank: int,
|
||||
torch_distributed_backend: str | Backend,
|
||||
torch_distributed_backend: Union[str, Backend],
|
||||
use_device_communicator: bool,
|
||||
use_message_queue_broadcaster: bool = False,
|
||||
group_name: str | None = None,
|
||||
group_name: Optional[str] = None,
|
||||
):
|
||||
group_name = group_name or "anonymous"
|
||||
self.unique_name = _get_unique_name(group_name)
|
||||
@@ -244,8 +243,8 @@ class GroupCoordinator:
|
||||
return self.ranks[(rank_in_group - 1) % world_size]
|
||||
|
||||
@contextmanager
|
||||
def graph_capture(self,
|
||||
graph_capture_context: GraphCaptureContext | None = None):
|
||||
def graph_capture(
|
||||
self, graph_capture_context: Optional[GraphCaptureContext] = None):
|
||||
if graph_capture_context is None:
|
||||
stream = torch.cuda.Stream()
|
||||
graph_capture_context = GraphCaptureContext(stream)
|
||||
@@ -261,7 +260,11 @@ class GroupCoordinator:
|
||||
with torch.cuda.stream(stream):
|
||||
yield graph_capture_context
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
def all_reduce(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
User-facing all-reduce function before we actually call the
|
||||
all-reduce operation.
|
||||
@@ -284,10 +287,14 @@ class GroupCoordinator:
|
||||
return torch.ops.vllm.all_reduce(input_,
|
||||
group_name=self.unique_name)
|
||||
else:
|
||||
return self._all_reduce_out_place(input_)
|
||||
return self._all_reduce_out_place(input_, op=op)
|
||||
|
||||
def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
return self.device_communicator.all_reduce(input_)
|
||||
def _all_reduce_out_place(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
return self.device_communicator.all_reduce(input_, op=op)
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
world_size = self.world_size
|
||||
@@ -302,7 +309,7 @@ class GroupCoordinator:
|
||||
def gather(self,
|
||||
input_: torch.Tensor,
|
||||
dst: int = 0,
|
||||
dim: int = -1) -> torch.Tensor | None:
|
||||
dim: int = -1) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
NOTE: We assume that the input tensor is on the same device across
|
||||
all the ranks.
|
||||
@@ -338,7 +345,7 @@ class GroupCoordinator:
|
||||
group=self.device_group)
|
||||
return input_
|
||||
|
||||
def broadcast_object(self, obj: Any | None = None, src: int = 0):
|
||||
def broadcast_object(self, obj: Optional[Any] = None, src: int = 0):
|
||||
"""Broadcast the input object.
|
||||
NOTE: `src` is the local rank of the source rank.
|
||||
"""
|
||||
@@ -363,9 +370,9 @@ class GroupCoordinator:
|
||||
return recv[0]
|
||||
|
||||
def broadcast_object_list(self,
|
||||
obj_list: list[Any],
|
||||
obj_list: List[Any],
|
||||
src: int = 0,
|
||||
group: ProcessGroup | None = None):
|
||||
group: Optional[ProcessGroup] = None):
|
||||
"""Broadcast the input object list.
|
||||
NOTE: `src` is the local rank of the source rank.
|
||||
"""
|
||||
@@ -445,11 +452,11 @@ class GroupCoordinator:
|
||||
|
||||
def broadcast_tensor_dict(
|
||||
self,
|
||||
tensor_dict: dict[str, torch.Tensor | Any] | None = None,
|
||||
tensor_dict: Optional[Dict[str, Union[torch.Tensor, Any]]] = None,
|
||||
src: int = 0,
|
||||
group: ProcessGroup | None = None,
|
||||
metadata_group: ProcessGroup | None = None
|
||||
) -> dict[str, torch.Tensor | Any] | None:
|
||||
group: Optional[ProcessGroup] = None,
|
||||
metadata_group: Optional[ProcessGroup] = None
|
||||
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
|
||||
"""Broadcast the input tensor dictionary.
|
||||
NOTE: `src` is the local rank of the source rank.
|
||||
"""
|
||||
@@ -463,7 +470,7 @@ class GroupCoordinator:
|
||||
|
||||
rank_in_group = self.rank_in_group
|
||||
if rank_in_group == src:
|
||||
metadata_list: list[tuple[Any, Any]] = []
|
||||
metadata_list: List[Tuple[Any, Any]] = []
|
||||
assert isinstance(
|
||||
tensor_dict,
|
||||
dict), (f"Expecting a dictionary, got {type(tensor_dict)}")
|
||||
@@ -530,10 +537,10 @@ class GroupCoordinator:
|
||||
|
||||
def send_tensor_dict(
|
||||
self,
|
||||
tensor_dict: dict[str, torch.Tensor | Any],
|
||||
dst: int | None = None,
|
||||
tensor_dict: Dict[str, Union[torch.Tensor, Any]],
|
||||
dst: Optional[int] = None,
|
||||
all_gather_group: Optional["GroupCoordinator"] = None,
|
||||
) -> dict[str, torch.Tensor | Any] | None:
|
||||
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
|
||||
"""Send the input tensor dictionary.
|
||||
NOTE: `dst` is the local rank of the source rank.
|
||||
"""
|
||||
@@ -553,7 +560,7 @@ class GroupCoordinator:
|
||||
dst = (self.rank_in_group + 1) % self.world_size
|
||||
assert dst < self.world_size, f"Invalid dst rank ({dst})"
|
||||
|
||||
metadata_list: list[tuple[Any, Any]] = []
|
||||
metadata_list: List[Tuple[Any, Any]] = []
|
||||
assert isinstance(
|
||||
tensor_dict,
|
||||
dict), f"Expecting a dictionary, got {type(tensor_dict)}"
|
||||
@@ -584,9 +591,9 @@ class GroupCoordinator:
|
||||
|
||||
def recv_tensor_dict(
|
||||
self,
|
||||
src: int | None = None,
|
||||
src: Optional[int] = None,
|
||||
all_gather_group: Optional["GroupCoordinator"] = None,
|
||||
) -> dict[str, torch.Tensor | Any] | None:
|
||||
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
|
||||
"""Recv the input tensor dictionary.
|
||||
NOTE: `src` is the local rank of the source rank.
|
||||
"""
|
||||
@@ -607,7 +614,7 @@ class GroupCoordinator:
|
||||
assert src < self.world_size, f"Invalid src rank ({src})"
|
||||
|
||||
recv_metadata_list = self.recv_object(src=src)
|
||||
tensor_dict: dict[str, Any] = {}
|
||||
tensor_dict: Dict[str, Any] = {}
|
||||
for key, value in recv_metadata_list:
|
||||
if isinstance(value, TensorMetadata):
|
||||
tensor = torch.empty(value.size,
|
||||
@@ -657,7 +664,7 @@ class GroupCoordinator:
|
||||
"""
|
||||
torch.distributed.barrier(group=self.cpu_group)
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
self.device_communicator.send(tensor, dst)
|
||||
@@ -665,7 +672,7 @@ class GroupCoordinator:
|
||||
def recv(self,
|
||||
size: torch.Size,
|
||||
dtype: torch.dtype,
|
||||
src: int | None = None) -> torch.Tensor:
|
||||
src: Optional[int] = None) -> torch.Tensor:
|
||||
"""Receives a tensor from the source rank."""
|
||||
"""NOTE: `src` is the local rank of the source rank."""
|
||||
return self.device_communicator.recv(size, dtype, src)
|
||||
@@ -683,7 +690,7 @@ class GroupCoordinator:
|
||||
self.mq_broadcaster = None
|
||||
|
||||
|
||||
_WORLD: GroupCoordinator | None = None
|
||||
_WORLD: Optional[GroupCoordinator] = None
|
||||
|
||||
|
||||
def get_world_group() -> GroupCoordinator:
|
||||
@@ -691,7 +698,7 @@ def get_world_group() -> GroupCoordinator:
|
||||
return _WORLD
|
||||
|
||||
|
||||
def init_world_group(ranks: list[int], local_rank: int,
|
||||
def init_world_group(ranks: List[int], local_rank: int,
|
||||
backend: str) -> GroupCoordinator:
|
||||
return GroupCoordinator(
|
||||
group_ranks=[ranks],
|
||||
@@ -703,11 +710,11 @@ def init_world_group(ranks: list[int], local_rank: int,
|
||||
|
||||
|
||||
def init_model_parallel_group(
|
||||
group_ranks: list[list[int]],
|
||||
group_ranks: List[List[int]],
|
||||
local_rank: int,
|
||||
backend: str,
|
||||
use_message_queue_broadcaster: bool = False,
|
||||
group_name: str | None = None,
|
||||
group_name: Optional[str] = None,
|
||||
) -> GroupCoordinator:
|
||||
|
||||
return GroupCoordinator(
|
||||
@@ -720,7 +727,7 @@ def init_model_parallel_group(
|
||||
)
|
||||
|
||||
|
||||
_TP: GroupCoordinator | None = None
|
||||
_TP: Optional[GroupCoordinator] = None
|
||||
|
||||
|
||||
def get_tp_group() -> GroupCoordinator:
|
||||
@@ -779,7 +786,7 @@ def init_distributed_environment(
|
||||
"world group already initialized with a different world size")
|
||||
|
||||
|
||||
_SP: GroupCoordinator | None = None
|
||||
_SP: Optional[GroupCoordinator] = None
|
||||
|
||||
|
||||
def get_sp_group() -> GroupCoordinator:
|
||||
@@ -790,7 +797,7 @@ def get_sp_group() -> GroupCoordinator:
|
||||
def initialize_model_parallel(
|
||||
tensor_model_parallel_size: int = 1,
|
||||
sequence_model_parallel_size: int = 1,
|
||||
backend: str | None = None,
|
||||
backend: Optional[str] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize model parallel groups.
|
||||
@@ -859,7 +866,7 @@ def get_sequence_model_parallel_rank() -> int:
|
||||
def ensure_model_parallel_initialized(
|
||||
tensor_model_parallel_size: int,
|
||||
sequence_model_parallel_size: int,
|
||||
backend: str | None = None,
|
||||
backend: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Helper to initialize model parallel groups if they are not initialized,
|
||||
or ensure tensor-parallel, sequence-parallel sizes
|
||||
@@ -970,8 +977,8 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
"torch._C._host_emptyCache() only available in Pytorch >=2.5")
|
||||
|
||||
|
||||
def in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
|
||||
source_rank: int = 0) -> list[bool]:
|
||||
def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
source_rank: int = 0) -> List[bool]:
|
||||
"""
|
||||
This is a collective operation that returns if each rank is in the same node
|
||||
as the source rank. It tests if processes are attached to the same
|
||||
@@ -1057,7 +1064,7 @@ def in_the_same_node_as(pg: ProcessGroup | StatelessProcessGroup,
|
||||
|
||||
def initialize_tensor_parallel_group(
|
||||
tensor_model_parallel_size: int = 1,
|
||||
backend: str | None = None,
|
||||
backend: Optional[str] = None,
|
||||
group_name_suffix: str = "") -> GroupCoordinator:
|
||||
"""Initialize a tensor parallel group for a specific model.
|
||||
|
||||
@@ -1121,7 +1128,7 @@ def initialize_tensor_parallel_group(
|
||||
|
||||
def initialize_sequence_parallel_group(
|
||||
sequence_model_parallel_size: int = 1,
|
||||
backend: str | None = None,
|
||||
backend: Optional[str] = None,
|
||||
group_name_suffix: str = "") -> GroupCoordinator:
|
||||
"""Initialize a sequence parallel group for a specific model.
|
||||
|
||||
|
||||
@@ -9,8 +9,7 @@ import dataclasses
|
||||
import pickle
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
from torch.distributed import TCPStore
|
||||
@@ -73,15 +72,15 @@ class StatelessProcessGroup:
|
||||
data_expiration_seconds: int = 3600 # 1 hour
|
||||
|
||||
# dst rank -> counter
|
||||
send_dst_counter: dict[int, int] = dataclasses.field(default_factory=dict)
|
||||
send_dst_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
|
||||
# src rank -> counter
|
||||
recv_src_counter: dict[int, int] = dataclasses.field(default_factory=dict)
|
||||
recv_src_counter: Dict[int, int] = dataclasses.field(default_factory=dict)
|
||||
broadcast_send_counter: int = 0
|
||||
broadcast_recv_src_counter: dict[int, int] = dataclasses.field(
|
||||
broadcast_recv_src_counter: Dict[int, int] = dataclasses.field(
|
||||
default_factory=dict)
|
||||
|
||||
# A deque to store the data entries, with key and timestamp.
|
||||
entries: deque[tuple[str, float]] = dataclasses.field(default_factory=deque)
|
||||
entries: Deque[Tuple[str, float]] = dataclasses.field(default_factory=deque)
|
||||
|
||||
def __post_init__(self):
|
||||
assert self.rank < self.world_size
|
||||
@@ -115,7 +114,7 @@ class StatelessProcessGroup:
|
||||
self.recv_src_counter[src] += 1
|
||||
return obj
|
||||
|
||||
def broadcast_obj(self, obj: Any | None, src: int) -> Any:
|
||||
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any:
|
||||
"""Broadcast an object from a source rank to all other ranks.
|
||||
It does not clean up after all ranks have received the object.
|
||||
Use it for limited times, e.g., for initialization.
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
import argparse
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import Any, cast
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
|
||||
from fastvideo import PipelineConfig, VideoGenerator
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
@@ -26,11 +26,11 @@ class GenerateSubcommand(CLISubcommand):
|
||||
self.init_arg_names = self._get_init_arg_names()
|
||||
self.generation_arg_names = self._get_generation_arg_names()
|
||||
|
||||
def _get_init_arg_names(self) -> list[str]:
|
||||
def _get_init_arg_names(self) -> List[str]:
|
||||
"""Get names of arguments for VideoGenerator initialization"""
|
||||
return ["num_gpus", "tp_size", "sp_size", "model_path"]
|
||||
|
||||
def _get_generation_arg_names(self) -> list[str]:
|
||||
def _get_generation_arg_names(self) -> List[str]:
|
||||
"""Get names of arguments for generate_video method"""
|
||||
return [field.name for field in dataclasses.fields(SamplingParam)]
|
||||
|
||||
@@ -130,13 +130,13 @@ class GenerateSubcommand(CLISubcommand):
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
return [GenerateSubcommand()]
|
||||
|
||||
|
||||
def update_config_from_args(config: Any,
|
||||
args_dict: dict[str, Any],
|
||||
prefix: str | None = None) -> None:
|
||||
args_dict: Dict[str, Any],
|
||||
prefix: Optional[str] = None) -> None:
|
||||
"""
|
||||
Update configuration object from arguments dictionary.
|
||||
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
|
||||
|
||||
from typing import List
|
||||
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
"""Initialize all commands from separate modules"""
|
||||
commands = []
|
||||
commands.extend(generate_cmd_init())
|
||||
|
||||
@@ -4,6 +4,7 @@ import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -18,8 +19,8 @@ class RaiseNotImplementedAction(argparse.Action):
|
||||
|
||||
|
||||
def launch_distributed(num_gpus: int,
|
||||
args: list[str],
|
||||
master_port: int | None = None) -> int:
|
||||
args: List[str],
|
||||
master_port: Optional[int] = None) -> int:
|
||||
"""
|
||||
Launch a distributed job with the given arguments
|
||||
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
from fastvideo.v1.pipelines.wan.wan_latent_pipeline import WanLatentPipeline
|
||||
|
||||
|
||||
def main():
|
||||
print("Starting data preprocessor")
|
||||
pipeline = WanLatentPipeline.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
|
||||
train_dataset = getdataset(args)
|
||||
sampler = DistributedSampler(train_dataset,
|
||||
rank=local_rank,
|
||||
num_replicas=world_size,
|
||||
shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
for batch in train_dataloader:
|
||||
pipeline(batch)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -10,7 +10,7 @@ import gc
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
@@ -53,9 +53,11 @@ class VideoGenerator:
|
||||
@classmethod
|
||||
def from_pretrained(cls,
|
||||
model_path: str,
|
||||
device: str | None = None,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
pipeline_config: str | PipelineConfig | None = None,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
@@ -126,9 +128,9 @@ class VideoGenerator:
|
||||
def generate_video(
|
||||
self,
|
||||
prompt: str,
|
||||
sampling_param: SamplingParam | None = None,
|
||||
sampling_param: Optional[SamplingParam] = None,
|
||||
**kwargs,
|
||||
) -> dict[str, Any] | list[np.ndarray]:
|
||||
) -> Union[Dict[str, Any], List[np.ndarray]]:
|
||||
"""
|
||||
Generate a video based on the given prompt.
|
||||
|
||||
|
||||
+12
-13
@@ -2,29 +2,28 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py
|
||||
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
|
||||
FASTVIDEO_NCCL_SO_PATH: str | None = None
|
||||
LD_LIBRARY_PATH: str | None = None
|
||||
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
|
||||
LD_LIBRARY_PATH: Optional[str] = None
|
||||
LOCAL_RANK: int = 0
|
||||
CUDA_VISIBLE_DEVICES: str | None = None
|
||||
CUDA_VISIBLE_DEVICES: Optional[str] = None
|
||||
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
|
||||
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
|
||||
FASTVIDEO_CONFIGURE_LOGGING: int = 1
|
||||
FASTVIDEO_LOGGING_LEVEL: str = "INFO"
|
||||
FASTVIDEO_LOGGING_PREFIX: str = ""
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_ATTENTION_CONFIG: str | None = None
|
||||
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
|
||||
FASTVIDEO_ATTENTION_CONFIG: Optional[str] = None
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
NVCC_THREADS: str | None = None
|
||||
CMAKE_BUILD_TYPE: str | None = None
|
||||
MAX_JOBS: Optional[str] = None
|
||||
NVCC_THREADS: Optional[str] = None
|
||||
CMAKE_BUILD_TYPE: Optional[str] = None
|
||||
VERBOSE: bool = False
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
|
||||
@@ -43,7 +42,7 @@ def get_default_config_root() -> str:
|
||||
)
|
||||
|
||||
|
||||
def maybe_convert_int(value: str | None) -> int | None:
|
||||
def maybe_convert_int(value: Optional[str]) -> Optional[int]:
|
||||
if value is None:
|
||||
return None
|
||||
return int(value)
|
||||
@@ -54,7 +53,7 @@ def maybe_convert_int(value: str | None) -> int | None:
|
||||
|
||||
# begin-env-vars-definition
|
||||
|
||||
environment_variables: dict[str, Callable[[], Any]] = {
|
||||
environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
|
||||
# ================== Installation Time Env Vars ==================
|
||||
|
||||
|
||||
+374
-18
@@ -4,10 +4,9 @@
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
from collections.abc import Callable
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import field
|
||||
from typing import Any
|
||||
from typing import Any, Callable, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -39,17 +38,17 @@ class FastVideoArgs:
|
||||
|
||||
# HuggingFace specific parameters
|
||||
trust_remote_code: bool = False
|
||||
revision: str | None = None
|
||||
revision: Optional[str] = None
|
||||
|
||||
# Parallelism
|
||||
num_gpus: int = 1
|
||||
tp_size: int | None = None
|
||||
sp_size: int | None = None
|
||||
dist_timeout: int | None = None # timeout for torch.distributed
|
||||
tp_size: Optional[int] = None
|
||||
sp_size: Optional[int] = None
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
flow_shift: Optional[float] = None
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
@@ -71,36 +70,40 @@ class FastVideoArgs:
|
||||
# Text encoder configuration
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = (
|
||||
"fp16",
|
||||
"fp16",
|
||||
# "fp16",
|
||||
)
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: tuple[Callable[[Any], Any], ...] = field(
|
||||
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
|
||||
default_factory=lambda: (postprocess_text, ))
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: str | None = None
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: str | None = None
|
||||
neg_magic: str | None = None
|
||||
timesteps_scale: bool | None = None
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
|
||||
# Logging
|
||||
log_level: str = "info"
|
||||
|
||||
# Inference parameters
|
||||
device_str: str | None = None
|
||||
device_str: Optional[str] = None
|
||||
device = None
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
|
||||
def __post_init__(self):
|
||||
pass
|
||||
|
||||
@@ -133,6 +136,13 @@ class FastVideoArgs:
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--inference-mode",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.inference_mode,
|
||||
help="Whether to use inference mode",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
@@ -377,7 +387,7 @@ class FastVideoArgs:
|
||||
_current_fastvideo_args = None
|
||||
|
||||
|
||||
def prepare_fastvideo_args(argv: list[str]) -> FastVideoArgs:
|
||||
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
|
||||
"""
|
||||
Prepare the inference arguments from the command line arguments.
|
||||
|
||||
@@ -424,3 +434,349 @@ def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
# TODO(will): may need to handle this for CI.
|
||||
raise ValueError("Current fastvideo args is not set.")
|
||||
return _current_fastvideo_args
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TrainingArgs(FastVideoArgs):
|
||||
data_path: str = ""
|
||||
dataloader_num_workers: int = 0
|
||||
num_height: int = 0
|
||||
num_width: int = 0
|
||||
num_frames: int = 0
|
||||
|
||||
train_batch_size: int = 0
|
||||
num_latent_t: int = 0
|
||||
group_frame: bool = False
|
||||
group_resolution: bool = False
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
cache_dir: str = ""
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
cfg: float = 0.0
|
||||
precondition_outputs: bool = False
|
||||
|
||||
# validation & logs
|
||||
validation_prompt_dir: str = ""
|
||||
validation_sampling_steps: str = ""
|
||||
validation_guidance_scale: str = ""
|
||||
validation_steps: float = 0.0
|
||||
log_validation: bool = False
|
||||
tracker_project_name: str = ""
|
||||
# seed: int
|
||||
|
||||
# output
|
||||
output_dir: str = ""
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: str = ""
|
||||
resume_from_lora_checkpoint: str = ""
|
||||
logging_dir: str = ""
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
max_train_steps: int = 0
|
||||
gradient_accumulation_steps: int = 0
|
||||
learning_rate: float = 0.0
|
||||
scale_lr: bool = False
|
||||
lr_scheduler: str = ""
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
gradient_checkpointing: bool = False
|
||||
selective_checkpointing: float = 0.0
|
||||
allow_tf32: bool = False
|
||||
mixed_precision: str = ""
|
||||
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 = ""
|
||||
logit_mean: float = 0.0
|
||||
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
|
||||
not_apply_cfg_solver: bool = False
|
||||
distill_cfg: float = 0.0
|
||||
scheduler_type: str = ""
|
||||
linear_quadratic_threshold: float = 0.0
|
||||
linear_range: float = 0.0
|
||||
weight_decay: float = 0.0
|
||||
use_ema: bool = False
|
||||
multi_phased_distill_schedule: str = ""
|
||||
pred_decay_weight: float = 0.0
|
||||
pred_decay_type: str = ""
|
||||
hunyuan_teacher_disable_cfg: bool = False
|
||||
|
||||
# master_weight_type
|
||||
master_weight_type: str = ""
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
parser.add_argument("--data-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to parquet files")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of workers for dataloader")
|
||||
parser.add_argument("--num-height",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of heights")
|
||||
parser.add_argument("--num-width",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of widths")
|
||||
parser.add_argument("--num-frames",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of frames")
|
||||
|
||||
# Training batch and model configuration
|
||||
parser.add_argument("--train-batch-size",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Training batch size")
|
||||
parser.add_argument("--num-latent-t",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of latent time steps")
|
||||
parser.add_argument("--group-frame",
|
||||
action=StoreBoolean,
|
||||
help="Whether to group frames during training")
|
||||
parser.add_argument("--group-resolution",
|
||||
action=StoreBoolean,
|
||||
help="Whether to group resolutions during training")
|
||||
|
||||
# Model paths
|
||||
parser.add_argument("--pretrained-model-name-or-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to pretrained model or model name")
|
||||
parser.add_argument("--dit-model-name-or-path",
|
||||
type=str,
|
||||
required=False,
|
||||
help="Path to DiT model or model name")
|
||||
parser.add_argument("--cache-dir",
|
||||
type=str,
|
||||
help="Directory to cache models")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
default=0.999,
|
||||
help="EMA decay rate")
|
||||
parser.add_argument("--ema-start-step",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Step to start EMA")
|
||||
parser.add_argument("--cfg",
|
||||
type=float,
|
||||
help="Classifier-free guidance scale")
|
||||
parser.add_argument(
|
||||
"--precondition-outputs",
|
||||
action=StoreBoolean,
|
||||
help="Whether to precondition the outputs of the model")
|
||||
|
||||
# Validation and logging
|
||||
parser.add_argument("--validation-prompt-dir",
|
||||
type=str,
|
||||
help="Directory containing validation prompts")
|
||||
parser.add_argument("--validation-sampling-steps",
|
||||
type=str,
|
||||
help="Validation sampling steps")
|
||||
parser.add_argument("--validation-guidance-scale",
|
||||
type=str,
|
||||
help="Validation guidance scale")
|
||||
parser.add_argument("--validation-steps",
|
||||
type=float,
|
||||
help="Number of validation steps")
|
||||
parser.add_argument("--log-validation",
|
||||
action=StoreBoolean,
|
||||
help="Whether to log validation results")
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
|
||||
# Output configuration
|
||||
parser.add_argument("--output-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Output directory for checkpoints and logs")
|
||||
parser.add_argument("--checkpoints-total-limit",
|
||||
type=int,
|
||||
help="Maximum number of checkpoints to keep")
|
||||
parser.add_argument("--checkpointing-steps",
|
||||
type=int,
|
||||
help="Steps between checkpoints")
|
||||
parser.add_argument("--resume-from-checkpoint",
|
||||
type=str,
|
||||
help="Path to checkpoint to resume from")
|
||||
parser.add_argument("--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")
|
||||
|
||||
# Training configuration
|
||||
parser.add_argument("--num-train-epochs",
|
||||
type=int,
|
||||
help="Number of training epochs")
|
||||
parser.add_argument("--max-train-steps",
|
||||
type=int,
|
||||
help="Maximum number of training steps")
|
||||
parser.add_argument("--gradient-accumulation-steps",
|
||||
type=int,
|
||||
help="Number of steps to accumulate gradients")
|
||||
parser.add_argument("--learning-rate",
|
||||
type=float,
|
||||
required=True,
|
||||
help="Learning rate")
|
||||
parser.add_argument("--scale-lr",
|
||||
action=StoreBoolean,
|
||||
help="Whether to scale learning rate")
|
||||
parser.add_argument("--lr-scheduler",
|
||||
type=str,
|
||||
default="constant",
|
||||
help="Learning rate scheduler type")
|
||||
parser.add_argument("--lr-warmup-steps",
|
||||
type=int,
|
||||
default=10,
|
||||
help="Number of warmup steps for learning rate")
|
||||
parser.add_argument("--max-grad-norm",
|
||||
type=float,
|
||||
help="Maximum gradient norm")
|
||||
parser.add_argument("--gradient-checkpointing",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use gradient checkpointing")
|
||||
parser.add_argument("--selective-checkpointing",
|
||||
type=float,
|
||||
help="Selective checkpointing threshold")
|
||||
parser.add_argument("--allow-tf32",
|
||||
action=StoreBoolean,
|
||||
help="Whether to allow TF32")
|
||||
parser.add_argument("--mixed-precision",
|
||||
type=str,
|
||||
help="Mixed precision training type")
|
||||
parser.add_argument("--train-sp-batch-size",
|
||||
type=int,
|
||||
help="Training spatial parallelism batch size")
|
||||
|
||||
# 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")
|
||||
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
default="uniform",
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="mean to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_std",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="std to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode_scale",
|
||||
type=float,
|
||||
default=1.29,
|
||||
help=
|
||||
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
)
|
||||
|
||||
# Additional training parameters
|
||||
parser.add_argument("--num-euler-timesteps",
|
||||
type=int,
|
||||
help="Number of Euler timesteps")
|
||||
parser.add_argument("--lr-num-cycles",
|
||||
type=int,
|
||||
help="Number of learning rate cycles")
|
||||
parser.add_argument("--lr-power",
|
||||
type=float,
|
||||
help="Learning rate power")
|
||||
parser.add_argument("--not-apply-cfg-solver",
|
||||
action=StoreBoolean,
|
||||
help="Whether to not apply CFG solver")
|
||||
parser.add_argument("--distill-cfg",
|
||||
type=float,
|
||||
help="Distillation CFG scale")
|
||||
parser.add_argument("--scheduler-type", type=str, help="Scheduler type")
|
||||
parser.add_argument("--linear-quadratic-threshold",
|
||||
type=float,
|
||||
help="Linear quadratic threshold")
|
||||
parser.add_argument("--linear-range", type=float, help="Linear range")
|
||||
parser.add_argument("--weight-decay", type=float, help="Weight decay")
|
||||
parser.add_argument("--use-ema",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA")
|
||||
parser.add_argument("--multi-phased-distill-schedule",
|
||||
type=str,
|
||||
help="Multi-phased distillation schedule")
|
||||
parser.add_argument("--pred-decay-weight",
|
||||
type=float,
|
||||
help="Prediction decay weight")
|
||||
parser.add_argument("--pred-decay-type",
|
||||
type=str,
|
||||
help="Prediction decay type")
|
||||
parser.add_argument("--hunyuan-teacher-disable-cfg",
|
||||
action=StoreBoolean,
|
||||
help="Whether to disable CFG for Hunyuan teacher")
|
||||
parser.add_argument("--master-weight-type",
|
||||
type=str,
|
||||
help="Master weight type")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -5,7 +5,7 @@ import time
|
||||
from collections import defaultdict
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -37,10 +37,10 @@ class ForwardContext:
|
||||
# attn_layers: Dict[str, Any]
|
||||
# TODO: extend to support per-layer dynamic forward context
|
||||
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
|
||||
forward_batch: ForwardBatch | None = None
|
||||
forward_batch: Optional[ForwardBatch] = None
|
||||
|
||||
|
||||
_forward_context: ForwardContext | None = None
|
||||
_forward_context: Optional[ForwardContext] = None
|
||||
|
||||
|
||||
def get_forward_context() -> ForwardContext:
|
||||
@@ -55,8 +55,8 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
forward_batch: ForwardBatch | None = None,
|
||||
fastvideo_args: FastVideoArgs | None = None):
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
Here we can inject common logic for every model forward pass.
|
||||
|
||||
@@ -8,7 +8,7 @@ This module provides classes and functions for running inference with diffusion
|
||||
"""
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
@@ -83,7 +83,7 @@ class InferenceEngine:
|
||||
self,
|
||||
prompt: str,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> dict[str, Any]:
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Run inference with the pipeline.
|
||||
|
||||
@@ -96,7 +96,7 @@ class InferenceEngine:
|
||||
Returns:
|
||||
A dictionary containing the generated videos and metadata.
|
||||
"""
|
||||
out_dict: dict[str, Any] = dict()
|
||||
out_dict: Dict[str, Any] = dict()
|
||||
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/custom_op.py
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
from typing import Any, Callable, Dict, Type
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -82,7 +81,7 @@ class CustomOp(nn.Module):
|
||||
# Examples:
|
||||
# - MyOp.enabled()
|
||||
# - op_registry["my_op"].enabled()
|
||||
op_registry: dict[str, type['CustomOp']] = {}
|
||||
op_registry: Dict[str, Type['CustomOp']] = {}
|
||||
|
||||
# Decorator to register custom ops.
|
||||
@classmethod
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/layernorm.py
|
||||
"""Custom normalization layers."""
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -21,7 +22,7 @@ class RMSNorm(CustomOp):
|
||||
hidden_size: int,
|
||||
eps: float = 1e-6,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
var_hidden_size: int | None = None,
|
||||
var_hidden_size: Optional[int] = None,
|
||||
has_weight: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -39,8 +40,8 @@ class RMSNorm(CustomOp):
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: torch.Tensor | None = None,
|
||||
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
orig_dtype = x.dtype
|
||||
x = x.to(torch.float32)
|
||||
@@ -129,7 +130,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
gate: torch.Tensor, shift: torch.Tensor,
|
||||
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
scale: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply gated residual connection, followed by layernorm and
|
||||
scale/shift in a single fused operation.
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/linear.py
|
||||
|
||||
from abc import abstractmethod
|
||||
from typing import Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -39,7 +40,7 @@ WEIGHT_LOADER_V2_SUPPORTED = [
|
||||
|
||||
def adjust_scalar_to_fused_array(
|
||||
param: torch.Tensor, loaded_weight: torch.Tensor,
|
||||
shard_id: str | int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
shard_id: Union[str, int]) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""For fused modules (QKV and MLP) we have an array of length
|
||||
N that holds 1 scale for each "logical" matrix. So the param
|
||||
is an array of length N. The loaded_weight corresponds to
|
||||
@@ -90,7 +91,7 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
"""Apply the weights in layer to the input tensor.
|
||||
Expects create_weights to have been called before on the layer."""
|
||||
raise NotImplementedError
|
||||
@@ -115,7 +116,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
|
||||
return F.linear(x, layer.weight, bias)
|
||||
|
||||
@@ -137,8 +138,8 @@ class LinearBase(torch.nn.Module):
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -151,13 +152,14 @@ class LinearBase(torch.nn.Module):
|
||||
params_dtype = torch.get_default_dtype()
|
||||
self.params_dtype = params_dtype
|
||||
if quant_config is None:
|
||||
self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
|
||||
)
|
||||
self.quant_method: Optional[
|
||||
QuantizeMethodBase] = UnquantizedLinearMethod()
|
||||
else:
|
||||
self.quant_method = quant_config.get_quant_method(self,
|
||||
prefix=prefix)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
|
||||
def forward(self,
|
||||
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -180,8 +182,8 @@ class ReplicatedLinear(LinearBase):
|
||||
output_size: int,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__(input_size,
|
||||
output_size,
|
||||
@@ -221,7 +223,8 @@ class ReplicatedLinear(LinearBase):
|
||||
f"to a parameter of size {param.size()}")
|
||||
param.data.copy_(loaded_weight)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
|
||||
def forward(self,
|
||||
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
|
||||
bias = self.bias if not self.skip_bias_add else None
|
||||
assert self.quant_method is not None
|
||||
output = self.quant_method.apply(self, x, bias)
|
||||
@@ -265,9 +268,9 @@ class ColumnParallelLinear(LinearBase):
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
output_sizes: list[int] | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
output_sizes: Optional[list[int]] = None,
|
||||
prefix: str = ""):
|
||||
# Divide the weight matrix along the last dimension.
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
@@ -342,8 +345,9 @@ class ColumnParallelLinear(LinearBase):
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
param.load_column_parallel_weight(loaded_weight=loaded_weight)
|
||||
|
||||
def forward(self,
|
||||
input_: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
|
||||
def forward(
|
||||
self,
|
||||
input_: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
|
||||
bias = self.bias if not self.skip_bias_add else None
|
||||
|
||||
# Matrix multiply.
|
||||
@@ -395,8 +399,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
bias: bool = True,
|
||||
gather_output: bool = False,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
self.output_sizes = output_sizes
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
@@ -413,7 +417,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None) -> None:
|
||||
loaded_shard_id: Optional[int] = None) -> None:
|
||||
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
@@ -506,8 +510,10 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
if isinstance(
|
||||
param,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
@@ -519,7 +525,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: int | None = None) -> None:
|
||||
loaded_shard_id: Optional[int] = None) -> None:
|
||||
if loaded_shard_id is None:
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight,
|
||||
@@ -592,11 +598,11 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
hidden_size: int,
|
||||
head_size: int,
|
||||
total_num_heads: int,
|
||||
total_num_kv_heads: int | None = None,
|
||||
total_num_kv_heads: Optional[int] = None,
|
||||
bias: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
self.hidden_size = hidden_size
|
||||
self.head_size = head_size
|
||||
@@ -631,7 +637,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
|
||||
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> Optional[int]:
|
||||
shard_offset_mapping = {
|
||||
"q": 0,
|
||||
"k": self.num_heads * self.head_size,
|
||||
@@ -640,7 +646,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
}
|
||||
return shard_offset_mapping.get(loaded_shard_id)
|
||||
|
||||
def _get_shard_size_mapping(self, loaded_shard_id: str) -> int | None:
|
||||
def _get_shard_size_mapping(self, loaded_shard_id: str) -> Optional[int]:
|
||||
shard_size_mapping = {
|
||||
"q": self.num_heads * self.head_size,
|
||||
"k": self.num_kv_heads * self.head_size,
|
||||
@@ -673,8 +679,10 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(param, PackedColumnParameter | PackedvLLMParameter
|
||||
) and param.packed_dim == param.output_dim:
|
||||
if isinstance(
|
||||
param,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
@@ -686,7 +694,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None):
|
||||
loaded_shard_id: Optional[str] = None):
|
||||
if loaded_shard_id is None: # special case for certain models
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_qkv_weight(loaded_weight=loaded_weight, shard_id=0)
|
||||
@@ -712,7 +720,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: str | None = None):
|
||||
loaded_shard_id: Optional[str] = None):
|
||||
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
@@ -837,9 +845,9 @@ class RowParallelLinear(LinearBase):
|
||||
bias: bool = True,
|
||||
input_is_parallel: bool = True,
|
||||
skip_bias_add: bool = False,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
reduce_results: bool = True,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
# Divide the weight matrix along the first dimension.
|
||||
self.tp_rank = get_tensor_model_parallel_rank()
|
||||
@@ -913,7 +921,7 @@ class RowParallelLinear(LinearBase):
|
||||
|
||||
param.load_row_parallel_weight(loaded_weight=loaded_weight)
|
||||
|
||||
def forward(self, input_) -> tuple[torch.Tensor, Parameter | None]:
|
||||
def forward(self, input_) -> tuple[torch.Tensor, Optional[Parameter]]:
|
||||
if self.input_is_parallel:
|
||||
input_parallel = input_
|
||||
else:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -16,10 +18,10 @@ class MLP(nn.Module):
|
||||
self,
|
||||
input_dim: int,
|
||||
mlp_hidden_dim: int,
|
||||
output_dim: int | None = None,
|
||||
output_dim: Optional[int] = None,
|
||||
bias: bool = True,
|
||||
act_type: str = "gelu_pytorch_tanh",
|
||||
dtype: torch.dtype | None = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
|
||||
import inspect
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -105,8 +105,8 @@ class QuantizationConfig(ABC):
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def override_quantization_method(cls, hf_quant_cfg,
|
||||
user_quant) -> QuantizationMethods | None:
|
||||
def override_quantization_method(
|
||||
cls, hf_quant_cfg, user_quant) -> Optional[QuantizationMethods]:
|
||||
"""
|
||||
Detects if this quantization method can support a given checkpoint
|
||||
format by overriding the user specified quantization method --
|
||||
@@ -135,7 +135,7 @@ class QuantizationConfig(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def get_quant_method(self, layer: torch.nn.Module,
|
||||
prefix: str) -> QuantizeMethodBase | None:
|
||||
prefix: str) -> Optional[QuantizeMethodBase]:
|
||||
"""Get the quantize method to use for the quantized layer.
|
||||
|
||||
Args:
|
||||
@@ -147,5 +147,5 @@ class QuantizationConfig(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def get_cache_scale(self, name: str) -> str | None:
|
||||
return None
|
||||
def get_cache_scale(self, name: str) -> Optional[str]:
|
||||
return None
|
||||
@@ -23,7 +23,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Rotary Positional Embeddings."""
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
max_position_embeddings: int,
|
||||
base: int | float,
|
||||
base: Union[int, float],
|
||||
is_neox_style: bool,
|
||||
dtype: torch.dtype,
|
||||
) -> None:
|
||||
@@ -101,7 +101,7 @@ class RotaryEmbedding(CustomOp):
|
||||
self.cos_sin_cache: torch.Tensor
|
||||
self.register_buffer("cos_sin_cache", cache, persistent=False)
|
||||
|
||||
def _compute_inv_freq(self, base: int | float) -> torch.Tensor:
|
||||
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
|
||||
"""Compute the inverse frequency."""
|
||||
# NOTE(woosuk): To exactly match the HF implementation, we need to
|
||||
# use CPU to compute the cache and then move it to GPU. However, we
|
||||
@@ -127,8 +127,8 @@ class RotaryEmbedding(CustomOp):
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""A PyTorch-native implementation of forward()."""
|
||||
if offsets is not None:
|
||||
positions = positions + offsets
|
||||
@@ -159,7 +159,7 @@ class RotaryEmbedding(CustomOp):
|
||||
return s
|
||||
|
||||
|
||||
def _to_tuple(x: int | tuple[int, ...], dim: int = 2) -> tuple[int, ...]:
|
||||
def _to_tuple(x: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
|
||||
if isinstance(x, int):
|
||||
return (x, ) * dim
|
||||
elif len(x) == dim:
|
||||
@@ -168,8 +168,8 @@ def _to_tuple(x: int | tuple[int, ...], dim: int = 2) -> tuple[int, ...]:
|
||||
raise ValueError(f"Expected length {dim} or int, but got {x}")
|
||||
|
||||
|
||||
def get_meshgrid_nd(start: int | tuple[int, ...],
|
||||
*args: int | tuple[int, ...],
|
||||
def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
|
||||
*args: Union[int, Tuple[int, ...]],
|
||||
dim: int = 2) -> torch.Tensor:
|
||||
"""
|
||||
Get n-D meshgrid with start, stop and num.
|
||||
@@ -217,12 +217,12 @@ def get_meshgrid_nd(start: int | tuple[int, ...],
|
||||
|
||||
def get_1d_rotary_pos_embed(
|
||||
dim: int,
|
||||
pos: torch.FloatTensor | int,
|
||||
pos: Union[torch.FloatTensor, int],
|
||||
theta: float = 10000.0,
|
||||
theta_rescale_factor: float = 1.0,
|
||||
interpolation_factor: float = 1.0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
|
||||
@@ -261,13 +261,13 @@ def get_nd_rotary_pos_embed(
|
||||
start,
|
||||
*args,
|
||||
theta=10000.0,
|
||||
theta_rescale_factor: float | list[float] = 1.0,
|
||||
interpolation_factor: float | list[float] = 1.0,
|
||||
theta_rescale_factor: Union[float, List[float]] = 1.0,
|
||||
interpolation_factor: Union[float, List[float]] = 1.0,
|
||||
shard_dim: int = 0,
|
||||
sp_rank: int = 0,
|
||||
sp_world_size: int = 1,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
||||
Supports sequence parallelism by allowing sharding of a specific dimension.
|
||||
@@ -324,7 +324,7 @@ def get_nd_rotary_pos_embed(
|
||||
else:
|
||||
grid = full_grid
|
||||
|
||||
if isinstance(theta_rescale_factor, int | float):
|
||||
if isinstance(theta_rescale_factor, (int, float)):
|
||||
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
||||
elif isinstance(theta_rescale_factor,
|
||||
list) and len(theta_rescale_factor) == 1:
|
||||
@@ -333,7 +333,7 @@ def get_nd_rotary_pos_embed(
|
||||
rope_dim_list
|
||||
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
if isinstance(interpolation_factor, int | float):
|
||||
if isinstance(interpolation_factor, (int, float)):
|
||||
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
||||
elif isinstance(interpolation_factor,
|
||||
list) and len(interpolation_factor) == 1:
|
||||
@@ -370,7 +370,7 @@ def get_rotary_pos_embed(
|
||||
interpolation_factor=1.0,
|
||||
shard_dim: int = 0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Generate rotary positional embeddings for the given sizes.
|
||||
|
||||
@@ -417,17 +417,17 @@ def get_rotary_pos_embed(
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
_ROPE_DICT: dict[tuple, RotaryEmbedding] = {}
|
||||
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {}
|
||||
|
||||
|
||||
def get_rope(
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
max_position: int,
|
||||
base: int | float,
|
||||
base: Union[int, float],
|
||||
is_neox_style: bool = True,
|
||||
rope_scaling: dict[str, Any] | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
rope_scaling: Optional[Dict[str, Any]] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
partial_rotary_factor: float = 1.0,
|
||||
) -> RotaryEmbedding:
|
||||
if dtype is None:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/utils.py
|
||||
"""Utility methods for model layers."""
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -9,7 +10,7 @@ def get_token_bin_counts_and_mask(
|
||||
tokens: torch.Tensor,
|
||||
vocab_size: int,
|
||||
num_seqs: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Compute the bin counts for the tokens.
|
||||
# vocab_size + 1 for padding.
|
||||
bin_counts = torch.zeros((num_seqs, vocab_size + 1),
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -35,7 +36,7 @@ class PatchEmbed(nn.Module):
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, list | tuple):
|
||||
if isinstance(patch_size, (list, tuple)):
|
||||
if len(patch_size) == 1:
|
||||
patch_size = (patch_size[0], patch_size[0])
|
||||
else:
|
||||
@@ -132,7 +133,7 @@ class ModulateProjection(nn.Module):
|
||||
hidden_size: int,
|
||||
factor: int = 2,
|
||||
act_layer: str = "silu",
|
||||
dtype: torch.dtype | None = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -24,7 +24,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_partition_sizes: List[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs):
|
||||
"""Create weights for embedding layer."""
|
||||
@@ -39,7 +39,7 @@ class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
||||
def apply(self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
return F.linear(x, layer.weight, bias)
|
||||
|
||||
def embedding(self, layer: torch.nn.Module,
|
||||
@@ -139,7 +139,7 @@ def get_masked_input_and_mask(
|
||||
input_: torch.Tensor, org_vocab_start_index: int,
|
||||
org_vocab_end_index: int, num_org_vocab_padding: int,
|
||||
added_vocab_start_index: int,
|
||||
added_vocab_end_index: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
added_vocab_end_index: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# torch.compile will fuse all of the pointwise ops below
|
||||
# into a single kernel, making it very fast
|
||||
org_vocab_mask = (input_ >= org_vocab_start_index) & (input_
|
||||
@@ -197,10 +197,10 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
def __init__(self,
|
||||
num_embeddings: int,
|
||||
embedding_dim: int,
|
||||
params_dtype: torch.dtype | None = None,
|
||||
org_num_embeddings: int | None = None,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
org_num_embeddings: Optional[int] = None,
|
||||
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
@@ -296,7 +296,7 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
org_vocab_start_index, org_vocab_end_index, added_vocab_start_index,
|
||||
added_vocab_end_index)
|
||||
|
||||
def get_sharded_to_full_mapping(self) -> list[int] | None:
|
||||
def get_sharded_to_full_mapping(self) -> Optional[List[int]]:
|
||||
"""Get a mapping that can be used to reindex the gathered
|
||||
logits for sampling.
|
||||
|
||||
@@ -310,9 +310,9 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
if self.tp_size < 2:
|
||||
return None
|
||||
|
||||
base_embeddings: list[int] = []
|
||||
added_embeddings: list[int] = []
|
||||
padding: list[int] = []
|
||||
base_embeddings: List[int] = []
|
||||
added_embeddings: List[int] = []
|
||||
padding: List[int] = []
|
||||
for tp_rank in range(self.tp_size):
|
||||
shard_indices = self._get_indices(self.num_embeddings_padded,
|
||||
self.org_vocab_size_padded,
|
||||
|
||||
@@ -11,7 +11,7 @@ from logging import Logger
|
||||
from logging.config import dictConfig
|
||||
from os import path
|
||||
from types import MethodType
|
||||
from typing import Any, cast
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
|
||||
@@ -278,7 +278,8 @@ def _trace_calls(log_path, root_dir, frame, event, arg=None):
|
||||
return partial(_trace_calls, log_path, root_dir)
|
||||
|
||||
|
||||
def enable_trace_function_call(log_file_path: str, root_dir: str | None = None):
|
||||
def enable_trace_function_call(log_file_path: str,
|
||||
root_dir: Optional[str] = None):
|
||||
"""
|
||||
Enable tracing of every function call in code under `root_dir`.
|
||||
This is useful for debugging hangs or crashes.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -18,7 +18,7 @@ class BaseDiT(nn.Module, ABC):
|
||||
num_attention_heads: int
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: tuple[
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init_subclass__(cls) -> None:
|
||||
@@ -33,9 +33,11 @@ class BaseDiT(nn.Module, ABC):
|
||||
f"Subclasses of BaseDiT must define '{attr}' class variable"
|
||||
)
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any],
|
||||
**kwargs) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.hf_config = hf_config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
@@ -44,10 +46,10 @@ class BaseDiT(nn.Module, ABC):
|
||||
@abstractmethod
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
pass
|
||||
@@ -63,7 +65,7 @@ class BaseDiT(nn.Module, ABC):
|
||||
)
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
@@ -81,7 +83,7 @@ class CachableDiT(BaseDiT):
|
||||
num_attention_heads: int
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: tuple[
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -94,8 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: torch.dtype | None = None,
|
||||
supported_attention_backends: tuple[_Backend, ...] | None = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -200,7 +202,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
txt: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Process modulation vectors
|
||||
img_mod_outputs = self.img_mod(vec)
|
||||
(
|
||||
@@ -301,8 +303,8 @@ class MMSingleStreamBlock(nn.Module):
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: torch.dtype | None = None,
|
||||
supported_attention_backends: tuple[_Backend, ...] | None = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -364,7 +366,7 @@ class MMSingleStreamBlock(nn.Module):
|
||||
x: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
txt_len: int,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
# Process modulation
|
||||
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
||||
@@ -440,8 +442,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
|
||||
|
||||
def __init__(self, config: HunyuanVideoConfig):
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
self.patch_size = [
|
||||
config.patch_size_t, config.patch_size, config.patch_size
|
||||
@@ -540,10 +542,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
|
||||
# TODO: change output to a dict
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs):
|
||||
"""
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
# ==============================================================================
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
@@ -54,7 +55,7 @@ class PatchEmbed2D(nn.Module):
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, list | tuple):
|
||||
if isinstance(patch_size, (list, tuple)):
|
||||
if len(patch_size) == 1:
|
||||
patch_size = (patch_size[0], patch_size[0])
|
||||
else:
|
||||
@@ -142,7 +143,7 @@ class SelfAttention(nn.Module):
|
||||
def __init__(self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: tuple[int, int, int] = (64, 32, 32),
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
@@ -189,10 +190,8 @@ class SelfAttention(nn.Module):
|
||||
|
||||
outs = []
|
||||
idx = 0
|
||||
for (chunk_size, cos_i, sin_i) in zip(self.rope_split,
|
||||
cos_splits,
|
||||
sin_splits,
|
||||
strict=False):
|
||||
for (chunk_size, cos_i, sin_i) in zip(self.rope_split, cos_splits,
|
||||
sin_splits):
|
||||
# slice the corresponding channels
|
||||
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
|
||||
idx += chunk_size
|
||||
@@ -332,8 +331,8 @@ class AdaLayerNormSingle(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
embedded_timestep = self.emb(timestep * self.time_step_rescale)
|
||||
|
||||
out, _ = self.linear(self.silu(embedded_timestep))
|
||||
@@ -378,7 +377,7 @@ class StepVideoTransformerBlock(nn.Module):
|
||||
dim: int,
|
||||
attention_head_dim: int,
|
||||
norm_eps: float = 1e-5,
|
||||
ff_inner_dim: int | None = None,
|
||||
ff_inner_dim: Optional[int] = None,
|
||||
ff_bias: bool = False,
|
||||
attention_type: str = 'torch'):
|
||||
super().__init__()
|
||||
@@ -418,7 +417,7 @@ class StepVideoTransformerBlock(nn.Module):
|
||||
kv: torch.Tensor,
|
||||
t_expand: torch.LongTensor,
|
||||
attn_mask=None,
|
||||
rope_positions: list | None = None,
|
||||
rope_positions: Optional[list] = None,
|
||||
cos_sin=None,
|
||||
mask_strategy=None) -> torch.Tensor:
|
||||
|
||||
@@ -463,8 +462,9 @@ class StepVideoModel(BaseDiT):
|
||||
_supported_attention_backends = StepVideoConfig(
|
||||
)._supported_attention_backends
|
||||
|
||||
def __init__(self, config: StepVideoConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.attention_head_dim = config.attention_head_dim
|
||||
self.in_channels = config.in_channels
|
||||
@@ -539,7 +539,7 @@ class StepVideoModel(BaseDiT):
|
||||
return hidden_states
|
||||
|
||||
def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states,
|
||||
q_seqlen) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
q_seqlen) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
kv_seqlens = encoder_attention_mask.sum(dim=1).int()
|
||||
mask = torch.zeros([len(kv_seqlens), q_seqlen,
|
||||
max(kv_seqlens)],
|
||||
@@ -594,12 +594,12 @@ class StepVideoModel(BaseDiT):
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | None = None,
|
||||
t_expand: torch.LongTensor | None = None,
|
||||
encoder_hidden_states_2: torch.Tensor | None = None,
|
||||
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
|
||||
encoder_attention_mask: torch.Tensor | None = None,
|
||||
fps: torch.Tensor | None = None,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
t_expand: Optional[torch.LongTensor] = None,
|
||||
encoder_hidden_states_2: Optional[torch.Tensor] = None,
|
||||
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None,
|
||||
fps: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
mask_strategy=None,
|
||||
guidance=None,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -52,7 +53,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
text_embed_dim: int,
|
||||
image_embed_dim: int | None = None,
|
||||
image_embed_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -75,7 +76,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: torch.Tensor | None = None,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
):
|
||||
temb = self.time_embedder(timestep)
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
@@ -172,7 +173,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
supported_attention_backends: tuple[_Backend, ...] | None = None
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
|
||||
) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps,
|
||||
supported_attention_backends)
|
||||
@@ -221,9 +222,9 @@ class WanTransformerBlock(nn.Module):
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: int | None = None,
|
||||
supported_attention_backends: tuple[_Backend, ...]
|
||||
| None = None,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
@@ -291,13 +292,13 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
assert orig_dtype != torch.float32
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
@@ -359,8 +360,9 @@ class WanTransformer3DModel(CachableDiT):
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = WanVideoConfig()._param_names_mapping
|
||||
|
||||
def __init__(self, config: WanVideoConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
@@ -416,10 +418,10 @@ class WanTransformer3DModel(CachableDiT):
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
forward_batch = get_forward_context().forward_batch
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -10,7 +11,7 @@ from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: tuple[
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
@@ -23,21 +24,21 @@ class TextEncoder(nn.Module, ABC):
|
||||
|
||||
@abstractmethod
|
||||
def forward(self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs) -> BaseEncoderOutput:
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
class ImageEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: tuple[
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: ImageEncoderConfig) -> None:
|
||||
@@ -54,5 +55,5 @@ class ImageEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
|
||||
"""Minimal implementation of CLIPVisionModel intended to be only used
|
||||
within a vision language model."""
|
||||
from collections.abc import Iterable
|
||||
from typing import Iterable, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -91,9 +91,9 @@ class CLIPTextEmbeddings(nn.Module):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.LongTensor | None = None,
|
||||
position_ids: torch.LongTensor | None = None,
|
||||
inputs_embeds: torch.FloatTensor | None = None,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if input_ids is not None:
|
||||
seq_length = input_ids.shape[-1]
|
||||
@@ -128,8 +128,8 @@ class CLIPAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig | CLIPTextConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
config: Union[CLIPVisionConfig, CLIPTextConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -209,8 +209,8 @@ class CLIPMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig | CLIPTextConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
config: Union[CLIPVisionConfig, CLIPTextConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -239,8 +239,8 @@ class CLIPEncoderLayer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPTextConfig | CLIPVisionConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
config: Union[CLIPTextConfig, CLIPVisionConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -284,9 +284,9 @@ class CLIPEncoder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig | CLIPTextConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
num_hidden_layers_override: int | None = None,
|
||||
config: Union[CLIPVisionConfig, CLIPTextConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -305,8 +305,8 @@ class CLIPEncoder(nn.Module):
|
||||
])
|
||||
|
||||
def forward(
|
||||
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
|
||||
) -> torch.Tensor | list[torch.Tensor]:
|
||||
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
|
||||
) -> Union[torch.Tensor, list[torch.Tensor]]:
|
||||
hidden_states_pool = [inputs_embeds]
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
@@ -325,8 +325,8 @@ class CLIPTextTransformer(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
num_hidden_layers_override: int | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
@@ -348,11 +348,11 @@ class CLIPTextTransformer(nn.Module):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
) -> BaseEncoderOutput:
|
||||
r"""
|
||||
Returns:
|
||||
@@ -440,11 +440,11 @@ class CLIPTextModel(TextEncoder):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
|
||||
@@ -456,8 +456,8 @@ class CLIPTextModel(TextEncoder):
|
||||
)
|
||||
return outputs
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
|
||||
# Define mapping for stacked parameters
|
||||
stacked_params_mapping = [
|
||||
@@ -467,7 +467,7 @@ class CLIPTextModel(TextEncoder):
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
# Handle q_proj, k_proj, v_proj -> qkv_proj mapping
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
@@ -498,9 +498,9 @@ class CLIPVisionTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
num_hidden_layers_override: int | None = None,
|
||||
require_post_norm: bool | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -540,7 +540,7 @@ class CLIPVisionTransformer(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
pixel_values: torch.Tensor,
|
||||
feature_sample_layers: list[int] | None = None,
|
||||
feature_sample_layers: Optional[list[int]] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
hidden_states = self.embeddings(pixel_values)
|
||||
@@ -582,7 +582,7 @@ class CLIPVisionModel(ImageEncoder):
|
||||
def forward(
|
||||
self,
|
||||
pixel_values: torch.Tensor,
|
||||
feature_sample_layers: list[int] | None = None,
|
||||
feature_sample_layers: Optional[list[int]] = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
last_hidden_state = self.vision_model(pixel_values,
|
||||
@@ -595,8 +595,8 @@ class CLIPVisionModel(ImageEncoder):
|
||||
|
||||
# (TODO) Add prefix argument for filtering out weights to be loaded
|
||||
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
@@ -604,7 +604,7 @@ class CLIPVisionModel(ImageEncoder):
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
loaded_params: Set[str] = set()
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
|
||||
@@ -23,8 +23,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Inference-only LLaMA model compatible with HuggingFace weights."""
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
from typing import Any, Dict, Iterable, Optional, Set, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -53,7 +52,7 @@ class LlamaMLP(nn.Module):
|
||||
hidden_size: int,
|
||||
intermediate_size: int,
|
||||
hidden_act: str,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
bias: bool = False,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
@@ -93,9 +92,9 @@ class LlamaAttention(nn.Module):
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
rope_theta: float = 10000,
|
||||
rope_scaling: dict[str, Any] | None = None,
|
||||
rope_scaling: Optional[Dict[str, Any]] = None,
|
||||
max_position_embeddings: int = 8192,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
bias: bool = False,
|
||||
bias_o_proj: bool = False,
|
||||
prefix: str = "") -> None:
|
||||
@@ -202,7 +201,7 @@ class LlamaDecoderLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: LlamaConfig,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -255,8 +254,8 @@ class LlamaDecoderLayer(nn.Module):
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
residual: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Self Attention
|
||||
if residual is None:
|
||||
residual = hidden_states
|
||||
@@ -319,11 +318,11 @@ class LlamaModel(TextEncoder):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
output_hidden_states = (output_hidden_states
|
||||
@@ -340,7 +339,7 @@ class LlamaModel(TextEncoder):
|
||||
0, hidden_states.shape[1],
|
||||
device=hidden_states.device).unsqueeze(0)
|
||||
|
||||
all_hidden_states: tuple[Any, ...] | None = (
|
||||
all_hidden_states: Optional[Tuple[Any, ...]] = (
|
||||
) if output_hidden_states else None
|
||||
for layer in self.layers:
|
||||
if all_hidden_states is not None:
|
||||
@@ -368,8 +367,8 @@ class LlamaModel(TextEncoder):
|
||||
|
||||
return output
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
@@ -379,7 +378,7 @@ class LlamaModel(TextEncoder):
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
if "rotary_emb.inv_freq" in name:
|
||||
continue
|
||||
@@ -401,7 +400,7 @@ class LlamaModel(TextEncoder):
|
||||
# continue
|
||||
if "scale" in name:
|
||||
# Remapping the name of FP8 kv-scale.
|
||||
kv_scale_name: str | None = maybe_remap_kv_scale_name(
|
||||
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
|
||||
name, params_dict)
|
||||
if kv_scale_name is None:
|
||||
continue
|
||||
|
||||
@@ -13,6 +13,7 @@
|
||||
# ==============================================================================
|
||||
import os
|
||||
from functools import wraps
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -178,10 +179,10 @@ class StepChatTokenizer:
|
||||
def vocab_size(self):
|
||||
return self._tokenizer.vocab_size()
|
||||
|
||||
def tokenize(self, text: str) -> list[int]:
|
||||
def tokenize(self, text: str) -> List[int]:
|
||||
return self._tokenizer.encode_as_ids(text)
|
||||
|
||||
def detokenize(self, token_ids: list[int]) -> str:
|
||||
def detokenize(self, token_ids: List[int]) -> str:
|
||||
return self._tokenizer.decode_ids(token_ids)
|
||||
|
||||
|
||||
@@ -346,9 +347,9 @@ class MultiQueryAttention(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
cu_seqlens: torch.Tensor | None,
|
||||
max_seq_len: torch.Tensor | None,
|
||||
mask: Optional[torch.Tensor],
|
||||
cu_seqlens: Optional[torch.Tensor],
|
||||
max_seq_len: Optional[torch.Tensor],
|
||||
):
|
||||
seqlen, bsz, dim = x.shape
|
||||
xqkv = self.wqkv(x)
|
||||
@@ -470,9 +471,9 @@ class TransformerBlock(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
cu_seqlens: torch.Tensor | None,
|
||||
max_seq_len: torch.Tensor | None,
|
||||
mask: Optional[torch.Tensor],
|
||||
cu_seqlens: Optional[torch.Tensor],
|
||||
max_seq_len: Optional[torch.Tensor],
|
||||
):
|
||||
residual = self.attention.forward(self.attention_norm(x), mask,
|
||||
cu_seqlens, max_seq_len)
|
||||
|
||||
@@ -20,8 +20,8 @@
|
||||
"""PyTorch T5 & UMT5 model."""
|
||||
|
||||
import math
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable, Optional, Set, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -64,7 +64,7 @@ class T5DenseActDense(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
quant_config: QuantizationConfig | None = None):
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False)
|
||||
@@ -85,7 +85,7 @@ class T5DenseGatedActDense(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
quant_config: QuantizationConfig | None = None):
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False,
|
||||
@@ -113,7 +113,7 @@ class T5LayerFF(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
quant_config: QuantizationConfig | None = None):
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
if config.is_gated_act:
|
||||
self.DenseReluDense = T5DenseGatedActDense(
|
||||
@@ -155,7 +155,7 @@ class T5Attention(nn.Module):
|
||||
config: T5Config,
|
||||
attn_type: str,
|
||||
has_relative_attention_bias=False,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.attn_type = attn_type
|
||||
@@ -294,7 +294,7 @@ class T5Attention(nn.Module):
|
||||
self,
|
||||
hidden_states: torch.Tensor, # (num_tokens, d_model)
|
||||
attention_mask: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata | None = None,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
) -> torch.Tensor:
|
||||
bs, seq_len, _ = hidden_states.shape
|
||||
num_seqs = bs
|
||||
@@ -344,7 +344,7 @@ class T5LayerSelfAttention(nn.Module):
|
||||
self,
|
||||
config,
|
||||
has_relative_attention_bias=False,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -361,7 +361,7 @@ class T5LayerSelfAttention(nn.Module):
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata | None = None,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
) -> torch.Tensor:
|
||||
normed_hidden_states = self.layer_norm.forward_native(hidden_states)
|
||||
attention_output = self.SelfAttention(
|
||||
@@ -377,7 +377,7 @@ class T5LayerCrossAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.EncDecAttention = T5Attention(config,
|
||||
@@ -390,7 +390,7 @@ class T5LayerCrossAttention(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata | None = None,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
) -> torch.Tensor:
|
||||
normed_hidden_states = self.layer_norm.forward_native(hidden_states)
|
||||
attention_output = self.EncDecAttention(
|
||||
@@ -407,7 +407,7 @@ class T5Block(nn.Module):
|
||||
config: T5Config,
|
||||
is_decoder: bool,
|
||||
has_relative_attention_bias=False,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.is_decoder = is_decoder
|
||||
@@ -431,7 +431,7 @@ class T5Block(nn.Module):
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata | None = None,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
hidden_states = self.layer[0](hidden_states=hidden_states,
|
||||
@@ -455,7 +455,7 @@ class T5Stack(nn.Module):
|
||||
is_decoder: bool,
|
||||
n_layers: int,
|
||||
embed_tokens=None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
is_umt5: bool = False):
|
||||
super().__init__()
|
||||
@@ -524,11 +524,11 @@ class T5EncoderModel(TextEncoder):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
attn_metadata = AttentionMetadata(None)
|
||||
@@ -540,8 +540,8 @@ class T5EncoderModel(TextEncoder):
|
||||
|
||||
return BaseEncoderOutput(last_hidden_state=hidden_states)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
@@ -549,7 +549,7 @@ class T5EncoderModel(TextEncoder):
|
||||
(".qkv_proj", ".v", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
loaded = False
|
||||
if "decoder" in name or "lm_head" in name:
|
||||
@@ -611,11 +611,11 @@ class UMT5EncoderModel(TextEncoder):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
attn_metadata = AttentionMetadata(None)
|
||||
@@ -630,8 +630,8 @@ class UMT5EncoderModel(TextEncoder):
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
@@ -639,7 +639,7 @@ class UMT5EncoderModel(TextEncoder):
|
||||
(".qkv_proj", ".v", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
loaded = False
|
||||
if "decoder" in name or "lm_head" in name:
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/vision.py
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Generic, TypeVar
|
||||
from typing import Generic, Optional, TypeVar, Union
|
||||
|
||||
import torch
|
||||
from transformers import PretrainedConfig
|
||||
@@ -48,9 +48,9 @@ class VisionEncoderInfo(ABC, Generic[_C]):
|
||||
|
||||
|
||||
def resolve_visual_encoder_outputs(
|
||||
encoder_outputs: torch.Tensor | list[torch.Tensor],
|
||||
feature_sample_layers: list[int] | None,
|
||||
post_layer_norm: torch.nn.LayerNorm | None,
|
||||
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
|
||||
feature_sample_layers: Optional[list[int]],
|
||||
post_layer_norm: Optional[torch.nn.LayerNorm],
|
||||
max_possible_layers: int,
|
||||
) -> torch.Tensor:
|
||||
"""Given the outputs a visual encoder module that may correspond to the
|
||||
|
||||
@@ -20,14 +20,14 @@ import contextlib
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Dict, Optional, Type, Union
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import AutoConfig, PretrainedConfig
|
||||
from transformers.models.auto.modeling_auto import (
|
||||
MODEL_FOR_CAUSAL_LM_MAPPING_NAMES)
|
||||
|
||||
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
|
||||
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
||||
# ChatGLMConfig.model_type: ChatGLMConfig,
|
||||
# DbrxConfig.model_type: DbrxConfig,
|
||||
# ExaoneConfig.model_type: ExaoneConfig,
|
||||
@@ -50,8 +50,8 @@ def download_from_hf(model_path: str):
|
||||
def get_hf_config(
|
||||
model: str,
|
||||
trust_remote_code: bool,
|
||||
revision: str | None = None,
|
||||
model_override_args: dict | None = None,
|
||||
revision: Optional[str] = None,
|
||||
model_override_args: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
is_gguf = check_gguf_file(model)
|
||||
@@ -83,8 +83,8 @@ def get_hf_config(
|
||||
|
||||
def get_diffusers_config(
|
||||
model: str,
|
||||
fastvideo_args: dict | None = None,
|
||||
) -> dict[str, Any]:
|
||||
fastvideo_args: Optional[dict] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Gets a configuration for the given diffusers model.
|
||||
|
||||
Args:
|
||||
@@ -104,7 +104,7 @@ def get_diffusers_config(
|
||||
try:
|
||||
# Load the config directly from the file
|
||||
with open(config_file) as f:
|
||||
config_dict: dict[str, Any] = json.load(f)
|
||||
config_dict: Dict[str, Any] = json.load(f)
|
||||
|
||||
# TODO(will): apply any overrides from inference args
|
||||
return config_dict
|
||||
@@ -139,7 +139,7 @@ def attach_additional_stop_token_ids(tokenizer):
|
||||
tokenizer.additional_stop_token_ids = None
|
||||
|
||||
|
||||
def check_gguf_file(model: str | os.PathLike) -> bool:
|
||||
def check_gguf_file(model: Union[str, os.PathLike]) -> bool:
|
||||
"""Check if the file is a GGUF model."""
|
||||
model = Path(model)
|
||||
if not model.is_file():
|
||||
|
||||
@@ -6,8 +6,8 @@ import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Generator, Iterable
|
||||
from typing import Any, cast
|
||||
from copy import deepcopy
|
||||
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -106,7 +106,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
fall_back_to_pt: bool = True
|
||||
"""Whether .pt weights can be used."""
|
||||
|
||||
allow_patterns_overrides: list[str] | None = None
|
||||
allow_patterns_overrides: Optional[list[str]] = None
|
||||
"""If defined, weights will load exclusively using these patterns."""
|
||||
|
||||
counter_before_loading_weights: float = 0.0
|
||||
@@ -116,8 +116,8 @@ class TextEncoderLoader(ComponentLoader):
|
||||
self,
|
||||
model_name_or_path: str,
|
||||
fall_back_to_pt: bool,
|
||||
allow_patterns_overrides: list[str] | None,
|
||||
) -> tuple[str, list[str], bool]:
|
||||
allow_patterns_overrides: Optional[list[str]],
|
||||
) -> Tuple[str, List[str], bool]:
|
||||
"""Prepare weights for the model.
|
||||
|
||||
If the model is not local, it will be downloaded."""
|
||||
@@ -139,7 +139,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
hf_folder = model_name_or_path
|
||||
|
||||
hf_weights_files: list[str] = []
|
||||
hf_weights_files: List[str] = []
|
||||
for pattern in allow_patterns:
|
||||
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
|
||||
if len(hf_weights_files) > 0:
|
||||
@@ -162,7 +162,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
def _get_weights_iterator(
|
||||
self, source: "Source"
|
||||
) -> Generator[tuple[str, torch.Tensor], None, None]:
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Get an iterator for the model weights based on the load format."""
|
||||
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
|
||||
source.model_or_path, source.fall_back_to_pt,
|
||||
@@ -182,7 +182,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
self,
|
||||
model_config: Any,
|
||||
model: nn.Module,
|
||||
) -> Generator[tuple[str, torch.Tensor], None, None]:
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
primary_weights = TextEncoderLoader.Source(
|
||||
model_config.model,
|
||||
prefix="",
|
||||
@@ -367,6 +367,7 @@ class TransformerLoader(ComponentLoader):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
hf_config = deepcopy(config)
|
||||
cls_name = config.pop("_class_name")
|
||||
if cls_name is None:
|
||||
raise ValueError(
|
||||
@@ -395,7 +396,10 @@ class TransformerLoader(ComponentLoader):
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params={"config": dit_config},
|
||||
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,
|
||||
|
||||
@@ -7,9 +7,9 @@
|
||||
import contextlib
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable, Generator, Hashable
|
||||
from itertools import chain
|
||||
from typing import Any
|
||||
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
|
||||
Optional, Tuple, Type)
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -52,7 +52,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
|
||||
|
||||
|
||||
def get_param_names_mapping(
|
||||
mapping_dict: dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
|
||||
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
|
||||
"""
|
||||
Creates a mapping function that transforms parameter names using regex patterns.
|
||||
|
||||
@@ -87,12 +87,12 @@ def get_param_names_mapping(
|
||||
|
||||
# TODO(PY): add compile option
|
||||
def load_fsdp_model(
|
||||
model_cls: type[nn.Module],
|
||||
init_params: dict[str, Any],
|
||||
weight_dir_list: list[str],
|
||||
model_cls: Type[nn.Module],
|
||||
init_params: Dict[str, Any],
|
||||
weight_dir_list: List[str],
|
||||
device: torch.device,
|
||||
cpu_offload: bool = False,
|
||||
default_dtype: torch.dtype | None = torch.bfloat16,
|
||||
default_dtype: Optional[torch.dtype] = torch.bfloat16,
|
||||
) -> torch.nn.Module:
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
@@ -121,6 +121,7 @@ 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
|
||||
|
||||
|
||||
@@ -129,7 +130,7 @@ def shard_model(
|
||||
*,
|
||||
cpu_offload: bool,
|
||||
reshard_after_forward: bool = True,
|
||||
dp_mesh: DeviceMesh | None = None,
|
||||
dp_mesh: Optional[DeviceMesh] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
|
||||
@@ -184,11 +185,11 @@ def shard_model(
|
||||
# TODO(PY): device mesh for cfg parallel
|
||||
def load_fsdp_model_from_full_model_state_dict(
|
||||
model: torch.nn.Module,
|
||||
full_sd_iterator: Generator[tuple[str, torch.Tensor], None, None],
|
||||
full_sd_iterator: Generator[Tuple[str, torch.Tensor], None, None],
|
||||
device: torch.device,
|
||||
strict: bool = False,
|
||||
cpu_offload: bool = False,
|
||||
param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None,
|
||||
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None,
|
||||
) -> _IncompatibleKeys:
|
||||
"""
|
||||
Converting full state dict into a sharded state dict
|
||||
@@ -212,7 +213,7 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
meta_sharded_sd = model.state_dict()
|
||||
|
||||
sharded_sd = {}
|
||||
to_merge_params: defaultdict[Hashable, dict[Any, Any]] = defaultdict(dict)
|
||||
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
|
||||
for source_param_name, full_tensor in full_sd_iterator:
|
||||
assert param_names_mapping is not None
|
||||
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
|
||||
|
||||
@@ -8,8 +8,8 @@ import os
|
||||
import tempfile
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from collections.abc import Generator
|
||||
from pathlib import Path
|
||||
from typing import Generator, List, Optional, Tuple, Union
|
||||
|
||||
import filelock
|
||||
import huggingface_hub.constants
|
||||
@@ -50,7 +50,8 @@ class DisabledTqdm(tqdm):
|
||||
super().__init__(*args, **kwargs, disable=True)
|
||||
|
||||
|
||||
def get_lock(model_name_or_path: str | Path, cache_dir: str | None = None):
|
||||
def get_lock(model_name_or_path: Union[str, Path],
|
||||
cache_dir: Optional[str] = None):
|
||||
lock_dir = cache_dir or temp_dir
|
||||
model_name_or_path = str(model_name_or_path)
|
||||
os.makedirs(os.path.dirname(lock_dir), exist_ok=True)
|
||||
@@ -76,10 +77,10 @@ def _shared_pointers(tensors):
|
||||
|
||||
def download_weights_from_hf(
|
||||
model_name_or_path: str,
|
||||
cache_dir: str | None,
|
||||
allow_patterns: list[str],
|
||||
revision: str | None = None,
|
||||
ignore_patterns: str | list[str] | None = None,
|
||||
cache_dir: Optional[str],
|
||||
allow_patterns: List[str],
|
||||
revision: Optional[str] = None,
|
||||
ignore_patterns: Optional[Union[str, List[str]]] = None,
|
||||
) -> str:
|
||||
"""Download model weights from Hugging Face Hub.
|
||||
|
||||
@@ -135,8 +136,8 @@ def download_weights_from_hf(
|
||||
def download_safetensors_index_file_from_hf(
|
||||
model_name_or_path: str,
|
||||
index_file: str,
|
||||
cache_dir: str | None,
|
||||
revision: str | None = None,
|
||||
cache_dir: Optional[str],
|
||||
revision: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Download hf safetensors index file from Hugging Face Hub.
|
||||
|
||||
@@ -171,9 +172,9 @@ def download_safetensors_index_file_from_hf(
|
||||
# Passing both of these to the weight loader functionality breaks.
|
||||
# So, we use the index_file to
|
||||
# look up which safetensors files should be used.
|
||||
def filter_duplicate_safetensors_files(hf_weights_files: list[str],
|
||||
def filter_duplicate_safetensors_files(hf_weights_files: List[str],
|
||||
hf_folder: str,
|
||||
index_file: str) -> list[str]:
|
||||
index_file: str) -> List[str]:
|
||||
# model.safetensors.index.json is a mapping from keys in the
|
||||
# torch state_dict to safetensors file holding that weight.
|
||||
index_file_name = os.path.join(hf_folder, index_file)
|
||||
@@ -196,7 +197,7 @@ def filter_duplicate_safetensors_files(hf_weights_files: list[str],
|
||||
|
||||
|
||||
def filter_files_not_needed_for_inference(
|
||||
hf_weights_files: list[str]) -> list[str]:
|
||||
hf_weights_files: List[str]) -> List[str]:
|
||||
"""
|
||||
Exclude files that are not needed for inference.
|
||||
|
||||
@@ -224,8 +225,8 @@ _BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elap
|
||||
|
||||
|
||||
def safetensors_weights_iterator(
|
||||
hf_weights_files: list[str]
|
||||
) -> Generator[tuple[str, torch.Tensor], None, None]:
|
||||
hf_weights_files: List[str]
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model safetensor files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
@@ -242,8 +243,8 @@ def safetensors_weights_iterator(
|
||||
|
||||
|
||||
def pt_weights_iterator(
|
||||
hf_weights_files: list[str]
|
||||
) -> Generator[tuple[str, torch.Tensor], None, None]:
|
||||
hf_weights_files: List[str]
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model bin/pt files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
@@ -279,7 +280,7 @@ def default_weight_loader(param: torch.Tensor,
|
||||
raise
|
||||
|
||||
|
||||
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> str | None:
|
||||
def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
|
||||
"""Remap the name of FP8 k/v_scale parameters.
|
||||
|
||||
This function handles the remapping of FP8 k/v_scale parameter names.
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
|
||||
|
||||
from collections.abc import Callable
|
||||
from fractions import Fraction
|
||||
from typing import Any
|
||||
from typing import Any, Callable, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Parameter
|
||||
@@ -113,8 +112,9 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
if shard_offset is None or shard_size is None:
|
||||
raise ValueError("shard_offset and shard_size must be provided")
|
||||
if isinstance(
|
||||
self, PackedColumnParameter
|
||||
| PackedvLLMParameter) and self.packed_dim == self.output_dim:
|
||||
self,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and self.packed_dim == self.output_dim:
|
||||
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
|
||||
shard_offset=shard_offset, shard_size=shard_size)
|
||||
|
||||
@@ -141,8 +141,9 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
assert num_heads is not None
|
||||
|
||||
if isinstance(
|
||||
self, PackedColumnParameter
|
||||
| PackedvLLMParameter) and self.output_dim == self.packed_dim:
|
||||
self,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and self.output_dim == self.packed_dim:
|
||||
shard_size, shard_offset = self.adjust_shard_indexes_for_packing(
|
||||
shard_offset=shard_offset, shard_size=shard_size)
|
||||
|
||||
@@ -229,7 +230,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
|
||||
self.qkv_idxs = {"q": 0, "k": 1, "v": 2}
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _shard_id_as_int(self, shard_id: str | int) -> int:
|
||||
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int:
|
||||
if isinstance(shard_id, int):
|
||||
return shard_id
|
||||
|
||||
@@ -254,7 +255,7 @@ class PerTensorScaleParameter(BasevLLMParameter):
|
||||
super().load_row_parallel_weight(*args, **kwargs)
|
||||
|
||||
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
|
||||
shard_id: str | int, **kwargs):
|
||||
shard_id: Union[str, int], **kwargs):
|
||||
"""
|
||||
Slice the parameter data based on the shard id for
|
||||
loading.
|
||||
@@ -281,7 +282,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
|
||||
for more details on the packed properties.
|
||||
"""
|
||||
|
||||
def __init__(self, packed_factor: int | Fraction, packed_dim: int,
|
||||
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
|
||||
**kwargs):
|
||||
self._packed_factor = packed_factor
|
||||
self._packed_dim = packed_dim
|
||||
@@ -296,7 +297,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
|
||||
return self._packed_factor
|
||||
|
||||
def adjust_shard_indexes_for_packing(self, shard_size,
|
||||
shard_offset) -> tuple[Any, Any]:
|
||||
shard_offset) -> Tuple[Any, Any]:
|
||||
return _adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
@@ -314,7 +315,7 @@ class PackedvLLMParameter(ModelWeightParameter):
|
||||
by accounting for packing and optionally, marlin tile size.
|
||||
"""
|
||||
|
||||
def __init__(self, packed_factor: int | Fraction, packed_dim: int,
|
||||
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
|
||||
**kwargs):
|
||||
self._packed_factor = packed_factor
|
||||
self._packed_dim = packed_dim
|
||||
@@ -403,7 +404,7 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_packing(shard_size, shard_offset,
|
||||
packed_factor) -> tuple[Any, Any]:
|
||||
packed_factor) -> Tuple[Any, Any]:
|
||||
shard_size = shard_size // packed_factor
|
||||
shard_offset = shard_offset // packed_factor
|
||||
return shard_size, shard_offset
|
||||
|
||||
@@ -8,10 +8,10 @@ import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable, Set
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
from typing import NoReturn, TypeVar, cast
|
||||
from typing import (AbstractSet, Callable, Dict, List, NoReturn, Optional,
|
||||
Tuple, Type, TypeVar, Union, cast)
|
||||
|
||||
import cloudpickle
|
||||
from torch import nn
|
||||
@@ -80,7 +80,7 @@ class _ModelInfo:
|
||||
architecture: str
|
||||
|
||||
@staticmethod
|
||||
def from_model_cls(model: type[nn.Module]) -> "_ModelInfo":
|
||||
def from_model_cls(model: Type[nn.Module]) -> "_ModelInfo":
|
||||
return _ModelInfo(architecture=model.__name__, )
|
||||
|
||||
|
||||
@@ -91,7 +91,7 @@ class _BaseRegisteredModel(ABC):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def load_model_cls(self) -> type[nn.Module]:
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -102,10 +102,10 @@ class _RegisteredModel(_BaseRegisteredModel):
|
||||
"""
|
||||
|
||||
interfaces: _ModelInfo
|
||||
model_cls: type[nn.Module]
|
||||
model_cls: Type[nn.Module]
|
||||
|
||||
@staticmethod
|
||||
def from_model_cls(model_cls: type[nn.Module]):
|
||||
def from_model_cls(model_cls: Type[nn.Module]):
|
||||
return _RegisteredModel(
|
||||
interfaces=_ModelInfo.from_model_cls(model_cls),
|
||||
model_cls=model_cls,
|
||||
@@ -114,7 +114,7 @@ class _RegisteredModel(_BaseRegisteredModel):
|
||||
def inspect_model_cls(self) -> _ModelInfo:
|
||||
return self.interfaces
|
||||
|
||||
def load_model_cls(self) -> type[nn.Module]:
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
return self.model_cls
|
||||
|
||||
|
||||
@@ -159,16 +159,16 @@ class _LazyRegisteredModel(_BaseRegisteredModel):
|
||||
return _run_in_subprocess(
|
||||
lambda: _ModelInfo.from_model_cls(self.load_model_cls()))
|
||||
|
||||
def load_model_cls(self) -> type[nn.Module]:
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
mod = importlib.import_module(self.module_name)
|
||||
return cast(type[nn.Module], getattr(mod, self.class_name))
|
||||
return cast(Type[nn.Module], getattr(mod, self.class_name))
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
def _try_load_model_cls(
|
||||
model_arch: str,
|
||||
model: _BaseRegisteredModel,
|
||||
) -> type[nn.Module] | None:
|
||||
) -> Optional[Type[nn.Module]]:
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
current_platform.verify_model_arch(model_arch)
|
||||
try:
|
||||
@@ -182,7 +182,7 @@ def _try_load_model_cls(
|
||||
def _try_inspect_model_cls(
|
||||
model_arch: str,
|
||||
model: _BaseRegisteredModel,
|
||||
) -> _ModelInfo | None:
|
||||
) -> Optional[_ModelInfo]:
|
||||
try:
|
||||
return model.inspect_model_cls()
|
||||
except Exception:
|
||||
@@ -194,15 +194,15 @@ def _try_inspect_model_cls(
|
||||
@dataclass
|
||||
class _ModelRegistry:
|
||||
# Keyed by model_arch
|
||||
models: dict[str, _BaseRegisteredModel] = field(default_factory=dict)
|
||||
models: Dict[str, _BaseRegisteredModel] = field(default_factory=dict)
|
||||
|
||||
def get_supported_archs(self) -> Set[str]:
|
||||
def get_supported_archs(self) -> AbstractSet[str]:
|
||||
return self.models.keys()
|
||||
|
||||
def register_model(
|
||||
self,
|
||||
model_arch: str,
|
||||
model_cls: type[nn.Module] | str,
|
||||
model_cls: Union[Type[nn.Module], str],
|
||||
) -> None:
|
||||
"""
|
||||
Register an external model to be used in vLLM.
|
||||
@@ -232,7 +232,7 @@ class _ModelRegistry:
|
||||
|
||||
self.models[model_arch] = model
|
||||
|
||||
def _raise_for_unsupported(self, architectures: list[str]) -> NoReturn:
|
||||
def _raise_for_unsupported(self, architectures: List[str]) -> NoReturn:
|
||||
all_supported_archs = self.get_supported_archs()
|
||||
|
||||
if any(arch in all_supported_archs for arch in architectures):
|
||||
@@ -244,13 +244,13 @@ class _ModelRegistry:
|
||||
f"Model architectures {architectures} are not supported for now. "
|
||||
f"Supported architectures: {all_supported_archs}")
|
||||
|
||||
def _try_load_model_cls(self, model_arch: str) -> type[nn.Module] | None:
|
||||
def _try_load_model_cls(self, model_arch: str) -> Optional[Type[nn.Module]]:
|
||||
if model_arch not in self.models:
|
||||
return None
|
||||
|
||||
return _try_load_model_cls(model_arch, self.models[model_arch])
|
||||
|
||||
def _try_inspect_model_cls(self, model_arch: str) -> _ModelInfo | None:
|
||||
def _try_inspect_model_cls(self, model_arch: str) -> Optional[_ModelInfo]:
|
||||
if model_arch not in self.models:
|
||||
return None
|
||||
|
||||
@@ -258,8 +258,8 @@ class _ModelRegistry:
|
||||
|
||||
def _normalize_archs(
|
||||
self,
|
||||
architectures: str | list[str],
|
||||
) -> list[str]:
|
||||
architectures: Union[str, List[str]],
|
||||
) -> List[str]:
|
||||
if isinstance(architectures, str):
|
||||
architectures = [architectures]
|
||||
if not architectures:
|
||||
@@ -274,8 +274,8 @@ class _ModelRegistry:
|
||||
|
||||
def inspect_model_cls(
|
||||
self,
|
||||
architectures: str | list[str],
|
||||
) -> tuple[_ModelInfo, str]:
|
||||
architectures: Union[str, List[str]],
|
||||
) -> Tuple[_ModelInfo, str]:
|
||||
architectures = self._normalize_archs(architectures)
|
||||
|
||||
for arch in architectures:
|
||||
@@ -287,8 +287,8 @@ class _ModelRegistry:
|
||||
|
||||
def resolve_model_cls(
|
||||
self,
|
||||
architectures: str | list[str],
|
||||
) -> tuple[type[nn.Module], str]:
|
||||
architectures: Union[str, List[str]],
|
||||
) -> Tuple[Type[nn.Module], str]:
|
||||
architectures = self._normalize_archs(architectures)
|
||||
|
||||
for arch in architectures:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.utils import BaseOutput
|
||||
@@ -31,15 +32,15 @@ class BaseScheduler(ABC):
|
||||
@abstractmethod
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: int | None = None) -> torch.Tensor:
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
timestep: int | torch.Tensor,
|
||||
timestep: Union[int, torch.Tensor],
|
||||
sample: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
) -> BaseOutput | tuple:
|
||||
) -> Union[BaseOutput, Tuple]:
|
||||
pass
|
||||
|
||||
@@ -20,7 +20,7 @@
|
||||
# ==============================================================================
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
@@ -75,7 +75,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
shift: float = 1.0,
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: int | None = None,
|
||||
n_tokens: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
@@ -130,7 +130,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
device: str | torch.device = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
n_tokens: int = 0,
|
||||
):
|
||||
"""
|
||||
@@ -193,7 +193,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: int | None = None) -> torch.Tensor:
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def sd3_time_shift(self, t: torch.Tensor):
|
||||
@@ -202,11 +202,11 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: float | torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
**kwargs,
|
||||
) -> FlowMatchDiscreteSchedulerOutput | tuple:
|
||||
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
@@ -232,7 +232,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if isinstance(timestep, int | torch.IntTensor | torch.LongTensor):
|
||||
if isinstance(timestep, (int, torch.IntTensor, torch.LongTensor)):
|
||||
raise ValueError((
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
|
||||
@@ -23,6 +23,7 @@
|
||||
# ==============================================================================
|
||||
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -202,7 +203,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
beta_start: float = 0.0001,
|
||||
beta_end: float = 0.02,
|
||||
beta_schedule: str = "linear",
|
||||
trained_betas: np.ndarray | list[float] | None = None,
|
||||
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
|
||||
solver_order: int = 2,
|
||||
prediction_type: str = "epsilon",
|
||||
thresholding: bool = False,
|
||||
@@ -211,16 +212,16 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
predict_x0: bool = True,
|
||||
solver_type: str = "bh2",
|
||||
lower_order_final: bool = True,
|
||||
disable_corrector: tuple[int, ...] = (),
|
||||
disable_corrector: Tuple[int, ...] = (),
|
||||
solver_p: SchedulerMixin = None,
|
||||
use_karras_sigmas: bool | None = False,
|
||||
use_exponential_sigmas: bool | None = False,
|
||||
use_beta_sigmas: bool | None = False,
|
||||
use_flow_sigmas: bool | None = False,
|
||||
flow_shift: float | None = 1.0,
|
||||
use_karras_sigmas: Optional[bool] = False,
|
||||
use_exponential_sigmas: Optional[bool] = False,
|
||||
use_beta_sigmas: Optional[bool] = False,
|
||||
use_flow_sigmas: Optional[bool] = False,
|
||||
flow_shift: Optional[float] = 1.0,
|
||||
timestep_spacing: str = "linspace",
|
||||
steps_offset: int = 0,
|
||||
final_sigmas_type: str | None = "zero", # "zero", "sigma_min"
|
||||
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
|
||||
rescale_betas_zero_snr: bool = False,
|
||||
):
|
||||
if self.config.use_beta_sigmas and not is_scipy_available():
|
||||
@@ -282,20 +283,21 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
|
||||
self.predict_x0 = predict_x0
|
||||
# setable values
|
||||
self.num_inference_steps: int | None = None
|
||||
self.num_inference_steps: Optional[int] = None
|
||||
timesteps = np.linspace(0,
|
||||
num_train_timesteps - 1,
|
||||
num_train_timesteps,
|
||||
dtype=np.float32)[::-1].copy()
|
||||
self.timesteps = torch.from_numpy(timesteps)
|
||||
self.model_outputs = [None] * solver_order
|
||||
self.timestep_list: list[int | torch.Tensor] = [None] * solver_order
|
||||
self.timestep_list: List[Union[int,
|
||||
torch.Tensor]] = [None] * solver_order
|
||||
self.lower_order_nums = 0
|
||||
self.disable_corrector = list(disable_corrector)
|
||||
self.solver_p = solver_p
|
||||
self.last_sample = None
|
||||
self._step_index: int | None = None
|
||||
self._begin_index: int | None = None
|
||||
self._step_index: Optional[int] = None
|
||||
self._begin_index: Optional[int] = None
|
||||
self.sigmas = self.sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
|
||||
@@ -331,7 +333,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
|
||||
def set_timesteps(self,
|
||||
num_inference_steps: int,
|
||||
device: str | torch.device = None):
|
||||
device: Union[str, torch.device] = None):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
@@ -535,7 +537,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._sigma_to_alpha_sigma_t
|
||||
def _sigma_to_alpha_sigma_t(
|
||||
self, sigma: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
self, sigma: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
if self.config.use_flow_sigmas:
|
||||
alpha_t = 1 - sigma
|
||||
sigma_t = sigma
|
||||
@@ -706,7 +708,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
model_output: torch.Tensor,
|
||||
*args,
|
||||
sample: torch.Tensor = None,
|
||||
order: int | None = None,
|
||||
order: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
@@ -806,7 +808,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
R_tensor: torch.Tensor = torch.stack(R)
|
||||
b = torch.tensor(b, device=device)
|
||||
|
||||
D1s_tensor: torch.Tensor | None = None
|
||||
D1s_tensor: Optional[torch.Tensor] = None
|
||||
if len(D1s) > 0:
|
||||
D1s_tensor = torch.stack(D1s, dim=1) # (B, K)
|
||||
# for order 2, we use a simplified version
|
||||
@@ -840,9 +842,9 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
self,
|
||||
this_model_output: torch.Tensor,
|
||||
*args,
|
||||
last_sample: torch.Tensor | None = None,
|
||||
this_sample: torch.Tensor | None = None,
|
||||
order: int | None = None,
|
||||
last_sample: Optional[torch.Tensor] = None,
|
||||
this_sample: Optional[torch.Tensor] = None,
|
||||
order: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
@@ -948,7 +950,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
R = torch.stack(R)
|
||||
b = torch.tensor(b, device=device)
|
||||
|
||||
D1s_tensor: torch.Tensor | None = torch.stack(
|
||||
D1s_tensor: Optional[torch.Tensor] = torch.stack(
|
||||
D1s, dim=1) if len(D1s) > 0 else None
|
||||
|
||||
# for order 1, we use a simplified version
|
||||
@@ -1014,10 +1016,10 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
timestep: int | torch.Tensor,
|
||||
timestep: Union[int, torch.Tensor],
|
||||
sample: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
) -> SchedulerOutput | tuple:
|
||||
) -> Union[SchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
|
||||
the multistep UniPC.
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py
|
||||
"""Utils for model executor."""
|
||||
from typing import Any
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -58,7 +58,7 @@ def set_random_seed(seed: int) -> None:
|
||||
|
||||
def set_weight_attrs(
|
||||
weight: torch.Tensor,
|
||||
weight_attrs: dict[str, Any] | None,
|
||||
weight_attrs: Optional[Dict[str, Any]],
|
||||
):
|
||||
"""Set attributes on a weight tensor.
|
||||
|
||||
@@ -109,7 +109,7 @@ def extract_layer_index(layer_name: str) -> int:
|
||||
- "model.encoder.layers.0.sub.1" -> ValueError
|
||||
"""
|
||||
subnames = layer_name.split(".")
|
||||
int_vals: list[int] = []
|
||||
int_vals: List[int] = []
|
||||
for subname in subnames:
|
||||
try:
|
||||
int_vals.append(int(subname))
|
||||
@@ -121,8 +121,8 @@ def extract_layer_index(layer_name: str) -> int:
|
||||
|
||||
|
||||
def modulate(x: torch.Tensor,
|
||||
shift: torch.Tensor | None = None,
|
||||
scale: torch.Tensor | None = None) -> torch.Tensor:
|
||||
shift: Optional[torch.Tensor] = None,
|
||||
scale: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
"""modulate by shift and scale
|
||||
|
||||
Args:
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Iterator
|
||||
from math import prod
|
||||
from typing import Optional, cast
|
||||
from typing import Iterator, Optional, Tuple, Union, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -40,9 +39,6 @@ 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)
|
||||
@@ -52,8 +48,8 @@ class ParallelTiledVAE(ABC):
|
||||
return cast(int, self.config.spatial_compression_ratio)
|
||||
|
||||
@property
|
||||
def scaling_factor(self) -> float | torch.Tensor:
|
||||
return cast(float | torch.Tensor, self.config.scaling_factor)
|
||||
def scaling_factor(self) -> Union[float, torch.tensor]:
|
||||
return cast(Union[float, torch.tensor], self.config.scaling_factor)
|
||||
|
||||
@abstractmethod
|
||||
def _encode(self, *args, **kwargs) -> torch.Tensor:
|
||||
@@ -161,7 +157,7 @@ class ParallelTiledVAE(ABC):
|
||||
|
||||
def _parallel_data_generator(
|
||||
self, gathered_results,
|
||||
gathered_dim_metadata) -> Iterator[tuple[torch.Tensor, int]]:
|
||||
gathered_dim_metadata) -> Iterator[Tuple[torch.Tensor, int]]:
|
||||
global_idx = 0
|
||||
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
|
||||
_start_shape = 0
|
||||
@@ -413,16 +409,16 @@ class ParallelTiledVAE(ABC):
|
||||
|
||||
def enable_tiling(
|
||||
self,
|
||||
tile_sample_min_height: int | None = None,
|
||||
tile_sample_min_width: int | None = None,
|
||||
tile_sample_min_num_frames: int | None = None,
|
||||
tile_sample_stride_height: int | None = None,
|
||||
tile_sample_stride_width: int | None = None,
|
||||
tile_sample_stride_num_frames: int | None = None,
|
||||
blend_num_frames: int | None = None,
|
||||
use_tiling: bool | None = None,
|
||||
use_temporal_tiling: bool | None = None,
|
||||
use_parallel_tiling: bool | None = None,
|
||||
tile_sample_min_height: Optional[int] = None,
|
||||
tile_sample_min_width: Optional[int] = None,
|
||||
tile_sample_min_num_frames: Optional[int] = None,
|
||||
tile_sample_stride_height: Optional[int] = None,
|
||||
tile_sample_stride_width: Optional[int] = None,
|
||||
tile_sample_stride_num_frames: Optional[int] = None,
|
||||
blend_num_frames: Optional[int] = None,
|
||||
use_tiling: Optional[bool] = None,
|
||||
use_temporal_tiling: Optional[bool] = None,
|
||||
use_parallel_tiling: Optional[bool] = None,
|
||||
) -> None:
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
@@ -486,7 +482,8 @@ class DiagonalGaussianDistribution:
|
||||
device=self.parameters.device,
|
||||
dtype=self.parameters.dtype)
|
||||
|
||||
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
def sample(self,
|
||||
generator: Optional[torch.Generator] = None) -> torch.Tensor:
|
||||
# make sure sample is on the same device as the parameters and has same dtype
|
||||
sample = randn_tensor(
|
||||
self.mean.shape,
|
||||
@@ -517,7 +514,7 @@ class DiagonalGaussianDistribution:
|
||||
|
||||
def nll(
|
||||
self, sample: torch.Tensor,
|
||||
dims: tuple[int, ...] = (1, 2, 3)) -> torch.Tensor:
|
||||
dims: Tuple[int, ...] = (1, 2, 3)) -> torch.Tensor:
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user