Compare commits

..
Author SHA1 Message Date
Will Lin 6b5b2dc6e7 mp training 2025-05-27 11:38:47 -07:00
Will Lin 91b7cc1be8 move utils into training_utils 2025-05-25 18:23:47 -07:00
Zihang-He d6365373b4 added gradient clipping 2025-05-25 23:54:10 +00:00
Will Lin 50fb94b902 gradient checking 2025-05-25 15:20:41 -07:00
Will Lin 2caa0d4d0b cleanup 2025-05-24 17:07:17 -07:00
Will Lin 24db823998 add validation 2025-05-24 16:27:05 -07:00
Will Lin 8c4704edf5 update train dependecies 2025-05-23 12:48:06 -07:00
JerryZhou54 dfba7ec833 Small fix 2025-05-23 19:03:58 +00:00
JerryZhou54 7a2e171f1b Small fix 2025-05-23 19:01:59 +00:00
JerryZhou54 a8aac6090a Integrate the new parquet dataloader into training pipeline 2025-05-23 18:27:52 +00:00
191d1be3b4 Will/training (#425)
Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-05-23 02:06:54 -04:00
JerryZhou54 338ea1e5f2 Add script to upload preprocessed dataset to HF 2025-05-23 06:05:40 +00:00
JerryZhou54 42a2f272d5 Preprocessing Stage: 1. Save to parquets periodically 2. Allow resuming from the middle 3. Doesn't support multi-gpu for now. Data Loader Stage: 1. Allow multi-gpu dataloader 2. Doesn't drop files even if number of parquet files is not divisible by num_gpus 3. Able to resume training 2025-05-23 00:57:39 +00:00
JerryZhou54 982bfcfdc8 Fix small issues when loading mask 2025-05-21 22:46:32 +00:00
JerryZhou54 2ac06379a7 Finish data preprocessing and loading 2025-05-21 22:34:03 +00:00
JerryZhou54 baaa1673f7 Add data preprocessing script for WAN 2025-05-19 15:09:26 +00:00
133 changed files with 48449 additions and 50304 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/env_utils.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 -4
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
@@ -143,8 +141,8 @@ jobs:
fail-fast: false
matrix:
python-version: [
# {version: "3.10", tag: "latest"},
# {version: "3.11", tag: "py3.11-latest"},
{version: "3.10", tag: "latest"},
{version: "3.11", tag: "py3.11-latest"},
{version: "3.12", tag: "py3.12-latest"}
]
uses: ./.github/workflows/runpod-test.yml
+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
@@ -70,8 +70,6 @@ DEFAULT_CONDA_PATTERNS = {
"optree",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
@@ -87,8 +85,6 @@ DEFAULT_PIP_PATTERNS = {
"onnx",
"nccl",
"transformers",
"accelerate",
"peft",
"zmq",
"nvidia",
"pynvml",
+2 -1
View File
@@ -2,6 +2,7 @@ import torch
from flex_sta_ref import get_sliding_tile_attention_mask
from st_attn import sliding_tile_attention
from torch.nn.attention.flex_attention import flex_attention
# from flash_attn_interface import flash_attn_func
from tqdm import tqdm
flex_attention = torch.compile(flex_attention, dynamic=False)
@@ -22,7 +23,7 @@ def h100_fwd_kernel_test(Q, K, V, kernel_size):
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
+46
View File
@@ -0,0 +1,46 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=./data
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# --gradient_checkpointing\
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/v1/pipelines/training_pipeline.py\
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--train_batch_size=1\
--num_latent_t 1 \
--sp_size 4 \
--tp_size 4 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16"
@@ -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)
+1 -2
View File
@@ -1,6 +1,5 @@
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
from fastvideo.version import __version__
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
+11 -9
View File
@@ -11,13 +11,16 @@ from fastvideo.v1.distributed import init_distributed_environment, initialize_mo
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
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
def main(args):
args.model_path = maybe_download_model(args.model_path)
# Assume using torchrun
local_rank = int(os.getenv("RANK", 0))
rank = int(os.environ.get("RANK", 0))
@@ -28,7 +31,7 @@ def main(args):
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
@@ -36,15 +39,15 @@ def main(args):
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=args.model_path,
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
num_gpus=world_size,
device_str="cuda",
**pipeline_config_args,
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -92,7 +95,6 @@ if __name__ == "__main__":
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)
@@ -117,4 +119,4 @@ if __name__ == "__main__":
)
args = parser.parse_args()
main(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,
@@ -0,0 +1,199 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from accelerate.logging import get_logger
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
# from fastvideo.utils.load import load_text_encoder, load_vae
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
logger = get_logger(__name__)
class T5dataset(Dataset):
def __init__(
self,
json_path,
vae_debug,
):
self.json_path = json_path
self.vae_debug = vae_debug
with open(self.json_path, "r") as f:
train_dataset = json.load(f)
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
def __getitem__(self, idx):
caption = self.train_dataset[idx]["caption"]
filename = self.train_dataset[idx]["latent_path"].split(".")[0]
length = self.train_dataset[idx]["length"]
if self.vae_debug:
latents = torch.load(
os.path.join(args.output_dir, "latent", self.train_dataset[idx]["latent_path"]),
map_location="cpu",
)
else:
latents = []
return dict(caption=caption, latents=latents, filename=filename, length=length)
def __len__(self):
return len(self.train_dataset)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
# device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# torch.cuda.set_device(local_rank)
# if not dist.is_initialized():
# dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
videoprocessor = VideoProcessor(vae_scale_factor=8)
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
vae_precision = "fp16"
text_encoder_precision = "fp32"
fastvideo_args = FastVideoArgs(model_path=args.model_path,
use_cpu_offload=False,
vae_precision=vae_precision,
text_encoder_precisions=(text_encoder_precision,))
fastvideo_args.device = device
fastvideo_args.device_str = f"cuda:{local_rank}"
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
fastvideo_args.text_encoder_configs = (T5Config(),)
# vae_loader = VAELoader()
# vae = vae_loader.load_vae()
text_encoder_loader = TextEncoderLoader()
tokenizer_loader = TokenizerLoader()
model_path = args.model_path
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
ENCODER_PATH = os.path.join(path, "text_encoder")
TOKENIZER_PATH = os.path.join(path, "tokenizer")
print(ENCODER_PATH)
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
train_dataset = T5dataset(latents_json_path, args.vae_debug)
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
# vae.enable_tiling()
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
# with torch.autocast("cuda", dtype=torch.float32):
print(data["caption"])
text_inputs = tokenizer(data["caption"], **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
post_process_func = t5_postprocess_text
prompt_embeds = post_process_func(outputs)
prompt_attention_mask = attention_mask
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.postprocess_video(video)
for idx, video_name in enumerate(data["filename"]):
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
video_name + ".pt")
# save latent
torch.save(prompt_embeds[idx], prompt_embed_path)
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
print(f"sample {video_name} saved")
if args.vae_debug:
export_to_video(video[idx], video_path, fps=16)
item = {}
item["length"] = int(data["length"][idx])
item["latent_path"] = video_name + ".pt"
item["prompt_embed_path"] = video_name + ".pt"
item["prompt_attention_mask"] = video_name + ".pt"
item["caption"] = data["caption"][idx]
json_data.append(item)
dist.barrier()
local_data = json_data
gathered_data = [None] * world_size
dist.all_gather_object(gathered_data, local_data)
if local_rank == 0:
# os.remove(latents_json_path)
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
# parser.add_argument("--model_type", type=str, default="mochi")
# text encoder & vae & diffusion model
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=1,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument("--vae_debug", action="store_true")
args = parser.parse_args()
main(args)
@@ -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)
@@ -0,0 +1,151 @@
import argparse
import json
import os
import torch
# import torch.distributed as dist
# from accelerate.logging import get_logger
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.dataset import getdataset
# from fastvideo.utils.load import load_vae
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
logger = init_logger(__name__)
model_path = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
VAE_PATH = os.path.join(path, "vae")
print(VAE_PATH)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
vae_precision = "fp16"
fastvideo_args = FastVideoArgs(model_path=VAE_PATH,
use_cpu_offload=False,
vae_precision=vae_precision)
fastvideo_args.device = device
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
# encoder_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# torch.cuda.set_device(local_rank)
# if not dist.is_initialized():
# dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
vae_loader = VAELoader()
vae = vae_loader.load(VAE_PATH, "vae", fastvideo_args)
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
# vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=torch.float16):
latents = vae.encode(data["pixel_values"].to(device)).sample()
for idx, video_path in enumerate(data["path"]):
video_name = os.path.basename(video_path).split(".")[0]
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
torch.save(latents[idx].to(torch.bfloat16), latent_path)
item = {}
item["length"] = latents[idx].shape[1]
item["latent_path"] = video_name + ".pt"
item["caption"] = data["text"][idx]
json_data.append(item)
print(f"{video_name} processed")
world_group.barrier()
local_data = json_data
gathered_data = [None] * world_size
for i in range(world_size):
if local_rank == i:
world_group.broadcast_object(local_data, src=i)
else:
gathered_data[i] = world_group.broadcast_object(None, src=i)
gathered_data[local_rank] = json_data
print(gathered_data)
if local_rank == 0:
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
# parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
@@ -0,0 +1,115 @@
import argparse
import os
import torch
# import torch.distributed as dist
from accelerate.logging import get_logger
# from fastvideo.utils.load import load_text_encoder
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
logger = get_logger(__name__)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
vae_precision = "fp16"
text_encoder_precision = "fp32"
fastvideo_args = FastVideoArgs(model_path=args.model_path,
use_cpu_offload=False,
vae_precision=vae_precision,
text_encoder_precisions=(text_encoder_precision,))
fastvideo_args.device = device
fastvideo_args.device_str = f"cuda:{local_rank}"
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
fastvideo_args.text_encoder_configs = (T5Config(),)
# vae_loader = VAELoader()
# vae = vae_loader.load_vae()
text_encoder_loader = TextEncoderLoader()
tokenizer_loader = TokenizerLoader()
model_path = args.model_path
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
ENCODER_PATH = os.path.join(path, "text_encoder")
TOKENIZER_PATH = os.path.join(path, "tokenizer")
print(ENCODER_PATH)
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
# autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
# output_dir/validation/prompt_attention_mask
# output_dir/validation/prompt_embed
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
os.makedirs(
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
exist_ok=True,
)
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True)
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
for prompt in prompts:
with torch.inference_mode():
# with torch.autocast("cuda", dtype=autocast_type):
text_inputs = tokenizer(prompt, **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
post_process_func = t5_postprocess_text
prompt_embeds = post_process_func(outputs)
prompt_attention_mask = attention_mask
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
prompt_attention_mask_path = os.path.join(
args.output_dir,
"validation",
"prompt_attention_mask",
f"{file_name}.pt",
)
torch.save(prompt_embeds[0], prompt_embed_path)
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
print(f"sample {file_name} saved")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
-97
View File
@@ -1,97 +0,0 @@
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.dataset.t2v_datasets import T2V_dataset
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
def getdataset(args):
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
]
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
]
transform = transforms.Compose([
# Normalize255(),
*resize,
])
transform_topcrop = transforms.Compose([
Normalize255(),
*resize_topcrop,
norm_fun,
])
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(
args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
)
raise NotImplementedError(args.dataset)
if __name__ == "__main__":
import random
from accelerate import Accelerator
from tqdm import tqdm
from fastvideo.dataset.t2v_datasets import dataset_prog
args = type(
"args",
(),
{
"ae": "CausalVAEModel_4x8x8",
"dataset": "t2v",
"attention_mode": "xformers",
"use_rope": True,
"text_max_length": 300,
"max_height": 320,
"max_width": 240,
"num_frames": 1,
"use_image_num": 0,
"interpolation_scale_t": 1,
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
"use_img_from_vid": False,
"speed_factor": 1.0,
"cfg": 0.1,
"text_encoder_name": "google/mt5-xxl",
"dataloader_num_workers": 10,
},
)
accelerator = Accelerator()
dataset = getdataset(args)
num = len(dataset_prog.img_cap_list)
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
print(e)
# import ipdb;ipdb.set_trace()
print(image_data)
zero += 1
continue
assert caps[0] is not None and len(caps[0]) > 0
print(num, zero)
import ipdb
ipdb.set_trace()
print("end")
-118
View File
@@ -1,118 +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,
):
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
self.json_path = json_path
self.cfg_rate = cfg_rate
self.datase_dir_path = os.path.dirname(json_path)
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
with open(self.json_path, "r") as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
self.num_latent_t = num_latent_t
# just zero embeddings [256, 4096]
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [data_item["length"] if "length" in data_item else 1 for data_item in self.data_anno]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
# load
latent = torch.load(
os.path.join(self.latent_dir, latent_file),
map_location="cpu",
weights_only=True,
)
latent = latent.squeeze(0)[:, -self.num_latent_t:]
if random.random() < self.cfg_rate:
prompt_embed = self.uncond_prompt_embed
prompt_attention_mask = self.uncond_prompt_mask
else:
prompt_embed = torch.load(
os.path.join(self.prompt_embed_dir, prompt_embed_file),
map_location="cpu",
weights_only=True,
)
prompt_attention_mask = torch.load(
os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file),
map_location="cpu",
weights_only=True,
)
return latent, prompt_embed, prompt_attention_mask
def __len__(self):
return len(self.data_anno)
def latent_collate_function(batch):
# return latent, prompt, latent_attn_mask, text_attn_mask
# latent_attn_mask: # b t h w
# text_attn_mask: b 1 l
# needs to check if the latent/prompt' size and apply padding & attn mask
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
# calculate max shape
max_t = max([latent.shape[1] for latent in latents])
max_h = max([latent.shape[2] for latent in latents])
max_w = max([latent.shape[3] for latent in latents])
# padding
latents = [
torch.nn.functional.pad(
latent,
(
0,
max_t - latent.shape[1],
0,
max_h - latent.shape[2],
0,
max_w - latent.shape[3],
),
) for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latents):
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
prompt_embeds = torch.stack(prompt_embeds, dim=0)
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
latents = torch.stack(latents, dim=0)
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
prompt_embed.shape,
latent_attn_mask.shape,
prompt_attention_mask.shape,
)
import pdb
pdb.set_trace()
-324
View File
@@ -1,324 +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 = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
class DataSetProg(metaclass=SingletonMeta):
def __init__(self):
self.cap_list = []
self.elements = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements = dict()
self.n_used_elements = dict()
def set_cap_list(self, num_workers, cap_list, n_elements):
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
self.elements = list(range(n_elements))
random.shuffle(self.elements)
print(f"n_elements: {len(self.elements)}", flush=True)
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info):
if work_info is None:
worker_id = 0
else:
worker_id = work_info.id
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
dataset_prog = DataSetProg()
def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
class T2V_dataset(Dataset):
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
self.use_image_num = args.use_image_num
self.transform = transform
self.transform_topcrop = transform_topcrop
self.temporal_sample = temporal_sample
self.tokenizer = tokenizer
self.text_max_length = args.text_max_length
self.cfg = args.cfg
self.speed_factor = args.speed_factor
self.max_height = args.max_height
self.max_width = args.max_width
self.drop_short_ratio = args.drop_short_ratio
assert self.speed_factor >= 1
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if "mt5" not in args.text_encoder_name:
self.support_Chinese = False
cap_list = self.get_cap_list()
assert len(cap_list) > 0
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
def set_checkpoint(self, n_used_elements):
for i in range(len(dataset_prog.n_used_elements)):
dataset_prog.n_used_elements[i] = n_used_elements
def __len__(self):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
def get_data(self, idx):
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx):
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
assert video.dtype == torch.uint8
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
text = dataset_prog.cap_list[idx]["cap"]
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(
pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
)
def get_image(self, idx):
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
# for i in image:
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images" in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
input_ids = text_tokens_and_mask["input_ids"] # 1, l
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
return dict(
pixel_values=image,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=image_data["path"],
)
def define_frame_index(self, cap_list):
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
cnt_too_short = 0
cnt_no_cap = 0
cnt_no_resolution = 0
cnt_resolution_mismatch = 0
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
if fps is None or duration is None:
continue
# ======resolution mismatch=====
resolution = i.get("resolution", None)
if resolution is None:
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"]["width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
if not is_pick:
print("resolution mismatch")
cnt_resolution_mismatch += 1
continue
# import ipdb;ipdb.set_trace()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps *
self.speed_factor): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"], frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(i["sample_frame_index"]) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
new_cap_list.append(i)
i["sample_num_frames"] = 1
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}")
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices):
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data):
cap_lists = []
with open(data, "r") as f:
folder_anno = [i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno, "r") as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
return cap_lists
def get_cap_list(self):
cap_lists = self.read_jsons(self.data)
return cap_lists
-608
View File
@@ -1,608 +0,0 @@
import numbers
import random
import torch
from PIL import Image
def _is_tensor_video_clip(clip):
if not torch.is_tensor(clip):
raise TypeError("clip should be Tensor. Got %s" % type(clip))
if not clip.ndimension() == 4:
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
return True
def center_crop_arr(pil_image, image_size):
"""
Center cropping implementation from ADM.
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
"""
while min(*pil_image.size) >= 2 * image_size:
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
scale = image_size / min(*pil_image.size)
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
arr = np.array(pil_image)
crop_y = (arr.shape[0] - image_size) // 2
crop_x = (arr.shape[1] - image_size) // 2
return Image.fromarray(arr[crop_y:crop_y + image_size, crop_x:crop_x + image_size])
def crop(clip, i, j, h, w):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
"""
if len(clip.size()) != 4:
raise ValueError("clip should be a 4D tensor")
return clip[..., i:i + h, j:j + w]
def resize(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
return torch.nn.functional.interpolate(
clip,
size=target_size,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
clip,
scale_factor=scale_,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
"""
Do spatial cropping and resizing to the video clip
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
i (int): i in (i,j) i.e coordinates of the upper left corner.
j (int): j in (i,j) i.e coordinates of the upper left corner.
h (int): Height of the cropped region.
w (int): Width of the cropped region.
size (tuple(int, int)): height and width of resized clip
Returns:
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
clip = crop(clip, i, j, h, w)
clip = resize(clip, size, interpolation_mode)
return clip
def center_crop(clip, crop_size):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
th, tw = crop_size
if h < th or w < tw:
raise ValueError("height and width must be no smaller than crop_size")
i = int(round((h - th) / 2.0))
j = int(round((w - tw) / 2.0))
return crop(clip, i, j, th, tw)
def center_crop_using_short_edge(clip):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h < w:
th, tw = h, h
i = 0
j = int(round((w - tw) / 2.0))
else:
th, tw = w, w
i = int(round((h - th) / 2.0))
j = 0
return crop(clip, i, j, th, tw)
def center_crop_th_tw(clip, th, tw, top_crop):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
# import ipdb;ipdb.set_trace()
h, w = clip.size(-2), clip.size(-1)
tr = th / tw
if h / w > tr:
new_h = int(w * tr)
new_w = w
else:
new_h = h
new_w = int(h / tr)
i = 0 if top_crop else int(round((h - new_h) / 2.0))
j = int(round((w - new_w) / 2.0))
return crop(clip, i, j, new_h, new_w)
def random_shift_crop(clip):
"""
Slide along the long edge, with the short edge as crop size
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h <= w:
short_edge = h
else:
short_edge = w
th, tw = short_edge, short_edge
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return crop(clip, i, j, th, tw)
def normalize_video(clip):
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
permute the dimensions of clip tensor
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
def normalize(clip, mean, std, inplace=False):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
mean (tuple): pixel RGB mean. Size is (3)
std (tuple): pixel standard deviation. Size is (3)
Returns:
normalized clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
if not inplace:
clip = clip.clone()
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
# print(mean)
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
return clip
def hflip(clip):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
Returns:
flipped clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
return clip.flip(-1)
class RandomCropVideo:
def __init__(self, size):
if isinstance(size, numbers.Number):
self.size = (int(size), int(size))
else:
self.size = size
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: randomly cropped video clip.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = self.size
if h < th or w < tw:
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
if w == tw and h == th:
return 0, 0, h, w
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return i, j, th, tw
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: cropped video clip by stride.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
return 0, 0, th, tw # from top-left
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class LongSideResizeVideo:
"""
First use the long side,
then resize to the specified size
"""
def __init__(
self,
size,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized video clip.
size is (T, C, 512, *) or (T, C, *, 512)
"""
_, _, h, w = clip.shape
if self.skip_low_resolution and max(h, w) <= self.size:
return clip
if h > w:
w = int(w * self.size / h)
h = self.size
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
return resize_clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class CenterCropResizeVideo:
"""
First use the short side for cropping length,
center crop video, then resize to the specified size
"""
def __init__(
self,
size,
top_crop=False,
interpolation_mode="bilinear",
):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
# clip_center_crop = center_crop_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
# import ipdb;ipdb.set_trace()
clip_center_crop_resize = resize(
clip_center_crop,
target_size=self.size,
interpolation_mode=self.interpolation_mode,
)
return clip_center_crop_resize
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class UCFCenterCropVideo:
"""
First scale to the specified size in equal proportion to the short edge,
then center cropping
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class KineticsRandomCropResizeVideo:
"""
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
return clip_resize
class CenterCropVideo:
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop(clip, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize:
"""
Normalize the video clip by mean subtraction and division by standard deviation
Args:
mean (3-tuple): pixel RGB mean
std (3-tuple): pixel RGB standard deviation
inplace (boolean): whether do in-place normalization
"""
def __init__(self, mean, std, inplace=False):
self.mean = mean
self.std = std
self.inplace = inplace
def __call__(self, clip):
"""
Args:
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
"""
return normalize(clip, self.mean, self.std, self.inplace)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self):
pass
def __call__(self, clip):
"""
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
Return:
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
"""
return normalize_video(clip)
def __repr__(self) -> str:
return self.__class__.__name__
class RandomHorizontalFlipVideo:
"""
Flip the video clip along the horizontal direction with a given probability
Args:
p (float): probability of the clip being flipped. Default value is 0.5
"""
def __init__(self, p=0.5):
self.p = p
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Size is (T, C, H, W)
Return:
clip (torch.tensor): Size is (T, C, H, W)
"""
if random.random() < self.p:
clip = hflip(clip)
return clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(p={self.p})"
# ------------------------------------------------------------
# --------------------- Sampling ---------------------------
# ------------------------------------------------------------
class TemporalRandomCrop(object):
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, size):
self.size = size
def __call__(self, total_frames):
rand_end = max(0, total_frames - self.size - 1)
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
class DynamicSampleDuration(object):
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, t_stride, extra_1):
self.t_stride = t_stride
self.extra_1 = extra_1
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
return 0, truncate_t
if __name__ == "__main__":
import os
import numpy as np
import torchvision.io as io
from torchvision import transforms
from torchvision.utils import save_image
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW")
trans = transforms.Compose([
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
])
target_video_len = 32
frame_interval = 1
total_frames = len(vframes)
print(total_frames)
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
# Sampling video frames
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
print(select_vframes.shape)
print(select_vframes.dtype)
select_vframes_trans = trans(select_vframes)
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
for i in range(target_video_len):
save_image(
select_vframes_trans[i],
os.path.join("./test000", "%04d.png" % i),
normalize=True,
value_range=(-1, 1),
)
+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,
+2 -1
View File
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -38,6 +38,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
linear_range=0.5,
):
if linear_quadratic:
raise NotImplementedError("Linear quadratic schedule is not implemented")
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
+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,
+870
View File
@@ -0,0 +1,870 @@
# !/bin/python3
# isort: skip_file
import argparse
import math
import os
import time
from collections import deque
import torch
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.utils.checkpoint import (save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
get_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.models.loader.component_loader import TransformerLoader, SchedulerLoader
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
SCHEDULER_PATH = os.path.join(MODEL_PATH, "scheduler")
def reshard_fsdp(model):
for m in FSDP.fsdp_modules(model):
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
gradient_accumulation_steps)
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
gradient_accumulation_steps)
absolute_mean = torch.mean(
torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(
torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
def distill_one_step(
transformer,
model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_teacher_disable_cfg,
):
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0, # codespell:ignore
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
}
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
# model_input = normalize_dit_input(model_type, latents)
model_input = latents
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
timesteps = (sigmas *
noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
noisy_model_input = noisy_model_input.to(torch.bfloat16)
forward_batch = ForwardBatch(data_type="video", enable_teacache=False)
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if hunyuan_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
with torch.autograd.graph.save_on_cpu(pin_memory=True):
model_pred = transformer(**teacher_kwargs)
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
).float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(
bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
).float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output,
index).to(torch.bfloat16)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
with torch.autograd.graph.save_on_cpu(pin_memory=True):
target_pred = transformer(
x_prev,
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
huber_c = 0.001
# loss = loss.mean()
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float()**2)) *
pred_decay_weight / gradient_accumulation_steps)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (torch.mean(model_pred.float()**2) *
pred_decay_weight /
gradient_accumulation_steps)
loss += pred_decay_loss
else:
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(model_pred.detach().float(), model_pred_norm,
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(ema_transformer.parameters(),
transformer.parameters()):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(),
1 - ema_decay))
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
grad_norm = torch.nn.utils.clip_grad_norm_(transformer.parameters(),
max_norm=max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item(), model_pred_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
torch.cuda.set_device(rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=args.sp_size,
sequence_model_parallel_size=args.sp_size)
fastvideo_args = FastVideoArgs(
model_path=MODEL_PATH,
num_gpus=world_size,
use_cpu_offload=False,
precision=args.master_weight_type,
dit_config=WanVideoConfig(),
device_str="cuda",
)
fastvideo_args.check_fastvideo_args()
device_str = f"cuda:{rank}"
device = torch.device(device_str)
fastvideo_args.device = device
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
logger.info("--> loading model from %s", TRANSFORMER_PATH)
fastvideo_args.device = device
transformer_loader = TransformerLoader()
transformer = transformer_loader.load(TRANSFORMER_PATH, "", fastvideo_args)
transformer = transformer.train()
transformer.requires_grad_(True)
teacher_loader = TransformerLoader()
teacher_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
fastvideo_args)
if args.use_ema:
ema_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
fastvideo_args)
else:
ema_transformer = None
logger.info(
" Total training parameters = %s M",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e6)
logger.info("--> model loaded")
teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
# scheduler
noise_scheduler_loader = SchedulerLoader()
noise_scheduler = noise_scheduler_loader.load(SCHEDULER_PATH, "",
fastvideo_args)
solver = EulerSolver(
noise_scheduler.sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(train_dataset))
logger.info(" Dataloader size = %s", len(train_dataloader))
logger.info(" Num Epochs = %s", args.num_train_epochs)
logger.info(" Resume training from step %s", init_steps)
logger.info(" Instantaneous batch size per device = %s",
args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s", args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e9)
# print dtype
logger.info(" Master weight dtype: %s",
transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
loss, grad_norm, pred_norm = distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.95)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=str, default="64")
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="Threshold for linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--use_ema",
action="store_true",
help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+2 -2
View File
@@ -237,7 +237,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
type=str,
default="540p",
choices=["540p", "720p"],
help="The resolution of the model.",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--load-key",
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
"--ring-degree",
type=int,
default=1,
help="Ring degree.",
help="Ulysses degree.",
)
return parser
+1 -1
View File
@@ -17,7 +17,7 @@ from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.utils.parallel_states import nccl_info
class Inference:
class Inference(object):
def __init__(
self,
+1 -1
View File
@@ -41,7 +41,7 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
elif mode == "Master":
prompt = master_mode_prompt.format(input=ori_prompt)
else:
raise Exception("Only supports Normal and Master mode, but got {}".format(mode))
raise Exception("Only supports Normal and Normal", mode)
return prompt
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
mochi_scaling_factor = 1.0
def normalize_dit_input(model_type, latents):
def normalize_dit_input(model_type, latents, args=None):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
@@ -41,5 +41,16 @@ def normalize_dit_input(model_type, latents):
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
vae_config = WanVAEConfig()
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
class STEP1TextEncoder(torch.nn.Module):
def __init__(self, model_dir, max_length=320):
super()
super(STEP1TextEncoder, self).__init__()
self.max_length = max_length
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
text_encoder = Step1Model.from_pretrained(model_dir)
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
@torch.no_grad
@torch.autocast(device_type='cuda', dtype=torch.bfloat16)
def forward(self, prompts, with_mask=True, max_length=None):
self.device = next(self.text_encoder.parameters()).device
if type(prompts) is str:
prompts = [prompts]
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
if type(prompts) is str:
prompts = [prompts]
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
-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,
+45 -7
View File
@@ -11,6 +11,7 @@ from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_
from torch.distributed.fsdp import FullOptimStateDictConfig, FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
import dataclasses
from fastvideo.utils.logging_ import main_print
@@ -32,7 +33,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)
@@ -44,13 +45,50 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
optimizer_path = os.path.join(save_dir, "optimizer.pt")
torch.save(optim_state, optimizer_path)
else:
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
weight_path = os.path.join(save_dstate_dictir, "discriminator_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
torch.save(optim_state, optimizer_path)
main_print(f"--> checkpoint saved at step {step}")
def save_checkpoint_v1(transformer, rank, output_dir, step):
# from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
# from torch.distributed.fsdp import StateDictType, FullStateDictConfig
# Configure FSDP to save full state dict
FSDP.set_state_dict_type(
transformer,
state_dict_type=StateDictType.FULL_STATE_DICT,
state_dict_config=FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
)
# Now get the state dict
cpu_state = transformer.state_dict()
# Save it (only on rank 0 since we used rank0_only=True)
# if torch.distributed.get_rank() == 0:
# torch.save(state_dict, "model_checkpoint.pt")
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
# weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
print(weight_path)
# save_file(cpu_state, weight_path)
torch.save(cpu_state, weight_path)
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
main_print(f"--> checkpoint saved at step {step}")
def save_checkpoint(transformer, rank, output_dir, step):
main_print(f"--> saving checkpoint at step {step}")
with FSDP.state_dict_type(
@@ -60,7 +98,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 +136,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 +177,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 +216,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 +279,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)
+38
View File
@@ -0,0 +1,38 @@
import platform
import accelerate
import peft
import torch
import transformers
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
VERSION = "1.2.0"
if __name__ == "__main__":
info = {
"FastVideo version": VERSION,
"Platform": platform.platform(),
"Python version": platform.python_version(),
"PyTorch version": torch.__version__,
"Transformers version": transformers.__version__,
"Accelerate version": accelerate.__version__,
"PEFT version": peft.__version__,
}
if is_torch_cuda_available():
info["PyTorch version"] += " (GPU)"
info["GPU type"] = torch.cuda.get_device_name()
if is_torch_npu_available():
info["PyTorch version"] += " (NPU)"
info["NPU type"] = torch.npu.get_device_name()
info["CANN version"] = torch.version.cann # codespell:ignore
try:
import bitsandbytes
info["Bitsandbytes version"] = bitsandbytes.__version__
except Exception:
pass
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
-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
@@ -1,6 +1,6 @@
import json
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Type
from typing import List, Optional, Type
import torch
from einops import rearrange
@@ -13,7 +13,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 +20,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[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
@@ -45,14 +42,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")
@@ -85,8 +82,6 @@ class SlidingTileAttentionBackend(AttentionBackend):
@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 +98,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):
@@ -129,12 +120,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 +205,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 +227,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
+2 -3
View File
@@ -13,7 +13,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):
@@ -39,7 +38,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,
@@ -156,7 +155,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,
+1 -3
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, List, Optional, Tuple
from typing import Any, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
@@ -11,7 +11,6 @@ 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,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_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:
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from typing import Optional, Tuple
import torch
@@ -163,8 +163,6 @@ class HunyuanVideoArchConfig(DiTArchConfig):
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__()
@@ -51,7 +51,6 @@ class StepVideoArchConfig(DiTArchConfig):
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: [])
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
+1 -19
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
from typing import Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -51,23 +51,6 @@ 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)
text_len = 512
@@ -85,7 +68,6 @@ class WanVideoArchConfig(DiTArchConfig):
image_dim: Optional[int] = None
added_kv_proj_dim: Optional[int] = None
rope_max_seq_len: int = 1024
exclude_lora_layers: List[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
+1 -1
View File
@@ -63,7 +63,7 @@ class WanVAEArchConfig(VAEArchConfig):
@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
+1 -2
View File
@@ -27,6 +27,7 @@ class PipelineConfig:
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
@@ -54,8 +55,6 @@ class PipelineConfig:
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: str = "STA_inference"
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
@@ -68,6 +68,9 @@ 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(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
@@ -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
+3
View File
@@ -37,6 +37,9 @@ class WanT2V480PConfig(PipelineConfig):
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 3
+66 -5
View File
@@ -1,5 +1,3 @@
import os
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
@@ -9,7 +7,7 @@ from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
def getdataset(args, start_idx=0) -> T2V_dataset:
def getdataset(args, start_idx=0):
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
@@ -27,8 +25,8 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
*resize_topcrop,
norm_fun,
])
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(args,
@@ -39,3 +37,66 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
start_idx=start_idx)
raise NotImplementedError(args.dataset)
if __name__ == "__main__":
import random
from accelerate import Accelerator
from tqdm import tqdm
from fastvideo.v1.dataset.t2v_datasets import dataset_prog
args = type(
"args",
(),
{
"ae": "CausalVAEModel_4x8x8",
"dataset": "t2v",
"attention_mode": "xformers",
"use_rope": True,
"text_max_length": 300,
"max_height": 320,
"max_width": 240,
"num_frames": 1,
"use_image_num": 0,
"interpolation_scale_t": 1,
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data":
"/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
"use_img_from_vid": False,
"speed_factor": 1.0,
"cfg": 0.1,
"text_encoder_name": "google/mt5-xxl",
"dataloader_num_workers": 10,
},
)
accelerator = Accelerator()
dataset = getdataset(args)
num = len(dataset_prog.img_cap_list)
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [
i["cap"] if isinstance(i["cap"], list) else [i["cap"]]
for i in image_data
]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
print(e)
# import ipdb;ipdb.set_trace()
print(image_data)
zero += 1
continue
assert caps[0] is not None and len(caps[0]) > 0
print(num, zero)
import ipdb
ipdb.set_trace()
print("end")
+1 -39
View File
@@ -9,7 +9,7 @@ frameworks that can handle parquet or lance file.
import pyarrow as pa
pyarrow_schema_i2v = pa.schema([
pyarrow_schema = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
@@ -30,10 +30,6 @@ pyarrow_schema_i2v = pa.schema([
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()),
@@ -46,37 +42,3 @@ pyarrow_schema_i2v = pa.schema([
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)
+28 -8
View File
@@ -13,7 +13,7 @@ class LatentDataset(Dataset):
json_path,
num_latent_t,
cfg_rate,
) -> None:
):
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
self.json_path = json_path
self.cfg_rate = cfg_rate
@@ -29,12 +29,13 @@ class LatentDataset(Dataset):
# json.load(f) already keeps the order
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
self.num_latent_t = num_latent_t
# just zero embeddings [256, 4096]
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item.get("length", 1) for data_item in self.data_anno
data_item["length"] if "length" in data_item else 1
for data_item in self.data_anno
]
def __getitem__(self, idx):
@@ -82,7 +83,7 @@ def latent_collate_function(batch):
max_w = max([latent.shape[3] for latent in latents])
# padding
latent_list: list[torch.Tensor] = [
latents = [
torch.nn.functional.pad(
latent,
(
@@ -96,14 +97,33 @@ def latent_collate_function(batch):
) for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latent_list), max_t, max_h, max_w)
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latent_list):
for i, latent in enumerate(latents):
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
prompt_embeds = torch.stack(prompt_embeds, dim=0)
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
latents = torch.stack(latent_list, dim=0)
latents = torch.stack(latents, dim=0)
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
num_latent_t=28)
dataloader = torch.utils.data.DataLoader(dataset,
batch_size=2,
shuffle=False,
collate_fn=latent_collate_function)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
prompt_embed.shape,
latent_attn_mask.shape,
prompt_attention_mask.shape,
)
import pdb
pdb.set_trace()
+209 -261
View File
@@ -1,29 +1,21 @@
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 torch.utils.data import IterableDataset, get_worker_info
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__)
# Path to your dataset
dataset_path = "/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/train/"
class ParquetVideoTextDataset(Dataset):
class ParquetVideoTextDataset(IterableDataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
def __init__(self,
@@ -32,292 +24,245 @@ class ParquetVideoTextDataset(Dataset):
rank: int = 0,
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2,
seed: int = 0,
validation: bool = False):
num_latent_t: int = 2):
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.world_size = world_size
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
# Find all parquet files recursively
print(f"Scanning for parquet files in {self.path}")
self.parquet_files = []
for root, _, files in os.walk(self.path):
for file in files:
if file.endswith('.parquet'):
self.parquet_files.append(os.path.join(root, file))
# Sort files for consistent ordering
self.parquet_files.sort()
self.plan_output_dir = os.path.join(
self.path,
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.json"
# Distribute files among workers
# drop last unenven files
print(f"Total files: {len(self.parquet_files)}")
total_files = len(self.parquet_files)
base_count = total_files // world_size
extra_files = total_files % world_size
if rank < extra_files:
start_idx = rank * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = rank * base_count + extra_files
end_idx = start_idx + base_count
self.parquet_files = self.parquet_files[start_idx:end_idx]
print(f"Files assigned to rank {rank}: {len(self.parquet_files)}")
if len(self.parquet_files) > 0:
print(f"First file: {self.parquet_files[0]}")
print(f"Last file: {self.parquet_files[-1]}")
# Initialize current file index
self.current_file_idx = 0
self.current_reader = None
self.current_batches = None
self.total_samples = 0
def _open_next_file(self):
"""Open the next parquet file for reading."""
num_workers = get_worker_info().num_workers
worker_id = get_worker_info().id
total_files = len(self.parquet_files)
base_count = total_files // num_workers
extra_files = total_files % num_workers
if worker_id < extra_files:
start_idx = worker_id * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = worker_id * base_count + extra_files
end_idx = start_idx + base_count
worker_parquet_files = self.parquet_files[start_idx:end_idx]
if self.current_file_idx >= len(worker_parquet_files):
print(
f"Rank {self.rank}, Worker {worker_id}: No more files to open (current_idx={self.current_file_idx}, total_files={len(worker_parquet_files)})"
)
return False
if self.current_reader is not None:
self.current_reader.close()
file_path = worker_parquet_files[self.current_file_idx]
print(
f"Rank {self.rank}, Worker {worker_id}: Opening file {self.current_file_idx + 1}/{len(worker_parquet_files)}: {file_path}"
)
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)
self.current_reader = pq.ParquetFile(file_path)
self.current_batches = self.current_reader.iter_batches(
batch_size=self.batch_size)
self.current_file_idx += 1
return True
except Exception as e:
logger.error("Failed to load negative prompt: %s", e)
self.cached_neg_prompt = None
print(f"Error opening file {file_path}: {str(e)}")
return False
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"
)
def __iter__(self):
"""Iterate over the dataset in a streaming fashion."""
print(f"Rank {self.rank}: Starting iteration")
# 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()
# First try to open a file
if not self._open_next_file():
print(f"Rank {self.rank}: Failed to open first file")
return
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:
while True:
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)
# Get next batch from current file
batch = next(self.current_batches)
batch_dict = batch.to_pydict()
processed = self._process_batch(batch_dict)
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)
# Update sample count
batch_size = len(processed["latents"])
self.total_samples += batch_size
# 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
# Print progress
if self.total_samples % 1000 == 0:
print(
f"Rank {self.rank}: Processed {self.total_samples} samples"
)
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
# Yield each item in the batch
for lat, emb, mask, info in zip(processed["latents"],
processed["embeddings"],
processed["masks"],
processed["info"]):
if lat.numel() == 0: # Split is validation
yield lat, emb, mask, info
else:
yield lat[:, -self.num_latent_t:], emb, mask, info
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
except StopIteration:
# Current file is exhausted, try next file
print(
f"Rank {self.rank}: Current file exhausted, trying next file"
)
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: No more files to process. Total samples: {self.total_samples}"
)
break
except Exception as e:
print(f"Error processing batch: {str(e)}")
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: Failed to open next file after error"
)
break
def _process_row(self, row) -> Dict[str, Any]:
# Clean up
if self.current_reader is not None:
self.current_reader.close()
def _process_batch(self, batch):
"""Process a PyArrow batch into tensors."""
out = {"lat": [], "emb": [], "msk": [], "info": []}
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"]
for i in range(len(batch["vae_latent_bytes"])):
vae_latent_bytes = batch["vae_latent_bytes"][i]
vae_latent_shape = batch["vae_latent_shape"][i]
text_embedding_bytes = batch["text_embedding_bytes"][i]
text_embedding_shape = batch["text_embedding_shape"][i]
text_attention_mask_bytes = batch["text_attention_mask_bytes"][i]
text_attention_mask_shape = batch["text_attention_mask_shape"][i]
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(
text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
# to string
file_name = str(batch["file_name"][i])
# Collect metadata
info = {
"width": batch["width"][i],
"height": batch["height"][i],
"num_frames": batch["num_frames"][i],
"duration_sec": batch["duration_sec"][i],
"fps": batch["fps"][i],
"file_name": batch["file_name"][i],
"caption": batch["caption"][i],
}
# Collect metadata
info = {
"width": row["width"],
"height": row["height"],
"num_frames": row["num_frames"],
"duration_sec": row["duration_sec"],
"fps": row["fps"],
"file_name": row["file_name"],
"caption": row["caption"],
}
out["lat"].append(torch.from_numpy(lat))
out["emb"].append(torch.from_numpy(emb))
out["msk"].append(torch.from_numpy(msk))
out["info"].append(info)
return {
"latents": torch.from_numpy(lat),
"embeddings": torch.from_numpy(emb),
"masks": torch.from_numpy(msk),
"info": info
"latents": torch.stack(out["lat"]) if out["lat"] else None,
"embeddings": torch.stack(out["emb"]) if out["emb"] else None,
"masks": torch.stack(out["msk"]) if out["msk"] else None,
"info": out["info"]
}
def bind_cpu_cores(local_rank, cpu_per_process=16):
"""根据local_rank绑定固定cpu核。"""
start = local_rank * cpu_per_process
end = start + cpu_per_process
cores = list(range(start, end))
print(f"[Rank {local_rank}] Binding to CPU cores: {cores}")
os.sched_setaffinity(0, cores)
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description='Benchmark Parquet dataset loading speed')
parser.add_argument('--path',
type=str,
default="your/dataset/path",
default=dataset_path,
help='Path to Parquet dataset')
parser.add_argument('--batch_size',
type=int,
@@ -352,6 +297,9 @@ if __name__ == "__main__":
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
)
# Bind CPU cores after distributed initialization
# bind_cpu_cores(local_rank, cpu_per_process=16)
# Create dataset
dataset = ParquetVideoTextDataset(
args.path,
+32 -29
View File
@@ -17,7 +17,7 @@ from fastvideo.utils.logging_ import main_print
class SingletonMeta(type):
_instances: dict[type, 'SingletonMeta'] = {}
_instances = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
@@ -28,15 +28,15 @@ class SingletonMeta(type):
class DataSetProg(metaclass=SingletonMeta):
def __init__(self) -> None:
self.cap_list: list[dict] = []
self.elements: list[int] = []
def __init__(self):
self.cap_list = []
self.elements = []
self.num_workers = 1
self.n_elements = 0
self.worker_elements: dict[int, list[int]] = {}
self.n_used_elements: dict[int, int] = {}
self.worker_elements = dict()
self.n_used_elements = dict()
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
def set_cap_list(self, num_workers, cap_list, n_elements):
self.num_workers = num_workers
self.cap_list = cap_list
self.n_elements = n_elements
@@ -52,8 +52,11 @@ class DataSetProg(metaclass=SingletonMeta):
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
def get_item(self, work_info) -> int:
worker_id = 0 if work_info is None else work_info.id
def get_item(self, work_info):
if work_info is None:
worker_id = 0
else:
worker_id = work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
@@ -65,11 +68,13 @@ class DataSetProg(metaclass=SingletonMeta):
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
def filter_resolution(h,
w,
max_h_div_w_ratio=17 / 16,
min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
class T2V_dataset(Dataset):
@@ -80,7 +85,7 @@ class T2V_dataset(Dataset):
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0) -> None:
start_idx=0):
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
@@ -127,18 +132,17 @@ class T2V_dataset(Dataset):
data = self.get_data(idx)
return data
def get_data(self, idx) -> dict:
def get_data(self, idx):
path = dataset_prog.cap_list[idx]["path"]
if path.endswith(".mp4"):
return self.get_video(idx)
else:
return self.get_image(idx)
def get_video(self, idx) -> dict:
def get_video(self, idx):
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
@@ -179,7 +183,7 @@ class T2V_dataset(Dataset):
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx) -> dict:
def get_image(self, idx):
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
@@ -197,14 +201,14 @@ class T2V_dataset(Dataset):
image = image.float() / 127.5 - 1.0
caps: list[str] = (image_data["cap"] if isinstance(
image_data["cap"], list) else [image_data["cap"]])
caps = (image_data["cap"]
if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
single_text = text[0] if random.random() > self.cfg else ""
text = text[0] if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
single_text,
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
@@ -222,7 +226,7 @@ class T2V_dataset(Dataset):
path=image_data["path"],
)
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
def define_frame_index(self, cap_list):
new_cap_list = []
sample_num_frames = []
cnt_too_long = 0
@@ -271,8 +275,7 @@ class T2V_dataset(Dataset):
cnt_resolution_mismatch += 1
continue
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
# from IPython import embed; embed()
# import ipdb;ipdb.set_trace()
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
@@ -323,14 +326,14 @@ class T2V_dataset(Dataset):
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices) -> torch.Tensor:
def decord_read(self, path, frame_indices):
decord_vr = self.v_decoder(path)
video_data = decord_vr.get_batch(frame_indices).asnumpy()
video_data = torch.from_numpy(video_data)
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
return video_data
def read_jsons(self, data) -> list[dict]:
def read_jsons(self, data):
cap_lists = []
with open(data) as f:
folder_anno = [
@@ -346,6 +349,6 @@ class T2V_dataset(Dataset):
cap_lists += sub_list
return cap_lists
def get_cap_list(self) -> list:
def get_cap_list(self):
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
+504 -15
View File
@@ -1,19 +1,42 @@
import numbers
import random
import torch
from PIL import Image
def _is_tensor_video_clip(clip) -> bool:
def _is_tensor_video_clip(clip):
if not torch.is_tensor(clip):
raise TypeError(f"clip should be Tensor. Got {type(clip)}")
raise TypeError("clip should be Tensor. Got %s" % type(clip))
if not clip.ndimension() == 4:
raise ValueError(f"clip should be 4D. Got {clip.dim()}D")
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
return True
def crop(clip, i, j, h, w) -> torch.Tensor:
def center_crop_arr(pil_image, image_size):
"""
Center cropping implementation from ADM.
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
"""
while min(*pil_image.size) >= 2 * image_size:
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size),
resample=Image.BOX)
scale = image_size / min(*pil_image.size)
pil_image = pil_image.resize(tuple(
round(x * scale) for x in pil_image.size),
resample=Image.BICUBIC)
arr = np.array(pil_image)
crop_y = (arr.shape[0] - image_size) // 2
crop_x = (arr.shape[1] - image_size) // 2
return Image.fromarray(arr[crop_y:crop_y + image_size,
crop_x:crop_x + image_size])
def crop(clip, i, j, h, w):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
@@ -23,7 +46,7 @@ def crop(clip, i, j, h, w) -> torch.Tensor:
return clip[..., i:i + h, j:j + w]
def resize(clip, target_size, interpolation_mode) -> torch.Tensor:
def resize(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
@@ -37,7 +60,71 @@ def resize(clip, target_size, interpolation_mode) -> torch.Tensor:
)
def center_crop_th_tw(clip, th, tw, top_crop) -> torch.Tensor:
def resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
clip,
scale_factor=scale_,
mode=interpolation_mode,
align_corners=True,
antialias=True,
)
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
"""
Do spatial cropping and resizing to the video clip
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
i (int): i in (i,j) i.e coordinates of the upper left corner.
j (int): j in (i,j) i.e coordinates of the upper left corner.
h (int): Height of the cropped region.
w (int): Width of the cropped region.
size (tuple(int, int)): height and width of resized clip
Returns:
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
clip = crop(clip, i, j, h, w)
clip = resize(clip, size, interpolation_mode)
return clip
def center_crop(clip, crop_size):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
th, tw = crop_size
if h < th or w < tw:
raise ValueError("height and width must be no smaller than crop_size")
i = int(round((h - th) / 2.0))
j = int(round((w - tw) / 2.0))
return crop(clip, i, j, th, tw)
def center_crop_using_short_edge(clip):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h < w:
th, tw = h, h
i = 0
j = int(round((w - tw) / 2.0))
else:
th, tw = w, w
i = int(round((h - th) / 2.0))
j = 0
return crop(clip, i, j, th, tw)
def center_crop_th_tw(clip, th, tw, top_crop):
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
@@ -56,7 +143,27 @@ def center_crop_th_tw(clip, th, tw, top_crop) -> torch.Tensor:
return crop(clip, i, j, new_h, new_w)
def normalize_video(clip) -> torch.Tensor:
def random_shift_crop(clip):
"""
Slide along the long edge, with the short edge as crop size
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
h, w = clip.size(-2), clip.size(-1)
if h <= w:
short_edge = h
else:
short_edge = w
th, tw = short_edge, short_edge
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return crop(clip, i, j, th, tw)
def normalize_video(clip):
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
permute the dimensions of clip tensor
@@ -67,12 +174,153 @@ def normalize_video(clip) -> torch.Tensor:
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError(
f"clip tensor should have data type uint8. Got {clip.dtype}")
raise TypeError("clip tensor should have data type uint8. Got %s" %
str(clip.dtype))
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
def normalize(clip, mean, std, inplace=False):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
mean (tuple): pixel RGB mean. Size is (3)
std (tuple): pixel standard deviation. Size is (3)
Returns:
normalized clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
if not inplace:
clip = clip.clone()
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
# print(mean)
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
return clip
def hflip(clip):
"""
Args:
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
Returns:
flipped clip (torch.tensor): Size is (T, C, H, W)
"""
if not _is_tensor_video_clip(clip):
raise ValueError("clip should be a 4D torch.tensor")
return clip.flip(-1)
class RandomCropVideo:
def __init__(self, size):
if isinstance(size, numbers.Number):
self.size = (int(size), int(size))
else:
self.size = size
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: randomly cropped video clip.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = self.size
if h < th or w < tw:
raise ValueError(
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
)
if w == tw and h == th:
return 0, 0, h, w
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return i, j, th, tw
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: cropped video clip by stride.
size is (T, C, OH, OW)
"""
i, j, h, w = self.get_params(clip)
return crop(clip, i, j, h, w)
def get_params(self, clip):
h, w = clip.shape[-2:]
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
return 0, 0, th, tw # from top-left
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size})"
class LongSideResizeVideo:
"""
First use the long side,
then resize to the specified size
"""
def __init__(
self,
size,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized video clip.
size is (T, C, 512, *) or (T, C, *, 512)
"""
_, _, h, w = clip.shape
if self.skip_low_resolution and max(h, w) <= self.size:
return clip
if h > w:
w = int(w * self.size / h)
h = self.size
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(clip,
target_size=(h, w),
interpolation_mode=self.interpolation_mode)
return resize_clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class CenterCropResizeVideo:
"""
First use the short side for cropping length,
@@ -84,7 +332,7 @@ class CenterCropResizeVideo:
size,
top_crop=False,
interpolation_mode="bilinear",
) -> None:
):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
@@ -92,7 +340,7 @@ class CenterCropResizeVideo:
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
def __call__(self, clip) -> torch.Tensor:
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
@@ -100,10 +348,12 @@ class CenterCropResizeVideo:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
# clip_center_crop = center_crop_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(clip,
self.size[0],
self.size[1],
top_crop=self.top_crop)
# import ipdb;ipdb.set_trace()
clip_center_crop_resize = resize(
clip_center_crop,
target_size=self.size,
@@ -115,15 +365,138 @@ class CenterCropResizeVideo:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class UCFCenterCropVideo:
"""
First scale to the specified size in equal proportion to the short edge,
then center cropping
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_resize = resize_scale(clip=clip,
target_size=self.size,
interpolation_mode=self.interpolation_mode)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class KineticsRandomCropResizeVideo:
"""
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
"""
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size,
self.interpolation_mode)
return clip_resize
class CenterCropVideo:
def __init__(
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
self.interpolation_mode = interpolation_mode
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
Returns:
torch.tensor: center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_center_crop = center_crop(clip, self.size)
return clip_center_crop
def __repr__(self) -> str:
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
class Normalize:
"""
Normalize the video clip by mean subtraction and division by standard deviation
Args:
mean (3-tuple): pixel RGB mean
std (3-tuple): pixel RGB standard deviation
inplace (boolean): whether do in-place normalization
"""
def __init__(self, mean, std, inplace=False):
self.mean = mean
self.std = std
self.inplace = inplace
def __call__(self, clip):
"""
Args:
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
"""
return normalize(clip, self.mean, self.std, self.inplace)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
class Normalize255:
"""
Convert tensor data type from uint8 to float, divide value by 255.0 and
"""
def __init__(self) -> None:
def __init__(self):
pass
def __call__(self, clip) -> torch.Tensor:
def __call__(self, clip):
"""
Args:
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
@@ -136,6 +509,34 @@ class Normalize255:
return self.__class__.__name__
class RandomHorizontalFlipVideo:
"""
Flip the video clip along the horizontal direction with a given probability
Args:
p (float): probability of the clip being flipped. Default value is 0.5
"""
def __init__(self, p=0.5):
self.p = p
def __call__(self, clip):
"""
Args:
clip (torch.tensor): Size is (T, C, H, W)
Return:
clip (torch.tensor): Size is (T, C, H, W)
"""
if random.random() < self.p:
clip = hflip(clip)
return clip
def __repr__(self) -> str:
return f"{self.__class__.__name__}(p={self.p})"
# ------------------------------------------------------------
# --------------------- Sampling ---------------------------
# ------------------------------------------------------------
class TemporalRandomCrop:
"""Temporally crop the given frame indices at a random location.
@@ -143,11 +544,99 @@ class TemporalRandomCrop:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, size) -> None:
def __init__(self, size):
self.size = size
def __call__(self, total_frames) -> tuple[int, int]:
def __call__(self, total_frames):
rand_end = max(0, total_frames - self.size - 1)
begin_index = random.randint(0, rand_end)
end_index = min(begin_index + self.size, total_frames)
return begin_index, end_index
class DynamicSampleDuration:
"""Temporally crop the given frame indices at a random location.
Args:
size (int): Desired length of frames will be seen in the model.
"""
def __init__(self, t_stride, extra_1):
self.t_stride = t_stride
self.extra_1 = extra_1
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(
range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
return 0, truncate_t
if __name__ == "__main__":
import os
import numpy as np
import torchvision.io as io
from torchvision import transforms
from torchvision.utils import save_image
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi",
pts_unit="sec",
output_format="TCHW")
trans = transforms.Compose([
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5],
std=[0.5, 0.5, 0.5],
inplace=True),
])
target_video_len = 32
frame_interval = 1
total_frames = len(vframes)
print(total_frames)
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
# Sampling video frames
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(start_frame_ind,
end_frame_ind - 1,
target_video_len,
dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
print(select_vframes.shape)
print(select_vframes.dtype)
select_vframes_trans = trans(select_vframes)
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) *
255).to(dtype=torch.uint8)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi",
select_vframes_trans_int.permute(0, 2, 3, 1),
fps=8)
for i in range(target_video_len):
save_image(
select_vframes_trans[i],
os.path.join("./test000", "%04d.png" % i),
normalize=True,
value_range=(-1, 1),
)
+2 -8
View File
@@ -2,10 +2,8 @@
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)
@@ -14,15 +12,11 @@ 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",
]
+7 -49
View File
@@ -655,7 +655,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
@@ -704,7 +704,7 @@ def init_world_group(ranks: List[int], local_rank: int,
group_ranks=[ranks],
local_rank=local_rank,
torch_distributed_backend=backend,
use_device_communicator=True,
use_device_communicator=False,
group_name="world",
)
@@ -747,10 +747,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(
@@ -794,18 +794,9 @@ 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,
) -> None:
"""
@@ -861,22 +852,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,20 +863,9 @@ 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,
) -> None:
"""Helper to initialize model parallel groups if they are not initialized,
@@ -912,8 +876,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 +895,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 +948,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
@@ -0,0 +1,29 @@
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from fastvideo.v1.pipelines.wan.wan_latent_pipeline import WanLatentPipeline
def main():
print("Starting data preprocessor")
pipeline = WanLatentPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
for batch in train_dataloader:
pipeline(batch)
if __name__ == "__main__":
main()
+6 -8
View File
@@ -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)
@@ -118,6 +118,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,
@@ -275,10 +276,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 +295,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.
+24 -83
View File
@@ -44,8 +44,6 @@ class FastVideoArgs:
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
# Video generation parameters
@@ -57,8 +55,6 @@ class FastVideoArgs:
# 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"
@@ -86,19 +82,10 @@ class FastVideoArgs:
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
enable_torch_compile: bool = False
use_cpu_offload: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
@@ -192,20 +179,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 +254,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 +269,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 +332,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 +343,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)
@@ -502,11 +438,6 @@ def get_current_fastvideo_args() -> FastVideoArgs:
@dataclasses.dataclass
class TrainingArgs(FastVideoArgs):
"""
Training arguments. Inherits from FastVideoArgs and adds training-specific
arguments. If there are any conflicts, the training arguments will take
precedence.
"""
data_path: str = ""
dataloader_num_workers: int = 0
num_height: int = 0
@@ -536,13 +467,14 @@ class TrainingArgs(FastVideoArgs):
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
seed: Optional[int] = None
# seed: int
# output
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: bool = False
resume_from_checkpoint: str = ""
resume_from_lora_checkpoint: str = ""
logging_dir: str = ""
# optimizer & scheduler
@@ -558,7 +490,13 @@ class TrainingArgs(FastVideoArgs):
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
use_cpu_offload: bool = False
# fp16_full_eval: bool
# fp16_backend: str
train_sp_batch_size: int = 0
use_lora: bool = False
lora_alpha: int = 0
lora_rank: int = 0
fsdp_sharding_startegy: str = ""
weighting_scheme: str = ""
@@ -566,6 +504,8 @@ class TrainingArgs(FastVideoArgs):
logit_std: float = 1.0
mode_scale: float = 0.0
# lr_scheduler
lr_scheduler: str = ""
num_euler_timesteps: int = 0
lr_num_cycles: int = 0
lr_power: float = 0.0
@@ -584,9 +524,6 @@ class TrainingArgs(FastVideoArgs):
# 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
@@ -602,10 +539,6 @@ class TrainingArgs(FastVideoArgs):
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)
@@ -701,9 +634,6 @@ class TrainingArgs(FastVideoArgs):
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",
@@ -719,6 +649,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument("--resume-from-lora-checkpoint",
type=str,
help="Path to LoRA checkpoint to resume from")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
@@ -767,6 +700,14 @@ class TrainingArgs(FastVideoArgs):
type=int,
help="Training spatial parallelism batch size")
# LoRA configuration
parser.add_argument("--use-lora",
action=StoreBoolean,
help="Whether to use LoRA")
parser.add_argument("--lora-alpha",
type=int,
help="LoRA alpha parameter")
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--fsdp-sharding-strategy",
type=str,
help="FSDP sharding strategy")
+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, Dict
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
-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
-1
View File
@@ -78,7 +78,6 @@ 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
-1
View File
@@ -441,7 +441,6 @@ 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)
-1
View File
@@ -459,7 +459,6 @@ 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
-2
View File
@@ -319,7 +319,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,7 +359,6 @@ 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:
+27 -30
View File
@@ -15,10 +15,10 @@ from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
from fastvideo.v1.models.loader.weight_utils import (
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
@@ -391,36 +391,33 @@ class TransformerLoader(ComponentLoader):
len(safetensors_list), model_path)
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
if fastvideo_args.training_mode:
assert isinstance(
fastvideo_args, TrainingArgs
), "fastvideo_args must be a TrainingArgs object when training_mode is True"
default_dtype = PRECISION_TO_TYPE[fastvideo_args.master_weight_type]
else:
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
default_dtype)
assert fastvideo_args.dp_shards is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
"config": dit_config,
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
data_parallel_size=fastvideo_args.dp_size,
data_parallel_shards=fastvideo_args.dp_shards,
cpu_offload=fastvideo_args.use_cpu_offload,
fsdp_inference=fastvideo_args.use_fsdp_inference,
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
training_mode=fastvideo_args.training_mode)
logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype)
# model = load_fsdp_model(model_cls=model_cls,
# init_params={
# "config": dit_config,
# "hf_config": hf_config
# },
# weight_dir_list=safetensors_list,
# device=fastvideo_args.device,
# cpu_offload=fastvideo_args.use_cpu_offload,
# default_dtype=default_dtype)
model = load_fsdp_model(model_cls=model_cls,
init_params={
"config": dit_config,
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
)
if fastvideo_args.enable_torch_compile:
logger.info("Torch Compile enabled for DiT")
for n, m in reversed(list(model.named_modules())):
+79 -59
View File
@@ -5,23 +5,23 @@
# Copyright 2025 The FastVideo Authors.
import contextlib
import re
from collections import defaultdict
from itertools import chain
from typing import (Any, Callable, DefaultDict, Dict, Generator, List, Optional,
Tuple, Type, Union)
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
Optional, Tuple, Type)
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
from torch.distributed._tensor import distribute_tensor
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
MixedPrecisionPolicy, fully_shard)
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.utils import get_param_names_mapping
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.utils import set_mixed_precision_policy
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@@ -54,60 +54,80 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
torch.set_default_dtype(old_dtype)
def get_param_names_mapping(
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
Args:
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
param_name (str): The parameter name to be transformed
Returns:
Callable[[str], str]: A function that maps parameter names from source to target format
"""
def mapping_fn(name: str) -> tuple[str, Any, Any]:
# Try to match and transform the name using the regex patterns in mapping_dict
for pattern, replacement in mapping_dict.items():
match = re.match(pattern, name)
if match:
merge_index = None
total_splitted_params = None
if isinstance(replacement, tuple):
merge_index = replacement[1]
total_splitted_params = replacement[2]
replacement = replacement[0]
name = re.sub(pattern, replacement, name)
return name, merge_index, total_splitted_params
# If no pattern matches, return the original name
return name, None, None
return mapping_fn
# TODO(PY): add compile option
def maybe_load_fsdp_model(
# param_dtype: torch.dtype,
# reduce_dtype: torch.dtype,
# output_dtype: torch.dtype,
# pp_enabled: bool = False,
# cpu_offload: bool = False,
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
data_parallel_size: int,
data_parallel_shards: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
fsdp_inference: bool = False,
output_dtype: Optional[torch.dtype] = None,
training_mode: bool = True,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
"""
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
# manually casting the inputs to the model
mp_policy = MixedPrecisionPolicy(param_dtype,
reduce_dtype,
output_dtype,
cast_forward_inputs=False)
set_mixed_precision_policy(master_dtype=default_dtype,
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
output_dtype=output_dtype)
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
# with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
dp_size = data_parallel_size if fsdp_inference or training_mode else 1
device_mesh = init_device_mesh(
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(dp_size, data_parallel_shards),
mesh_dim_names=("dp", "sp"),
mesh_shape=(get_sequence_model_parallel_world_size(), ),
mesh_dim_names=("dp", ),
)
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
mesh=device_mesh)
dp_mesh=device_mesh["dp"])
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
load_model_from_full_model_state_dict(
load_fsdp_model_from_full_model_state_dict(
model,
weight_iterator,
device,
param_dtype,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
@@ -116,9 +136,9 @@ def maybe_load_fsdp_model(
if p.is_meta:
raise RuntimeError(
f"Unexpected param or buffer {n} on meta device.")
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
for p in model.parameters():
p.requires_grad = False
# set_state_dict(model, StateDictType.LOCAL_STATE_DICT)
return model
@@ -129,7 +149,6 @@ def shard_model(
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
mesh: Optional[DeviceMesh] = None,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
@@ -156,7 +175,7 @@ def shard_model(
"""
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": mesh,
"mesh": dp_mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
@@ -185,28 +204,25 @@ def shard_model(
# TODO(PY): device mesh for cfg parallel
def load_model_from_full_model_state_dict(
model: Union[FSDPModule, torch.nn.Module],
def load_fsdp_model_from_full_model_state_dict(
model: torch.nn.Module,
full_sd_iterator: Generator[Tuple[str, torch.Tensor], None, None],
device: torch.device,
param_dtype: torch.dtype,
strict: bool = False,
cpu_offload: bool = False,
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None,
training_mode: bool = True,
) -> _IncompatibleKeys:
"""
Converting full state dict into a sharded state dict
and loading it into FSDP model (if training) or normal huggingface model
and loading it into FSDP model
Args:
model (Union[FSDPModule, torch.nn.Module]): Model to generate fully qualified names for cpu_state_dict
model (FSDPModule): Model to generate fully qualified names for cpu_state_dict
full_sd_iterator (Generator): an iterator yielding (param_name, tensor) pairs
device (torch.device): device used to move full state dict tensors
param_dtype (torch.dtype): dtype used to move full state dict tensors
strict (bool): flag to check if to load the model in strict mode
cpu_offload (bool): flag to check if FSDP offload is enabled
cpu_offload (bool): flag to check if offload to CPU is enabled
param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name
training_mode (bool): apply FSDP only for training
Returns:
``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields:
* **missing_keys** is a list of str containing the missing keys
@@ -215,18 +231,23 @@ def load_model_from_full_model_state_dict(
Raises:
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sd = model.state_dict()
meta_sharded_sd = model.state_dict()
# s = fully_shard.state(model)
# logger.info(f"type(s): {type(s)}")
# logger.info(f"s: {s}")
# import pdb; pdb.set_trace()
sharded_sd = {}
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
# cat at output dim according to the merge_index order
# cat at dim=1 according to the merge_index order
sorted_tensors = [
to_merge_params[target_param_name][i]
for i in range(num_params_to_merge)
@@ -236,25 +257,24 @@ def load_model_from_full_model_state_dict(
else:
continue
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
sharded_meta_param = meta_sharded_sd.get(target_param_name)
if sharded_meta_param is None:
raise ValueError(
f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict"
)
full_tensor = full_tensor.to(sharded_meta_param.dtype).to(device)
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
if not hasattr(sharded_meta_param, "device_mesh"):
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
sharded_tensor = full_tensor
else:
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
sharded_tensor = distribute_tensor(
full_tensor,
meta_sharded_param.device_mesh,
meta_sharded_param.placements,
sharded_meta_param.device_mesh,
sharded_meta_param.placements,
)
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
# choose `assign=True` since we cannot call `copy_` on meta tensor
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
-36
View File
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""Utilities for selecting and loading models."""
import contextlib
import re
from typing import Any, Callable, Dict
import torch
@@ -18,37 +16,3 @@ def set_default_torch_dtype(dtype: torch.dtype):
torch.set_default_dtype(dtype)
yield
torch.set_default_dtype(old_dtype)
def get_param_names_mapping(
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
Args:
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
param_name (str): The parameter name to be transformed
Returns:
Callable[[str], str]: A function that maps parameter names from source to target format
"""
def mapping_fn(name: str) -> tuple[str, Any, Any]:
# Try to match and transform the name using the regex patterns in mapping_dict
for pattern, replacement in mapping_dict.items():
match = re.match(pattern, name)
if match:
merge_index = None
total_splitted_params = None
if isinstance(replacement, tuple):
merge_index = replacement[1]
total_splitted_params = replacement[2]
replacement = replacement[0]
name = re.sub(pattern, replacement, name)
return name, merge_index, total_splitted_params
# If no pattern matches, return the original name
return name, None, None
return mapping_fn
@@ -1,817 +0,0 @@
# Copied from https://github.com/huggingface/diffusers/blob/v0.31.0/src/diffusers/schedulers/scheduling_unipc_multistep.py
# Convert unipc for flow matching
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import math
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import (KarrasDiffusionSchedulers,
SchedulerMixin,
SchedulerOutput)
from diffusers.utils import deprecate, is_scipy_available
from fastvideo.v1.models.schedulers.base import BaseScheduler
if is_scipy_available():
pass
class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
`UniPCMultistepScheduler` is a training-free framework designed for the fast sampling of diffusion models.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
solver_order (`int`, default `2`):
The UniPC order which can be any positive integer. The effective order of accuracy is `solver_order + 1`
due to the UniC. It is recommended to use `solver_order=2` for guided sampling, and `solver_order=3` for
unconditional sampling.
prediction_type (`str`, defaults to "flow_prediction"):
Prediction type of the scheduler function; must be `flow_prediction` for this scheduler, which predicts
the flow of the diffusion process.
thresholding (`bool`, defaults to `False`):
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
as Stable Diffusion.
dynamic_thresholding_ratio (`float`, defaults to 0.995):
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
sample_max_value (`float`, defaults to 1.0):
The threshold value for dynamic thresholding. Valid only when `thresholding=True` and `predict_x0=True`.
predict_x0 (`bool`, defaults to `True`):
Whether to use the updating algorithm on the predicted x0.
solver_type (`str`, default `bh2`):
Solver type for UniPC. It is recommended to use `bh1` for unconditional sampling when steps < 10, and `bh2`
otherwise.
lower_order_final (`bool`, default `True`):
Whether to use lower-order solvers in the final steps. Only valid for < 15 inference steps. This can
stabilize the sampling of DPMSolver for steps < 15, especially for steps <= 10.
disable_corrector (`list`, default `[]`):
Decides which step to disable the corrector to mitigate the misalignment between `epsilon_theta(x_t, c)`
and `epsilon_theta(x_t^c, c)` which can influence convergence for a large guidance scale. Corrector is
usually disabled during the first few steps.
solver_p (`SchedulerMixin`, default `None`):
Any other scheduler that if specified, the algorithm becomes `solver_p + UniC`.
use_karras_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use Karras sigmas for step sizes in the noise schedule during the sampling process. If `True`,
the sigmas are determined according to a sequence of noise levels {σi}.
use_exponential_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use exponential sigmas for step sizes in the noise schedule during the sampling process.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
steps_offset (`int`, defaults to 0):
An offset added to the inference steps, as required by some model families.
final_sigmas_type (`str`, defaults to `"zero"`):
The final `sigma` value for the noise schedule during the sampling process. If `"sigma_min"`, the final
sigma is the same as the last sigma in the training schedule. If `zero`, the final sigma is set to 0.
"""
_compatibles = [e.name for e in KarrasDiffusionSchedulers]
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
solver_order: int = 2,
prediction_type: str = "flow_prediction",
shift: Optional[float] = 1.0,
use_dynamic_shifting=False,
thresholding: bool = False,
dynamic_thresholding_ratio: float = 0.995,
sample_max_value: float = 1.0,
predict_x0: bool = True,
solver_type: str = "bh2",
lower_order_final: bool = True,
disable_corrector: Tuple = (),
solver_p: SchedulerMixin = None,
timestep_spacing: str = "linspace",
steps_offset: int = 0,
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
**kwargs):
if solver_type not in ["bh1", "bh2"]:
if solver_type in ["midpoint", "heun", "logrho"]:
self.register_to_config(solver_type="bh2")
else:
raise NotImplementedError(
f"{solver_type} is not implemented for {self.__class__}")
self.predict_x0 = predict_x0
# setable values
self.num_inference_steps: Optional[int] = None
alphas = np.linspace(1, 1 / num_train_timesteps,
num_train_timesteps)[::-1].copy()
sigmas = 1.0 - alphas
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32)
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
assert shift is not None
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
self.sigmas = sigmas
self.timesteps = sigmas * num_train_timesteps
self.model_outputs = [None] * solver_order
self.timestep_list: List[Optional[Any]] = [None] * solver_order
self.lower_order_nums = 0
self.disable_corrector = list(disable_corrector)
self.solver_p = solver_p
self.last_sample = None
self._step_index: Optional[int] = None
self._begin_index: Optional[int] = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
def set_shift(self, shift: float) -> None:
self.config.shift = shift
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
# Modified from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler.set_timesteps
def set_timesteps(
self,
num_inference_steps: Union[int, None] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[Union[float, None]] = None,
shift: Optional[Union[float, None]] = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
Total number of the spacing of the time steps.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError(
" you have to pass a value for `mu` when `use_dynamic_shifting` is set to be `True`"
)
if sigmas is None:
assert num_inference_steps is not None
sigmas = np.linspace(self.sigma_max, self.sigma_min,
num_inference_steps +
1).copy()[:-1] # pyright: ignore
if self.config.use_dynamic_shifting:
assert mu is not None
sigmas = self.time_shift(mu, 1.0, sigmas) # pyright: ignore
else:
if shift is None:
shift = self.config.shift
assert isinstance(sigmas, np.ndarray)
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
if self.config.final_sigmas_type == "sigma_min":
sigma_last = ((1 - self.alphas_cumprod[0]) /
self.alphas_cumprod[0])**0.5
elif self.config.final_sigmas_type == "zero":
sigma_last = 0
else:
raise ValueError(
f"`final_sigmas_type` must be one of 'zero', or 'sigma_min', but got {self.config.final_sigmas_type}"
)
timesteps = sigmas * self.config.num_train_timesteps
sigmas = np.concatenate([sigmas, [sigma_last]
]).astype(np.float32) # pyright: ignore
self.sigmas = torch.from_numpy(sigmas)
self.timesteps = torch.from_numpy(timesteps).to(device=device,
dtype=torch.int64)
self.num_inference_steps = len(timesteps)
self.model_outputs = [
None,
] * self.config.solver_order
self.lower_order_nums = 0
self.last_sample = None
if self.solver_p:
self.solver_p.set_timesteps(self.num_inference_steps, device=device)
# add an index counter for schedulers that allow duplicated timesteps
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor:
"""
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
photorealism as well as better image-text alignment, especially when using very large guidance weights."
https://arxiv.org/abs/2205.11487
"""
dtype = sample.dtype
batch_size, channels, *remaining_dims = sample.shape
if dtype not in (torch.float32, torch.float64):
sample = sample.float(
) # upcast for quantile calculation, and clamp not implemented for cpu half
# Flatten sample for doing quantile calculation along each image
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
s = torch.quantile(abs_sample,
self.config.dynamic_thresholding_ratio,
dim=1)
s = torch.clamp(
s, min=1, max=self.config.sample_max_value
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
s = s.unsqueeze(
1) # (batch_size, 1) because clamp will broadcast along dim=0
sample = torch.clamp(
sample, -s, s
) / s # "we threshold xt0 to the range [-s, s] and then divide by s"
sample = sample.reshape(batch_size, channels, *remaining_dims)
sample = sample.to(dtype)
return sample
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler._sigma_to_t
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def _sigma_to_alpha_sigma_t(self, sigma) -> Tuple[Any, Any]:
return 1 - sigma, sigma
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.set_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
def convert_model_output(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
**kwargs,
) -> torch.Tensor:
r"""
Convert the model output to the corresponding type the UniPC algorithm needs.
Args:
model_output (`torch.Tensor`):
The direct output from the learned diffusion model.
timestep (`int`):
The current discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
Returns:
`torch.Tensor`:
The converted model output.
"""
timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None)
if sample is None:
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
"missing `sample` as a required keyword argument")
if timestep is not None:
deprecate(
"timesteps",
"1.0.0",
"Passing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
)
sigma = self.sigmas[self.step_index]
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
if self.predict_x0:
if self.config.prediction_type == "flow_prediction":
sigma_t = self.sigmas[self.step_index]
x0_pred = sample - sigma_t * model_output
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
)
if self.config.thresholding:
x0_pred = self._threshold_sample(x0_pred)
return x0_pred
else:
if self.config.prediction_type == "flow_prediction":
sigma_t = self.sigmas[self.step_index]
epsilon = sample - (1 - sigma_t) * model_output
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
)
if self.config.thresholding:
sigma_t = self.sigmas[self.step_index]
x0_pred = sample - sigma_t * model_output
x0_pred = self._threshold_sample(x0_pred)
epsilon = model_output + x0_pred
return epsilon
def multistep_uni_p_bh_update(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
order: Optional[int] = None, # pyright: ignore
**kwargs,
) -> torch.Tensor:
"""
One step for the UniP (B(h) version). Alternatively, `self.solver_p` is used if is specified.
Args:
model_output (`torch.Tensor`):
The direct output from the learned diffusion model at the current timestep.
prev_timestep (`int`):
The previous discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
order (`int`):
The order of UniP at this timestep (corresponds to the *p* in UniPC-p).
Returns:
`torch.Tensor`:
The sample tensor at the previous timestep.
"""
prev_timestep = args[0] if len(args) > 0 else kwargs.pop(
"prev_timestep", None)
if sample is None:
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
" missing `sample` as a required keyword argument")
if order is None:
if len(args) > 2:
order = args[2]
else:
raise ValueError(
" missing `order` as a required keyword argument")
if prev_timestep is not None:
deprecate(
"prev_timestep",
"1.0.0",
"Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
)
model_output_list = self.model_outputs
s0 = self.timestep_list[-1]
m0 = model_output_list[-1]
x = sample
if self.solver_p:
x_t = self.solver_p.step(model_output, s0, x).prev_sample
return x_t
sigma_t, sigma_s0 = self.sigmas[self.step_index + 1], self.sigmas[
self.step_index] # pyright: ignore
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = sample.device
rks = []
D1s: Optional[List[Any]] = []
for i in range(1, order):
si = self.step_index - i # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
D1s.append((mi - m0) / rk) # pyright: ignore
rks.append(1.0)
rks = torch.tensor(rks, device=device)
R = []
b = []
hh = -h if self.predict_x0 else h
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
h_phi_k = h_phi_1 / hh - 1
factorial_i = 1
if self.config.solver_type == "bh1":
B_h = hh
elif self.config.solver_type == "bh2":
B_h = torch.expm1(hh)
else:
raise NotImplementedError()
for i in range(1, order + 1):
R.append(torch.pow(rks, i - 1))
b.append(h_phi_k * factorial_i / B_h)
factorial_i *= i + 1
h_phi_k = h_phi_k / hh - 1 / factorial_i
R = torch.stack(R)
b = torch.tensor(b, device=device)
if D1s is not None and len(D1s) > 0:
D1s = torch.stack(D1s, dim=1) # (B, K)
# for order 2, we use a simplified version
if order == 2:
rhos_p = torch.tensor([0.5], dtype=x.dtype, device=device)
else:
assert isinstance(R, torch.Tensor)
rhos_p = torch.linalg.solve(R[:-1, :-1],
b[:-1]).to(device).to(x.dtype)
else:
D1s = None
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
D1s) # pyright: ignore
else:
pred_res = 0
x_t = x_t_ - alpha_t * B_h * pred_res
else:
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
if D1s is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
D1s) # pyright: ignore
else:
pred_res = 0
x_t = x_t_ - sigma_t * B_h * pred_res
x_t = x_t.to(x.dtype)
return x_t
def multistep_uni_c_bh_update(
self,
this_model_output: torch.Tensor,
*args,
last_sample: torch.Tensor = None,
this_sample: torch.Tensor = None,
order: Optional[int] = None, # pyright: ignore
**kwargs,
) -> torch.Tensor:
"""
One step for the UniC (B(h) version).
Args:
this_model_output (`torch.Tensor`):
The model outputs at `x_t`.
this_timestep (`int`):
The current timestep `t`.
last_sample (`torch.Tensor`):
The generated sample before the last predictor `x_{t-1}`.
this_sample (`torch.Tensor`):
The generated sample after the last predictor `x_{t}`.
order (`int`):
The `p` of UniC-p at this step. The effective order of accuracy should be `order + 1`.
Returns:
`torch.Tensor`:
The corrected sample tensor at the current timestep.
"""
this_timestep = args[0] if len(args) > 0 else kwargs.pop(
"this_timestep", None)
if last_sample is None:
if len(args) > 1:
last_sample = args[1]
else:
raise ValueError(
" missing`last_sample` as a required keyword argument")
if this_sample is None:
if len(args) > 2:
this_sample = args[2]
else:
raise ValueError(
" missing`this_sample` as a required keyword argument")
if order is None:
if len(args) > 3:
order = args[3]
else:
raise ValueError(
" missing`order` as a required keyword argument")
if this_timestep is not None:
deprecate(
"this_timestep",
"1.0.0",
"Passing `this_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
)
model_output_list = self.model_outputs
m0 = model_output_list[-1]
x = last_sample
x_t = this_sample
model_t = this_model_output
sigma_t, sigma_s0 = self.sigmas[self.step_index], self.sigmas[
self.step_index - 1] # pyright: ignore
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = this_sample.device
rks = []
D1s: Optional[List[Any]] = []
for i in range(1, order):
si = self.step_index - (i + 1) # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
D1s.append((mi - m0) / rk) # pyright: ignore
rks.append(1.0)
rks = torch.tensor(rks, device=device)
R = []
b = []
hh = -h if self.predict_x0 else h
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
h_phi_k = h_phi_1 / hh - 1
factorial_i = 1
if self.config.solver_type == "bh1":
B_h = hh
elif self.config.solver_type == "bh2":
B_h = torch.expm1(hh)
else:
raise NotImplementedError()
for i in range(1, order + 1):
R.append(torch.pow(rks, i - 1))
b.append(h_phi_k * factorial_i / B_h)
factorial_i *= i + 1
h_phi_k = h_phi_k / hh - 1 / factorial_i
R = torch.stack(R)
b = torch.tensor(b, device=device)
if D1s is not None and len(D1s) > 0:
D1s = torch.stack(D1s, dim=1)
else:
D1s = None
# for order 1, we use a simplified version
if order == 1:
rhos_c = torch.tensor([0.5], dtype=x.dtype, device=device)
else:
rhos_c = torch.linalg.solve(R, b).to(device).to(x.dtype)
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
else:
corr_res = 0
D1_t = model_t - m0
x_t = x_t_ - alpha_t * B_h * (corr_res + rhos_c[-1] * D1_t)
else:
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
if D1s is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
else:
corr_res = 0
D1_t = model_t - m0
x_t = x_t_ - sigma_t * B_h * (corr_res + rhos_c[-1] * D1_t)
x_t = x_t.to(x.dtype)
return x_t
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
step_index: int = indices[pos].item()
return step_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._init_step_index
def _init_step_index(self, timestep) -> None:
"""
Initialize the step_index counter for the scheduler.
"""
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def step(self,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
generator=None) -> Union[SchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC.
Args:
model_output (`torch.Tensor`):
The direct output from learned diffusion model.
timestep (`int`):
The current discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`.
Returns:
[`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a
tuple is returned where the first element is the sample tensor.
"""
if self.num_inference_steps is None:
raise ValueError(
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
)
if self.step_index is None:
self._init_step_index(timestep)
use_corrector = (
self.step_index > 0
and self.step_index - 1 not in self.disable_corrector
and self.last_sample is not None # pyright: ignore
)
model_output_convert = self.convert_model_output(model_output,
sample=sample)
if use_corrector:
sample = self.multistep_uni_c_bh_update(
this_model_output=model_output_convert,
last_sample=self.last_sample,
this_sample=sample,
order=self.this_order,
)
for i in range(self.config.solver_order - 1):
self.model_outputs[i] = self.model_outputs[i + 1]
self.timestep_list[i] = self.timestep_list[i + 1]
self.model_outputs[-1] = model_output_convert
self.timestep_list[-1] = timestep # pyright: ignore
if self.config.lower_order_final:
this_order = min(self.config.solver_order,
len(self.timesteps) -
self.step_index) # pyright: ignore
else:
this_order = self.config.solver_order
self.this_order: int = min(this_order, self.lower_order_nums +
1) # warmup for multistep
assert self.this_order > 0
self.last_sample = sample
prev_sample = self.multistep_uni_p_bh_update(
model_output=
model_output, # pass the original non-converted model output, in case solver-p is used
sample=sample,
order=self.this_order,
)
if self.lower_order_nums < self.config.solver_order:
self.lower_order_nums += 1
# upon completion increase step index by one
assert self._step_index is not None
self._step_index += 1 # pyright: ignore
if not return_dict:
return (prev_sample, )
return SchedulerOutput(prev_sample=prev_sample)
def scale_model_input(self, sample: torch.Tensor, *args,
**kwargs) -> torch.Tensor:
"""
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
current timestep.
Args:
sample (`torch.Tensor`):
The input sample.
Returns:
`torch.Tensor`:
A scaled input sample.
"""
return sample
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.add_noise
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=original_samples.device,
dtype=original_samples.dtype)
if original_samples.device.type == "mps" and torch.is_floating_point(
timesteps):
# mps does not support float64
schedule_timesteps = self.timesteps.to(original_samples.device,
dtype=torch.float32)
timesteps = timesteps.to(original_samples.device,
dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(original_samples.device)
timesteps = timesteps.to(original_samples.device)
# begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [
self.index_for_timestep(t, schedule_timesteps)
for t in timesteps
]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timesteps.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timesteps.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(original_samples.shape):
sigma = sigma.unsqueeze(-1)
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
noisy_samples = alpha_t * original_samples + sigma_t * noise
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
-3
View File
@@ -39,9 +39,6 @@ class ParallelTiledVAE(ABC):
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
def to(self, device) -> 'ParallelTiledVAE':
return self
@property
def temporal_compression_ratio(self) -> int:
return cast(int, self.config.temporal_compression_ratio)
+2 -11
View File
@@ -5,12 +5,9 @@ Diffusion pipelines for fastvideo.v1.
This package contains diffusion pipelines for generating videos and images.
"""
from typing import cast
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
from fastvideo.v1.utils import (maybe_download_model,
@@ -19,12 +16,7 @@ from fastvideo.v1.utils import (maybe_download_model,
logger = init_logger(__name__)
class PipelineWithLoRA(LoRAPipeline, ComposedPipelineBase):
"""Type for a pipeline that has both ComposedPipelineBase and LoRAPipeline functionality."""
pass
def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
@@ -53,7 +45,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
return cast(PipelineWithLoRA, pipeline)
return pipeline
__all__ = [
@@ -62,5 +54,4 @@ __all__ = [
"ComposedPipelineBase",
"PipelineRegistry",
"ForwardBatch",
"LoRAPipeline",
]
@@ -40,9 +40,6 @@ class ComposedPipelineBase(ABC):
is_video_pipeline: bool = False # To be overridden by video pipelines
_required_config_modules: List[str] = []
training_args: Optional[TrainingArgs] = None
fastvideo_args: Optional[FastVideoArgs] = None
modules: Dict[str, torch.nn.Module] = {}
# TODO(will): args should support both inference args and training args
def __init__(self,
@@ -54,15 +51,7 @@ class ComposedPipelineBase(ABC):
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
if fastvideo_args.training_mode:
assert isinstance(fastvideo_args, TrainingArgs)
self.training_args = fastvideo_args
assert self.training_args is not None
else:
self.fastvideo_args = fastvideo_args
assert self.fastvideo_args is not None
self.fastvideo_args = fastvideo_args
self.model_path = model_path
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
@@ -88,22 +77,27 @@ class ComposedPipelineBase(ABC):
self.modules = self.load_modules(fastvideo_args)
if fastvideo_args.training_mode:
assert self.training_args is not None
if self.training_args.log_validation:
self.initialize_validation_pipeline(self.training_args)
self.initialize_training_pipeline(self.training_args)
if fastvideo_args.log_validation:
self.initialize_validation_pipeline(fastvideo_args)
self.initialize_training_pipeline(fastvideo_args)
self.initialize_pipeline(fastvideo_args)
if not fastvideo_args.training_mode:
# logger.info("Creating pipeline stages...")
# self.create_pipeline_stages(fastvideo_args)
if fastvideo_args.training_mode:
logger.info("Creating training pipeline stages...")
self.create_training_stages(fastvideo_args)
else:
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
def initialize_training_pipeline(self, training_args: TrainingArgs):
def initialize_training_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if training_mode is True, the pipeline must implement this method")
def initialize_validation_pipeline(self, training_args: TrainingArgs):
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if log_validation is True, the pipeline must implement this method"
)
@@ -140,7 +134,7 @@ class ComposedPipelineBase(ABC):
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args is None or args.inference_mode:
if args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path,
device_str=device or "cuda" if
torch.cuda.is_available() else "cpu",
@@ -161,7 +155,7 @@ class ComposedPipelineBase(ABC):
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
fastvideo_args.num_gpus = int(os.environ.get("WORLD_SIZE", 1))
# we use cpu offload for training
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
@@ -169,13 +163,12 @@ class ComposedPipelineBase(ABC):
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
# fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.master_weight_type == 'fp32', 'only fp32 is supported for training'
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
return cls(model_path,
fastvideo_args,
@@ -190,19 +183,16 @@ class ComposedPipelineBase(ABC):
if local_rank == -1 or world_size == -1 or rank == -1:
raise ValueError(
"Local rank, world size, and rank must be set. Use torchrun to launch the script or pass rank to the worker process."
"Local rank, world size, and rank must be set. Use torchrun to launch the script."
)
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
assert fastvideo_args.tp_size is not None, "tp_size must be set"
assert fastvideo_args.sp_size is not None, "sp_size must be set"
initialize_model_parallel(
tensor_model_parallel_size=fastvideo_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size,
data_parallel_size=fastvideo_args.dp_size)
sequence_model_parallel_size=fastvideo_args.sp_size)
device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device = device
@@ -250,11 +240,18 @@ class ComposedPipelineBase(ABC):
@abstractmethod
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the inference pipeline stages.
Create the pipeline stages.
"""
raise NotImplementedError
def create_training_stages(self, training_args: TrainingArgs):
# @abstractmethod
# def create_validation_stages(self, fastvideo_args: FastVideoArgs):
# """
# Create the validation pipeline stages.
# """
# raise NotImplementedError
def create_training_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the training pipeline stages.
"""
+303
View File
@@ -0,0 +1,303 @@
import os
import sys
import time
from collections import deque
import torch
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.utils.checkpoint import save_checkpoint
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper
from fastvideo.utils.validation import log_validation
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class WanTrainingPipeline(ComposedPipelineBase): # == distill_one_step
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
device = fastvideo_args.device
local_rank = int(os.environ.get("LOCAL_RANK", -1))
rank = int(os.environ.get("RANK", -1))
assert rank != -1
assert local_rank != -1
sp_group = get_sp_group()
world_size = sp_group.world_size
rank = sp_group.rank
args = fastvideo_args
transformer = self.get_module("transformer")
teacher_transformer = self.get_module("teacher_transformer")
ema_transformer = None
assert not fastvideo_args.use_ema, "ema is not supported now"
assert teacher_transformer is not None
assert transformer is not None
train_dataset = self.train_dataset
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
solver = self.solver
noise_random_generator = None
uncond_prompt_embed = self.uncond_prompt_embed
uncond_prompt_mask = self.uncond_prompt_mask
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(train_dataset))
logger.info(" Dataloader size = %s", len(train_dataloader))
logger.info(" Num Epochs = %s", args.num_train_epochs)
logger.info(" Resume training from step %s", init_steps)
logger.info(" Instantaneous batch size per device = %s",
args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s", args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e9)
# print dtype
logger.info(" Master weight dtype: %s",
transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
# loader = self.get_module("train_dataloader")
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule,
step)
loss, grad_norm, pred_norm = self.distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
raise NotImplementedError("lora is not supported now")
# save_lora_checkpoint(transformer, optimizer, rank,
# args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
raise NotImplementedError("ema is not supported now")
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir,
step)
sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.
linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
raise NotImplementedError("lora is not supported now")
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers", args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
print(args)
main(args)
-149
View File
@@ -1,149 +0,0 @@
import logging
from collections import defaultdict
from typing import Any, DefaultDict, Dict, Hashable, List, Optional
import torch
import torch.distributed as dist
from safetensors.torch import load_file
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.layers.lora.linear import (BaseLayerWithLoRA, get_lora_layer,
replace_submodule)
from fastvideo.v1.models.loader.utils import get_param_names_mapping
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.utils import maybe_download_lora
logger = logging.getLogger(__name__)
class LoRAPipeline(ComposedPipelineBase):
"""
Pipeline that supports injecting LoRA adapters into the diffusion transformer.
TODO: support training.
"""
lora_adapters: Dict[str, Dict[str, torch.Tensor]] = defaultdict(
dict) # state dicts of loaded lora adapters
cur_adapter_name: str = ""
lora_layers: Dict[str, BaseLayerWithLoRA] = {}
fastvideo_args: FastVideoArgs
exclude_lora_layers: List[str] = []
device: torch.device = torch.device(f"cuda:{torch.cuda.current_device()}")
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.exclude_lora_layers = self.modules[
"transformer"].config.arch_config.exclude_lora_layers
self.convert_to_lora_layers()
if self.fastvideo_args.lora_path is not None:
self.set_lora_adapter(
self.fastvideo_args.lora_nickname, # type: ignore
self.fastvideo_args.lora_path)
def is_target_layer(self, module_name: str) -> bool:
if self.fastvideo_args.lora_target_names is None:
return True
return any(target_name in module_name
for target_name in self.fastvideo_args.lora_target_names)
def convert_to_lora_layers(self) -> None:
"""
Converts the transformer to a LoRA transformer.
"""
for name, layer in self.modules["transformer"].named_modules():
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers:
if exclude_layer in name:
excluded = True
break
if excluded:
continue
layer = get_lora_layer(layer)
if layer is not None:
self.lora_layers[name] = layer
replace_submodule(self.modules["transformer"], name, layer)
def set_lora_adapter(self,
lora_nickname: str,
lora_path: Optional[str] = None): # type: ignore
"""
Loads a LoRA adapter into the pipeline and applies it to the transformer.
Args:
lora_nickname: The "nick name" of the adapter when referenced in the pipeline.
lora_path: The path to the adapter, either a local path or a Hugging Face repo id.
"""
if lora_nickname not in self.lora_adapters and lora_path is None:
raise ValueError(
f"Adapter {lora_nickname} not found in the pipeline. Please provide lora_path to load it."
)
adapter_updated = False
rank = dist.get_rank()
if lora_path is not None:
lora_local_path = maybe_download_lora(lora_path)
lora_state_dict = load_file(lora_local_path)
# Map the hf layer names to our custom layer names
param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"]._param_names_mapping)
lora_param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"]._lora_param_names_mapping)
to_merge_params: DefaultDict[Hashable,
Dict[Any, Any]] = defaultdict(dict)
for name, weight in lora_state_dict.items():
name = ".".join(
name.split(".")
[1:-1]) # remove the transformer prefix and .weight suffix
name, _, _ = lora_param_names_mapping_fn(name)
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
name)
# for (in_dim, r) @ (r, out_dim), we only merge (r, out_dim * n) where n is the number of linear layers to fuse
# see param mapping in HunyuanVideoArchConfig
if merge_index is not None and "lora_B" in name:
to_merge_params[target_name][merge_index] = weight
if len(to_merge_params[target_name]) == num_params_to_merge:
# cat at output dim according to the merge_index order
sorted_tensors = [
to_merge_params[target_name][i]
for i in range(num_params_to_merge)
]
weight = torch.cat(sorted_tensors, dim=1)
del to_merge_params[target_name]
else:
continue
self.lora_adapters[lora_nickname][target_name] = weight.to(
self.device)
adapter_updated = True
logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path)
if not adapter_updated and lora_nickname == self.cur_adapter_name:
return
# Merge the new adapter
adapted_count = 0
for name, layer in self.lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
if layer.merged:
layer.unmerge_lora_weights()
layer.set_lora_weights(
self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name],
training_mode=self.fastvideo_args.training_mode)
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path, name)
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
self.cur_adapter_name = lora_nickname
+1 -12
View File
@@ -7,8 +7,7 @@ This module defines the dataclasses used to pass state between pipeline componen
in a functional manner, reducing the need for explicit parameter passing.
"""
import pprint
from dataclasses import asdict, dataclass, field
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union
import torch
@@ -115,20 +114,10 @@ class ForwardBatch:
enable_teacache: bool = False
teacache_params: Optional[TeaCacheParams | WanTeaCacheParams] = None
# STA parameters
STA_param: Optional[List] = None
is_cfg_negative: bool = False
mask_search_final_result_pos: Optional[List[List]] = None
mask_search_final_result_neg: Optional[List[List]] = None
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
# Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
def __str__(self):
return pprint.pformat(asdict(self), indent=2, width=120)
+2 -3
View File
@@ -6,11 +6,10 @@ import importlib
import pkgutil
from dataclasses import dataclass, field
from functools import lru_cache
from typing import AbstractSet, Dict, Optional, Tuple, Type, Union
from typing import AbstractSet, Dict, Optional, Tuple, Type
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
logger = init_logger(__name__)
@@ -34,7 +33,7 @@ class _PipelineRegistry:
def resolve_pipeline_cls(
self,
architecture: str,
) -> Tuple[Union[Type[ComposedPipelineBase], Type[LoRAPipeline]], str]:
) -> Tuple[Type[ComposedPipelineBase], str]:
if not architecture:
logger.warning("No pipeline architecture is specified")
@@ -1,92 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
I2V Data Preprocessing pipeline implementation.
This module contains an implementation of the I2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
from typing import Any, Dict, List, Optional
import numpy as np
import torch
from PIL import Image
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
BasePreprocessPipeline)
class PreprocessPipeline_I2V(BasePreprocessPipeline):
"""I2V preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
]
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for I2V pipeline."""
return [f.name for f in pyarrow_schema_i2v]
def get_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
processed_images = []
for frame in first_frame:
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
processed_images.append(processed_img)
# Get CLIP features
pixel_values = torch.cat(
[img['pixel_values'] for img in processed_images],
dim=0).to(fastvideo_args.device)
with torch.no_grad():
image_inputs = {'pixel_values': pixel_values}
with set_forward_context(current_timestep=0, attn_metadata=None):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
return {"clip_feature": clip_features}
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
record.update({
"clip_feature_bytes": clip_feature.tobytes(),
"clip_feature_shape": list(clip_feature.shape),
"clip_feature_dtype": str(clip_feature.dtype),
})
else:
record.update({
"clip_feature_bytes": b"",
"clip_feature_shape": [],
"clip_feature_dtype": "",
})
return record
EntryClass = PreprocessPipeline_I2V
@@ -1,23 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
BasePreprocessPipeline)
class PreprocessPipeline_T2V(BasePreprocessPipeline):
"""T2V preprocessing pipeline implementation."""
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def get_schema_fields(self):
"""Get the schema fields for T2V pipeline."""
return [f.name for f in pyarrow_schema_t2v]
EntryClass = PreprocessPipeline_T2V
@@ -0,0 +1,567 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class PreprocessPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
args,
):
# Initialize class variables for data sharing
self.video_data = {} # Store video metadata and paths
self.latent_data = {} # Store latent tensors
self.preprocess_validation_text(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def preprocess_video_and_text(self, fastvideo_args: FastVideoArgs, args):
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
for root, _, files in os.walk(combined_parquet_dir):
for file in files:
if file.endswith('.parquet'):
table = pq.read_table(os.path.join(root, file))
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
num_processed_samples = 0
# Add progress bar for video preprocessing
pbar = tqdm(train_dataloader,
desc="Processing videos",
unit="batch",
disable=local_rank != 0)
for batch_idx, data in enumerate(pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
# Get corresponding captions for this batch
batch_captions = valid_data["text"]
batch = ForwardBatch(
data_type="video",
prompt=batch_captions,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Remove padding from prompt_embeds using attention mask for all batches
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Create a list to store non-padded embeddings and masks
non_padded_embeds = []
non_padded_masks = []
# Process each item in the batch
for i in range(prompt_embeds.size(0)):
seq_len = seq_lens[i].item()
# Slice the embeddings and masks to keep only non-padding parts
non_padded_embeds.append(prompt_embeds[i, :seq_len])
non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
prompt_embeds = non_padded_embeds
prompt_attention_mask = non_padded_masks
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
height, width = valid_data["pixel_values"][idx].shape[-2:]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# Create record for Parquet dataset
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape":
list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "video",
"width": width,
"height": height,
"num_frames": latents[idx].shape[1],
"duration_sec": float(valid_data["duration"][idx]),
"fps": float(valid_data["fps"][idx]),
}
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array(
[record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["vae_latent_dtype"] for record in batch_data]),
pa.array([
record["text_embedding_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_embedding_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_embedding_dtype"] for record in batch_data
]),
pa.array([
record["text_attention_mask_bytes"]
for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"]
for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"]
for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(
arrays, names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info(f"Collected batch with {len(table)} samples")
if num_processed_samples >= args.flush_frequency:
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(
num_processed_samples // args.samples_per_file, 1)
print(
f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(),
total_chunks))
chunks_per_worker = (total_chunks + num_workers -
1) // num_workers
print(
f"Using {num_workers} workers to process {total_chunks} chunks"
)
logger.info(f"Chunks per worker: {chunks_per_worker}")
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info(
f"Processed chunk with {written} samples")
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially"
)
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(
work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total samples written: {total_written}")
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
# Initialize Parquet dataset
validation_parquet_path = os.path.join(validation_parquet_dir,
"data.parquet")
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
with torch.inference_mode():
# Text Encoder
batch = ForwardBatch(
data_type="video",
prompt=prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds = result_batch.prompt_embeds[0]
prompt_attention_mask = result_batch.prompt_attention_mask[0]
file_name = prompt.split(".")[0]
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
# Slice the embeddings to keep only the non-padding parts
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
# Log the shapes after removing padding
logger.info(
f"Shape after removing padding - Embeddings: {text_embedding.shape}, Mask: {text_attention_mask.shape}"
)
# Create record for Parquet dataset
record = {
"id": file_name,
"vae_latent_bytes": b"", # Not available for validation
"vae_latent_shape": [],
"vae_latent_dtype": "",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": file_name,
"caption": prompt,
"media_type": "video",
"width": 0, # Not available for validation
"height": 0, # Not available for validation
"num_frames": 0, # Not available for validation
"duration_sec": 0.0, # Not available for validation
"fps": 0.0, # Not available for validation
}
batch_data.append(record)
logger.info(f"Saved validation sample: {file_name}")
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array([record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array([record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array([record["vae_latent_dtype"] for record in batch_data]),
pa.array(
[record["text_embedding_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["text_embedding_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["text_embedding_dtype"] for record in batch_data]),
pa.array([
record["text_attention_mask_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"] for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(arrays,
names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
logger.info(f"Total validation samples: {len(table)}")
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=1) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
}
for future in tqdm(futures, desc="Processing chunks"):
try:
total_written += future.result()
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially")
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total validation samples written: {total_written}")
# Clear memory
del table
gc.collect() # Force garbage collection
@staticmethod
def process_chunk_range(args):
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
num_samples = len(table)
# Create worker-specific subdirectory
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Check how many files there are already in the dir, and update i accordingly
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_idx, end_idx):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
chunk = table.slice(start_sample, end_sample - start_sample)
# Create chunk file in worker's directory
chunk_path = os.path.join(
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
# Write to temporary file
pq.write_table(chunk, temp_path, compression='zstd')
# Rename temporary file to final file
if os.path.exists(chunk_path):
os.remove(
chunk_path) # Remove existing file if it exists
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception as e:
# Clean up temporary file if it exists
if os.path.exists(temp_path):
os.remove(temp_path)
raise e
return total_written
except Exception as e:
logger.error(
f"Error processing chunks {start_idx}-{end_idx} for worker {worker_id}: {str(e)}"
)
raise
EntryClass = PreprocessPipeline
@@ -1,539 +0,0 @@
import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any, Dict, List, Optional
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
logger = init_logger(__name__)
class BasePreprocessPipeline(ComposedPipelineBase):
"""Base class for preprocessing pipelines that handles common functionality."""
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
args,
):
# Initialize class variables for data sharing
self.video_data: Dict[str, Any] = {} # Store video metadata and paths
self.latent_data: Dict[str, Any] = {} # Store latent tensors
self.preprocess_validation_text(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def get_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""Get additional features specific to the pipeline type. Override in subclasses."""
return {}
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for the pipeline type. Override in subclasses."""
raise NotImplementedError
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset."""
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx] if valid_data else "",
"media_type": "video",
"width":
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
"height":
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
"num_frames":
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
"duration_sec":
float(valid_data["duration"][idx]) if valid_data else 0.0,
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
}
if extra_features:
record.update(extra_features)
return record
def preprocess_video_and_text(self, fastvideo_args: FastVideoArgs, args):
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
for root, _, files in os.walk(combined_parquet_dir):
for file in files:
if file.endswith('.parquet'):
table = pq.read_table(os.path.join(root, file))
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
num_processed_samples = 0
# Add progress bar for video preprocessing
pbar = tqdm(train_dataloader,
desc="Processing videos",
unit="batch",
disable=local_rank != 0)
for batch_idx, data in enumerate(pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
# Get extra features if needed
extra_features = self.get_extra_features(
valid_data, fastvideo_args)
batch_captions = valid_data["text"]
batch = ForwardBatch(
data_type="video",
prompt=batch_captions,
prompt_embeds=[],
prompt_attention_mask=[],
)
assert hasattr(self, "prompt_encoding_stage")
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
non_padded_embeds = []
non_padded_masks = []
# Process each item in the batch
for i in range(prompt_embeds.size(0)):
seq_len = seq_lens[i].item()
# Slice the embeddings and masks to keep only non-padding parts
non_padded_embeds.append(prompt_embeds[i, :seq_len])
non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
prompt_embeds = non_padded_embeds
prompt_attention_mask = non_padded_masks
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu().numpy(
)
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset
record = self.create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if num_processed_samples >= args.flush_frequency:
self._flush_tables(num_processed_samples, args,
combined_parquet_dir)
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
"""Process validation text prompts and save them to parquet files.
This base implementation handles the common validation text processing logic.
Subclasses can override this method to add pipeline-specific features.
"""
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
prompts = [sampling_param.negative_prompt] + prompts
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
with torch.inference_mode():
# Text Encoder
batch = ForwardBatch(
data_type="video",
prompt=prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
assert hasattr(self, "prompt_encoding_stage")
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds = result_batch.prompt_embeds[0]
prompt_attention_mask = result_batch.prompt_attention_mask[0]
file_name = prompt.split(".")[0]
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
# Log the shapes after removing padding
logger.info(
"Shape after removing padding - Embeddings: %s, Mask: %s",
text_embedding.shape, text_attention_mask.shape)
# Create record for Parquet dataset
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=None,
idx=0,
extra_features=None)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays, names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
logger.info("Total validation samples: %s", len(table))
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=1) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
}
for future in tqdm(futures, desc="Processing chunks"):
try:
total_written += future.result()
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info("Total validation samples written: %s", total_written)
# Clear memory
del table
gc.collect() # Force garbage collection
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(num_processed_samples // args.samples_per_file, 1)
print(f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(), total_chunks))
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
print(f"Using {num_workers} workers to process {total_chunks} chunks")
logger.info("Chunks per worker: %s", chunks_per_worker)
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info("Processed chunk with %s samples", written)
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
# Retry failed ranges sequentially
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info("Total samples written: %s", total_written)
@staticmethod
def process_chunk_range(args: Any) -> int:
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
num_samples = len(table)
# Create worker-specific subdirectory
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Check how many files there are already in the dir, and update i accordingly
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_idx, end_idx):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
chunk = table.slice(start_sample, end_sample - start_sample)
# Create chunk file in worker's directory
chunk_path = os.path.join(
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
# Write to temporary file
pq.write_table(chunk, temp_path, compression='zstd')
# Rename temporary file to final file
if os.path.exists(chunk_path):
os.remove(
chunk_path) # Remove existing file if it exists
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception as e:
# Clean up temporary file if it exists
if os.path.exists(temp_path):
os.remove(temp_path)
raise e
return total_written
except Exception as e:
logger.error("Error processing chunks %s-%s for worker %s: %s",
start_idx, end_idx, worker_id, str(e))
raise
@@ -36,7 +36,6 @@ class ConditioningStage(PipelineStage):
Returns:
The batch with applied conditioning.
"""
# TODO!!
if not batch.do_classifier_free_guidance:
return batch
else:
+1 -2
View File
@@ -7,7 +7,6 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
@@ -24,7 +23,7 @@ class DecodingStage(PipelineStage):
"""
def __init__(self, vae) -> None:
self.vae: ParallelTiledVAE = vae
self.vae = vae
def forward(
self,
+12 -162
View File
@@ -23,6 +23,7 @@ from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.utils import PRECISION_TO_TYPE
st_attn_available = False
spec = importlib.util.find_spec("st_attn")
@@ -47,17 +48,6 @@ class DenoisingStage(PipelineStage):
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
# when used for validation, transformer is None as it is taking from the
# training loop
if transformer is not None:
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
def forward(
self,
@@ -84,7 +74,6 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
# TODO(will): make the precision configurable for inference
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
@@ -95,6 +84,7 @@ class DenoisingStage(PipelineStage):
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
# b c t h w -> b t n s h w
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
@@ -162,10 +152,6 @@ class DenoisingStage(PipelineStage):
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
@@ -204,9 +190,18 @@ class DenoisingStage(PipelineStage):
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
dtype=torch.bfloat16,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
@@ -228,7 +223,6 @@ class DenoisingStage(PipelineStage):
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# fastvideo_args or training_args, and attn_metadata.
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
@@ -247,7 +241,6 @@ class DenoisingStage(PipelineStage):
# Apply guidance
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
@@ -297,10 +290,6 @@ class DenoisingStage(PipelineStage):
# Update batch with final latents
batch.latents = latents
# Save STA mask search results if needed
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == 'STA_searching':
self.save_sta_search_results(batch)
if fastvideo_args.use_cpu_offload:
self.transformer.to('cpu')
torch.cuda.empty_cache()
@@ -372,142 +361,3 @@ class DenoisingStage(PipelineStage):
noise_cfg = (guidance_rescale * noise_pred_rescaled +
(1 - guidance_rescale) * noise_cfg)
return noise_cfg
def prepare_sta_param(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs):
"""
Prepare Sliding Tile Attention (STA) parameters and settings.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
"""
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
from fastvideo.v1.STA_configuration import configure_sta
STA_mode = fastvideo_args.STA_mode
skip_time_steps = fastvideo_args.skip_time_steps
if batch.timesteps is None:
raise ValueError("Timesteps must be provided")
timesteps_num = batch.timesteps.shape[0]
logger.info("STA_mode: %s", STA_mode)
if (batch.num_frames, batch.height,
batch.width) != (69, 768, 1280) and STA_mode != "STA_inference":
raise NotImplementedError(
"STA mask search/tuning is not supported for this resolution")
if STA_mode == "STA_searching" or STA_mode == "STA_tuning" or STA_mode == "STA_tuning_cfg":
size = (batch.width, batch.height)
if size == (1280, 768):
# TODO: make it configurable
sparse_mask_candidates_searching = [
"3, 1, 10", "1, 5, 7", "3, 3, 3", "1, 6, 5", "1, 3, 10",
"3, 6, 1"
]
sparse_mask_candidates_tuning = [
"3, 1, 10", "1, 5, 7", "3, 3, 3", "1, 6, 5", "1, 3, 10",
"3, 6, 1"
]
full_mask = ["3,6,10"]
else:
raise NotImplementedError(
"STA mask search is not supported for this resolution")
layer_num = self.transformer.config.num_layers
# specific for HunyuanVideo
if hasattr(self.transformer.config, "num_single_layers"):
layer_num += self.transformer.config.num_single_layers
head_num = self.transformer.config.num_attention_heads
if STA_mode == "STA_searching":
STA_param = configure_sta(
mode='STA_searching',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_candidates=sparse_mask_candidates_searching +
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
)
elif STA_mode == 'STA_tuning':
STA_param = configure_sta(
mode='STA_tuning',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_search_files_path=
f'output/mask_search_result_pos_{size[0]}x{size[1]}/',
mask_candidates=sparse_mask_candidates_tuning,
full_attention_mask=[int(x) for x in full_mask[0].split(',')],
skip_time_steps=
skip_time_steps, # Use full attention for first 12 steps
save_dir=
f'output/mask_search_strategy_{size[0]}x{size[1]}/', # Custom save directory
timesteps=timesteps_num)
elif STA_mode == 'STA_tuning_cfg':
STA_param = configure_sta(
mode='STA_tuning_cfg',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_search_files_path_pos=
f'output/mask_search_result_pos_{size[0]}x{size[1]}/',
mask_search_files_path_neg=
f'output/mask_search_result_neg_{size[0]}x{size[1]}/',
mask_candidates=sparse_mask_candidates_tuning,
full_attention_mask=[int(x) for x in full_mask[0].split(',')],
skip_time_steps=skip_time_steps,
save_dir=f'output/mask_search_strategy_{size[0]}x{size[1]}/',
timesteps=timesteps_num)
elif STA_mode == 'STA_inference':
import fastvideo.v1.envs as envs
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
if config_file is None:
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
STA_param = configure_sta(mode='STA_inference',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
load_path=config_file)
batch.STA_param = STA_param
batch.mask_search_final_result_pos = [[] for _ in range(timesteps_num)]
batch.mask_search_final_result_neg = [[] for _ in range(timesteps_num)]
def save_sta_search_results(self, batch: ForwardBatch):
"""
Save the STA mask search results.
Args:
batch: The current batch information.
"""
size = (batch.width, batch.height)
if size == (1280, 768):
# TODO: make it configurable
sparse_mask_candidates_searching = [
"3, 1, 10", "1, 5, 7", "3, 3, 3", "1, 6, 5", "1, 3, 10",
"3, 6, 1"
]
else:
raise NotImplementedError(
"STA mask search is not supported for this resolution")
from fastvideo.v1.STA_configuration import save_mask_search_results
if batch.mask_search_final_result_pos is not None and batch.prompt is not None:
save_mask_search_results(
[
dict(layer_data)
for layer_data in batch.mask_search_final_result_pos
],
prompt=str(batch.prompt),
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_pos_{size[0]}x{size[1]}/'
)
if batch.mask_search_final_result_neg is not None and batch.prompt is not None:
save_mask_search_results(
[
dict(layer_data)
for layer_data in batch.mask_search_final_result_neg
],
prompt=str(batch.prompt),
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_neg_{size[0]}x{size[1]}/'
)
@@ -46,8 +46,6 @@ class EncodingStage(PipelineStage):
Returns:
The batch with encoded outputs.
"""
self.vae = self.vae.to(fastvideo_args.device)
image_path = batch.image_path
# TODO(will): remove this once we add input/output validation for stages
if image_path is None:
@@ -20,7 +20,6 @@ from fastvideo.v1.models.encoders.bert import HunyuanClip # type: ignore
from fastvideo.v1.models.encoders.stepllm import STEP1TextEncoder
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.v1.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
@@ -30,7 +29,7 @@ from fastvideo.v1.pipelines.stages import (DecodingStage, DenoisingStage,
logger = init_logger(__name__)
class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
class StepVideoPipeline(ComposedPipelineBase):
_required_config_modules = ["transformer", "scheduler", "vae"]
+824
View File
@@ -0,0 +1,824 @@
import gc
import os
import sys
import time
import traceback
from abc import ABC, abstractmethod
from collections import deque
from copy import deepcopy
import imageio
import numpy as np
import torch
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.checkpoint import save_checkpoint_v1
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.training_utils import (
_clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas)
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
GRADIENT_CHECK_DTYPE = torch.bfloat16
class TrainingPipeline(ComposedPipelineBase, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
All reusable components and code should be implemented in this class.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_training_pipeline(self, fastvideo_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = fastvideo_args.device
self.sp_group = get_sp_group()
self.world_size = self.sp_group.world_size
self.rank = self.sp_group.rank
self.local_rank = self.sp_group.local_rank
self.transformer = self.get_module("transformer")
assert self.transformer is not None
self.transformer.requires_grad_(True)
self.transformer.train()
args = fastvideo_args
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * self.world_size,
num_training_steps=args.max_train_steps * self.world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = ParquetVideoTextDataset(
args.data_path,
batch_size=args.train_batch_size,
rank=self.rank,
world_size=self.world_size,
cfg_rate=args.cfg,
num_latent_t=args.num_latent_t)
train_dataloader = StatefulDataLoader(
train_dataset,
batch_size=args.train_batch_size,
num_workers=args.
dataloader_num_workers, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
self.lr_scheduler = lr_scheduler
self.train_dataset = train_dataset
self.train_dataloader = train_dataloader
self.init_steps = init_steps
self.optimizer = optimizer
self.noise_scheduler = noise_scheduler
# self.noise_random_generator = noise_random_generator
# num_update_steps_per_epoch = math.ceil(
# len(train_dataloader) / args.gradient_accumulation_steps *
# args.sp_size / args.train_sp_batch_size)
# args.num_train_epochs = math.ceil(args.max_train_steps /
# num_update_steps_per_epoch)
if self.rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@abstractmethod
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"Training pipelines must implement this method")
@abstractmethod
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
loader, noise_scheduler, noise_random_generator,
gradient_accumulation_steps, sp_size,
precondition_outputs, max_grad_norm, weighting_scheme,
logit_mean, logit_std, mode_scale):
"""
Train one step of the model.
"""
raise NotImplementedError(
"Training pipeline must implement this method")
def log_validation(self, transformer, fastvideo_args, global_step):
fastvideo_args.inference_mode = True
fastvideo_args.use_cpu_offload = False
if not fastvideo_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
# Prepare validation prompts
print('fastvideo_args.validation_prompt_dir',
fastvideo_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
fastvideo_args.validation_prompt_dir,
batch_size=1,
rank=0,
world_size=1,
cfg_rate=0,
num_latent_t=args.num_latent_t)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=False)
transformer.requires_grad_(False)
for p in transformer.parameters():
p.requires_grad = False
transformer.eval()
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
self.validation_pipeline.latent_preparation_stage.transformer = transformer
self.validation_pipeline.denoising_stage.transformer = transformer
# Process each validation prompt
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
logger.info(f"infos: {infos}")
caption = infos['caption']
captions.append(caption)
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
logger.info('embed dtype', prompt_embeds.dtype)
# Prepare batch for validation
# print('shape of embeddings', prompt_embeds.shape)
batch = ForwardBatch(
# **shallow_asdict(sampling_param),
data_type="video",
latents=None,
# seed=sampling_param.seed,
# data_type="video",
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=args.num_height,
width=args.num_width,
num_frames=args.num_frames,
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=50,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
extra={},
)
# Run validation inference
with torch.autocast("cuda", dtype=torch.bfloat16):
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
# Log validation results
rank = int(os.environ.get("RANK", 0))
if rank == 0:
video_filenames = []
video_captions = []
for i, video in enumerate(videos):
caption = captions[i]
filename = os.path.join(
fastvideo_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_captions.append(
caption) # Store the caption for each video
logs = {
"validation_videos": [
wandb.Video(filename,
caption=caption) for filename, caption in zip(
video_filenames, video_captions)
]
}
wandb.log(logs, step=global_step)
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def gradient_check_parameters(self,
transformer,
latents,
encoder_hidden_states,
encoder_attention_mask,
timesteps,
target,
eps=5e-2,
max_params_to_check=2000):
"""
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
"""
# Move all inputs to CPU and clear GPU memory
inputs_cpu = {
'latents': latents.cpu(),
'encoder_hidden_states': encoder_hidden_states.cpu(),
'encoder_attention_mask': encoder_attention_mask.cpu(),
'timesteps': timesteps.cpu(),
'target': target.cpu()
}
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
torch.cuda.empty_cache()
def compute_loss():
# Move inputs to GPU, compute loss, cleanup
inputs_gpu = {
k:
v.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE
if k != 'encoder_attention_mask' else None)
for k, v in inputs_cpu.items()
}
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
# with torch.autocast(enabled=False, device_type="cuda"):
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
with set_forward_context(
current_timestep=inputs_gpu['timesteps'],
attn_metadata=None):
model_pred = transformer(
hidden_states=inputs_gpu['latents'],
encoder_hidden_states=inputs_gpu[
'encoder_hidden_states'],
timestep=inputs_gpu['timesteps'],
encoder_attention_mask=inputs_gpu[
'encoder_attention_mask'],
return_dict=False)[0]
if self.fastvideo_args.precondition_outputs:
sigmas = get_sigmas(self.noise_scheduler,
inputs_gpu['latents'].device,
inputs_gpu['timesteps'],
n_dim=inputs_gpu['latents'].ndim,
dtype=inputs_gpu['latents'].dtype)
model_pred = inputs_gpu['latents'] - model_pred * sigmas
target_adjusted = inputs_gpu['target']
else:
target_adjusted = inputs_gpu['target']
loss = torch.mean((model_pred - target_adjusted)**2)
# Cleanup and return
loss_cpu = loss.cpu()
del inputs_gpu, model_pred, target_adjusted
if 'sigmas' in locals(): del sigmas
torch.cuda.empty_cache()
return loss_cpu.to(self.fastvideo_args.device)
try:
# Get analytical gradients
transformer.zero_grad()
analytical_loss = compute_loss()
analytical_loss.backward()
# Check gradients for selected parameters
absolute_errors = []
param_count = 0
for name, param in transformer.named_parameters():
if not (param.requires_grad and param.grad is not None
and param_count < max_params_to_check
and param.grad.abs().max() > 5e-4):
continue
# Get local parameter and gradient tensors
local_param = param._local_tensor if hasattr(
param, '_local_tensor') else param
local_grad = param.grad._local_tensor if hasattr(
param.grad, '_local_tensor') else param.grad
# Find first significant gradient element
flat_param = local_param.data.view(-1)
flat_grad = local_grad.view(-1)
check_idx = next((i for i in range(min(10, flat_param.numel()))
if abs(flat_grad[i]) > 1e-4), 0)
# Store original values
orig_value = flat_param[check_idx].item()
analytical_grad = flat_grad[check_idx].item()
# Compute numerical gradient
for delta in [eps, -eps]:
with torch.no_grad():
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0: loss_plus = loss.item()
else: loss_minus = loss.item()
# Restore parameter and compute error
with torch.no_grad():
flat_param[check_idx] = orig_value
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
abs_error = abs(analytical_grad - numerical_grad)
rel_error = abs_error / max(abs(analytical_grad),
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
logger.info(
f"{name}[{check_idx}]: analytical={analytical_grad:.6f}, "
f"numerical={numerical_grad:.6f}, abs_error={abs_error:.2e}, rel_error={rel_error:.2%}"
)
# param_count += 1
# Compute and log statistics
if absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
logger.info(
f"Gradient check stats: min={min_err:.2e}, max={max_err:.2e}, mean={mean_err:.2e}"
)
if self.rank <= 0:
wandb.log({
"grad_check/min_abs_error":
min_err,
"grad_check/max_abs_error":
max_err,
"grad_check/mean_abs_error":
mean_err,
"grad_check/analytical_loss":
analytical_loss.item(),
})
return max_err
return float('inf')
except Exception as e:
logger.error(f"Gradient check failed: {e}")
traceback.print_exc()
return float('inf')
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
noise_random_generator):
"""
Setup and perform gradient check on a fresh batch.
Args:
args: Training arguments
loader_iter: Data loader iterator
noise_scheduler: Noise scheduler for diffusion
noise_random_generator: Random number generator for noise
Returns:
float or None: Maximum gradient error or None if check is disabled/fails
"""
if not ENABLE_GRADIENT_CHECK:
return None
try:
# Get a fresh batch and process it exactly like train_one_step
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
loader_iter)
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
check_latents = check_latents.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE)
check_encoder_hidden_states = check_encoder_hidden_states.to(
self.fastvideo_args.device, dtype=GRADIENT_CHECK_DTYPE)
check_latents = normalize_dit_input("wan", check_latents)
batch_size = check_latents.shape[0]
check_noise = torch.randn_like(check_latents)
check_u = compute_density_for_timestep_sampling(
weighting_scheme=args.weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=args.logit_mean,
logit_std=args.logit_std,
mode_scale=args.mode_scale,
)
check_indices = (check_u *
noise_scheduler.config.num_train_timesteps).long()
check_timesteps = noise_scheduler.timesteps[check_indices].to(
device=check_latents.device)
check_sigmas = get_sigmas(
noise_scheduler,
check_latents.device,
check_timesteps,
n_dim=check_latents.ndim,
dtype=check_latents.dtype,
)
check_noisy_model_input = (
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
# Compute target exactly like train_one_step
if args.precondition_outputs:
check_target = check_latents
else:
check_target = check_noise - check_latents
# Perform gradient check with the exact same inputs as training
max_grad_error = self.gradient_check_parameters(
transformer=self.transformer,
latents=
check_noisy_model_input, # Use noisy input like in training
encoder_hidden_states=check_encoder_hidden_states,
encoder_attention_mask=check_encoder_attention_mask,
timesteps=check_timesteps,
target=check_target,
max_params_to_check=100 # Check more parameters
)
if max_grad_error > 5e-2:
logger.error(
f"❌ Large gradient error detected: {max_grad_error:.2e}")
else:
logger.info(
f"✅ Gradient check passed: max error {max_grad_error:.2e}")
return max_grad_error
except Exception as e:
logger.error(f"Gradient check setup failed: {e}")
traceback.print_exc()
return None
class WanTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def create_training_stages(self, fastvideo_args: FastVideoArgs):
pass
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(fastvideo_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
# TODO(will): clean this up
args_copy.precision = "bf16"
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=args_copy)
self.validation_pipeline = validation_pipeline
def train_one_step(
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
):
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
encoder_attention_mask,
infos,
) = next(loader_iter)
latents = latents.to(self.fastvideo_args.device,
dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
self.fastvideo_args.device, dtype=torch.bfloat16)
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
print('device before forward ',
next(transformer.named_parameters())[1].device)
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)[0]
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
if precondition_outputs:
target = latents
else:
target = noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
print('device before backwardin context',
next(transformer.named_parameters())[1].device)
print('device before backward out context',
next(transformer.named_parameters())[1].device)
loss.backward()
print('device after backward out context',
next(transformer.named_parameters())[1].device)
avg_loss = loss.detach().clone()
sp_group = get_sp_group()
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
model_parts = [self.transformer]
grad_norm = _clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
optimizer.step()
print('device after optimizer step',
next(transformer.named_parameters())[1].device)
lr_scheduler.step()
print('device after scheduler step',
next(transformer.named_parameters())[1].device)
return total_loss, grad_norm.item()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
args = fastvideo_args
self.fastvideo_args = args
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
noise_random_generator = None
from diffusers import FlowMatchEulerDiscreteScheduler
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
total_batch_size = (self.world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
# logger.info(f" Num examples = {len(train_dataset)}")
# logger.info(f" Dataloader size = {len(train_dataloader)}")
# logger.info(f" Num Epochs = {args.num_train_epochs}")
logger.info(f" Resume training from step {init_steps}")
logger.info(
f" Instantaneous batch size per device = {args.train_batch_size}")
logger.info(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
logger.info(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}"
)
logger.info(f" Total optimization steps = {args.max_train_steps}")
logger.info(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in self.transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
logger.info(
f" Master weight dtype: {self.transformer.parameters().__next__().dtype}"
)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
loader_iter = iter(train_dataloader)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader_iter)
# get gpu memory usage
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage before train_one_step: {gpu_memory_usage} MB")
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
args.weighting_scheme,
args.logit_mean,
args.logit_std,
args.mode_scale,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage after train_one_step: {gpu_memory_usage} MB")
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info(f"Performing gradient check at step {step}")
self.setup_gradient_check(args, loader_iter, noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
raise NotImplementedError("LoRA is not supported now")
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step, pipe)
else:
# Your existing checkpoint saving code
save_checkpoint_v1(self.transformer, self.rank,
args.output_dir, step)
self.transformer.train()
self.sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
self.log_validation(self.transformer, args, step)
if args.use_lora:
raise NotImplementedError("LoRA is not supported now")
# save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps, pipe)
else:
save_checkpoint_v1(self.transformer, self.rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
print(args)
main(args)
@@ -1,64 +1,22 @@
import json
import math
import os
import time
from typing import Any, Dict, List, Optional, Tuple, Union
from typing import List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
import torch.distributed.checkpoint.stateful
from safetensors.torch import save_file
import torch.distributed.tensor
from fastvideo.v1.logger import init_logger
from fastvideo.v1.training.checkpointing_utils import (ModelWrapper,
OptimizerWrapper,
RandomStateWrapper,
SchedulerWrapper)
logger = init_logger(__name__)
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
def gather_state_dict_on_cpu_rank0(
model,
device: Optional[torch.device] = None,
) -> Dict[str, Any]:
rank = dist.get_rank()
cpu_state_dict = {}
sharded_sd = model.state_dict()
for param_name, param in sharded_sd.items():
if hasattr(param, "_local_tensor"):
# DTensor case
if param.is_cpu:
# Gather directly on CPU
param = param.full_tensor()
else:
if device is not None:
param = param.to(device)
param = param.full_tensor()
else:
# Regular tensor case
if param.is_cpu:
pass
else:
if device is not None:
param = param.to(device)
if rank == 0:
cpu_state_dict[param_name] = param.cpu()
return cpu_state_dict
def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
generator,
logit_mean: Optional[float] = None,
logit_std: Optional[float] = None,
mode_scale: Optional[float] = None,
logit_mean: float = None,
logit_std: float = None,
mode_scale: float = None,
):
"""
Compute the density for sampling the timesteps when doing SD3 training.
@@ -89,7 +47,7 @@ def get_sigmas(noise_scheduler,
device,
timesteps,
n_dim=4,
dtype=torch.float32) -> torch.Tensor:
dtype=torch.float32):
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
schedule_timesteps = noise_scheduler.timesteps.to(device)
timesteps = timesteps.to(device)
@@ -102,156 +60,10 @@ def get_sigmas(noise_scheduler,
return sigma
def save_checkpoint(transformer,
rank,
output_dir,
step,
optimizer=None,
dataloader=None,
scheduler=None,
noise_generator=None) -> None:
"""
Save checkpoint following finetrainer's distributed checkpoint approach.
Saves both distributed checkpoint and consolidated model weights.
"""
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
states = {
"model": ModelWrapper(transformer),
"random_state": RandomStateWrapper(noise_generator),
}
if optimizer is not None:
states["optimizer"] = OptimizerWrapper(transformer, optimizer)
if dataloader is not None:
states["dataloader"] = dataloader
if scheduler is not None:
states["scheduler"] = SchedulerWrapper(scheduler)
dcp_dir = os.path.join(save_dir, "distributed_checkpoint")
logger.info("rank: %s, saving distributed checkpoint to %s",
rank,
dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.save(states, checkpoint_id=dcp_dir)
end_time = time.perf_counter()
logger.info("rank: %s, distributed checkpoint saved in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
cpu_state = gather_state_dict_on_cpu_rank0(transformer, device=None)
if rank == 0:
# Save model weights (consolidated)
weight_path = os.path.join(save_dir,
"diffusion_pytorch_model.safetensors")
logger.info("rank: %s, saving consolidated checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
save_file(cpu_state, weight_path)
logger.info("rank: %s, consolidated checkpoint saved to %s",
rank,
weight_path,
local_main_process_only=False)
# Save model config
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
logger = init_logger(__name__)
def load_checkpoint(transformer,
rank,
checkpoint_path,
optimizer=None,
dataloader=None,
scheduler=None,
noise_generator=None) -> int:
"""
Load checkpoint following finetrainer's distributed checkpoint approach.
Returns the step number from which training should resume.
"""
if not os.path.exists(checkpoint_path):
logger.warning("Checkpoint path %s does not exist", checkpoint_path)
return 0
# Extract step number from checkpoint path
step = int(os.path.basename(checkpoint_path).split('-')[-1])
if rank == 0:
logger.info("Loading checkpoint from step %s", step)
dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint")
if not os.path.exists(dcp_dir):
logger.warning("Distributed checkpoint directory %s does not exist",
dcp_dir)
return 0
states = {
"model": ModelWrapper(transformer),
"random_state": RandomStateWrapper(noise_generator),
}
if optimizer is not None:
states["optimizer"] = OptimizerWrapper(transformer, optimizer)
if dataloader is not None:
states["dataloader"] = dataloader
if scheduler is not None:
states["scheduler"] = SchedulerWrapper(scheduler)
logger.info("rank: %s, loading distributed checkpoint from %s",
rank,
dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.load(states, checkpoint_id=dcp_dir)
end_time = time.perf_counter()
logger.info("rank: %s, distributed checkpoint loaded in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
logger.info("--> checkpoint loaded from step %s", step)
return step
def normalize_dit_input(model_type, latents, args=None) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
vae_config = WanVAEConfig()
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
latents_mean = latents_mean.view(1, -1, 1, 1,
1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
def clip_grad_norm_while_handling_failing_dtensor_cases(
def _clip_grad_norm_while_handling_failing_dtensor_cases(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
@@ -275,8 +87,8 @@ def clip_grad_norm_while_handling_failing_dtensor_cases(
)
except Exception as e:
logger.warning(
"An error occurred while clipping gradients: %s. Gradient clipping will be skipped and gradient "
"norm will not be logged.", e)
f"An error occurred while clipping gradients: {e}. Gradient clipping will be skipped and gradient "
f"norm will not be logged.")
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = True
return None
@@ -336,7 +148,6 @@ def clip_grad_norm_(
total_norm = total_norm.full_tensor()
if pp_mesh is not None:
raise NotImplementedError("Pipeline parallel is not supported")
if math.isinf(norm_type):
dist.all_reduce(total_norm,
op=dist.ReduceOp.MAX,
@@ -396,7 +207,10 @@ def _get_total_norm(
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
) -> torch.Tensor:
tensors = [tensors] if isinstance(tensors, torch.Tensor) else list(tensors)
if isinstance(tensors, torch.Tensor):
tensors = [tensors]
else:
tensors = list(tensors)
norm_type = float(norm_type)
if len(tensors) == 0:
return torch.tensor(0.0)
@@ -449,8 +263,8 @@ def _group_tensors_by_device_and_dtype(
with_indices: bool = False,
) -> dict[tuple[torch.device, torch.dtype], tuple[
List[List[Optional[torch.Tensor]]], List[int]]]:
return torch._C._group_tensors_by_device_and_dtype( # type: ignore[no-any-return]
tensorlistlist, with_indices)
return torch._C._group_tensors_by_device_and_dtype(tensorlistlist,
with_indices)
def _device_has_foreach_support(device: torch.device) -> bool:
@@ -9,7 +9,6 @@ using the modular pipeline architecture.
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.v1.pipelines.stages import (
@@ -17,23 +16,17 @@ from fastvideo.v1.pipelines.stages import (
EncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
class WanImageToVideoPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
@@ -0,0 +1,19 @@
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
class WanLatentPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs):
logger.info("WAN Latent Pipeline forward")
pass

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