Compare commits

..
Author SHA1 Message Date
SolitaryThinker d0c53871a4 fix tensor type hint 2025-05-23 14:42:23 -07:00
SolitaryThinker 0d5306f61f update min python to 3.10 2025-05-23 14:42:23 -07:00
162 changed files with 44096 additions and 51978 deletions
+8 -8
View File
@@ -4,6 +4,14 @@ title: "[Bug] "
labels: ['Bug']
body:
- type: textarea
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.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
- type: textarea
attributes:
label: Describe the bug
@@ -17,13 +25,5 @@ body:
What command or script did you run? Which **model** are you using?
placeholder: |
A placeholder for the command.
validations:
required: true
- type: textarea
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.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
-2
View File
@@ -77,8 +77,6 @@ jobs:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
encoder-test:
needs: change-filter
+1 -1
View File
@@ -10,7 +10,7 @@ jobs:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
python-version: "3.10"
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
- uses: pre-commit/action@v3.0.1
+1
View File
@@ -27,6 +27,7 @@ env
**/build/
**.pyc
**.txt
**.json
# Distribution / packaging
build/
+2 -2
View File
@@ -33,7 +33,7 @@ repos:
args: [--in-place, --verbose]
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.11.12
rev: v0.11.4
hooks:
- id: ruff
args: [--output-format, github, --fix]
@@ -48,7 +48,7 @@ repos:
hooks:
- id: isort
- repo: https://github.com/jackdewinter/pymarkdown
rev: v0.9.30
rev: v0.9.29
hooks:
- id: pymarkdown
args: [fix]
+42491 -42491
View File
File diff suppressed because it is too large Load Diff
+1 -3
View File
@@ -18,7 +18,6 @@ import os
import re
import sys
from pathlib import Path
from typing import Optional
import requests
@@ -168,8 +167,7 @@ _cached_base: str = ""
_cached_branch: str = ""
def get_repo_base_and_branch(
pr_number: str) -> tuple[Optional[str], Optional[str]]:
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
global _cached_base, _cached_branch
if _cached_base and _cached_branch:
return _cached_base, _cached_branch
+1 -2
View File
@@ -5,7 +5,6 @@ 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 = '../../../..'
@@ -89,7 +88,7 @@ class Example:
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: Optional[str] = None
category: str | None = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
@@ -57,9 +57,8 @@ Run the script with:
python example.py
```
The generated video will be saved in the current directory under `my_videos/`
The generated video will be saved in the current directory under `my_videos/`.
More inference example scripts can be found in `scripts/inference/`
## Available Models
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
@@ -80,6 +79,7 @@ def main():
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
+4 -6
View File
@@ -2,7 +2,7 @@ from fastvideo import VideoGenerator
# from fastvideo.v1.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -11,9 +11,7 @@ def main():
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=2,
use_fsdp_inference=True,
use_cpu_offload=False
num_gpus=1,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
@@ -25,7 +23,7 @@ def main():
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
video = generator.generate_video(prompt)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
@@ -36,7 +34,7 @@ def main():
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
video2 = generator.generate_video(prompt2)
if __name__ == "__main__":
+1 -1
View File
@@ -1,5 +1,5 @@
from fastvideo import VideoGenerator
from fastvideo.v1.configs.pipelines.base import PipelineConfig
def main():
@@ -1,45 +0,0 @@
from fastvideo import VideoGenerator
from fastvideo.v1.configs.sample import SamplingParam
OUTPUT_PATH = "./lora"
def main():
# Initialize VideoGenerator with the Wan model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=2,
lora_path="benjamin-paine/steamboat-willie-1.3b",
lora_nickname="steamboat"
)
kwargs = {
"height": 480,
"width": 832,
"num_frames": 81,
"guidance_scale": 5.0,
"num_inference_steps": 32,
}
# Generate video with LoRA style
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
video = generator.generate_video(
prompt,
# sampling_param=sampling_param,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt=negative_prompt,
**kwargs
)
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
video = generator.generate_video(
prompt,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt=negative_prompt,
**kwargs
)
if __name__ == "__main__":
main()
@@ -1,5 +0,0 @@
# STA Mask Search Examples
```bash
bash examples/inference/sta_mask_search/inference_wan_sta.sh
```
@@ -1,39 +0,0 @@
#!/bin/bash
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
base_port=29503
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
gpu_ids=$(seq 0 $((num_gpu-1)))
skip_time_steps=12
output_path="inference_results/sta/mask_search_full"
STA_mode="STA_searching"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode &
sleep 1
done
wait
echo "STA searching completed"
output_path="inference_results/sta/mask_search_sparse"
STA_mode="STA_tuning"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode \
--skip_time_steps $skip_time_steps &
sleep 1
done
wait
echo "STA tuning completed"
echo "All jobs completed"
@@ -1,63 +0,0 @@
import os
import argparse
from fastvideo import VideoGenerator, SamplingParam
def main(args):
os.makedirs(args.output_path, exist_ok=True)
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
num_gpus=args.num_gpus, # Adjust based on your hardware
STA_mode=args.STA_mode,
skip_time_steps=args.skip_time_steps
)
# Prompts for your video
prompt = args.prompt
prompt_path = args.prompt_path
negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
if prompt_path is not None:
with open(prompt_path, "r") as f:
prompts = f.readlines()
else:
prompts = [prompt]
params = SamplingParam(
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
fps=args.fps,
guidance_scale=args.guidance_scale,
seed=args.seed,
return_frames=True, # Also return frames from this call (defaults to False)
output_path=args.output_path, # Controls where videos are saved
save_video=True,
negative_prompt=negative_prompt
)
# Generate the video
for prompt in prompts:
video = generator.generate_video(
prompt,
sampling_param=params,
)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--prompt", type=str, default="A man is dancing.")
parser.add_argument("--prompt_path", type=str, default=None)
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1280)
parser.add_argument("--num_frames", type=int, default=69)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--guidance_scale", type=float, default=5.0)
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--output_path", type=str, default="my_videos/")
parser.add_argument("--num_gpus", type=int, default=1)
parser.add_argument("--STA_mode", type=str, default="STA_searching")
parser.add_argument("--skip_time_steps", type=int, default=12)
args = parser.parse_args()
main(args)
-120
View File
@@ -1,120 +0,0 @@
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.preprocess_pipeline_i2v import PreprocessPipeline_I2V
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
logger = init_logger(__name__)
def main(args):
args.model_path = maybe_download_model(args.model_path)
# Assume using torchrun
local_rank = int(os.getenv("RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=args.model_path,
num_gpus=world_size,
device_str="cuda",
**pipeline_config_args,
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--preprocess_video_batch_size",
type=int,
default=2,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--preprocess_text_batch_size",
type=int,
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--samples_per_file",
type=int,
default=64
)
parser.add_argument(
"--flush_frequency",
type=int,
default=256,
help="how often to save to parquet files"
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--preprocess_task", type=str, 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)
@@ -68,8 +68,7 @@ def main(args):
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)
if args.model_type != "wan":
vae.enable_tiling()
vae.enable_tiling()
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
@@ -33,8 +33,7 @@ def main(args):
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
if args.model_type != "wan":
vae.enable_tiling()
vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
+40 -74
View File
@@ -12,7 +12,6 @@ import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers import FlowMatchEulerDiscreteScheduler
from fastvideo.distill.solver import PCMFMScheduler
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from peft import LoraConfig
@@ -24,7 +23,7 @@ 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.latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
@@ -124,21 +123,13 @@ def distill_one_step(
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
if args.model_type == "wan":
teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"return_dict": True,
}
else:
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,
}
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,
@@ -150,70 +141,47 @@ def distill_one_step(
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
if args.model_type == "wan":
cond_teacher_kwargs ={
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"return_dict": True,
}
else:
cond_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,
}
cond_teacher_output = teacher_transformer(**cond_teacher_kwargs)[0].float()
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
if args.model_type == "wan":
uncond_teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
"timestep": timesteps,
"return_dict": True,
}
else:
uncond_teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
"timestep": timesteps,
"encoder_attention_mask": uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
"return_dict": False,
}
uncond_teacher_output = teacher_transformer(**uncond_teacher_kwargs)[0].float()
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if args.model_type == "wan":
target_pred_kwargs = {
"hidden_states": x_prev.float(),
"encoder_hidden_states": encoder_hidden_states,
"timestep":timesteps_prev,
"return_dict":True,
}
else:
target_pred_kwargs = {
"hidden_states": x_prev.float(),
"encoder_hidden_states": encoder_hidden_states,
"timestep":timesteps_prev,
"encoder_attention_mask":encoder_attention_mask,
"return_dict":False,
}
if ema_transformer is not None:
target_pred = ema_transformer(**target_pred_kwargs)[0]
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
target_pred = transformer(**target_pred_kwargs)[0]
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
@@ -274,7 +242,7 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank == 0 and args.output_dir is not None:
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
@@ -351,9 +319,7 @@ def main(args):
teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler()
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
sigmas = linear_quadratic_schedule(
@@ -425,7 +391,7 @@ def main(args):
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:
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@@ -527,7 +493,7 @@ def main(args):
"phases": num_phases,
})
progress_bar.update(1)
if rank == 0:
if rank <= 0:
wandb.log(
{
"train_loss": loss,
+4 -4
View File
@@ -23,7 +23,7 @@ from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.utils.latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
save_lora_checkpoint)
@@ -296,7 +296,7 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank == 0 and args.output_dir is not None:
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
@@ -462,7 +462,7 @@ def main(args):
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:
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@@ -559,7 +559,7 @@ def main(args):
"step_time": f"{step_time:.2f}s",
})
progress_bar.update(1)
if rank == 0:
if rank <= 0:
wandb.log(
{
"generator_loss": generator_loss,
-486
View File
@@ -1,486 +0,0 @@
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from typing import Any, Dict, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
from diffusers.models.attention import FeedForward
from diffusers.models.attention_processor import Attention
from diffusers.models.cache_utils import CacheMixin
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import FP32LayerNorm
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class WanAttnProcessor2_0:
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("WanAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
encoder_hidden_states_img = None
if attn.add_k_proj is not None:
# 512 is the context length of the text encoder, hardcoded for now
image_context_length = encoder_hidden_states.shape[1] - 512
encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length]
encoder_hidden_states = encoder_hidden_states[:, image_context_length:]
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
query = attn.to_q(hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
if rotary_emb is not None:
def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor):
x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2)))
x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4)
return x_out.type_as(hidden_states)
query = apply_rotary_emb(query, rotary_emb)
key = apply_rotary_emb(key, rotary_emb)
# I2V task
hidden_states_img = None
if encoder_hidden_states_img is not None:
key_img = attn.add_k_proj(encoder_hidden_states_img)
key_img = attn.norm_added_k(key_img)
value_img = attn.add_v_proj(encoder_hidden_states_img)
key_img = key_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
value_img = value_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
hidden_states_img = F.scaled_dot_product_attention(
query, key_img, value_img, attn_mask=None, dropout_p=0.0, is_causal=False
)
hidden_states_img = hidden_states_img.transpose(1, 2).flatten(2, 3)
hidden_states_img = hidden_states_img.type_as(query)
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
hidden_states = hidden_states.type_as(query)
if hidden_states_img is not None:
hidden_states = hidden_states + hidden_states_img
hidden_states = attn.to_out[0](hidden_states)
hidden_states = attn.to_out[1](hidden_states)
return hidden_states
class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None):
super().__init__()
self.norm1 = FP32LayerNorm(in_features)
self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu")
self.norm2 = FP32LayerNorm(out_features)
if pos_embed_seq_len is not None:
self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features))
else:
self.pos_embed = None
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
if self.pos_embed is not None:
batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape
encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim)
encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states)
return hidden_states
class WanTimeTextImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
time_proj_dim: int,
text_embed_dim: int,
image_embed_dim: Optional[int] = None,
pos_embed_seq_len: Optional[int] = None,
):
super().__init__()
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim)
self.act_fn = nn.SiLU()
self.time_proj = nn.Linear(dim, time_proj_dim)
self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len)
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
):
timestep = self.timesteps_proj(timestep)
time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype
if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8:
timestep = timestep.to(time_embedder_dtype)
temb = self.time_embedder(timestep).type_as(encoder_hidden_states)
timestep_proj = self.time_proj(self.act_fn(temb))
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
if encoder_hidden_states_image is not None:
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
class WanRotaryPosEmbed(nn.Module):
def __init__(
self, attention_head_dim: int, patch_size: Tuple[int, int, int], max_seq_len: int, theta: float = 10000.0
):
super().__init__()
self.attention_head_dim = attention_head_dim
self.patch_size = patch_size
self.max_seq_len = max_seq_len
h_dim = w_dim = 2 * (attention_head_dim // 6)
t_dim = attention_head_dim - h_dim - w_dim
freqs = []
for dim in [t_dim, h_dim, w_dim]:
freq = get_1d_rotary_pos_embed(
dim, max_seq_len, theta, use_real=False, repeat_interleave_real=False, freqs_dtype=torch.float64
)
freqs.append(freq)
self.freqs = torch.cat(freqs, dim=1)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w
freqs = self.freqs.to(hidden_states.device)
freqs = freqs.split_with_sizes(
[
self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6),
self.attention_head_dim // 6,
self.attention_head_dim // 6,
],
dim=1,
)
freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1)
freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1)
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
return freqs
class WanTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.attn1 = Attention(
query_dim=dim,
heads=num_heads,
kv_heads=num_heads,
dim_head=dim // num_heads,
qk_norm=qk_norm,
eps=eps,
bias=True,
cross_attention_dim=None,
out_bias=True,
processor=WanAttnProcessor2_0(),
)
# 2. Cross-attention
self.attn2 = Attention(
query_dim=dim,
heads=num_heads,
kv_heads=num_heads,
dim_head=dim // num_heads,
qk_norm=qk_norm,
eps=eps,
bias=True,
cross_attention_dim=None,
out_bias=True,
added_kv_proj_dim=added_kv_proj_dim,
added_proj_bias=True,
processor=WanAttnProcessor2_0(),
)
self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
# 3. Feed-forward
self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate")
self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
rotary_emb: torch.Tensor,
) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table + temb.float()
).chunk(6, dim=1)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states)
attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb)
hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states)
# 2. Cross-attention
norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states)
attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states)
hidden_states = hidden_states + attn_output
# 3. Feed-forward
norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as(
hidden_states
)
ff_output = self.ffn(norm_hidden_states)
hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states)
return hidden_states
class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin):
r"""
A Transformer model for video-like data used in the Wan model.
Args:
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
num_attention_heads (`int`, defaults to `40`):
Fixed length for text embeddings.
attention_head_dim (`int`, defaults to `128`):
The number of channels in each head.
in_channels (`int`, defaults to `16`):
The number of channels in the input.
out_channels (`int`, defaults to `16`):
The number of channels in the output.
text_dim (`int`, defaults to `512`):
Input dimension for text embeddings.
freq_dim (`int`, defaults to `256`):
Dimension for sinusoidal time embeddings.
ffn_dim (`int`, defaults to `13824`):
Intermediate dimension in feed-forward network.
num_layers (`int`, defaults to `40`):
The number of layers of transformer blocks to use.
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
Window size for local attention (-1 indicates global attention).
cross_attn_norm (`bool`, defaults to `True`):
Enable cross-attention normalization.
qk_norm (`bool`, defaults to `True`):
Enable query/key normalization.
eps (`float`, defaults to `1e-6`):
Epsilon value for normalization layers.
add_img_emb (`bool`, defaults to `False`):
Whether to use img_emb.
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the added key and value projections. If `None`, no projection is used.
"""
_supports_gradient_checkpointing = True
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
_no_split_modules = ["WanTransformerBlock"]
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
@register_to_config
def __init__(
self,
patch_size: Tuple[int] = (1, 2, 2),
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: Optional[str] = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
pos_embed_seq_len: Optional[int] = None,
) -> None:
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
out_channels = out_channels or in_channels
# 1. Patch & position embedding
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size)
# 2. Condition embeddings
# image_embedding_dim=1280 for I2V model
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
time_proj_dim=inner_dim * 6,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
pos_embed_seq_len=pos_embed_seq_len,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList(
[
WanTransformerBlock(
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
)
for _ in range(num_layers)
]
)
# 4. Output norm & projection
self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False)
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
lora_scale = attention_kwargs.pop("scale", 1.0)
else:
lora_scale = 1.0
if USE_PEFT_BACKEND:
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
rotary_emb = self.rope(hidden_states)
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image
)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb
)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
# Move the shift and scale tensors to the same device as hidden_states.
# When using multi-GPU inference via accelerate these will be on the
# first device rather than the last device, which hidden_states ends up
# on.
shift = shift.to(hidden_states.device)
scale = scale.to(hidden_states.device)
hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(
batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1
)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
if not return_dict:
return (output,)
return Transformer2DModelOutput(sample=output)
-609
View File
@@ -1,609 +0,0 @@
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import html
from typing import Any, Callable, Dict, List, Optional, Union
import regex as re
import torch
from transformers import AutoTokenizer, UMT5EncoderModel
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.loaders import WanLoraLoaderMixin
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput
from einops import rearrange
from transformers import UMT5EncoderModel, T5TokenizerFast
from fastvideo.models.mochi_hf.modeling_wan import WanTransformer3DModel
from fastvideo.utils.communications import all_gather
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
XLA_AVAILABLE = True
else:
XLA_AVAILABLE = False
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
if is_ftfy_available():
import ftfy
EXAMPLE_DOC_STRING = """
Examples:
```python
>>> import torch
>>> from diffusers.utils import export_to_video
>>> from diffusers import AutoencoderKLWan, WanPipeline
>>> from diffusers.schedulers.scheduling_unipc_multistep import UniPCMultistepScheduler
>>> # Available models: Wan-AI/Wan2.1-T2V-14B-Diffusers, Wan-AI/Wan2.1-T2V-1.3B-Diffusers
>>> model_id = "Wan-AI/Wan2.1-T2V-14B-Diffusers"
>>> vae = AutoencoderKLWan.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32)
>>> pipe = WanPipeline.from_pretrained(model_id, vae=vae, torch_dtype=torch.bfloat16)
>>> flow_shift = 5.0 # 5.0 for 720P, 3.0 for 480P
>>> pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config, flow_shift=flow_shift)
>>> pipe.to("cuda")
>>> prompt = "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
>>> negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
>>> output = pipe(
... prompt=prompt,
... negative_prompt=negative_prompt,
... height=720,
... width=1280,
... num_frames=81,
... guidance_scale=5.0,
... ).frames[0]
>>> export_to_video(output, "output.mp4", fps=16)
```
"""
def basic_clean(text):
text = ftfy.fix_text(text)
text = html.unescape(html.unescape(text))
return text.strip()
def whitespace_clean(text):
text = re.sub(r"\s+", " ", text)
text = text.strip()
return text
def prompt_clean(text):
text = whitespace_clean(basic_clean(text))
return text
class WanPipeline(DiffusionPipeline, WanLoraLoaderMixin):
r"""
Pipeline for text-to-video generation using Wan.
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
Args:
tokenizer ([`T5Tokenizer`]):
Tokenizer from [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5Tokenizer),
specifically the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
text_encoder ([`T5EncoderModel`]):
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
transformer ([`WanTransformer3DModel`]):
Conditional Transformer to denoise the input latents.
scheduler ([`UniPCMultistepScheduler`]):
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
vae ([`AutoencoderKLWan`]):
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
"""
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(
self,
tokenizer: AutoTokenizer,
text_encoder: UMT5EncoderModel,
transformer: WanTransformer3DModel,
vae: AutoencoderKLWan,
scheduler: FlowMatchEulerDiscreteScheduler,
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer,
scheduler=scheduler,
)
self.vae_scale_factor_temporal = 2 ** sum(self.vae.temperal_downsample) if getattr(self, "vae", None) else 4
self.vae_scale_factor_spatial = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_videos_per_prompt: int = 1,
max_sequence_length: int = 226,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
prompt = [prompt_clean(u) for u in prompt]
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
seq_lens = mask.gt(0).sum(dim=1).long()
prompt_embeds = self.text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
prompt_embeds = torch.stack(
[torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0
)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
return prompt_embeds
def encode_prompt(
self,
prompt: Union[str, List[str]],
negative_prompt: Optional[Union[str, List[str]]] = None,
do_classifier_free_guidance: bool = True,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
max_sequence_length: int = 226,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
r"""
Encodes the prompt into text encoder hidden states.
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
less than `1`).
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
Whether to use classifier free guidance or not.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
device: (`torch.device`, *optional*):
torch device
dtype: (`torch.dtype`, *optional*):
torch dtype
"""
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
if prompt is not None:
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds = self._get_t5_prompt_embeds(
prompt=prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
if do_classifier_free_guidance and negative_prompt_embeds is None:
negative_prompt = negative_prompt or ""
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
if prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
" the batch size of `prompt`."
)
negative_prompt_embeds = self._get_t5_prompt_embeds(
prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
return prompt_embeds, negative_prompt_embeds
def check_inputs(
self,
prompt,
negative_prompt,
height,
width,
prompt_embeds=None,
negative_prompt_embeds=None,
callback_on_step_end_tensor_inputs=None,
):
if height % 16 != 0 or width % 16 != 0:
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
elif negative_prompt is not None and (
not isinstance(negative_prompt, str) and not isinstance(negative_prompt, list)
):
raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")
def prepare_latents(
self,
batch_size: int,
num_channels_latents: int = 16,
height: int = 480,
width: int = 832,
num_frames: int = 81,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if latents is not None:
return latents.to(device=device, dtype=dtype)
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
shape = (
batch_size,
num_channels_latents,
num_latent_frames,
int(height) // self.vae_scale_factor_spatial,
int(width) // self.vae_scale_factor_spatial,
)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
@property
def guidance_scale(self):
return self._guidance_scale
@property
def do_classifier_free_guidance(self):
return self._guidance_scale > 1.0
@property
def num_timesteps(self):
return self._num_timesteps
@property
def current_timestep(self):
return self._current_timestep
@property
def interrupt(self):
return self._interrupt
@property
def attention_kwargs(self):
return self._attention_kwargs
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]] = None,
negative_prompt: Union[str, List[str]] = None,
height: int = 480,
width: int = 832,
num_frames: int = 81,
num_inference_steps: int = 50,
guidance_scale: float = 5.0,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
output_type: Optional[str] = "np",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 512,
):
r"""
The call function to the pipeline for generation.
Args:
prompt (`str` or `List[str]`, *optional*):
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
instead.
height (`int`, defaults to `480`):
The height in pixels of the generated image.
width (`int`, defaults to `832`):
The width in pixels of the generated image.
num_frames (`int`, defaults to `81`):
The number of frames in the generated video.
num_inference_steps (`int`, defaults to `50`):
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
expense of slower inference.
guidance_scale (`float`, defaults to `5.0`):
Guidance scale as defined in [Classifier-Free Diffusion
Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2.
of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting
`guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to
the text `prompt`, usually at the expense of lower image quality.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
The number of images to generate per prompt.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
generation deterministic.
latents (`torch.Tensor`, *optional*):
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
tensor is generated by sampling using the supplied random `generator`.
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
provided, text embeddings are generated from the `prompt` input argument.
output_type (`str`, *optional*, defaults to `"np"`):
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`WanPipelineOutput`] instead of a plain tuple.
attention_kwargs (`dict`, *optional*):
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
`self.processor` in
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
callback_on_step_end_tensor_inputs (`List`, *optional*):
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
`._callback_tensor_inputs` attribute of your pipeline class.
autocast_dtype (`torch.dtype`, *optional*, defaults to `torch.bfloat16`):
The dtype to use for the torch.amp.autocast.
Examples:
Returns:
[`~WanPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`WanPipelineOutput`] is returned, otherwise a `tuple` is returned where
the first element is a list with the generated images and the second element is a list of `bool`s
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
# 1. Check inputs. Raise error if not correct
self.check_inputs(
prompt,
negative_prompt,
height,
width,
prompt_embeds,
negative_prompt_embeds,
callback_on_step_end_tensor_inputs,
)
if num_frames % self.vae_scale_factor_temporal != 1:
logger.warning(
f"`num_frames - 1` has to be divisible by {self.vae_scale_factor_temporal}. Rounding to the nearest number."
)
num_frames = num_frames // self.vae_scale_factor_temporal * self.vae_scale_factor_temporal + 1
num_frames = max(num_frames, 1)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._current_timestep = None
self._interrupt = False
device = self._execution_device
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
# 3. Encode input prompt
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
prompt=prompt,
negative_prompt=negative_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
num_videos_per_prompt=num_videos_per_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
max_sequence_length=max_sequence_length,
device=device,
)
transformer_dtype = self.transformer.dtype
prompt_embeds = prompt_embeds.to(transformer_dtype)
if negative_prompt_embeds is not None:
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
# 4. Prepare timesteps
self.scheduler.set_timesteps(num_inference_steps, device=device)
timesteps = self.scheduler.timesteps
# 5. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
num_channels_latents,
height,
width,
num_frames,
torch.float32,
device,
generator,
latents,
)
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
if get_sequence_parallel_state():
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
# 6. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
self._current_timestep = t
latent_model_input = latents.to(transformer_dtype)
timestep = t.expand(latents.shape[0])
noise_pred = self.transformer(
hidden_states=latent_model_input,
timestep=timestep,
encoder_hidden_states=prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
if self.do_classifier_free_guidance:
noise_uncond = self.transformer(
hidden_states=latent_model_input,
timestep=timestep,
encoder_hidden_states=negative_prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
xm.mark_step()
if get_sequence_parallel_state():
latents = all_gather(latents, dim=2)
self._current_timestep = None
if not output_type == "latent":
latents = latents.to(self.vae.dtype)
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
latents.device, latents.dtype
)
latents = latents / latents_std + latents_mean
video = self.vae.decode(latents, return_dict=False)[0]
video = self.video_processor.postprocess_video(video, output_type=output_type)
else:
video = latents
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return (video,)
return WanPipelineOutput(frames=video)
+2 -2
View File
@@ -86,7 +86,7 @@ def inference(args):
num_inference_steps=args.num_inference_steps,
generator=generator,
).frames
if nccl_info.global_rank == 0:
if nccl_info.global_rank <= 0:
os.makedirs(args.output_path, exist_ok=True)
suffix = prompt.split(".")[0]
export_to_video(
@@ -107,7 +107,7 @@ def inference(args):
generator=generator,
).frames
if nccl_info.global_rank == 0:
if nccl_info.global_rank <= 0:
export_to_video(videos[0], args.output_path + ".mp4", fps=24)
+2 -2
View File
@@ -94,7 +94,7 @@ def main(args):
guidance_scale=args.guidance_scale,
generator=generator,
).frames
if nccl_info.global_rank == 0:
if nccl_info.global_rank <= 0:
os.makedirs(args.output_path, exist_ok=True)
suffix = prompt.split(".")[0]
export_to_video(
@@ -116,7 +116,7 @@ def main(args):
generator=generator,
).frames
if nccl_info.global_rank == 0:
if nccl_info.global_rank <= 0:
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
+4 -4
View File
@@ -20,7 +20,7 @@ from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.utils.latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
@@ -185,7 +185,7 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank == 0 and args.output_dir is not None:
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
@@ -316,7 +316,7 @@ def main(args):
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:
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@@ -393,7 +393,7 @@ def main(args):
"grad_norm": grad_norm,
})
progress_bar.update(1)
if rank == 0:
if rank <= 0:
wandb.log(
{
"train_loss": loss,
+6 -6
View File
@@ -32,7 +32,7 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
if rank == 0 and not discriminator:
if rank <= 0 and not discriminator:
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(model.config)
@@ -60,7 +60,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
):
cpu_state = transformer.state_dict()
# todo move to get_state_dict
if rank == 0:
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
@@ -98,7 +98,7 @@ def save_checkpoint_generator_discriminator(
hf_weight_dir = os.path.join(save_dir, "hf_weights")
os.makedirs(hf_weight_dir, exist_ok=True)
# save using safetensors
if rank == 0:
if rank <= 0:
config_dict = dict(model.config)
config_path = os.path.join(hf_weight_dir, "config.json")
# save dict as json
@@ -139,7 +139,7 @@ def save_checkpoint_generator_discriminator(
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
model_state = discriminator.state_dict()
state_dict = {"optimizer": optim_state, "model": model_state}
if rank == 0:
if rank <= 0:
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
torch.save(state_dict, discriminator_fsdp_state_fil)
@@ -178,7 +178,7 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
):
discriminator_state = torch.load(checkpoint_file)
model_state = discriminator_state["model"]
if rank == 0:
if rank <= 0:
optim_state = discriminator_state["optimizer"]
else:
optim_state = None
@@ -241,7 +241,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipelin
optimizer,
)
if rank == 0:
if rank <= 0:
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
-88
View File
@@ -1,88 +0,0 @@
import torch
mochi_latents_mean = torch.tensor([
-0.06730895953510081,
-0.038011381506090416,
-0.07477820912866141,
-0.05565264470995561,
0.012767231469026969,
-0.04703542746246419,
0.043896967884726704,
-0.09346305707025976,
-0.09918314763016893,
-0.008729793427399178,
-0.011931556316503654,
-0.0321993391887285,
]).view(1, 12, 1, 1, 1)
mochi_latents_std = torch.tensor([
0.9263795028493863,
0.9248894543193766,
0.9393059390890617,
0.959253732819592,
0.8244560132752793,
0.917259975397747,
0.9294154431013696,
1.3720942357788521,
0.881393668867029,
0.9168315692124348,
0.9185249279345552,
0.9274757570805041,
]).view(1, 12, 1, 1, 1)
mochi_scaling_factor = 1.0
wan_latents_mean = torch.tensor([
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
]).view(1, 16, 1, 1, 1)
wan_latents_std = torch.tensor([
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.916,
]).view(1, 16, 1, 1, 1)
def normalize_dit_input(model_type, latents):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
latents = (latents - latents_mean) / latents_std
return latents
elif model_type == "hunyuan_hf":
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
latents_mean = wan_latents_mean.to(latents.device, latents.dtype)
latents_std = wan_latents_std.to(latents.device, latents.dtype)
latents = (latents - latents_mean) / latents_std
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
+2 -69
View File
@@ -3,9 +3,9 @@ from pathlib import Path
import torch
import torch.nn.functional as F
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi, AutoencoderKLWan
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
from torch import nn
from transformers import AutoTokenizer, T5EncoderModel, UMT5EncoderModel
from transformers import AutoTokenizer, T5EncoderModel
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
MMSingleStreamBlock)
@@ -14,7 +14,6 @@ from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLC
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (HunyuanVideoSingleTransformerBlock,
HunyuanVideoTransformer3DModel, HunyuanVideoTransformerBlock)
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
from fastvideo.models.wan_hf.modeling_wan import WanTransformer3DModel, WanTransformerBlock
from fastvideo.utils.logging_ import main_print
hunyuan_config = {
@@ -201,48 +200,6 @@ class MochiTextEncoderWrapper(nn.Module):
return prompt_embeds, prompt_attention_mask
class WanTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
self.text_encoder = UMT5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path,
"text_encoder")).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "tokenizer"))
self.max_sequence_length = 256
def encode_prompt(self, prompt):
device = self.text_encoder.device
dtype = self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.max_sequence_length - 1:-1])
main_print(f"Truncated text input: {prompt} to: {removed_text} for model input.")
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
return prompt_embeds, prompt_attention_mask
def load_hunyuan_state_dict(model, dit_model_name_or_path):
load_key = "module"
@@ -283,20 +240,6 @@ def load_transformer(
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
elif model_type == "wan":
if dit_model_name_or_path:
transformer = WanTransformer3DModel.from_pretrained(
dit_model_name_or_path,
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
else:
transformer = WanTransformer3DModel.from_pretrained(
pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
elif model_type == "hunyuan_hf":
if dit_model_name_or_path:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
@@ -340,12 +283,6 @@ def load_vae(model_type, pretrained_model_name_or_path):
torch_dtype=weight_dtype).to("cuda")
autocast_type = torch.bfloat16
fps = 24
elif model_type == "wan":
vae = AutoencoderKLWan.from_pretrained(pretrained_model_name_or_path,
subfolder="vae",
torch_dtype=weight_dtype).to("cuda")
autocast_type = torch.bfloat16
fps = 24
elif model_type == "hunyuan":
vae_precision = torch.float32
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
@@ -374,8 +311,6 @@ def load_vae(model_type, pretrained_model_name_or_path):
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
if model_type == "mochi":
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "wan":
text_encoder = WanTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "hunyuan" or "hunyuan_hf":
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
else:
@@ -387,8 +322,6 @@ def get_no_split_modules(transformer):
# if of type MochiTransformer3DModel
if isinstance(transformer, MochiTransformer3DModel):
return (MochiTransformerBlock, )
elif isinstance(transformer, WanTransformer3DModel):
return (WanTransformerBlock, )
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
return (HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformerBlock)
elif isinstance(transformer, HYVideoDiffusionTransformer):
+11 -23
View File
@@ -129,22 +129,13 @@ def sample_validation_video(
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
with torch.autocast("cuda", dtype=torch.bfloat16):
if model_type == "wan":
pred_kwargs = {
"hidden_states": latent_model_input,
"encoder_hidden_states": prompt_embeds,
"timestep":timestep,
"return_dict":False,
}
else:
pred_kwargs = {
"hidden_states": latent_model_input,
"encoder_hidden_states": prompt_embeds,
"timestep":timestep,
"encoder_attention_mask":prompt_attention_mask,
"return_dict":False,
}
noise_pred = transformer(**pred_kwargs)[0]
noise_pred = transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask,
return_dict=False,
)[0]
# Mochi CFG + Sampling runs in FP32
noise_pred = noise_pred.to(torch.float32)
@@ -175,12 +166,10 @@ def sample_validation_video(
# denormalize with the mean and std if available and not None
has_latents_mean = (hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None)
has_latents_std = (hasattr(vae.config, "latents_std") and vae.config.latents_std is not None)
if model_type == "wan":
vae.config.scaling_factor = 1
if has_latents_mean and has_latents_std:
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, num_channels_latents, 1, 1,
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1,
1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(vae.config.latents_std).view(1, num_channels_latents, 1, 1, 1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype))
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
else:
latents = latents / vae.config.scaling_factor
@@ -213,15 +202,14 @@ def log_validation(
vae_spatial_scale_factor = 8
vae_temporal_scale_factor = 6
num_channels_latents = 12
elif args.model_type == "hunyuan" or "hunyuan_hf" or "wan":
elif args.model_type == "hunyuan" or "hunyuan_hf":
vae_spatial_scale_factor = 8
vae_temporal_scale_factor = 4
num_channels_latents = 16
else:
raise ValueError(f"Model type {args.model_type} not supported")
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
if args.model_type != "wan":
vae.enable_tiling()
vae.enable_tiling()
if scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
else:
-419
View File
@@ -1,419 +0,0 @@
import json
import os
from collections import defaultdict
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
def configure_sta(mode: str = 'STA_searching',
layer_num: int = 40,
time_step_num: int = 50,
head_num: int = 40,
**kwargs) -> List[List[List[Any]]]:
"""
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
Parameters:
----------
mode : str
The STA mode to use. Options are:
- 'STA_searching': Generate a set of mask candidates for initial search
- 'STA_tuning': Select best mask strategy based on previously saved results
- 'STA_inference': Load and use a previously tuned mask strategy
layer_num: int, number of layers
time_step_num: int, number of timesteps
head_num: int, number of heads
**kwargs : dict
Mode-specific parameters:
For 'STA_searching':
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
For 'STA_tuning':
- mask_search_files_path: str, required, path to mask search results
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
- skip_time_steps: int, optional, number of time steps to use full attention (default 12)
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
For 'STA_inference':
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
"""
valid_modes = [
'STA_searching', 'STA_tuning', 'STA_inference', 'STA_tuning_cfg'
]
if mode not in valid_modes:
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
if mode == 'STA_searching':
# Get parameters with defaults
mask_candidates: Optional[List[str]] = kwargs.get('mask_candidates')
if mask_candidates is None:
raise ValueError(
"mask_candidates is required for STA_searching mode")
mask_selected: List[int] = kwargs.get('mask_selected',
list(range(len(mask_candidates))))
# Parse selected masks
selected_masks: List[List[int]] = []
for index in mask_selected:
mask = mask_candidates[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Create 3D mask structure with fixed dimensions (t=50, l=60)
masks_3d: List[List[List[List[int]]]] = []
for i in range(time_step_num): # Fixed t dimension = 50
row = []
for j in range(layer_num): # Fixed l dimension = 60
row.append(selected_masks) # Add all masks at each position
masks_3d.append(row)
return masks_3d
elif mode == 'STA_tuning':
# Get required parameters
mask_search_files_path: Optional[str] = kwargs.get(
'mask_search_files_path')
if not mask_search_files_path:
raise ValueError(
"mask_search_files_path is required for STA_tuning mode")
# Get optional parameters with defaults
mask_candidates_tuning: Optional[List[str]] = kwargs.get(
'mask_candidates')
if mask_candidates_tuning is None:
raise ValueError("mask_candidates is required for STA_tuning mode")
mask_selected_tuning: List[int] = kwargs.get(
'mask_selected', list(range(len(mask_candidates_tuning))))
skip_time_steps_tuning: Optional[int] = kwargs.get('skip_time_steps')
save_dir_tuning: Optional[str] = kwargs.get('save_dir',
"mask_candidates")
# Parse selected masks
selected_masks_tuning: List[List[int]] = []
for index in mask_selected_tuning:
mask = mask_candidates_tuning[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks_tuning.append(masks_list)
# Read JSON results
results = read_specific_json_files(mask_search_files_path)
averaged_results = average_head_losses(results, selected_masks_tuning)
# Add full attention mask for specific cases
full_attention_mask_tuning: Optional[List[int]] = kwargs.get(
'full_attention_mask')
if full_attention_mask_tuning is not None:
selected_masks_tuning.append(full_attention_mask_tuning)
# Select best mask strategy
timesteps_tuning: int = kwargs.get('timesteps', time_step_num)
if skip_time_steps_tuning is None:
skip_time_steps_tuning = 12
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
averaged_results, selected_masks_tuning, skip_time_steps_tuning,
timesteps_tuning, head_num)
# Save mask strategy
if save_dir_tuning is not None:
os.makedirs(save_dir_tuning, exist_ok=True)
file_path = os.path.join(
save_dir_tuning,
f'mask_strategy_s{skip_time_steps_tuning}.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
for strategy, count in strategy_counts.items():
print(
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
elif mode == 'STA_tuning_cfg':
# Get required parameters for both positive and negative paths
mask_search_files_path_pos: Optional[str] = kwargs.get(
'mask_search_files_path_pos')
mask_search_files_path_neg: Optional[str] = kwargs.get(
'mask_search_files_path_neg')
save_dir_cfg: Optional[str] = kwargs.get('save_dir')
if not mask_search_files_path_pos or not mask_search_files_path_neg or not save_dir_cfg:
raise ValueError(
"mask_search_files_path_pos, mask_search_files_path_neg, and save_dir are required for STA_tuning_cfg mode"
)
# Get optional parameters with defaults
mask_candidates_cfg: Optional[List[str]] = kwargs.get('mask_candidates')
if mask_candidates_cfg is None:
raise ValueError(
"mask_candidates is required for STA_tuning_cfg mode")
mask_selected_cfg: List[int] = kwargs.get(
'mask_selected', list(range(len(mask_candidates_cfg))))
skip_time_steps_cfg: Optional[int] = kwargs.get('skip_time_steps')
# Parse selected masks
selected_masks_cfg: List[List[int]] = []
for index in mask_selected_cfg:
mask = mask_candidates_cfg[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks_cfg.append(masks_list)
# Read JSON results for both positive and negative paths
pos_results = read_specific_json_files(mask_search_files_path_pos)
neg_results = read_specific_json_files(mask_search_files_path_neg)
# Combine positive and negative results into one list
combined_results = pos_results + neg_results
# Average the combined results
averaged_results = average_head_losses(combined_results,
selected_masks_cfg)
# Add full attention mask for specific cases
full_attention_mask_cfg: Optional[List[int]] = kwargs.get(
'full_attention_mask')
if full_attention_mask_cfg is not None:
selected_masks_cfg.append(full_attention_mask_cfg)
timesteps_cfg: int = kwargs.get('timesteps', time_step_num)
if skip_time_steps_cfg is None:
skip_time_steps_cfg = 12
# Select best mask strategy using combined results
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
averaged_results, selected_masks_cfg, skip_time_steps_cfg,
timesteps_cfg, head_num)
# Save mask strategy
os.makedirs(save_dir_cfg, exist_ok=True)
file_path = os.path.join(save_dir_cfg,
f'mask_strategy_s{skip_time_steps_cfg}.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
for strategy, count in strategy_counts.items():
print(
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
else: # STA_inference
# Get parameters with defaults
load_path: Optional[str] = kwargs.get(
'load_path', "mask_candidates/mask_strategy.json")
if load_path is None:
raise ValueError("load_path is required for STA_inference mode")
# Load previously saved mask strategy
with open(load_path) as f:
mask_strategy = json.load(f)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
# Helper functions
def read_specific_json_files(folder_path: str) -> List[Dict[str, Any]]:
"""Read and parse JSON files containing mask search results."""
json_contents: List[Dict[str, Any]] = []
# List files only in the current directory (no walk)
files = os.listdir(folder_path)
# Filter files
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
print(f"Found {len(matching_files)} matching files: {matching_files}")
for file_name in matching_files:
file_path = os.path.join(folder_path, file_name)
with open(file_path) as file:
data = json.load(file)
json_contents.append(data)
return json_contents
def average_head_losses(
results: List[Dict[str, Any]],
selected_masks: List[List[int]]) -> Dict[str, Dict[str, np.ndarray]]:
"""Average losses across all prompts for each mask strategy."""
# Initialize a dictionary to store the averaged results
averaged_losses: Dict[str, Dict[str, np.ndarray]] = {}
loss_type = 'L2_loss'
# Get all loss types (e.g., 'L2_loss')
averaged_losses[loss_type] = {}
for mask in selected_masks:
mask_str = str(mask)
data_shape = np.array(results[0][loss_type][mask_str]).shape
accumulated_data = np.zeros(data_shape)
# Sum across all prompts
for prompt_result in results:
accumulated_data += np.array(prompt_result[loss_type][mask_str])
# Average by dividing by number of prompts
averaged_data = accumulated_data / len(results)
averaged_losses[loss_type][mask_str] = averaged_data
return averaged_losses
def select_best_mask_strategy(
averaged_results: Dict[str, Dict[str, np.ndarray]],
selected_masks: List[List[int]],
skip_time_steps: int = 12,
timesteps: int = 50,
head_num: int = 40
) -> Tuple[Dict[str, List[int]], float, Dict[str, int]]:
"""Select the best mask strategy for each head based on loss minimization."""
best_mask_strategy: Dict[str, List[int]] = {}
loss_type = 'L2_loss'
# Get the shape of time steps and layers
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
# Counter for sparsity calculation
total_tokens = 0 # total number of masked tokens
total_length = 0 # total sequence length
strategy_counts: Dict[str, int] = {
str(strategy): 0
for strategy in selected_masks
}
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
for t in range(timesteps):
for layer_idx in range(layers):
for h in range(head_num):
if t < skip_time_steps: # First steps use full attention
strategy = full_attn_strategy
else:
# Get losses for this head across all strategies
head_losses = []
for strategy in selected_masks[:
-1]: # Exclude full attention
head_losses.append(averaged_results[loss_type][str(
strategy)][t][layer_idx][h])
# Find which strategy gives minimum loss
best_strategy_idx = np.argmin(head_losses)
strategy = selected_masks[best_strategy_idx]
best_mask_strategy[f'{t}_{layer_idx}_{h}'] = strategy
# Calculate sparsity
nums = strategy # strategy is already a list of numbers
total_tokens += nums[0] * nums[1] * nums[
2] # masked tokens for chosen strategy
total_length += full_attn_strategy[0] * full_attn_strategy[
1] * full_attn_strategy[2]
# Count strategy usage
strategy_counts[str(strategy)] += 1
overall_sparsity = 1 - total_tokens / total_length
return best_mask_strategy, overall_sparsity, strategy_counts
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
t_max: int = 50,
l_max: int = 60,
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
result: List[List[List[Optional[List[int]]]]] = [[[
None for _ in range(h_max)
] for _ in range(l_max)] for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, layer_idx, h = map(int, key.split('_'))
result[t][layer_idx][h] = value
return result
def save_mask_search_results(
mask_search_final_result: List[Dict[str, List[float]]],
prompt: str,
mask_strategies: List[str],
output_dir: str = 'output/mask_search_result/') -> Optional[str]:
if not mask_search_final_result:
print("No mask search results to save")
return None
# Create result dictionary with defaultdict for nested lists
mask_search_dict: Dict[str, Dict[str, List[List[float]]]] = {
"L2_loss": defaultdict(list),
"L1_loss": defaultdict(list)
}
mask_selected = list(range(len(mask_strategies)))
selected_masks: List[List[int]] = []
for index in mask_selected:
mask = mask_strategies[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Process each mask strategy
for i, mask_strategy in enumerate(selected_masks):
mask_strategy_str = str(mask_strategy)
# Process L2 loss
step_results: List[List[float]] = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L2_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
step_results.append(layer_losses)
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
step_results = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L1_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
step_results.append(layer_losses)
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
# Create the output directory if it doesn't exist
os.makedirs(output_dir, exist_ok=True)
# Create a filename based on the first 20 characters of the prompt
filename = prompt[:50].replace(" ", "_")
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
# Save the results to a JSON file
with open(filepath, 'w') as f:
json.dump(mask_search_dict, f, indent=4)
print(f"Successfully saved mask research results to {filepath}")
return filepath
+6 -8
View File
@@ -3,8 +3,7 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, fields
from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
Type, TypeVar)
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
if TYPE_CHECKING:
from fastvideo.v1.fastvideo_args import FastVideoArgs
@@ -27,12 +26,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
@@ -46,7 +45,7 @@ class AttentionBackend(ABC):
@staticmethod
@abstractmethod
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
@@ -57,8 +56,7 @@ class AttentionMetadata:
current_timestep: int
def asdict_zerocopy(self,
skip_fields: Optional[Set[str]] = None
) -> Dict[str, Any]:
skip_fields: set[str] | None = None) -> dict[str, Any]:
"""Similar to dataclasses.asdict, but avoids deepcopying."""
if skip_fields is None:
skip_fields = set()
@@ -124,7 +122,7 @@ class AttentionImpl(ABC, Generic[T]):
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -1,7 +1,5 @@
# 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
@@ -28,7 +26,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
@@ -36,15 +34,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
@@ -56,7 +54,7 @@ class FlashAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+3 -5
View File
@@ -1,5 +1,3 @@
from typing import List, Optional, Type
import torch
from sageattention import sageattn
@@ -17,7 +15,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
@@ -25,7 +23,7 @@ class SageAttentionBackend(AttentionBackend):
return "SAGE_ATTN"
@staticmethod
def get_impl_cls() -> Type["SageAttentionImpl"]:
def get_impl_cls() -> type["SageAttentionImpl"]:
return SageAttentionImpl
# @staticmethod
@@ -41,7 +39,7 @@ class SageAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
+3 -5
View File
@@ -1,5 +1,3 @@
from typing import List, Optional, Type
import torch
from fastvideo.v1.attention.backends.abstract import (
@@ -16,7 +14,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
@@ -24,7 +22,7 @@ class SDPABackend(AttentionBackend):
return "SDPA"
@staticmethod
def get_impl_cls() -> Type["SDPAImpl"]:
def get_impl_cls() -> type["SDPAImpl"]:
return SDPAImpl
# @staticmethod
@@ -40,7 +38,7 @@ class SDPAImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -1,6 +1,5 @@
import json
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Type
import torch
from einops import rearrange
@@ -13,7 +12,6 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -21,9 +19,7 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(
mask_strategy: Dict[str,
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
def dict_to_3d_list(mask_strategy) -> list[list[list[torch.Tensor | None]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
@@ -45,14 +41,14 @@ def dict_to_3d_list(
class RangeDict(dict):
def __getitem__(self, item: int) -> str:
def __getitem__(self, item):
for key in self.keys():
if isinstance(key, tuple):
low, high = key
if low <= item <= high:
return str(super().__getitem__(key))
return super().__getitem__(key)
elif key == item:
return str(super().__getitem__(key))
return super().__getitem__(key)
raise KeyError(f"seq_len {item} not supported for STA")
@@ -61,7 +57,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]
@@ -70,23 +66,21 @@ 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
@dataclass
class SlidingTileAttentionMetadata(AttentionMetadata):
current_timestep: int
STA_param: List[List[
Any]] # each timestep with one metadata, shape [num_layers, num_heads]
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
@@ -103,12 +97,8 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> SlidingTileAttentionMetadata:
param = forward_batch.STA_param
if param is None:
return SlidingTileAttentionMetadata(
current_timestep=current_timestep, STA_param=[])
return SlidingTileAttentionMetadata(current_timestep=current_timestep,
STA_param=param[current_timestep])
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
class SlidingTileAttentionImpl(AttentionImpl):
@@ -119,7 +109,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
@@ -129,12 +119,12 @@ class SlidingTileAttentionImpl(AttentionImpl):
if config_file is None:
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
# TODO(kevin): get mask strategy for different STA modes
with open(config_file) as f:
mask_strategy = json.load(f)
self.mask_strategy = dict_to_3d_list(mask_strategy)
mask_strategy = dict_to_3d_list(mask_strategy)
self.prefix = prefix
self.mask_strategy = mask_strategy
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
# STA config
@@ -214,24 +204,16 @@ class SlidingTileAttentionImpl(AttentionImpl):
v: torch.Tensor,
attn_metadata: SlidingTileAttentionMetadata,
) -> torch.Tensor:
if self.mask_strategy is None:
raise ValueError(
"mask_strategy cannot be None for SlidingTileAttention")
if self.mask_strategy[0] is None:
raise ValueError(
"mask_strategy[0] cannot be None for SlidingTileAttention")
assert self.mask_strategy is not None, "mask_strategy cannot be None for SlidingTileAttention"
assert self.mask_strategy[
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
timestep = attn_metadata.current_timestep
forward_context: ForwardContext = get_forward_context()
forward_batch = forward_context.forward_batch
if forward_batch is None:
raise ValueError("forward_batch cannot be None")
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
layer_idx = int(self.prefix.split('.')[-3])
if attn_metadata.STA_param is None or len(
attn_metadata.STA_param) <= layer_idx:
raise ValueError("Invalid STA_param")
STA_param = attn_metadata.STA_param[layer_idx]
# TODO: remove hardcode
text_length = q.shape[1] - self.img_seq_length
has_text = text_length > 0
@@ -244,62 +226,15 @@ class SlidingTileAttentionImpl(AttentionImpl):
sp_group = get_sp_group()
current_rank = sp_group.rank_in_group
start_head = current_rank * head_num
# searching or tuning mode
if len(STA_param) < head_num * sp_group.world_size:
sparse_attn_hidden_states_all = []
full_mask_window = STA_param[-1]
for window_size in STA_param[:-1]:
sparse_hidden_states = sliding_tile_attention(
query, key, value, [window_size] * head_num, text_length,
has_text, self.img_latent_shape_str).transpose(1, 2)
sparse_attn_hidden_states_all.append(sparse_hidden_states)
hidden_states = sliding_tile_attention(
query, key, value, [full_mask_window] * head_num, text_length,
has_text, self.img_latent_shape_str).transpose(1, 2)
attn_L2_loss = []
attn_L1_loss = []
# average loss across all heads
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
# L2 loss
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() -
hidden_states.float())**2,
dim=[0, 1, 3]).cpu().numpy()
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
attn_L2_loss.append(attn_L2_loss_)
# L1 loss
attn_L1_loss_ = torch.mean(
torch.abs(sparse_attn_hidden_states.float() -
hidden_states.float()),
dim=[0, 1, 3]).cpu().numpy()
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
attn_L1_loss.append(attn_L1_loss_)
layer_loss_save = {"L2_loss": attn_L2_loss, "L1_loss": attn_L1_loss}
if forward_batch.is_cfg_negative:
if forward_batch.mask_search_final_result_neg is not None:
forward_batch.mask_search_final_result_neg[timestep].append(
layer_loss_save)
else:
if forward_batch.mask_search_final_result_pos is not None:
forward_batch.mask_search_final_result_pos[timestep].append(
layer_loss_save)
else:
# windows = [
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
# for head_idx in range(head_num)
# ]
windows = [
STA_param[head_idx + start_head] for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
windows = [
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
return hidden_states
+14 -17
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional, Tuple
import torch
import torch.nn as nn
@@ -13,7 +11,6 @@ from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.utils import get_compute_dtype
class DistributedAttention(nn.Module):
@@ -23,11 +20,11 @@ class DistributedAttention(nn.Module):
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -39,7 +36,7 @@ class DistributedAttention(nn.Module):
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = get_compute_dtype()
dtype = torch.get_default_dtype()
attn_backend = get_attn_backend(
head_size,
dtype,
@@ -63,10 +60,10 @@ class DistributedAttention(nn.Module):
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
replicated_q: Optional[torch.Tensor] = None,
replicated_k: Optional[torch.Tensor] = None,
replicated_v: Optional[torch.Tensor] = None,
) -> tuple[torch.Tensor, Optional[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]:
"""Forward pass for distributed attention.
Args:
@@ -142,11 +139,11 @@ class LocalAttention(nn.Module):
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
@@ -156,7 +153,7 @@ class LocalAttention(nn.Module):
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = get_compute_dtype()
dtype = torch.get_default_dtype()
attn_backend = get_attn_backend(
head_size,
dtype,
+14 -13
View File
@@ -2,9 +2,10 @@
# 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 Generator, Optional, Tuple, Type, cast
from typing import cast
import torch
@@ -17,7 +18,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) -> Optional[_Backend]:
def backend_name_to_enum(backend_name: str) -> _Backend | None:
"""
Convert a string backend name to a _Backend enum value.
@@ -31,7 +32,7 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
None
def get_env_variable_attn_backend() -> Optional[_Backend]:
def get_env_variable_attn_backend() -> _Backend | None:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
@@ -53,10 +54,10 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: Optional[_Backend] = None
forced_attn_backend: _Backend | None = None
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
def global_force_attn_backend(attn_backend: _Backend | None) -> None:
'''
Force all attention operations to use a specified backend.
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_global_forced_attn_backend() -> _Backend | None:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
@@ -82,8 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@@ -92,8 +93,8 @@ def get_attn_backend(
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
supported_attention_backends: tuple[_Backend, ...] | None = None,
) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
@@ -102,13 +103,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: Optional[_Backend] = (
backend_by_global_setting: _Backend | None = (
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: Optional[str] = envs.FASTVIDEO_ATTENTION_BACKEND
backend_by_env_var: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
@@ -120,7 +121,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
+3 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field, fields
from typing import Any, Dict
from typing import Any
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)}
+3 -5
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, List, Optional, Tuple
from typing import Any
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
@@ -11,8 +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)
_lora_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,
@@ -21,7 +20,6 @@ class DiTArchConfig(ArchConfig):
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
exclude_lora_layers: List[str] = field(default_factory=list)
def __post_init__(self) -> None:
if not self._compile_conditions:
@@ -34,7 +32,7 @@ class DiTConfig(ModelConfig):
# FastVideoDiT-specific parameters
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
quant_config: QuantizationConfig | None = None
@staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
import torch
@@ -156,15 +155,13 @@ 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: Optional[torch.dtype] = None
dtype: torch.dtype | None = None
text_embed_dim: int = 4096
pooled_projection_dim: int = 768
rope_theta: int = 256
qk_norm: str = "rms_norm"
exclude_lora_layers: List[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple, Union
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -40,18 +39,17 @@ class StepVideoArchConfig(DiTArchConfig):
num_attention_heads: int = 48
attention_head_dim: int = 128
in_channels: int = 64
out_channels: Optional[int] = 64
out_channels: int | None = 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: Optional[Union[int, List[int], Tuple[int, ...]]] = field(
caption_channels: int | list[int] | tuple[int, ...] | None = field(
default_factory=lambda: [6144, 1024])
attention_type: Optional[str] = "torch"
use_additional_conditions: Optional[bool] = False
exclude_lora_layers: List[str] = field(default_factory=lambda: [])
attention_type: str | None = "torch"
use_additional_conditions: bool | None = False
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
+3 -22
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -51,25 +50,8 @@ class WanVideoArchConfig(DiTArchConfig):
r"blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
_lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\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
@@ -82,10 +64,9 @@ class WanVideoArchConfig(DiTArchConfig):
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: Optional[int] = None
added_kv_proj_dim: Optional[int] = None
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
exclude_lora_layers: List[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
+11 -11
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
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: 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
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
@dataclass
@@ -61,8 +61,8 @@ class EncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
lora_config: Optional[Any] = None
quant_config: QuantizationConfig | None = None
lora_config: Any | None = None
@dataclass
+4 -5
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
@@ -51,8 +50,8 @@ class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=CLIPTextArchConfig)
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
prefix: str = "clip"
@@ -61,6 +60,6 @@ class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
prefix: str = "clip"
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@@ -12,7 +11,7 @@ class LlamaArchConfig(TextEncoderArchConfig):
intermediate_size: int = 11008
num_hidden_layers: int = 32
num_attention_heads: int = 32
num_key_value_heads: Optional[int] = None
num_key_value_heads: int | None = None
hidden_act: str = "silu"
max_position_embeddings: int = 2048
initializer_range: float = 0.02
@@ -24,11 +23,11 @@ class LlamaArchConfig(TextEncoderArchConfig):
pretraining_tp: int = 1
tie_word_embeddings: bool = False
rope_theta: float = 10000.0
rope_scaling: Optional[float] = None
rope_scaling: float | None = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: Optional[int] = None
head_dim: int | None = None
hidden_state_skip_layer: int = 2
text_len: int = 256
+1 -2
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@@ -12,7 +11,7 @@ class T5ArchConfig(TextEncoderArchConfig):
d_kv: int = 64
d_ff: int = 2048
num_layers: int = 6
num_decoder_layers: Optional[int] = None
num_decoder_layers: int | None = None
num_heads: int = 8
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
+2 -2
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, Union
from typing import Any
import torch
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import StoreBoolean
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0
scaling_factor: float | torch.Tensor = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@@ -9,19 +8,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
+8 -9
View File
@@ -1,5 +1,4 @@
from dataclasses import dataclass, field
from typing import Tuple
import torch
@@ -10,12 +9,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,
@@ -33,7 +32,7 @@ class WanVAEArchConfig(VAEArchConfig):
0.2503,
-0.2921,
)
latents_std: Tuple[float, ...] = (
latents_std: tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
@@ -55,15 +54,15 @@ 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)
@dataclass
class WanVAEConfig(VAEConfig):
arch_config: WanVAEArchConfig = field(default_factory=WanVAEArchConfig)
arch_config: VAEArchConfig = field(default_factory=WanVAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
+14 -13
View File
@@ -1,6 +1,7 @@
import json
from collections.abc import Callable
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
from typing import Any, cast
import torch
@@ -17,7 +18,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
@@ -26,7 +27,8 @@ class PipelineConfig:
"""Base configuration for all pipeline architectures."""
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
flow_shift: float | None = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
@@ -42,20 +44,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: Optional[str] = None
STA_mode: str = "STA_inference"
skip_time_steps: int = 15
mask_strategy_file_path: str | None = None
# Compilation
enable_torch_compile: bool = False
@@ -108,7 +108,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,8 +124,9 @@ 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):
for target_config, source_config in zip(current_value,
new_value,
strict=False):
target_config.update_model_config(source_config)
else:
setattr(self, key, new_value)
+15 -11
View File
@@ -1,5 +1,6 @@
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Callable, Tuple, TypedDict
from typing import TypedDict
import torch
@@ -35,11 +36,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:]
@@ -50,8 +51,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
@@ -68,20 +69,23 @@ class HunyuanConfig(PipelineConfig):
embedded_cfg_scale: int = 6
flow_shift: int = 7
# Video parameters
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):
+5 -5
View File
@@ -1,7 +1,7 @@
"""Registry for pipeline weight-specific configurations."""
import os
from typing import Callable, Dict, Optional, Type
from collections.abc import Callable
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) -> Optional[type[PipelineConfig]]:
pipeline_name_or_path: str) -> type[PipelineConfig] | None:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
@@ -18,6 +18,9 @@ class StepVideoT2VConfig(PipelineConfig):
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 13
timesteps_scale: bool = False
+14 -9
View File
@@ -1,5 +1,5 @@
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch
@@ -11,13 +11,15 @@ 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)]
prompt_embeds_tensor: torch.tensor = torch.stack([
prompt_embeds = [
u[:v] for u, v in zip(hidden_state, seq_lens, strict=False)
]
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
],
@@ -37,20 +39,23 @@ class WanT2V480PConfig(PipelineConfig):
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
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
+6 -6
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
from typing import Any
from fastvideo.v1.logger import init_logger
@@ -15,12 +15,12 @@ class SamplingParam:
data_type: str = "video"
# Image inputs
image_path: Optional[str] = None
image_path: str | None = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[str] = None
prompt_path: Optional[str] = None
prompt: str | list[str] | None = None
negative_prompt: str | None = None
prompt_path: str | None = 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)
+6 -6
View File
@@ -1,5 +1,6 @@
import os
from typing import Any, Callable, Dict, Optional
from collections.abc import Callable
from typing import Any
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
@@ -14,7 +15,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,
@@ -26,7 +27,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(),
@@ -35,7 +36,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":
@@ -46,8 +47,7 @@ SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
}
def get_sampling_param_cls_for_name(
pipeline_name_or_path: str) -> Optional[Any]:
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
"""Get the appropriate sampling param for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
-41
View File
@@ -1,41 +0,0 @@
import os
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
def getdataset(args, start_idx=0) -> T2V_dataset:
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
]
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
]
transform = transforms.Compose([
# Normalize255(),
*resize,
])
transform_topcrop = transforms.Compose([
Normalize255(),
*resize_topcrop,
norm_fun,
])
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
start_idx=start_idx)
raise NotImplementedError(args.dataset)
-82
View File
@@ -1,82 +0,0 @@
# 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_i2v = 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()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_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()),
])
pyarrow_schema_t2v = 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()),
])
-136
View File
@@ -1,136 +0,0 @@
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import torchvision
from tqdm import tqdm
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"num_frames": num_frames
}
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
prepare_dataset_json(args.folder, args.output, args.workers)
-109
View File
@@ -1,109 +0,0 @@
import json
import os
import random
import torch
from torch.utils.data import Dataset
class LatentDataset(Dataset):
def __init__(
self,
json_path,
num_latent_t,
cfg_rate,
) -> None:
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
self.json_path = json_path
self.cfg_rate = cfg_rate
self.datase_dir_path = os.path.dirname(json_path)
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path,
"prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
"prompt_attention_mask")
with open(self.json_path) as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
self.num_latent_t = num_latent_t
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item.get("length", 1) for data_item in self.data_anno
]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
prompt_attention_mask_file = self.data_anno[idx][
"prompt_attention_mask"]
# load
latent = torch.load(
os.path.join(self.latent_dir, latent_file),
map_location="cpu",
weights_only=True,
)
latent = latent.squeeze(0)[:, -self.num_latent_t:]
if random.random() < self.cfg_rate:
prompt_embed = self.uncond_prompt_embed
prompt_attention_mask = self.uncond_prompt_mask
else:
prompt_embed = torch.load(
os.path.join(self.prompt_embed_dir, prompt_embed_file),
map_location="cpu",
weights_only=True,
)
prompt_attention_mask = torch.load(
os.path.join(self.prompt_attention_mask_dir,
prompt_attention_mask_file),
map_location="cpu",
weights_only=True,
)
return latent, prompt_embed, prompt_attention_mask
def __len__(self):
return len(self.data_anno)
def latent_collate_function(batch):
# return latent, prompt, latent_attn_mask, text_attn_mask
# latent_attn_mask: # b t h w
# text_attn_mask: b 1 l
# needs to check if the latent/prompt' size and apply padding & attn mask
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
# calculate max shape
max_t = max([latent.shape[1] for latent in latents])
max_h = max([latent.shape[2] for latent in latents])
max_w = max([latent.shape[3] for latent in latents])
# padding
latent_list: list[torch.Tensor] = [
torch.nn.functional.pad(
latent,
(
0,
max_t - latent.shape[1],
0,
max_h - latent.shape[2],
0,
max_w - latent.shape[3],
),
) for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latent_list), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latent_list):
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
prompt_embeds = torch.stack(prompt_embeds, dim=0)
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
latents = torch.stack(latent_list, dim=0)
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
-470
View File
@@ -1,470 +0,0 @@
import argparse
import json
import os
import random
import time
from collections import defaultdict
from typing import Any, Dict, List
import numpy as np
import pyarrow.parquet as pq
import torch
import tqdm
from einops import rearrange
from torch import distributed as dist
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.distributed import (get_dp_group,
get_sequence_model_parallel_rank,
get_sp_group)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class ParquetVideoTextDataset(Dataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
def __init__(self,
path: str,
batch_size: int = 1024,
rank: int = 0,
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2,
seed: int = 0,
validation: bool = False):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.local_rank = get_sequence_model_parallel_rank()
self.sp_group = get_sp_group()
self.dp_group = get_dp_group()
self.dp_world_size = self.dp_group.world_size
self.sp_world_size = self.sp_group.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.validation = validation
# Negative prompt caching
self.neg_metadata = None
self.cached_neg_prompt: Dict[str, Any] | None = None
self.plan_output_dir = os.path.join(
self.path,
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.json"
)
ranks = get_sp_group().ranks
group_ranks: List[List] = [[] for _ in range(self.world_size)]
torch.distributed.all_gather_object(group_ranks, ranks)
if rank == 0:
# If a plan already exists, then skip creating a new plan
# This will be useful when resume training
if os.path.exists(self.plan_output_dir):
print(f"Using existing plan from {self.plan_output_dir}")
else:
print(f"Creating new plan for {self.plan_output_dir}")
# 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))
# the negative prompt is always the first row in the first
# parquet file
if validation:
self.neg_metadata = metadatas[0]
metadatas = metadatas[1:]
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
# Get all sp groups
# e.g. if num_gpus = 4, sp_size = 2
# group_ranks = [(0, 1), (2, 3)]
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
group_ranks_list: List[Any] = list(
set(tuple(r) for r in group_ranks))
num_sp_groups = len(group_ranks_list)
plan = defaultdict(list)
for idx, metadata in enumerate(metadatas):
sp_group_idx = idx % num_sp_groups
for global_rank in group_ranks_list[sp_group_idx]:
plan[global_rank].append(metadata)
if validation:
assert self.neg_metadata is not None
plan["negative_prompt"] = [self.neg_metadata]
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
else:
pass
dist.barrier()
if validation:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.neg_metadata = plan["negative_prompt"][0]
def _load_and_cache_negative_prompt(self) -> None:
"""Load and cache the negative prompt. Only rank 0 in each SP group should call this."""
if not self.validation or self.neg_metadata is None:
return
if self.cached_neg_prompt is not None:
return
# Only rank 0 in each SP group should read the negative prompt
try:
file_path, row_idx = self.neg_metadata
parquet_file = pq.ParquetFile(file_path)
# Since negative prompt is always the first row (row_idx = 0),
# it's always in the first row group
row_group_index = 0
local_index = row_idx # This will be 0 for the negative prompt
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
# Process the negative prompt row
self.cached_neg_prompt = self._process_row(row_dict)
except Exception as e:
logger.error("Failed to load negative prompt: %s", e)
self.cached_neg_prompt = None
def get_validation_negative_prompt(
self
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Dict[str, Any]]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
Returns the processed negative prompt data (latents, embeddings, masks, info).
"""
if not self.validation:
raise ValueError(
"get_validation_negative_prompt() can only be called in validation mode"
)
# Load and cache if needed (only rank 0 in SP group will actually load)
if self.cached_neg_prompt is None:
self._load_and_cache_negative_prompt()
if self.cached_neg_prompt is None:
raise RuntimeError(
f"Rank {self.rank} (SP rank {self.local_rank}): Could not retrieve negative prompt data"
)
# Extract the components
lat, emb, mask, info = (self.cached_neg_prompt["latents"],
self.cached_neg_prompt["embeddings"],
self.cached_neg_prompt["masks"],
self.cached_neg_prompt["info"])
# Apply the same processing as in __getitem__
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 __len__(self):
if self.local_indices is None:
try:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.local_indices = plan[str(self.rank)]
except Exception as err:
raise Exception(
"The data plan hasn't been created yet") from err
assert self.local_indices is not None
return len(self.local_indices)
def __getitem__(self, idx):
if self.local_indices is None:
try:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.local_indices = plan[self.rank]
except Exception as err:
raise Exception(
"The data plan hasn't been created yet") from err
assert self.local_indices is not None
file_path, row_idx = self.local_indices[idx]
parquet_file = pq.ParquetFile(file_path)
# Calculate the row group to read into memory and the local idx
# This way we can avoid reading in the entire parquet file
cumulative = 0
for i in range(parquet_file.num_row_groups):
num_rows = parquet_file.metadata.row_group(i).num_rows
if cumulative + num_rows > row_idx:
row_group_index = i
local_index = row_idx - cumulative
break
cumulative += num_rows
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
processed = self._process_row(row_dict)
lat, emb, mask, info = processed["latents"], processed[
"embeddings"], processed["masks"], processed["info"]
if lat.numel() == 0: # Validation parquet
return lat, emb, mask, info
else:
lat = lat[:, -self.num_latent_t:]
if self.sp_world_size > 1:
lat = rearrange(lat,
"t (n s) h w -> t n s h w",
n=self.sp_world_size).contiguous()
lat = lat[:, self.local_rank, :, :, :]
return lat, emb, mask, info
def _process_row(self, row) -> Dict[str, Any]:
"""Process a PyArrow batch into tensors."""
vae_latent_bytes = row["vae_latent_bytes"]
vae_latent_shape = row["vae_latent_shape"]
text_embedding_bytes = row["text_embedding_bytes"]
text_embedding_shape = row["text_embedding_shape"]
text_attention_mask_bytes = row["text_attention_mask_bytes"]
text_attention_mask_shape = row["text_attention_mask_shape"]
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
# Collect metadata
info = {
"width": row["width"],
"height": row["height"],
"num_frames": row["num_frames"],
"duration_sec": row["duration_sec"],
"fps": row["fps"],
"file_name": row["file_name"],
"caption": row["caption"],
}
return {
"latents": torch.from_numpy(lat),
"embeddings": torch.from_numpy(emb),
"masks": torch.from_numpy(msk),
"info": info
}
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description='Benchmark Parquet dataset loading speed')
parser.add_argument('--path',
type=str,
default="your/dataset/path",
help='Path to Parquet dataset')
parser.add_argument('--batch_size',
type=int,
default=4,
help='Batch size for DataLoader')
parser.add_argument('--num_batches',
type=int,
default=100,
help='Number of batches to benchmark')
parser.add_argument('--vae_debug', action="store_true")
args = parser.parse_args()
# Initialize distributed training
local_rank = int(os.environ.get("LOCAL_RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
rank = int(os.environ.get("RANK", 0))
# Initialize CUDA device first
if torch.cuda.is_available():
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}")
else:
device = torch.device("cpu")
# Initialize distributed training
if world_size > 1:
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=rank)
print(
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
)
# Create dataset
dataset = ParquetVideoTextDataset(
args.path,
batch_size=args.batch_size,
rank=rank,
world_size=world_size,
)
# Create DataLoader with proper settings
dataloader = StatefulDataLoader(
dataset,
batch_size=args.batch_size,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
# Example of how to load dataloader state
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
# dataloader.load_state_dict(dataloader_state[rank])
# Warm-up with synchronization
if rank == 0:
print("Warming up...")
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
# Example of how to save dataloader state
# if i == 30:
# dist.barrier()
# local_data = {rank: dataloader.state_dict()}
# gathered_data = [None] * world_size
# dist.all_gather_object(gathered_data, local_data)
# if rank == 0:
# global_state_dict = {}
# for d in gathered_data:
# global_state_dict.update(d)
# torch.save(global_state_dict, "dataloader_state.pt")
assert torch.sum(masks[0]).item() == torch.count_nonzero(
embeddings[0]).item() // 4096
if args.vae_debug:
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.models.loader.component_loader import VAELoader
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
fastvideo_args = FastVideoArgs(
model_path=VAE_PATH,
vae_config=WanVAEConfig(load_encoder=False),
vae_precision="fp32")
fastvideo_args.device = device
vae_loader = VAELoader()
vae = vae_loader.load(model_path=VAE_PATH,
architecture="",
fastvideo_args=fastvideo_args)
videoprocessor = VideoProcessor(vae_scale_factor=8)
with torch.inference_mode():
video = vae.decode(latents[0].unsqueeze(0).to(device))
video = videoprocessor.postprocess_video(video)
video_path = os.path.join("/workspace/FastVideo/debug_videos",
infos["caption"][0][:50] + ".mp4")
export_to_video(video[0], video_path, fps=16)
# Move data to device
# latents = latents.to(device)
# embeddings = embeddings.to(device)
if world_size > 1:
dist.barrier()
# Benchmark
if rank == 0:
print(f"Benchmarking with batch_size={args.batch_size}")
start_time = time.time()
total_samples = 0
for i, (latents, embeddings, masks,
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
if i >= args.num_batches:
break
# Move data to device
latents = latents.to(device)
embeddings = embeddings.to(device)
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
# Print progress only from rank 0
if rank == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
print(
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
)
# Final statistics
if world_size > 1:
dist.barrier()
if rank == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
print("\nBenchmark Results:")
print(f"Total time: {elapsed:.2f} seconds")
print(f"Total samples: {total_samples}")
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
if world_size > 1:
dist.destroy_process_group()
-351
View File
@@ -1,351 +0,0 @@
import json
import math
import os
import random
from collections import Counter
from os.path import join as opj
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from torch.utils.data import Dataset
from fastvideo.utils.dataset_utils import DecordInit
from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances: dict[type, 'SingletonMeta'] = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
class DataSetProg(metaclass=SingletonMeta):
def __init__(self) -> None:
self.cap_list: list[dict] = []
self.elements: list[int] = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements: dict[int, list[int]] = {}
self.n_used_elements: dict[int, int] = {}
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
self.elements = list(range(n_elements))
random.shuffle(self.elements)
print(f"n_elements: {len(self.elements)}", flush=True)
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(
math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info) -> int:
worker_id = 0 if work_info is None else work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
dataset_prog = DataSetProg()
def filter_resolution(h: int,
w: int,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16) -> bool:
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class T2V_dataset(Dataset):
def __init__(self,
args,
transform,
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0) -> None:
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
self.use_image_num = args.use_image_num
self.transform = transform
self.transform_topcrop = transform_topcrop
self.temporal_sample = temporal_sample
self.tokenizer = tokenizer
self.text_max_length = args.text_max_length
self.cfg = args.cfg
self.speed_factor = args.speed_factor
self.max_height = args.max_height
self.max_width = args.max_width
self.drop_short_ratio = args.drop_short_ratio
assert self.speed_factor >= 1
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if "mt5" not in args.text_encoder_name:
self.support_Chinese = False
cap_list = self.get_cap_list()
assert len(cap_list) > 0
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
def set_checkpoint(self, n_used_elements):
for i in range(len(dataset_prog.n_used_elements)):
dataset_prog.n_used_elements[i] = n_used_elements
def __len__(self):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
def get_data(self, idx) -> dict:
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx) -> dict:
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
assert video.dtype == torch.uint8
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
text = dataset_prog.cap_list[idx]["cap"]
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx) -> dict:
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
# for i in image:
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images"
in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps: list[str] = (image_data["cap"] if isinstance(
image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
single_text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
single_text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"] # 1, l
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
return dict(
pixel_values=image,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=image_data["path"],
)
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
cnt_too_short = 0
cnt_no_cap = 0
cnt_no_resolution = 0
cnt_resolution_mismatch = 0
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
if fps is None or duration is None:
continue
# ======resolution mismatch=====
resolution = i.get("resolution", None)
if resolution is None:
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None
or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"][
"width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
if not is_pick:
print("resolution mismatch")
cnt_resolution_mismatch += 1
continue
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
# from IPython import embed; embed()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps * self.speed_factor
): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"],
frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(
len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
new_cap_list.append(i)
i["sample_num_frames"] = 1
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices) -> torch.Tensor:
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data) -> list[dict]:
cap_lists = []
with open(data) as f:
folder_anno = [
i.strip().split(",") for i in f.readlines()
if len(i.strip()) > 0
]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno) as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
return cap_lists
def get_cap_list(self) -> list:
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
-153
View File
@@ -1,153 +0,0 @@
import random
import torch
def _is_tensor_video_clip(clip) -> bool:
if not torch.is_tensor(clip):
raise TypeError(f"clip should be Tensor. Got {type(clip)}")
if not clip.ndimension() == 4:
raise ValueError(f"clip should be 4D. Got {clip.dim()}D")
return True
def crop(clip, i, j, h, w) -> torch.Tensor:
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
"""
if len(clip.size()) != 4:
raise ValueError("clip should be a 4D tensor")
return clip[..., i:i + h, j:j + w]
def resize(clip, target_size, interpolation_mode) -> torch.Tensor:
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
return torch.nn.functional.interpolate(
clip,
size=target_size,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def center_crop_th_tw(clip, th, tw, top_crop) -> torch.Tensor:
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
# import ipdb;ipdb.set_trace()
h, w = clip.size(-2), clip.size(-1)
tr = th / tw
if h / w > tr:
new_h = int(w * tr)
new_w = w
else:
new_h = h
new_w = int(h / tr)
i = 0 if top_crop else int(round((h - new_h) / 2.0))
j = int(round((w - new_w) / 2.0))
return crop(clip, i, j, new_h, new_w)
def normalize_video(clip) -> torch.Tensor:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
permute the dimensions of clip tensor
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError(
f"clip tensor should have data type uint8. Got {clip.dtype}")
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
class CenterCropResizeVideo:
"""
First use the short side for cropping length,
center crop video, then resize to the specified size
"""
def __init__(
self,
size,
top_crop=False,
interpolation_mode="bilinear",
) -> None:
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
def __call__(self, clip) -> torch.Tensor:
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop_th_tw(clip,
self.size[0],
self.size[1],
top_crop=self.top_crop)
clip_center_crop_resize = resize(
clip_center_crop,
target_size=self.size,
interpolation_mode=self.interpolation_mode,
)
return clip_center_crop_resize
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self) -> None:
pass
def __call__(self, clip) -> torch.Tensor:
"""
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
return normalize_video(clip)
def __repr__(self) -> str:
return self.__class__.__name__
class TemporalRandomCrop:
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, size) -> None:
self.size = size
def __call__(self, total_frames) -> tuple[int, int]:
rand_end = max(0, total_frames - self.size - 1)
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
-10
View File
@@ -1,10 +0,0 @@
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="")
+3 -11
View File
@@ -2,27 +2,19 @@
from fastvideo.v1.distributed.communication_op import *
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_data_parallel_rank,
get_data_parallel_world_size, get_dp_group,
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size,
get_sp_group, get_tensor_model_parallel_rank,
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,
model_parallel_is_initialized)
init_distributed_environment, initialize_model_parallel)
from fastvideo.v1.distributed.utils import *
__all__ = [
"init_distributed_environment",
"initialize_model_parallel",
"get_data_parallel_world_size",
"get_data_parallel_rank",
"get_sequence_model_parallel_rank",
"get_sequence_model_parallel_world_size",
"get_tensor_model_parallel_rank",
"get_tensor_model_parallel_world_size",
"cleanup_dist_env_and_memory",
"get_world_group",
"get_dp_group",
"get_sp_group",
"model_parallel_is_initialized",
]
@@ -1,182 +1,14 @@
# 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 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
from torch.distributed import ProcessGroup
class DeviceCommunicatorBase:
"""
Base class for device-specific communicator with autograd support.
Base class for device-specific communicator.
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.
@@ -184,8 +16,8 @@ class DeviceCommunicatorBase:
def __init__(self,
cpu_group: ProcessGroup,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
unique_name: str = ""):
self.device = device or torch.device("cpu")
self.cpu_group = cpu_group
@@ -199,33 +31,40 @@ class DeviceCommunicatorBase:
self.rank_in_group = dist.get_group_rank(self.cpu_group,
self.global_rank)
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_reduce(self, input_: torch.Tensor) -> torch.Tensor:
dist.all_reduce(input_, group=self.device_group)
return input_
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()
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)
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
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]:
dim: int = -1) -> torch.Tensor | None:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
@@ -254,7 +93,82 @@ class DeviceCommunicatorBase:
output_tensor = None
return output_tensor
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
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:
"""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:
@@ -264,7 +178,7 @@ class DeviceCommunicatorBase:
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
src: int | None = 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,8 +1,6 @@
# 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
@@ -14,37 +12,35 @@ class CudaCommunicator(DeviceCommunicatorBase):
def __init__(self,
cpu_group: ProcessGroup,
device: Optional[torch.device] = None,
device_group: Optional[ProcessGroup] = None,
device: torch.device | None = None,
device_group: ProcessGroup | None = 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: Optional[PyNcclCommunicator] = None
self.pynccl_comm: PyNcclCommunicator | None = None
if self.world_size > 1:
self.pynccl_comm = PyNcclCommunicator(
group=self.cpu_group,
device=self.device,
)
def all_reduce(self,
input_,
op: Optional[torch.distributed.ReduceOp] = None):
def all_reduce(self, input_):
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
out = pynccl_comm.all_reduce(input_, op=op)
out = pynccl_comm.all_reduce(input_)
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, op=op)
torch.distributed.all_reduce(out, group=self.device_group)
return out
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
def send(self, tensor: torch.Tensor, dst: int | None = 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:
@@ -59,7 +55,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
src: int | None = 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,8 +1,6 @@
# 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
@@ -22,9 +20,9 @@ class PyNcclCommunicator:
def __init__(
self,
group: Union[ProcessGroup, StatelessProcessGroup],
device: Union[int, str, torch.device],
library_path: Optional[str] = None,
group: ProcessGroup | StatelessProcessGroup,
device: int | str | torch.device,
library_path: str | None = None,
):
"""
Args:
@@ -27,7 +27,7 @@
import ctypes
import platform
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any
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: Optional[str] = None):
def __init__(self, so_file: str | None = 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
+57 -106
View File
@@ -27,15 +27,16 @@ 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, Callable, Dict, List, Optional, Tuple, Union
from typing import Any, Optional
from unittest.mock import patch
import torch
import torch.distributed
from torch.distributed import Backend, ProcessGroup, ReduceOp
from torch.distributed import Backend, ProcessGroup
import fastvideo.v1.envs as envs
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
@@ -57,15 +58,15 @@ TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
def _split_tensor_dict(
tensor_dict: Dict[str, Union[torch.Tensor, Any]]
) -> Tuple[List[Tuple[str, Any]], List[torch.Tensor]]:
tensor_dict: dict[str, 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,
@@ -81,7 +82,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:
@@ -97,7 +98,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:
@@ -128,7 +129,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:
@@ -143,16 +144,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: Optional[Any] # shared memory broadcaster
mq_broadcaster: Any | None # shared memory broadcaster
def __init__(
self,
group_ranks: List[List[int]],
group_ranks: list[list[int]],
local_rank: int,
torch_distributed_backend: Union[str, Backend],
torch_distributed_backend: str | Backend,
use_device_communicator: bool,
use_message_queue_broadcaster: bool = False,
group_name: Optional[str] = None,
group_name: str | None = None,
):
group_name = group_name or "anonymous"
self.unique_name = _get_unique_name(group_name)
@@ -243,8 +244,8 @@ class GroupCoordinator:
return self.ranks[(rank_in_group - 1) % world_size]
@contextmanager
def graph_capture(
self, graph_capture_context: Optional[GraphCaptureContext] = None):
def graph_capture(self,
graph_capture_context: GraphCaptureContext | None = None):
if graph_capture_context is None:
stream = torch.cuda.Stream()
graph_capture_context = GraphCaptureContext(stream)
@@ -260,11 +261,7 @@ class GroupCoordinator:
with torch.cuda.stream(stream):
yield graph_capture_context
def all_reduce(
self,
input_: torch.Tensor,
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
) -> torch.Tensor:
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
"""
User-facing all-reduce function before we actually call the
all-reduce operation.
@@ -287,14 +284,10 @@ class GroupCoordinator:
return torch.ops.vllm.all_reduce(input_,
group_name=self.unique_name)
else:
return self._all_reduce_out_place(input_, op=op)
return self._all_reduce_out_place(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_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
return self.device_communicator.all_reduce(input_)
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
world_size = self.world_size
@@ -309,7 +302,7 @@ class GroupCoordinator:
def gather(self,
input_: torch.Tensor,
dst: int = 0,
dim: int = -1) -> Optional[torch.Tensor]:
dim: int = -1) -> torch.Tensor | None:
"""
NOTE: We assume that the input tensor is on the same device across
all the ranks.
@@ -345,7 +338,7 @@ class GroupCoordinator:
group=self.device_group)
return input_
def broadcast_object(self, obj: Optional[Any] = None, src: int = 0):
def broadcast_object(self, obj: Any | None = None, src: int = 0):
"""Broadcast the input object.
NOTE: `src` is the local rank of the source rank.
"""
@@ -370,9 +363,9 @@ class GroupCoordinator:
return recv[0]
def broadcast_object_list(self,
obj_list: List[Any],
obj_list: list[Any],
src: int = 0,
group: Optional[ProcessGroup] = None):
group: ProcessGroup | None = None):
"""Broadcast the input object list.
NOTE: `src` is the local rank of the source rank.
"""
@@ -452,11 +445,11 @@ class GroupCoordinator:
def broadcast_tensor_dict(
self,
tensor_dict: Optional[Dict[str, Union[torch.Tensor, Any]]] = None,
tensor_dict: dict[str, torch.Tensor | Any] | None = None,
src: int = 0,
group: Optional[ProcessGroup] = None,
metadata_group: Optional[ProcessGroup] = None
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
group: ProcessGroup | None = None,
metadata_group: ProcessGroup | None = None
) -> dict[str, torch.Tensor | Any] | None:
"""Broadcast the input tensor dictionary.
NOTE: `src` is the local rank of the source rank.
"""
@@ -470,7 +463,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)}")
@@ -537,10 +530,10 @@ class GroupCoordinator:
def send_tensor_dict(
self,
tensor_dict: Dict[str, Union[torch.Tensor, Any]],
dst: Optional[int] = None,
tensor_dict: dict[str, torch.Tensor | Any],
dst: int | None = None,
all_gather_group: Optional["GroupCoordinator"] = None,
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
) -> dict[str, torch.Tensor | Any] | None:
"""Send the input tensor dictionary.
NOTE: `dst` is the local rank of the source rank.
"""
@@ -560,7 +553,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)}"
@@ -591,9 +584,9 @@ class GroupCoordinator:
def recv_tensor_dict(
self,
src: Optional[int] = None,
src: int | None = None,
all_gather_group: Optional["GroupCoordinator"] = None,
) -> Optional[Dict[str, Union[torch.Tensor, Any]]]:
) -> dict[str, torch.Tensor | Any] | None:
"""Recv the input tensor dictionary.
NOTE: `src` is the local rank of the source rank.
"""
@@ -614,7 +607,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,
@@ -655,7 +648,7 @@ class GroupCoordinator:
tensor_dict[key] = value
return tensor_dict
def barrier(self) -> None:
def barrier(self):
"""Barrier synchronization among the group.
NOTE: don't use `device_group` here! `barrier` in NCCL is
terrible because it is internally a broadcast operation with
@@ -664,7 +657,7 @@ class GroupCoordinator:
"""
torch.distributed.barrier(group=self.cpu_group)
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
def send(self, tensor: torch.Tensor, dst: int | None = 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)
@@ -672,7 +665,7 @@ class GroupCoordinator:
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: Optional[int] = None) -> torch.Tensor:
src: int | None = 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)
@@ -690,7 +683,7 @@ class GroupCoordinator:
self.mq_broadcaster = None
_WORLD: Optional[GroupCoordinator] = None
_WORLD: GroupCoordinator | None = None
def get_world_group() -> GroupCoordinator:
@@ -698,23 +691,23 @@ 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],
local_rank=local_rank,
torch_distributed_backend=backend,
use_device_communicator=True,
use_device_communicator=False,
group_name="world",
)
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: Optional[str] = None,
group_name: str | None = None,
) -> GroupCoordinator:
return GroupCoordinator(
@@ -727,7 +720,7 @@ def init_model_parallel_group(
)
_TP: Optional[GroupCoordinator] = None
_TP: GroupCoordinator | None = None
def get_tp_group() -> GroupCoordinator:
@@ -747,10 +740,10 @@ def set_custom_all_reduce(enable: bool):
def init_distributed_environment(
world_size: int = 1,
rank: int = 0,
world_size: int = -1,
rank: int = -1,
distributed_init_method: str = "env://",
local_rank: int = 0,
local_rank: int = -1,
backend: str = "nccl",
):
logger.debug(
@@ -786,7 +779,7 @@ def init_distributed_environment(
"world group already initialized with a different world size")
_SP: Optional[GroupCoordinator] = None
_SP: GroupCoordinator | None = None
def get_sp_group() -> GroupCoordinator:
@@ -794,19 +787,10 @@ def get_sp_group() -> GroupCoordinator:
return _SP
_DP: Optional[GroupCoordinator] = None
def get_dp_group() -> GroupCoordinator:
assert _DP is not None, ("data parallel group is not initialized")
return _DP
def initialize_model_parallel(
tensor_model_parallel_size: int = 1,
sequence_model_parallel_size: int = 1,
data_parallel_size: int = 1,
backend: Optional[str] = None,
backend: str | None = None,
) -> None:
"""
Initialize model parallel groups.
@@ -861,22 +845,6 @@ def initialize_model_parallel(
backend,
group_name="sp")
# Build the data parallel groups.
num_data_parallel_groups: int = (world_size // data_parallel_size)
global _DP
assert _DP is None, ("data parallel group is already initialized")
group_ranks = []
for i in range(num_data_parallel_groups):
ranks = list(range(i * data_parallel_size,
(i + 1) * data_parallel_size))
group_ranks.append(ranks)
_DP = init_model_parallel_group(group_ranks,
get_world_group().local_rank,
backend,
group_name="dp")
def get_sequence_model_parallel_world_size() -> int:
"""Return world size for the sequence model parallel group."""
@@ -888,21 +856,10 @@ def get_sequence_model_parallel_rank() -> int:
return get_sp_group().rank_in_group
def get_data_parallel_world_size() -> int:
"""Return world size for the data parallel group."""
return get_dp_group().world_size
def get_data_parallel_rank() -> int:
"""Return my rank for the data parallel group."""
return get_dp_group().rank_in_group
def ensure_model_parallel_initialized(
tensor_model_parallel_size: int,
sequence_model_parallel_size: int,
data_parallel_size: int,
backend: Optional[str] = None,
backend: str | None = None,
) -> None:
"""Helper to initialize model parallel groups if they are not initialized,
or ensure tensor-parallel, sequence-parallel sizes
@@ -912,8 +869,7 @@ def ensure_model_parallel_initialized(
get_world_group().device_group)
if not model_parallel_is_initialized():
initialize_model_parallel(tensor_model_parallel_size,
sequence_model_parallel_size,
data_parallel_size, backend)
sequence_model_parallel_size, backend)
return
assert (
@@ -932,7 +888,7 @@ def ensure_model_parallel_initialized(
def model_parallel_is_initialized() -> bool:
"""Check if tensor, sequence parallel groups are initialized."""
return _TP is not None and _SP is not None and _DP is not None
return _TP is not None and _SP is not None
_TP_STATE_PATCHED = False
@@ -985,11 +941,6 @@ def destroy_model_parallel() -> None:
_SP.destroy()
_SP = None
global _DP
if _DP:
_DP.destroy()
_DP = None
def destroy_distributed_environment() -> None:
global _WORLD
@@ -1019,8 +970,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: Union[ProcessGroup, StatelessProcessGroup],
source_rank: int = 0) -> List[bool]:
def in_the_same_node_as(pg: 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
@@ -1106,7 +1057,7 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
def initialize_tensor_parallel_group(
tensor_model_parallel_size: int = 1,
backend: Optional[str] = None,
backend: str | None = None,
group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a tensor parallel group for a specific model.
@@ -1170,7 +1121,7 @@ def initialize_tensor_parallel_group(
def initialize_sequence_parallel_group(
sequence_model_parallel_size: int = 1,
backend: Optional[str] = None,
backend: str | None = None,
group_name_suffix: str = "") -> GroupCoordinator:
"""Initialize a sequence parallel group for a specific model.
+7 -6
View File
@@ -9,7 +9,8 @@ import dataclasses
import pickle
import time
from collections import deque
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
from collections.abc import Sequence
from typing import Any
import torch
from torch.distributed import TCPStore
@@ -72,15 +73,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
@@ -114,7 +115,7 @@ class StatelessProcessGroup:
self.recv_src_counter[src] += 1
return obj
def broadcast_obj(self, obj: Optional[Any], src: int) -> Any:
def broadcast_obj(self, obj: Any | None, 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.
+6 -6
View File
@@ -4,7 +4,7 @@
import argparse
import dataclasses
import os
from typing import Any, Dict, List, Optional, cast
from typing import Any, 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: Optional[str] = None) -> None:
args_dict: dict[str, Any],
prefix: str | None = None) -> None:
"""
Update configuration object from arguments dictionary.
+1 -3
View File
@@ -1,14 +1,12 @@
# 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())
+2 -3
View File
@@ -4,7 +4,6 @@ import argparse
import os
import subprocess
import sys
from typing import List, Optional
from fastvideo.v1.logger import init_logger
@@ -19,8 +18,8 @@ class RaiseNotImplementedAction(argparse.Action):
def launch_distributed(num_gpus: int,
args: List[str],
master_port: Optional[int] = None) -> int:
args: list[str],
master_port: int | None = None) -> int:
"""
Launch a distributed job with the given arguments
+12 -16
View File
@@ -10,7 +10,7 @@ import gc
import math
import os
import time
from typing import Any, Dict, List, Optional, Union
from typing import Any
import imageio
import numpy as np
@@ -44,8 +44,8 @@ class VideoGenerator:
Initialize the video generator.
Args:
pipeline: The pipeline to use for inference
fastvideo_args: The inference arguments
executor_class: The executor class to use for inference
"""
self.fastvideo_args = fastvideo_args
self.executor = executor_class(fastvideo_args)
@@ -53,11 +53,9 @@ class VideoGenerator:
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: str | PipelineConfig | None = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -118,6 +116,7 @@ class VideoGenerator:
# initialize_distributed_and_parallelism(fastvideo_args)
executor_class = Executor.get_class(fastvideo_args)
return cls(
fastvideo_args=fastvideo_args,
executor_class=executor_class,
@@ -127,9 +126,9 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
sampling_param: Optional[SamplingParam] = None,
sampling_param: SamplingParam | None = None,
**kwargs,
) -> Union[Dict[str, Any], List[np.ndarray]]:
) -> dict[str, Any] | list[np.ndarray]:
"""
Generate a video based on the given prompt.
@@ -275,10 +274,10 @@ class VideoGenerator:
# Save video if requested
if batch.save_video:
output_path = batch.output_path
if output_path:
os.makedirs(output_path, exist_ok=True)
video_path = os.path.join(output_path, f"{prompt[:100]}.mp4")
save_path = batch.output_path
if save_path:
os.makedirs(os.path.dirname(save_path), exist_ok=True)
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", video_path)
else:
@@ -294,9 +293,6 @@ class VideoGenerator:
"generation_time": gen_time
}
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
self.executor.set_lora_adapter(lora_nickname, lora_path)
def shutdown(self):
"""
Shutdown the video generator.
+13 -12
View File
@@ -2,28 +2,29 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py
import os
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
from collections.abc import Callable
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
LD_LIBRARY_PATH: Optional[str] = None
FASTVIDEO_NCCL_SO_PATH: str | None = None
LD_LIBRARY_PATH: str | None = None
LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: Optional[str] = None
CUDA_VISIBLE_DEVICES: str | None = 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: Optional[str] = None
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
FASTVIDEO_ATTENTION_CONFIG: Optional[str] = None
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: Optional[str] = None
NVCC_THREADS: Optional[str] = None
CMAKE_BUILD_TYPE: Optional[str] = None
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_SERVER_DEV_MODE: bool = False
@@ -42,7 +43,7 @@ def get_default_config_root() -> str:
)
def maybe_convert_int(value: Optional[str]) -> Optional[int]:
def maybe_convert_int(value: str | None) -> int | None:
if value is None:
return None
return int(value)
@@ -53,7 +54,7 @@ def maybe_convert_int(value: Optional[str]) -> Optional[int]:
# begin-env-vars-definition
environment_variables: Dict[str, Callable[[], Any]] = {
environment_variables: dict[str, Callable[[], Any]] = {
# ================== Installation Time Env Vars ==================
+20 -435
View File
@@ -4,9 +4,10 @@
import argparse
import dataclasses
from collections.abc import Callable
from contextlib import contextmanager
from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple
from typing import Any
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.logger import init_logger
@@ -38,27 +39,23 @@ class FastVideoArgs:
# HuggingFace specific parameters
trust_remote_code: bool = False
revision: Optional[str] = None
revision: str | None = None
# Parallelism
num_gpus: int = 1
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dp_size: int = 1
dp_shards: Optional[int] = None
dist_timeout: Optional[int] = None # timeout for torch.distributed
tp_size: int | None = None
sp_size: int | None = None
dist_timeout: int | None = None # timeout for torch.distributed
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
flow_shift: float | None = None
output_type: str = "pil"
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
precision: str = "bf16"
use_cpu_offload: bool = True
use_fsdp_inference: bool = True
# VAE configuration
vae_precision: str = "fp16"
@@ -74,49 +71,36 @@ 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
STA_mode: str = "STA_inference"
skip_time_steps: int = 15
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# STA parameters
mask_strategy_file_path: Optional[str] = None
mask_strategy_file_path: str | None = None
enable_torch_compile: bool = False
use_cpu_offload: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
pos_magic: str | None = None
neg_magic: str | None = None
timesteps_scale: bool | None = None
# Logging
log_level: str = "info"
# Inference parameters
device_str: Optional[str] = None
device_str: str | None = None
device = None
@property
def training_mode(self) -> bool:
return not self.inference_mode
def __post_init__(self):
pass
@@ -149,13 +133,6 @@ 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",
@@ -192,20 +169,6 @@ class FastVideoArgs:
default=FastVideoArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--data-parallel-size",
"--dp-size",
type=int,
default=FastVideoArgs.dp_size,
help="The data parallelism size.",
)
parser.add_argument(
"--data-parallel-shards",
"--dp-shards",
type=int,
default=FastVideoArgs.dp_shards,
help="The data parallelism shards.",
)
parser.add_argument(
"--dist-timeout",
type=int,
@@ -281,21 +244,6 @@ class FastVideoArgs:
)
# STA (Spatial-Temporal Attention) parameters
parser.add_argument(
"--STA-mode",
type=str,
default=FastVideoArgs.STA_mode,
choices=[
"STA_inference", "STA_searching", "STA_tuning", "STA_tuning_cfg"
],
help="STA mode",
)
parser.add_argument(
"--skip-time-steps",
type=int,
default=FastVideoArgs.skip_time_steps,
help="Number of time steps to warmup (full attention) for STA",
)
parser.add_argument(
"--mask-strategy-file-path",
type=str,
@@ -311,16 +259,8 @@ class FastVideoArgs:
parser.add_argument(
"--use-cpu-offload",
action=StoreBoolean,
help=
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
help="Use CPU offload for the model load",
)
parser.add_argument(
"--use-fsdp-inference",
action=StoreBoolean,
help=
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -382,10 +322,6 @@ class FastVideoArgs:
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
# Use getattr with default value from the dataclass for potentially missing attributes
@@ -397,20 +333,10 @@ class FastVideoArgs:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if not self.inference_mode:
assert self.dp_size is not None, "dp_size must be set for training"
assert self.dp_shards is not None, "dp_shards must be set for training"
assert self.sp_size is not None, "sp_size must be set for training"
if self.tp_size is None:
self.tp_size = self.num_gpus
if self.sp_size is None:
self.sp_size = self.num_gpus
if self.dp_shards is None:
self.dp_shards = self.num_gpus
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
assert self.dp_size <= self.num_gpus and self.num_gpus % self.dp_size == 0, "num_gpus must >= and be divisible by dp_size"
assert self.dp_shards <= self.num_gpus and self.num_gpus % self.dp_shards == 0, "num_gpus must >= and be divisible by dp_shards"
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
@@ -451,7 +377,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.
@@ -498,344 +424,3 @@ def get_current_fastvideo_args() -> FastVideoArgs:
# TODO(will): may need to handle this for CI.
raise ValueError("Current fastvideo args is not set.")
return _current_fastvideo_args
@dataclasses.dataclass
class TrainingArgs(FastVideoArgs):
"""
Training arguments. Inherits from FastVideoArgs and adds training-specific
arguments. If there are any conflicts, the training arguments will take
precedence.
"""
data_path: str = ""
dataloader_num_workers: int = 0
num_height: int = 0
num_width: int = 0
num_frames: int = 0
train_batch_size: int = 0
num_latent_t: int = 0
group_frame: bool = False
group_resolution: bool = False
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
cache_dir: str = ""
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
cfg: float = 0.0
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
seed: Optional[int] = None
# output
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: bool = False
logging_dir: str = ""
# optimizer & scheduler
num_train_epochs: int = 0
max_train_steps: int = 0
gradient_accumulation_steps: int = 0
learning_rate: float = 0.0
scale_lr: bool = False
lr_scheduler: str = ""
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
gradient_checkpointing: bool = False
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
train_sp_batch_size: int = 0
fsdp_sharding_startegy: str = ""
weighting_scheme: str = ""
logit_mean: float = 0.0
logit_std: float = 1.0
mode_scale: float = 0.0
num_euler_timesteps: int = 0
lr_num_cycles: int = 0
lr_power: float = 0.0
not_apply_cfg_solver: bool = False
distill_cfg: float = 0.0
scheduler_type: str = ""
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
pred_decay_type: str = ""
hunyuan_teacher_disable_cfg: bool = False
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
@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
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
# 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")
parser.add_argument("--seed",
type=int,
help="Seed for deterministic training")
# Output configuration
parser.add_argument("--output-dir",
type=str,
required=True,
help="Output directory for checkpoints and logs")
parser.add_argument("--checkpoints-total-limit",
type=int,
help="Maximum number of checkpoints to keep")
parser.add_argument("--checkpointing-steps",
type=int,
help="Steps between checkpoints")
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
# Training configuration
parser.add_argument("--num-train-epochs",
type=int,
help="Number of training epochs")
parser.add_argument("--max-train-steps",
type=int,
help="Maximum number of training steps")
parser.add_argument("--gradient-accumulation-steps",
type=int,
help="Number of steps to accumulate gradients")
parser.add_argument("--learning-rate",
type=float,
required=True,
help="Learning rate")
parser.add_argument("--scale-lr",
action=StoreBoolean,
help="Whether to scale learning rate")
parser.add_argument("--lr-scheduler",
type=str,
default="constant",
help="Learning rate scheduler type")
parser.add_argument("--lr-warmup-steps",
type=int,
default=10,
help="Number of warmup steps for learning rate")
parser.add_argument("--max-grad-norm",
type=float,
help="Maximum gradient norm")
parser.add_argument("--gradient-checkpointing",
action=StoreBoolean,
help="Whether to use gradient checkpointing")
parser.add_argument("--selective-checkpointing",
type=float,
help="Selective checkpointing threshold")
parser.add_argument("--allow-tf32",
action=StoreBoolean,
help="Whether to allow TF32")
parser.add_argument("--mixed-precision",
type=str,
help="Mixed precision training type")
parser.add_argument("--train-sp-batch-size",
type=int,
help="Training spatial parallelism batch size")
parser.add_argument("--fsdp-sharding-strategy",
type=str,
help="FSDP sharding strategy")
parser.add_argument(
"--weighting_scheme",
type=str,
default="uniform",
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
)
parser.add_argument(
"--logit_mean",
type=float,
default=0.0,
help="mean to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--logit_std",
type=float,
default=1.0,
help="std to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--mode_scale",
type=float,
default=1.29,
help=
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
# Additional training parameters
parser.add_argument("--num-euler-timesteps",
type=int,
help="Number of Euler timesteps")
parser.add_argument("--lr-num-cycles",
type=int,
help="Number of learning rate cycles")
parser.add_argument("--lr-power",
type=float,
help="Learning rate power")
parser.add_argument("--not-apply-cfg-solver",
action=StoreBoolean,
help="Whether to not apply CFG solver")
parser.add_argument("--distill-cfg",
type=float,
help="Distillation CFG scale")
parser.add_argument("--scheduler-type", type=str, help="Scheduler type")
parser.add_argument("--linear-quadratic-threshold",
type=float,
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
parser.add_argument("--multi-phased-distill-schedule",
type=str,
help="Multi-phased distillation schedule")
parser.add_argument("--pred-decay-weight",
type=float,
help="Prediction decay weight")
parser.add_argument("--pred-decay-type",
type=str,
help="Prediction decay type")
parser.add_argument("--hunyuan-teacher-disable-cfg",
action=StoreBoolean,
help="Whether to disable CFG for Hunyuan teacher")
parser.add_argument("--master-weight-type",
type=str,
help="Master weight type")
return parser
+5 -5
View File
@@ -5,7 +5,7 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING
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: Optional[ForwardBatch] = None
forward_batch: ForwardBatch | None = None
_forward_context: Optional[ForwardContext] = None
_forward_context: ForwardContext | None = 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: Optional[ForwardBatch] = None,
fastvideo_args: Optional[FastVideoArgs] = None):
forward_batch: ForwardBatch | None = None,
fastvideo_args: FastVideoArgs | None = 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.
+207
View File
@@ -0,0 +1,207 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
"""
Inference module for diffusion models.
This module provides classes and functions for running inference with diffusion models.
"""
import time
from typing import Any
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
build_pipeline)
# TODO(will): remove, check if this is hunyuan specific
from fastvideo.v1.utils import align_to
logger = init_logger(__name__)
class InferenceEngine:
"""
Engine for running inference with diffusion models.
"""
def __init__(
self,
pipeline: ComposedPipelineBase,
fastvideo_args: FastVideoArgs,
):
"""
Initialize the inference engine.
Args:
pipeline: The pipeline to use for inference.
fastvideo_args: The inference arguments.
default_negative_prompt: The default negative prompt to use.
"""
self.pipeline = pipeline
self.fastvideo_args = fastvideo_args
@classmethod
def create_engine(
cls,
fastvideo_args: FastVideoArgs,
) -> "InferenceEngine":
"""
Create an inference engine with the specified arguments.
Args:
fastvideo_args: The inference arguments.
model_loader_cls: The model loader class to use. If None, it will be
determined from the model type.
pipeline_type: The type of pipeline to create. If None, it will be
determined from the model type.
Returns:
The created inference engine.
Raises:
ValueError: If the model type is not recognized or if the pipeline type
is not recognized.
"""
logger.info("Building pipeline...")
# TODO(will): I don't really like this api.
# it should be something closer to pipeline_cls.from_pretrained(...)
# this way for training we can just do pipeline_cls.from_pretrained(
# checkpoint_path) and have it handle everything.
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
pipeline = build_pipeline(fastvideo_args)
logger.info("Pipeline Ready")
# Create the inference engine
return cls(pipeline, fastvideo_args)
def run(
self,
prompt: str,
fastvideo_args: FastVideoArgs,
) -> dict[str, Any]:
"""
Run inference with the pipeline.
Args:
prompt: The prompt to use for generation.
negative_prompt: The negative prompt to use. If None, the default will be used.
seed: The random seed to use. If None, a random seed will be used.
**kwargs: Additional arguments to pass to the pipeline.
Returns:
A dictionary containing the generated videos and metadata.
"""
out_dict: dict[str, Any] = dict()
num_videos_per_prompt = fastvideo_args.num_videos
seed = fastvideo_args.seed
height = fastvideo_args.height
width = fastvideo_args.width
video_length = fastvideo_args.num_frames
negative_prompt = fastvideo_args.neg_prompt
infer_steps = fastvideo_args.num_inference_steps
guidance_scale = fastvideo_args.guidance_scale
flow_shift = fastvideo_args.flow_shift
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
image_path = fastvideo_args.image_path
# ========================================================================
# Arguments: target_width, target_height, target_video_length
# ========================================================================
if width <= 0 or height <= 0 or video_length <= 0:
raise ValueError(
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
)
if (video_length - 1) % 4 != 0:
raise ValueError(
f"`video_length-1` must be a multiple of 4, got {video_length}")
target_height = align_to(height, 16)
target_width = align_to(width, 16)
target_video_length = video_length
out_dict["size"] = (target_height, target_width, target_video_length)
# ========================================================================
# Arguments: prompt, new_prompt, negative_prompt
# ========================================================================
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
# negative prompt
if negative_prompt is not None:
negative_prompt = negative_prompt.strip()
# TODO(PY): move to hunyuan stage
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# ========================================================================
# Print infer args
# ========================================================================
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {target_video_length}
prompt: {prompt}
neg_prompt: {negative_prompt}
seed: {seed}
infer_steps: {infer_steps}
num_videos_per_prompt: {num_videos_per_prompt}
guidance_scale: {guidance_scale}
n_tokens: {n_tokens}
flow_shift: {flow_shift}
embedded_guidance_scale: {embedded_guidance_scale}"""
logger.info(debug_str)
# return
# sp_group = get_sp_group()
# local_rank = sp_group.rank
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
image_path=image_path,
prompt=prompt,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
# generator=generator,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={}, # Any additional parameters
)
print('===============================================')
print(batch)
print('===============================================')
print('===============================================')
print(fastvideo_args)
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.perf_counter()
samples = self.pipeline.forward(
batch=batch,
fastvideo_args=fastvideo_args,
).output
# TODO(will): fix and move to hunyuan stage
# out_dict["seeds"] = batch.seeds
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.perf_counter() - start_time
logger.info("Success, time: %s", gen_time)
return out_dict
+3 -2
View File
@@ -1,7 +1,8 @@
# 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 typing import Any, Callable, Dict, Type
from collections.abc import Callable
from typing import Any
import torch.nn as nn
@@ -81,7 +82,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
+4 -5
View File
@@ -1,7 +1,6 @@
# 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
@@ -22,7 +21,7 @@ class RMSNorm(CustomOp):
hidden_size: int,
eps: float = 1e-6,
dtype: torch.dtype = torch.float32,
var_hidden_size: Optional[int] = None,
var_hidden_size: int | None = None,
has_weight: bool = True,
) -> None:
super().__init__()
@@ -40,8 +39,8 @@ class RMSNorm(CustomOp):
def forward_native(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""PyTorch-native implementation equivalent to forward()."""
orig_dtype = x.dtype
x = x.to(torch.float32)
@@ -130,7 +129,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.
+34 -42
View File
@@ -2,7 +2,6 @@
# 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
@@ -40,7 +39,7 @@ WEIGHT_LOADER_V2_SUPPORTED = [
def adjust_scalar_to_fused_array(
param: torch.Tensor, loaded_weight: torch.Tensor,
shard_id: Union[str, int]) -> tuple[torch.Tensor, torch.Tensor]:
shard_id: 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
@@ -91,7 +90,7 @@ class LinearMethodBase(QuantizeMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
bias: torch.Tensor | None = 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
@@ -116,7 +115,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
def apply(self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None) -> torch.Tensor:
bias: torch.Tensor | None = None) -> torch.Tensor:
return F.linear(x, layer.weight, bias)
@@ -138,8 +137,8 @@ class LinearBase(torch.nn.Module):
input_size: int,
output_size: int,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -152,14 +151,13 @@ class LinearBase(torch.nn.Module):
params_dtype = torch.get_default_dtype()
self.params_dtype = params_dtype
if quant_config is None:
self.quant_method: Optional[
QuantizeMethodBase] = UnquantizedLinearMethod()
self.quant_method: QuantizeMethodBase | None = UnquantizedLinearMethod(
)
else:
self.quant_method = quant_config.get_quant_method(self,
prefix=prefix)
def forward(self,
x: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
raise NotImplementedError
@@ -182,8 +180,8 @@ class ReplicatedLinear(LinearBase):
output_size: int,
bias: bool = True,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__(input_size,
output_size,
@@ -223,8 +221,7 @@ 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, Optional[Parameter]]:
def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
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)
@@ -268,9 +265,9 @@ class ColumnParallelLinear(LinearBase):
bias: bool = True,
gather_output: bool = False,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
output_sizes: Optional[list[int]] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
output_sizes: list[int] | None = None,
prefix: str = ""):
# Divide the weight matrix along the last dimension.
self.tp_size = get_tensor_model_parallel_world_size()
@@ -345,9 +342,8 @@ 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, Optional[Parameter]]:
def forward(self,
input_: torch.Tensor) -> tuple[torch.Tensor, Parameter | None]:
bias = self.bias if not self.skip_bias_add else None
# Matrix multiply.
@@ -399,8 +395,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
bias: bool = True,
gather_output: bool = False,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
self.output_sizes = output_sizes
tp_size = get_tensor_model_parallel_world_size()
@@ -417,7 +413,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader(self,
param: Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[int] = None) -> None:
loaded_shard_id: int | None = None) -> None:
param_data = param.data
output_dim = getattr(param, "output_dim", None)
@@ -510,10 +506,8 @@ 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)
@@ -525,7 +519,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self,
param: BasevLLMParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[int] = None) -> None:
loaded_shard_id: int | None = None) -> None:
if loaded_shard_id is None:
if isinstance(param, PerTensorScaleParameter):
param.load_merged_column_weight(loaded_weight=loaded_weight,
@@ -598,11 +592,11 @@ class QKVParallelLinear(ColumnParallelLinear):
hidden_size: int,
head_size: int,
total_num_heads: int,
total_num_kv_heads: Optional[int] = None,
total_num_kv_heads: int | None = None,
bias: bool = True,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
quant_config: Optional[QuantizationConfig] = None,
params_dtype: torch.dtype | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
self.hidden_size = hidden_size
self.head_size = head_size
@@ -637,7 +631,7 @@ class QKVParallelLinear(ColumnParallelLinear):
quant_config=quant_config,
prefix=prefix)
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> Optional[int]:
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> int | None:
shard_offset_mapping = {
"q": 0,
"k": self.num_heads * self.head_size,
@@ -646,7 +640,7 @@ class QKVParallelLinear(ColumnParallelLinear):
}
return shard_offset_mapping.get(loaded_shard_id)
def _get_shard_size_mapping(self, loaded_shard_id: str) -> Optional[int]:
def _get_shard_size_mapping(self, loaded_shard_id: str) -> int | None:
shard_size_mapping = {
"q": self.num_heads * self.head_size,
"k": self.num_kv_heads * self.head_size,
@@ -679,10 +673,8 @@ 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)
@@ -694,7 +686,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader_v2(self,
param: BasevLLMParameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[str] = None):
loaded_shard_id: str | None = 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)
@@ -720,7 +712,7 @@ class QKVParallelLinear(ColumnParallelLinear):
def weight_loader(self,
param: Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: Optional[str] = None):
loaded_shard_id: str | None = None):
param_data = param.data
output_dim = getattr(param, "output_dim", None)
@@ -845,9 +837,9 @@ class RowParallelLinear(LinearBase):
bias: bool = True,
input_is_parallel: bool = True,
skip_bias_add: bool = False,
params_dtype: Optional[torch.dtype] = None,
params_dtype: torch.dtype | None = None,
reduce_results: bool = True,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
# Divide the weight matrix along the first dimension.
self.tp_rank = get_tensor_model_parallel_rank()
@@ -921,7 +913,7 @@ class RowParallelLinear(LinearBase):
param.load_row_parallel_weight(loaded_weight=loaded_weight)
def forward(self, input_) -> tuple[torch.Tensor, Optional[Parameter]]:
def forward(self, input_) -> tuple[torch.Tensor, Parameter | None]:
if self.input_is_parallel:
input_parallel = input_
else:
-295
View File
@@ -1,295 +0,0 @@
# Code adapted from SGLang https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/lora/layers.py
from typing import Dict, List, Tuple, Type, Union
import torch
from torch import nn
from torch.distributed.tensor import DTensor, distribute_tensor
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.linear import (ColumnParallelLinear, LinearBase,
MergedColumnParallelLinear,
QKVParallelLinear, ReplicatedLinear,
RowParallelLinear)
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
class BaseLayerWithLoRA(nn.Module):
def __init__(
self,
base_layer: nn.Module,
):
super().__init__()
self.base_layer: nn.Module = base_layer
self.lora_A: torch.Tensor = None
self.lora_B: torch.Tensor = None
self.merged: bool = False
self.weight = base_layer.weight
self.cpu_weight = base_layer.weight.to("cpu")
self.unmerge_count = 0
# indicates adapter weights don't contain this layer
# (which shouldn't normally happen, but we want to separate it from the case of erroneous merging)
self.disable_lora: bool = False
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.base_layer.forward(x)
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
return B
def set_lora_weights(self,
A: torch.Tensor,
B: torch.Tensor,
training_mode: bool = False) -> None:
self.lora_A = A # share storage with weights in the pipeline
self.lora_B = B
self.disable_lora = False
if not training_mode:
self.merge_lora_weights()
@torch.no_grad()
def merge_lora_weights(self) -> None:
if self.disable_lora:
return
if self.merged:
raise ValueError(
"LoRA weights already merged. Please unmerge them first.")
assert self.lora_A is not None and self.lora_B is not None, "LoRA weights not set. Please set them first."
if isinstance(self.base_layer.weight, DTensor):
mesh = self.base_layer.weight.data.device_mesh
placements = self.base_layer.weight.data.placements
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}").full_tensor()
data += (self.slice_lora_b_weights(self.lora_B)
@ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placements).to(current_device)
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}")
data += \
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight.data = data.to(current_device)
self.merged = True
@torch.no_grad()
def unmerge_lora_weights(self) -> None:
if self.disable_lora:
return
if not self.merged:
raise ValueError(
"LoRA weights not merged. Please merge them first before unmerging."
)
self.unmerge_count += 1
# Avoid precision loss
if self.unmerge_count % 3 == 0:
self.base_layer.weight.data = self.cpu_weight.data.to(
self.base_layer.weight)
if isinstance(self.base_layer.weight, DTensor):
mesh = self.base_layer.weight.data.device_mesh
placement = self.base_layer.weight.data.placements
device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}").full_tensor()
data -= self.slice_lora_b_weights(
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placement).to(device)
else:
self.base_layer.weight.data -= \
self.slice_lora_b_weights(self.lora_B) @\
self.slice_lora_a_weights(self.lora_A)
self.merged = False
class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
"""
Vocab parallel embedding layer with support for LoRA (Low-Rank Adaptation).
Note: The current version does not yet implement the LoRA functionality.
This class behaves exactly the same as the base VocabParallelEmbedding.
Future versions will integrate LoRA functionality to support efficient parameter fine-tuning.
"""
def __init__(
self,
base_layer: VocabParallelEmbedding,
) -> None:
super().__init__(base_layer)
def forward(self, input_: torch.Tensor) -> torch.Tensor:
raise NotImplementedError(
"We don't support VocabParallelEmbeddingWithLoRA yet.")
class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: ColumnParallelLinear,
) -> None:
super().__init__(base_layer)
def forward(self, input_: torch.Tensor) -> torch.Tensor:
# duplicate the logic in ColumnParallelLinear
bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None
output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_, bias)
if self.base_layer.gather_output:
output = tensor_model_parallel_all_gather(output_parallel)
else:
output = output_parallel
output_bias = self.base_layer.bias if self.base_layer.skip_bias_add else None
return output, output_bias
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.base_layer.output_partition_sizes[0]
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
B = B[start_idx:end_idx, :]
return B
class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def __init__(
self,
base_layer: MergedColumnParallelLinear,
) -> None:
super().__init__(base_layer)
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A.to(self.base_layer.weight)
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
# Since the outputs for both gate and up are identical, we use a random one.
shard_size = self.base_layer.output_partition_sizes[0]
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
return B[:, start_idx:end_idx, :]
class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def __init__(
self,
base_layer: QKVParallelLinear,
) -> None:
super().__init__(base_layer)
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A
def slice_lora_b_weights(
self, B: List[torch.Tensor]) -> Tuple[torch.Tensor, torch.Tensor]:
tp_rank = get_tensor_model_parallel_rank()
B_q, B_kv = B
base_layer = self.base_layer
q_proj_shard_size = base_layer.q_proj_shard_size
kv_proj_shard_size = base_layer.kv_proj_shard_size
num_kv_head_replicas = base_layer.num_kv_head_replicas
q_start_idx = q_proj_shard_size * tp_rank
q_end_idx = q_start_idx + q_proj_shard_size
kv_shard_id = tp_rank // num_kv_head_replicas
kv_start_idx = kv_proj_shard_size * kv_shard_id
kv_end_idx = kv_start_idx + kv_proj_shard_size
return B_q[q_start_idx:q_end_idx, :], B_kv[:,
kv_start_idx:kv_end_idx, :]
class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: RowParallelLinear,
) -> None:
super().__init__(base_layer)
def forward(self, input_: torch.Tensor):
# duplicate the logic in RowParallelLinear
if self.base_layer.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.base_layer.tp_size)
input_parallel = splitted_input[tp_rank].contiguous()
output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_parallel)
if self.set_lora:
output_parallel = self.apply_lora(output_parallel, input_parallel)
if self.base_layer.reduce_results and self.base_layer.tp_size > 1:
output_ = tensor_model_parallel_all_reduce(output_parallel)
else:
output_ = output_parallel
if not self.base_layer.skip_bias_add:
output = (output_ + self.base_layer.bias
if self.base_layer.bias is not None else output_)
output_bias = None
else:
output = output_
output_bias = self.base_layer.bias
return output, output_bias
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.base_layer.input_size_per_partition
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
A = A[:, start_idx:end_idx].contiguous()
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
return B
def get_lora_layer(layer: nn.Module) -> Union[BaseLayerWithLoRA, None]:
supported_layer_types: Dict[Type[LinearBase], Type[BaseLayerWithLoRA]] = {
# the order matters
# VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
QKVParallelLinear: QKVParallelLinearWithLoRA,
MergedColumnParallelLinear: MergedColumnParallelLinearWithLoRA,
ColumnParallelLinear: ColumnParallelLinearWithLoRA,
RowParallelLinear: RowParallelLinearWithLoRA,
ReplicatedLinear: BaseLayerWithLoRA,
}
for src_layer_type, lora_layer_type in supported_layer_types.items():
if isinstance(layer, src_layer_type): # pylint: disable=unidiomatic-typecheck
ret = lora_layer_type(layer)
return ret
return None
# source: https://github.com/vllm-project/vllm/blob/93b38bea5dd03e1b140ca997dfaadef86f8f1855/vllm/lora/utils.py#L9
def replace_submodule(model: nn.Module, module_name: str,
new_module: nn.Module) -> nn.Module:
"""Replace a submodule in a model with a new module."""
parent = model.get_submodule(".".join(module_name.split(".")[:-1]))
target_name = module_name.split(".")[-1]
setattr(parent, target_name, new_module)
return new_module
+2 -4
View File
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
import torch
import torch.nn as nn
@@ -18,10 +16,10 @@ class MLP(nn.Module):
self,
input_dim: int,
mlp_hidden_dim: int,
output_dim: Optional[int] = None,
output_dim: int | None = None,
bias: bool = True,
act_type: str = "gelu_pytorch_tanh",
dtype: Optional[torch.dtype] = None,
dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
@@ -3,7 +3,7 @@
import inspect
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Optional
from typing import TYPE_CHECKING, Any
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) -> Optional[QuantizationMethods]:
def override_quantization_method(cls, hf_quant_cfg,
user_quant) -> QuantizationMethods | None:
"""
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) -> Optional[QuantizeMethodBase]:
prefix: str) -> QuantizeMethodBase | None:
"""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) -> Optional[str]:
return None
def get_cache_scale(self, name: str) -> str | None:
return None
+20 -20
View File
@@ -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, Dict, List, Optional, Tuple, Union
from typing import Any
import torch
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: Union[int, float],
base: 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: Union[int, float]) -> torch.Tensor:
def _compute_inv_freq(self, base: 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: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
offsets: torch.Tensor | None = 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: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
def _to_tuple(x: 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: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
*args: Union[int, Tuple[int, ...]],
def get_meshgrid_nd(start: int | tuple[int, ...],
*args: 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: Union[int, Tuple[int, ...]],
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[torch.FloatTensor, int],
pos: 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: Union[float, List[float]] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0,
theta_rescale_factor: float | list[float] = 1.0,
interpolation_factor: 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: Union[int, float],
base: int | float,
is_neox_style: bool = True,
rope_scaling: Optional[Dict[str, Any]] = None,
dtype: Optional[torch.dtype] = None,
rope_scaling: dict[str, Any] | None = None,
dtype: torch.dtype | None = None,
partial_rotary_factor: float = 1.0,
) -> RotaryEmbedding:
if dtype is None:
+1 -2
View File
@@ -1,7 +1,6 @@
# 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
@@ -10,7 +9,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),
+2 -3
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Optional
import torch
import torch.nn as nn
@@ -36,7 +35,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:
@@ -133,7 +132,7 @@ class ModulateProjection(nn.Module):
hidden_size: int,
factor: int = 2,
act_layer: str = "silu",
dtype: Optional[torch.dtype] = None,
dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
+11 -11
View File
@@ -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: Optional[torch.Tensor] = None) -> torch.Tensor:
bias: torch.Tensor | None = 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: Optional[torch.dtype] = None,
org_num_embeddings: Optional[int] = None,
params_dtype: torch.dtype | None = None,
org_num_embeddings: int | None = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = 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) -> Optional[List[int]]:
def get_sharded_to_full_mapping(self) -> list[int] | None:
"""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,
+2 -3
View File
@@ -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, Optional, cast
from typing import Any, cast
import fastvideo.v1.envs as envs
@@ -278,8 +278,7 @@ 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: Optional[str] = None):
def enable_trace_function_call(log_file_path: str, root_dir: str | None = None):
"""
Enable tracing of every function call in code under `root_dir`.
This is useful for debugging hangs or crashes.
+8 -11
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Any, List, Optional, Tuple, Union
from typing import Any
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,11 +33,9 @@ class BaseDiT(nn.Module, ABC):
f"Subclasses of BaseDiT must define '{attr}' class variable"
)
def __init__(self, config: DiTConfig, hf_config: dict[str, Any],
**kwargs) -> None:
def __init__(self, config: DiTConfig, **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"
@@ -46,10 +44,10 @@ class BaseDiT(nn.Module, ABC):
@abstractmethod
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
**kwargs) -> torch.Tensor:
pass
@@ -65,7 +63,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
@@ -78,13 +76,12 @@ class CachableDiT(BaseDiT):
# These are required class attributes that should be overridden by concrete implementations
_fsdp_shard_conditions = []
_param_names_mapping = {}
_lora_param_names_mapping: dict = {}
# Ensure these instance attributes are properly defined in subclasses
hidden_size: int
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:
+11 -14
View File
@@ -1,7 +1,5 @@
# 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
@@ -96,8 +94,8 @@ class MMDoubleStreamBlock(nn.Module):
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[_Backend, ...] | None = None,
prefix: str = "",
):
super().__init__()
@@ -202,7 +200,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)
(
@@ -303,8 +301,8 @@ class MMSingleStreamBlock(nn.Module):
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[_Backend, ...] | None = None,
prefix: str = "",
):
super().__init__()
@@ -366,7 +364,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)
@@ -441,10 +439,9 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
_supported_attention_backends = HunyuanVideoConfig(
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
_lora_param_names_mapping = HunyuanVideoConfig()._lora_param_names_mapping
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
def __init__(self, config: HunyuanVideoConfig):
super().__init__(config=config)
self.patch_size = [
config.patch_size_t, config.patch_size, config.patch_size
@@ -543,10 +540,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# TODO: change output to a dict
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
**kwargs):
"""
+19 -20
View File
@@ -10,7 +10,6 @@
# 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
@@ -55,7 +54,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:
@@ -143,7 +142,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,
@@ -190,8 +189,10 @@ class SelfAttention(nn.Module):
outs = []
idx = 0
for (chunk_size, cos_i, sin_i) in zip(self.rope_split, cos_splits,
sin_splits):
for (chunk_size, cos_i, sin_i) in zip(self.rope_split,
cos_splits,
sin_splits,
strict=False):
# slice the corresponding channels
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
idx += chunk_size
@@ -331,8 +332,8 @@ class AdaLayerNormSingle(nn.Module):
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale)
out, _ = self.linear(self.silu(embedded_timestep))
@@ -377,7 +378,7 @@ class StepVideoTransformerBlock(nn.Module):
dim: int,
attention_head_dim: int,
norm_eps: float = 1e-5,
ff_inner_dim: Optional[int] = None,
ff_inner_dim: int | None = None,
ff_bias: bool = False,
attention_type: str = 'torch'):
super().__init__()
@@ -417,7 +418,7 @@ class StepVideoTransformerBlock(nn.Module):
kv: torch.Tensor,
t_expand: torch.LongTensor,
attn_mask=None,
rope_positions: Optional[list] = None,
rope_positions: list | None = None,
cos_sin=None,
mask_strategy=None) -> torch.Tensor:
@@ -459,13 +460,11 @@ class StepVideoModel(BaseDiT):
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
_param_names_mapping = StepVideoConfig()._param_names_mapping
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
def __init__(self, config: StepVideoConfig) -> None:
super().__init__(config=config)
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
@@ -540,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)],
@@ -595,12 +594,12 @@ class StepVideoModel(BaseDiT):
def forward(
self,
hidden_states: torch.Tensor,
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,
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,
return_dict: bool = True,
mask_strategy=None,
guidance=None,
+13 -17
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
@@ -53,7 +52,7 @@ class WanTimeTextImageEmbedding(nn.Module):
dim: int,
time_freq_dim: int,
text_embed_dim: int,
image_embed_dim: Optional[int] = None,
image_embed_dim: int | None = None,
):
super().__init__()
@@ -76,7 +75,7 @@ class WanTimeTextImageEmbedding(nn.Module):
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
encoder_hidden_states_image: torch.Tensor | None = None,
):
temb = self.time_embedder(timestep)
timestep_proj = self.time_modulation(temb)
@@ -173,7 +172,7 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
supported_attention_backends: tuple[_Backend, ...] | None = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -222,9 +221,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: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[_Backend, ...]
| None = None,
prefix: str = ""):
super().__init__()
@@ -292,13 +291,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)
@@ -319,7 +318,6 @@ class WanTransformerBlock(nn.Module):
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
@@ -360,11 +358,9 @@ class WanTransformer3DModel(CachableDiT):
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
def __init__(self, config: WanVideoConfig) -> None:
super().__init__(config=config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
@@ -420,10 +416,10 @@ class WanTransformer3DModel(CachableDiT):
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
+9 -10
View File
@@ -1,5 +1,4 @@
from abc import ABC, abstractmethod
from typing import Optional, Tuple
import torch
from torch import nn
@@ -11,7 +10,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:
@@ -24,21 +23,21 @@ class TextEncoder(nn.Module, ABC):
@abstractmethod
def forward(self,
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,
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,
**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:
@@ -55,5 +54,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
+38 -38
View File
@@ -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 typing import Iterable, Optional, Set, Tuple, Union
from collections.abc import Iterable
import torch
import torch.nn as nn
@@ -91,9 +91,9 @@ class CLIPTextEmbeddings(nn.Module):
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
input_ids: torch.LongTensor | None = None,
position_ids: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = 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: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -209,8 +209,8 @@ class CLIPMLP(nn.Module):
def __init__(
self,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -239,8 +239,8 @@ class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: Union[CLIPTextConfig, CLIPVisionConfig],
quant_config: Optional[QuantizationConfig] = None,
config: CLIPTextConfig | CLIPVisionConfig,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -284,9 +284,9 @@ class CLIPEncoder(nn.Module):
def __init__(
self,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
config: CLIPVisionConfig | CLIPTextConfig,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = 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
) -> Union[torch.Tensor, list[torch.Tensor]]:
self, inputs_embeds: torch.Tensor, return_all_hidden_states: bool
) -> 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: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
prefix: str = ""):
super().__init__()
self.config = config
@@ -348,11 +348,11 @@ class CLIPTextTransformer(nn.Module):
def forward(
self,
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,
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,
) -> BaseEncoderOutput:
r"""
Returns:
@@ -440,11 +440,11 @@ class CLIPTextModel(TextEncoder):
def forward(
self,
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,
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,
**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: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
quant_config: QuantizationConfig | None = None,
num_hidden_layers_override: int | None = None,
require_post_norm: bool | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -540,7 +540,7 @@ class CLIPVisionTransformer(nn.Module):
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None,
feature_sample_layers: list[int] | None = 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: Optional[list[int]] = None,
feature_sample_layers: list[int] | None = 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:
+18 -17
View File
@@ -23,7 +23,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Inference-only LLaMA model compatible with HuggingFace weights."""
from typing import Any, Dict, Iterable, Optional, Set, Tuple
from collections.abc import Iterable
from typing import Any
import torch
from torch import nn
@@ -52,7 +53,7 @@ class LlamaMLP(nn.Module):
hidden_size: int,
intermediate_size: int,
hidden_act: str,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
bias: bool = False,
prefix: str = "",
) -> None:
@@ -92,9 +93,9 @@ class LlamaAttention(nn.Module):
num_heads: int,
num_kv_heads: int,
rope_theta: float = 10000,
rope_scaling: Optional[Dict[str, Any]] = None,
rope_scaling: dict[str, Any] | None = None,
max_position_embeddings: int = 8192,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
bias: bool = False,
bias_o_proj: bool = False,
prefix: str = "") -> None:
@@ -201,7 +202,7 @@ class LlamaDecoderLayer(nn.Module):
def __init__(
self,
config: LlamaConfig,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -254,8 +255,8 @@ class LlamaDecoderLayer(nn.Module):
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
residual: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
# Self Attention
if residual is None:
residual = hidden_states
@@ -318,11 +319,11 @@ class LlamaModel(TextEncoder):
def forward(
self,
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,
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,
**kwargs,
) -> BaseEncoderOutput:
output_hidden_states = (output_hidden_states
@@ -339,7 +340,7 @@ class LlamaModel(TextEncoder):
0, hidden_states.shape[1],
device=hidden_states.device).unsqueeze(0)
all_hidden_states: Optional[Tuple[Any, ...]] = (
all_hidden_states: tuple[Any, ...] | None = (
) if output_hidden_states else None
for layer in self.layers:
if all_hidden_states is not None:
@@ -367,8 +368,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"),
@@ -378,7 +379,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
@@ -400,7 +401,7 @@ class LlamaModel(TextEncoder):
# continue
if "scale" in name:
# Remapping the name of FP8 kv-scale.
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
kv_scale_name: str | None = maybe_remap_kv_scale_name(
name, params_dict)
if kv_scale_name is None:
continue
+8 -9
View File
@@ -13,7 +13,6 @@
# ==============================================================================
import os
from functools import wraps
from typing import List, Optional
import torch
import torch.nn as nn
@@ -179,10 +178,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)
@@ -347,9 +346,9 @@ class MultiQueryAttention(nn.Module):
def forward(
self,
x: torch.Tensor,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
):
seqlen, bsz, dim = x.shape
xqkv = self.wqkv(x)
@@ -471,9 +470,9 @@ class TransformerBlock(nn.Module):
def forward(
self,
x: torch.Tensor,
mask: Optional[torch.Tensor],
cu_seqlens: Optional[torch.Tensor],
max_seq_len: Optional[torch.Tensor],
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
):
residual = self.attention.forward(self.attention_norm(x), mask,
cu_seqlens, max_seq_len)
+29 -29
View File
@@ -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: Optional[QuantizationConfig] = None):
quant_config: QuantizationConfig | None = 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: Optional[QuantizationConfig] = None):
quant_config: QuantizationConfig | None = 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: Optional[QuantizationConfig] = None):
quant_config: QuantizationConfig | None = 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: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = 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: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = 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: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -361,7 +361,7 @@ class T5LayerSelfAttention(nn.Module):
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
attn_metadata: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = 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: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = 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: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = 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: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = 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: Optional[AttentionMetadata] = None,
attn_metadata: AttentionMetadata | None = 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: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
is_umt5: bool = False):
super().__init__()
@@ -524,11 +524,11 @@ class T5EncoderModel(TextEncoder):
def forward(
self,
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,
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,
**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: 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,
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,
**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:
+4 -4
View File
@@ -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, Optional, TypeVar, Union
from typing import Generic, TypeVar
import torch
from transformers import PretrainedConfig
@@ -48,9 +48,9 @@ class VisionEncoderInfo(ABC, Generic[_C]):
def resolve_visual_encoder_outputs(
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
feature_sample_layers: Optional[list[int]],
post_layer_norm: Optional[torch.nn.LayerNorm],
encoder_outputs: torch.Tensor | list[torch.Tensor],
feature_sample_layers: list[int] | None,
post_layer_norm: torch.nn.LayerNorm | None,
max_possible_layers: int,
) -> torch.Tensor:
"""Given the outputs a visual encoder module that may correspond to the

Some files were not shown because too many files have changed in this diff Show More