Compare commits
40
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8db5dff76f | ||
|
|
4388fa043d | ||
|
|
d6ef6c6ae4 | ||
|
|
da485fbe40 | ||
|
|
f657eb40dc | ||
|
|
19d75b9af3 | ||
|
|
2ecdc2bb8d | ||
|
|
43cb9075f2 | ||
|
|
2768c94977 | ||
|
|
2c35841a39 | ||
|
|
0f0285d1ee | ||
|
|
9f6b0ddc27 | ||
|
|
bb96fa2003 | ||
|
|
e3d0cbe185 | ||
|
|
d5ec468d43 | ||
|
|
e55fa6e5dc | ||
|
|
61b6ddeee1 | ||
|
|
a9a000f45d | ||
|
|
66b8b8561e | ||
|
|
6684872616 | ||
|
|
7f654e3332 | ||
|
|
8631c1b806 | ||
|
|
5357e12b5a | ||
|
|
bdfdf1dfee | ||
|
|
d156461785 | ||
|
|
6edf113838 | ||
|
|
dcf7738cbc | ||
|
|
b2ebaaf865 | ||
|
|
7768bb80f6 | ||
|
|
a335811869 | ||
|
|
357b0533fe | ||
|
|
2ec3732758 | ||
|
|
a004408a93 | ||
|
|
007e237e69 | ||
|
|
8e18dc9f71 | ||
|
|
7ab32539af | ||
|
|
6ef8fcb61d | ||
|
|
016e24da63 | ||
|
|
85b8717545 | ||
|
|
657fd745e1 |
@@ -4,14 +4,6 @@ title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
@@ -25,5 +17,13 @@ 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
|
||||
@@ -77,6 +77,8 @@ jobs:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
|
||||
@@ -10,7 +10,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
|
||||
@@ -27,7 +27,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
**.json
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
@@ -33,7 +33,7 @@ repos:
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
rev: v0.11.12
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
@@ -48,7 +48,7 @@ repos:
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
|
||||
+42906
-42906
File diff suppressed because it is too large
Load Diff
@@ -57,8 +57,9 @@ 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.
|
||||
@@ -79,7 +80,6 @@ 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"
|
||||
|
||||
@@ -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,7 +11,9 @@ 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=1,
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
@@ -23,7 +25,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)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
@@ -34,7 +36,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)
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
def main():
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
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()
|
||||
@@ -0,0 +1,5 @@
|
||||
# STA Mask Search Examples
|
||||
|
||||
```bash
|
||||
bash examples/inference/sta_mask_search/inference_wan_sta.sh
|
||||
```
|
||||
@@ -0,0 +1,39 @@
|
||||
#!/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"
|
||||
@@ -0,0 +1,63 @@
|
||||
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)
|
||||
@@ -0,0 +1,120 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import PreprocessPipeline_I2V
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def main(args):
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
# Assume using torchrun
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
"use_cpu_offload": False,
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
pipeline_config_args = shallow_asdict(pipeline_config)
|
||||
pipeline_config_args.update(kwargs)
|
||||
fastvideo_args = FastVideoArgs(model_path=args.model_path,
|
||||
num_gpus=world_size,
|
||||
device_str="cuda",
|
||||
**pipeline_config_args,
|
||||
)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
|
||||
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
|
||||
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_video_batch_size",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_text_batch_size",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--samples_per_file",
|
||||
type=int,
|
||||
default=64
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flush_frequency",
|
||||
type=int,
|
||||
default=256,
|
||||
help="how often to save to parquet files"
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -68,7 +68,8 @@ 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)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
|
||||
@@ -33,7 +33,8 @@ 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)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
|
||||
+75
-41
@@ -12,6 +12,7 @@ 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
|
||||
@@ -23,7 +24,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.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.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)
|
||||
@@ -123,13 +124,21 @@ 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):
|
||||
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 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,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor([1000.0],
|
||||
device=noisy_model_input.device,
|
||||
@@ -141,47 +150,70 @@ def distill_one_step(
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if 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()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
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()
|
||||
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
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 = transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
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]
|
||||
else:
|
||||
target_pred = transformer(**target_pred_kwargs)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
|
||||
|
||||
@@ -242,7 +274,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
|
||||
@@ -319,7 +351,9 @@ def main(args):
|
||||
teacher_transformer.requires_grad_(False)
|
||||
if args.use_ema:
|
||||
ema_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
@@ -391,7 +425,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)
|
||||
|
||||
@@ -493,7 +527,7 @@ def main(args):
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
|
||||
@@ -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.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.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,
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,609 @@
|
||||
# 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)
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.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,
|
||||
|
||||
@@ -32,7 +32,7 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0 and not discriminator:
|
||||
if rank == 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
@@ -60,7 +60,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
# todo move to get_state_dict
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
@@ -98,7 +98,7 @@ def save_checkpoint_generator_discriminator(
|
||||
hf_weight_dir = os.path.join(save_dir, "hf_weights")
|
||||
os.makedirs(hf_weight_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(hf_weight_dir, "config.json")
|
||||
# save dict as json
|
||||
@@ -139,7 +139,7 @@ def save_checkpoint_generator_discriminator(
|
||||
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
|
||||
model_state = discriminator.state_dict()
|
||||
state_dict = {"optimizer": optim_state, "model": model_state}
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
|
||||
torch.save(state_dict, discriminator_fsdp_state_fil)
|
||||
|
||||
@@ -178,7 +178,7 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
|
||||
):
|
||||
discriminator_state = torch.load(checkpoint_file)
|
||||
model_state = discriminator_state["model"]
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
optim_state = discriminator_state["optimizer"]
|
||||
else:
|
||||
optim_state = None
|
||||
@@ -241,7 +241,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipelin
|
||||
optimizer,
|
||||
)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
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")
|
||||
+69
-2
@@ -3,9 +3,9 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi, AutoencoderKLWan
|
||||
from torch import nn
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
from transformers import AutoTokenizer, T5EncoderModel, UMT5EncoderModel
|
||||
|
||||
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
|
||||
MMSingleStreamBlock)
|
||||
@@ -14,6 +14,7 @@ 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 = {
|
||||
@@ -200,6 +201,48 @@ 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"
|
||||
@@ -240,6 +283,20 @@ 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(
|
||||
@@ -283,6 +340,12 @@ 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")
|
||||
@@ -311,6 +374,8 @@ 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:
|
||||
@@ -322,6 +387,8 @@ 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):
|
||||
|
||||
@@ -129,13 +129,22 @@ 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):
|
||||
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]
|
||||
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]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
@@ -166,10 +175,12 @@ 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, 12, 1, 1,
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).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_std = (torch.tensor(vae.config.latents_std).view(1, num_channels_latents, 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
|
||||
@@ -202,14 +213,15 @@ 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":
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf" or "wan":
|
||||
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)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,419 @@
|
||||
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 List, Optional, Type
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -13,6 +13,7 @@ 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
|
||||
|
||||
@@ -20,7 +21,9 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: Dict[str,
|
||||
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
@@ -42,14 +45,14 @@ def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
|
||||
class RangeDict(dict):
|
||||
|
||||
def __getitem__(self, item):
|
||||
def __getitem__(self, item: int) -> str:
|
||||
for key in self.keys():
|
||||
if isinstance(key, tuple):
|
||||
low, high = key
|
||||
if low <= item <= high:
|
||||
return super().__getitem__(key)
|
||||
return str(super().__getitem__(key))
|
||||
elif key == item:
|
||||
return super().__getitem__(key)
|
||||
return str(super().__getitem__(key))
|
||||
raise KeyError(f"seq_len {item} not supported for STA")
|
||||
|
||||
|
||||
@@ -82,6 +85,8 @@ 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):
|
||||
@@ -98,8 +103,12 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
forward_batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
|
||||
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])
|
||||
|
||||
|
||||
class SlidingTileAttentionImpl(AttentionImpl):
|
||||
@@ -120,12 +129,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)
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
self.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
|
||||
@@ -205,16 +214,24 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
v: torch.Tensor,
|
||||
attn_metadata: SlidingTileAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
|
||||
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"
|
||||
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")
|
||||
|
||||
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])
|
||||
|
||||
# TODO: remove hardcode
|
||||
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]
|
||||
|
||||
text_length = q.shape[1] - self.img_seq_length
|
||||
has_text = text_length > 0
|
||||
@@ -227,15 +244,62 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
sp_group = get_sp_group()
|
||||
current_rank = sp_group.rank_in_group
|
||||
start_head = current_rank * head_num
|
||||
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)
|
||||
|
||||
# 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)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -13,6 +13,7 @@ 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):
|
||||
@@ -38,7 +39,7 @@ class DistributedAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
@@ -155,7 +156,7 @@ class LocalAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional, Tuple
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
@@ -11,6 +11,7 @@ class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
@@ -20,6 +21,7 @@ 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 Optional, Tuple
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -163,6 +163,8 @@ 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,6 +51,7 @@ 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,5 +1,5 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -51,6 +51,23 @@ 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
|
||||
@@ -68,6 +85,7 @@ 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__()
|
||||
|
||||
@@ -63,7 +63,7 @@ class WanVAEArchConfig(VAEArchConfig):
|
||||
|
||||
@dataclass
|
||||
class WanVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
arch_config: WanVAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
use_feature_cache: bool = True
|
||||
|
||||
use_tiling: bool = False
|
||||
|
||||
@@ -27,7 +27,6 @@ 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
|
||||
@@ -55,6 +54,8 @@ 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,9 +68,6 @@ 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,9 +18,6 @@ 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
|
||||
|
||||
@@ -37,9 +37,6 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 3
|
||||
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import os
|
||||
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
|
||||
|
||||
def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
|
||||
]
|
||||
resize = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width)),
|
||||
]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
start_idx=start_idx)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
@@ -0,0 +1,82 @@
|
||||
# schema.py
|
||||
"""
|
||||
Unified data schema and format for saving and loading image/video data after
|
||||
preprocessing.
|
||||
|
||||
It uses apache arrow in-memory format that can be consumed by modern data
|
||||
frameworks that can handle parquet or lance file.
|
||||
"""
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
@@ -0,0 +1,136 @@
|
||||
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)
|
||||
@@ -0,0 +1,109 @@
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
) -> None:
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
self.cfg_rate = cfg_rate
|
||||
self.datase_dir_path = os.path.dirname(json_path)
|
||||
self.video_dir = os.path.join(self.datase_dir_path, "video")
|
||||
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
|
||||
self.prompt_embed_dir = os.path.join(self.datase_dir_path,
|
||||
"prompt_embed")
|
||||
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
|
||||
"prompt_attention_mask")
|
||||
with open(self.json_path) as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
|
||||
self.num_latent_t = num_latent_t
|
||||
|
||||
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
|
||||
|
||||
self.uncond_prompt_mask = torch.zeros(256).bool()
|
||||
self.lengths = [
|
||||
data_item.get("length", 1) for data_item in self.data_anno
|
||||
]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
latent_file = self.data_anno[idx]["latent_path"]
|
||||
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
|
||||
prompt_attention_mask_file = self.data_anno[idx][
|
||||
"prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(
|
||||
os.path.join(self.latent_dir, latent_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
if random.random() < self.cfg_rate:
|
||||
prompt_embed = self.uncond_prompt_embed
|
||||
prompt_attention_mask = self.uncond_prompt_mask
|
||||
else:
|
||||
prompt_embed = torch.load(
|
||||
os.path.join(self.prompt_embed_dir, prompt_embed_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
prompt_attention_mask = torch.load(
|
||||
os.path.join(self.prompt_attention_mask_dir,
|
||||
prompt_attention_mask_file),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
return latent, prompt_embed, prompt_attention_mask
|
||||
|
||||
def __len__(self):
|
||||
return len(self.data_anno)
|
||||
|
||||
|
||||
def latent_collate_function(batch):
|
||||
# return latent, prompt, latent_attn_mask, text_attn_mask
|
||||
# latent_attn_mask: # b t h w
|
||||
# text_attn_mask: b 1 l
|
||||
# needs to check if the latent/prompt' size and apply padding & attn mask
|
||||
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
|
||||
# calculate max shape
|
||||
max_t = max([latent.shape[1] for latent in latents])
|
||||
max_h = max([latent.shape[2] for latent in latents])
|
||||
max_w = max([latent.shape[3] for latent in latents])
|
||||
|
||||
# padding
|
||||
latent_list: list[torch.Tensor] = [
|
||||
torch.nn.functional.pad(
|
||||
latent,
|
||||
(
|
||||
0,
|
||||
max_t - latent.shape[1],
|
||||
0,
|
||||
max_h - latent.shape[2],
|
||||
0,
|
||||
max_w - latent.shape[3],
|
||||
),
|
||||
) for latent in latents
|
||||
]
|
||||
# attn mask
|
||||
latent_attn_mask = torch.ones(len(latent_list), max_t, max_h, max_w)
|
||||
# set to 0 if padding
|
||||
for i, latent in enumerate(latent_list):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
latents = torch.stack(latent_list, dim=0)
|
||||
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
|
||||
@@ -0,0 +1,470 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from einops import rearrange
|
||||
from torch import distributed as dist
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.distributed import (get_dp_group,
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sp_group)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ParquetVideoTextDataset(Dataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
def __init__(self,
|
||||
path: str,
|
||||
batch_size: int = 1024,
|
||||
rank: int = 0,
|
||||
world_size: int = 1,
|
||||
cfg_rate: float = 0.0,
|
||||
num_latent_t: int = 2,
|
||||
seed: int = 0,
|
||||
validation: bool = False):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.rank = rank
|
||||
self.local_rank = get_sequence_model_parallel_rank()
|
||||
self.sp_group = get_sp_group()
|
||||
self.dp_group = get_dp_group()
|
||||
self.dp_world_size = self.dp_group.world_size
|
||||
self.sp_world_size = self.sp_group.world_size
|
||||
self.world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
self.cfg_rate = cfg_rate
|
||||
self.num_latent_t = num_latent_t
|
||||
self.local_indices = None
|
||||
self.validation = validation
|
||||
|
||||
# Negative prompt caching
|
||||
self.neg_metadata = None
|
||||
self.cached_neg_prompt: Dict[str, Any] | None = None
|
||||
|
||||
self.plan_output_dir = os.path.join(
|
||||
self.path,
|
||||
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.json"
|
||||
)
|
||||
|
||||
ranks = get_sp_group().ranks
|
||||
group_ranks: List[List] = [[] for _ in range(self.world_size)]
|
||||
torch.distributed.all_gather_object(group_ranks, ranks)
|
||||
|
||||
if rank == 0:
|
||||
# If a plan already exists, then skip creating a new plan
|
||||
# This will be useful when resume training
|
||||
if os.path.exists(self.plan_output_dir):
|
||||
print(f"Using existing plan from {self.plan_output_dir}")
|
||||
else:
|
||||
print(f"Creating new plan for {self.plan_output_dir}")
|
||||
# Find all parquet files recursively, and record num_rows for each file
|
||||
print(f"Scanning for parquet files in {self.path}")
|
||||
metadatas = []
|
||||
for root, _, files in os.walk(self.path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
num_rows = pq.ParquetFile(
|
||||
file_path).metadata.num_rows
|
||||
for row_idx in range(num_rows):
|
||||
metadatas.append((file_path, row_idx))
|
||||
|
||||
# the negative prompt is always the first row in the first
|
||||
# parquet file
|
||||
if validation:
|
||||
self.neg_metadata = metadatas[0]
|
||||
metadatas = metadatas[1:]
|
||||
|
||||
# Generate the plan that distribute rows among workers
|
||||
random.seed(seed)
|
||||
random.shuffle(metadatas)
|
||||
|
||||
# Get all sp groups
|
||||
# e.g. if num_gpus = 4, sp_size = 2
|
||||
# group_ranks = [(0, 1), (2, 3)]
|
||||
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
|
||||
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
|
||||
group_ranks_list: List[Any] = list(
|
||||
set(tuple(r) for r in group_ranks))
|
||||
num_sp_groups = len(group_ranks_list)
|
||||
plan = defaultdict(list)
|
||||
for idx, metadata in enumerate(metadatas):
|
||||
sp_group_idx = idx % num_sp_groups
|
||||
for global_rank in group_ranks_list[sp_group_idx]:
|
||||
plan[global_rank].append(metadata)
|
||||
|
||||
if validation:
|
||||
assert self.neg_metadata is not None
|
||||
plan["negative_prompt"] = [self.neg_metadata]
|
||||
with open(self.plan_output_dir, "w") as f:
|
||||
json.dump(plan, f)
|
||||
else:
|
||||
pass
|
||||
|
||||
dist.barrier()
|
||||
if validation:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.neg_metadata = plan["negative_prompt"][0]
|
||||
|
||||
def _load_and_cache_negative_prompt(self) -> None:
|
||||
"""Load and cache the negative prompt. Only rank 0 in each SP group should call this."""
|
||||
if not self.validation or self.neg_metadata is None:
|
||||
return
|
||||
|
||||
if self.cached_neg_prompt is not None:
|
||||
return
|
||||
|
||||
# Only rank 0 in each SP group should read the negative prompt
|
||||
try:
|
||||
file_path, row_idx = self.neg_metadata
|
||||
parquet_file = pq.ParquetFile(file_path)
|
||||
|
||||
# Since negative prompt is always the first row (row_idx = 0),
|
||||
# it's always in the first row group
|
||||
row_group_index = 0
|
||||
local_index = row_idx # This will be 0 for the negative prompt
|
||||
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
# Process the negative prompt row
|
||||
self.cached_neg_prompt = self._process_row(row_dict)
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to load negative prompt: %s", e)
|
||||
self.cached_neg_prompt = None
|
||||
|
||||
def get_validation_negative_prompt(
|
||||
self
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Dict[str, Any]]:
|
||||
"""
|
||||
Get the negative prompt for validation.
|
||||
This method ensures the negative prompt is loaded and cached properly.
|
||||
Returns the processed negative prompt data (latents, embeddings, masks, info).
|
||||
"""
|
||||
if not self.validation:
|
||||
raise ValueError(
|
||||
"get_validation_negative_prompt() can only be called in validation mode"
|
||||
)
|
||||
|
||||
# Load and cache if needed (only rank 0 in SP group will actually load)
|
||||
if self.cached_neg_prompt is None:
|
||||
self._load_and_cache_negative_prompt()
|
||||
|
||||
if self.cached_neg_prompt is None:
|
||||
raise RuntimeError(
|
||||
f"Rank {self.rank} (SP rank {self.local_rank}): Could not retrieve negative prompt data"
|
||||
)
|
||||
|
||||
# Extract the components
|
||||
lat, emb, mask, info = (self.cached_neg_prompt["latents"],
|
||||
self.cached_neg_prompt["embeddings"],
|
||||
self.cached_neg_prompt["masks"],
|
||||
self.cached_neg_prompt["info"])
|
||||
|
||||
# Apply the same processing as in __getitem__
|
||||
if lat.numel() == 0: # Validation parquet
|
||||
return lat, emb, mask, info
|
||||
else:
|
||||
lat = lat[:, -self.num_latent_t:]
|
||||
if self.sp_world_size > 1:
|
||||
lat = rearrange(lat,
|
||||
"t (n s) h w -> t n s h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
lat = lat[:, self.local_rank, :, :, :]
|
||||
return lat, emb, mask, info
|
||||
|
||||
def __len__(self):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[str(self.rank)]
|
||||
except Exception as err:
|
||||
raise Exception(
|
||||
"The data plan hasn't been created yet") from err
|
||||
assert self.local_indices is not None
|
||||
return len(self.local_indices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[self.rank]
|
||||
except Exception as err:
|
||||
raise Exception(
|
||||
"The data plan hasn't been created yet") from err
|
||||
assert self.local_indices is not None
|
||||
file_path, row_idx = self.local_indices[idx]
|
||||
parquet_file = pq.ParquetFile(file_path)
|
||||
|
||||
# Calculate the row group to read into memory and the local idx
|
||||
# This way we can avoid reading in the entire parquet file
|
||||
cumulative = 0
|
||||
for i in range(parquet_file.num_row_groups):
|
||||
num_rows = parquet_file.metadata.row_group(i).num_rows
|
||||
if cumulative + num_rows > row_idx:
|
||||
row_group_index = i
|
||||
local_index = row_idx - cumulative
|
||||
break
|
||||
cumulative += num_rows
|
||||
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
processed = self._process_row(row_dict)
|
||||
lat, emb, mask, info = processed["latents"], processed[
|
||||
"embeddings"], processed["masks"], processed["info"]
|
||||
if lat.numel() == 0: # Validation parquet
|
||||
return lat, emb, mask, info
|
||||
else:
|
||||
lat = lat[:, -self.num_latent_t:]
|
||||
if self.sp_world_size > 1:
|
||||
lat = rearrange(lat,
|
||||
"t (n s) h w -> t n s h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
lat = lat[:, self.local_rank, :, :, :]
|
||||
return lat, emb, mask, info
|
||||
|
||||
def _process_row(self, row) -> Dict[str, Any]:
|
||||
"""Process a PyArrow batch into tensors."""
|
||||
|
||||
vae_latent_bytes = row["vae_latent_bytes"]
|
||||
vae_latent_shape = row["vae_latent_shape"]
|
||||
text_embedding_bytes = row["text_embedding_bytes"]
|
||||
text_embedding_shape = row["text_embedding_shape"]
|
||||
text_attention_mask_bytes = row["text_attention_mask_bytes"]
|
||||
text_attention_mask_shape = row["text_attention_mask_shape"]
|
||||
|
||||
# Process latent
|
||||
if not vae_latent_shape: # No VAE latent is stored. Split is validation
|
||||
lat = np.array([])
|
||||
else:
|
||||
lat = np.frombuffer(vae_latent_bytes,
|
||||
dtype=np.float32).reshape(vae_latent_shape)
|
||||
# Make array writable
|
||||
lat = np.copy(lat)
|
||||
|
||||
if random.random() < self.cfg_rate:
|
||||
emb = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
emb = np.frombuffer(text_embedding_bytes,
|
||||
dtype=np.float32).reshape(text_embedding_shape)
|
||||
# Make array writable
|
||||
emb = np.copy(emb)
|
||||
if emb.shape[0] < 512:
|
||||
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
|
||||
padded_emb[:emb.shape[0], :] = emb
|
||||
emb = padded_emb
|
||||
elif emb.shape[0] > 512:
|
||||
emb = emb[:512, :]
|
||||
|
||||
# Process mask
|
||||
if len(text_attention_mask_bytes) > 0 and len(
|
||||
text_attention_mask_shape) > 0:
|
||||
msk = np.frombuffer(text_attention_mask_bytes,
|
||||
dtype=np.uint8).astype(np.bool_)
|
||||
msk = msk.reshape(1, -1)
|
||||
# Make array writable
|
||||
msk = np.copy(msk)
|
||||
if msk.shape[1] < 512:
|
||||
padded_msk = np.zeros((1, 512), dtype=np.bool_)
|
||||
padded_msk[:, :msk.shape[1]] = msk
|
||||
msk = padded_msk
|
||||
elif msk.shape[1] > 512:
|
||||
msk = msk[:, :512]
|
||||
else:
|
||||
msk = np.ones((1, 512), dtype=np.bool_)
|
||||
|
||||
# Collect metadata
|
||||
info = {
|
||||
"width": row["width"],
|
||||
"height": row["height"],
|
||||
"num_frames": row["num_frames"],
|
||||
"duration_sec": row["duration_sec"],
|
||||
"fps": row["fps"],
|
||||
"file_name": row["file_name"],
|
||||
"caption": row["caption"],
|
||||
}
|
||||
|
||||
return {
|
||||
"latents": torch.from_numpy(lat),
|
||||
"embeddings": torch.from_numpy(emb),
|
||||
"masks": torch.from_numpy(msk),
|
||||
"info": info
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Benchmark Parquet dataset loading speed')
|
||||
parser.add_argument('--path',
|
||||
type=str,
|
||||
default="your/dataset/path",
|
||||
help='Path to Parquet dataset')
|
||||
parser.add_argument('--batch_size',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Batch size for DataLoader')
|
||||
parser.add_argument('--num_batches',
|
||||
type=int,
|
||||
default=100,
|
||||
help='Number of batches to benchmark')
|
||||
parser.add_argument('--vae_debug', action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
# Initialize distributed training
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
# Initialize CUDA device first
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
|
||||
# Initialize distributed training
|
||||
if world_size > 1:
|
||||
dist.init_process_group(backend="nccl",
|
||||
init_method="env://",
|
||||
world_size=world_size,
|
||||
rank=rank)
|
||||
print(
|
||||
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
|
||||
)
|
||||
|
||||
# Create dataset
|
||||
dataset = ParquetVideoTextDataset(
|
||||
args.path,
|
||||
batch_size=args.batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataloader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_size=args.batch_size,
|
||||
num_workers=1, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
drop_last=True)
|
||||
|
||||
# Example of how to load dataloader state
|
||||
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
|
||||
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
|
||||
# dataloader.load_state_dict(dataloader_state[rank])
|
||||
|
||||
# Warm-up with synchronization
|
||||
if rank == 0:
|
||||
print("Warming up...")
|
||||
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
|
||||
# Example of how to save dataloader state
|
||||
# if i == 30:
|
||||
# dist.barrier()
|
||||
# local_data = {rank: dataloader.state_dict()}
|
||||
# gathered_data = [None] * world_size
|
||||
# dist.all_gather_object(gathered_data, local_data)
|
||||
# if rank == 0:
|
||||
# global_state_dict = {}
|
||||
# for d in gathered_data:
|
||||
# global_state_dict.update(d)
|
||||
# torch.save(global_state_dict, "dataloader_state.pt")
|
||||
assert torch.sum(masks[0]).item() == torch.count_nonzero(
|
||||
embeddings[0]).item() // 4096
|
||||
if args.vae_debug:
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=VAE_PATH,
|
||||
vae_config=WanVAEConfig(load_encoder=False),
|
||||
vae_precision="fp32")
|
||||
fastvideo_args.device = device
|
||||
vae_loader = VAELoader()
|
||||
vae = vae_loader.load(model_path=VAE_PATH,
|
||||
architecture="",
|
||||
fastvideo_args=fastvideo_args)
|
||||
|
||||
videoprocessor = VideoProcessor(vae_scale_factor=8)
|
||||
|
||||
with torch.inference_mode():
|
||||
video = vae.decode(latents[0].unsqueeze(0).to(device))
|
||||
video = videoprocessor.postprocess_video(video)
|
||||
video_path = os.path.join("/workspace/FastVideo/debug_videos",
|
||||
infos["caption"][0][:50] + ".mp4")
|
||||
export_to_video(video[0], video_path, fps=16)
|
||||
|
||||
# Move data to device
|
||||
# latents = latents.to(device)
|
||||
# embeddings = embeddings.to(device)
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
# Benchmark
|
||||
if rank == 0:
|
||||
print(f"Benchmarking with batch_size={args.batch_size}")
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
for i, (latents, embeddings, masks,
|
||||
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
|
||||
if i >= args.num_batches:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(device)
|
||||
embeddings = embeddings.to(device)
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
total_samples += batch_size
|
||||
|
||||
# Print progress only from rank 0
|
||||
if rank == 0 and (i + 1) % 10 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
print(
|
||||
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
|
||||
)
|
||||
|
||||
# Final statistics
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if rank == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
|
||||
print("\nBenchmark Results:")
|
||||
print(f"Total time: {elapsed:.2f} seconds")
|
||||
print(f"Total samples: {total_samples}")
|
||||
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
|
||||
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
|
||||
|
||||
if world_size > 1:
|
||||
dist.destroy_process_group()
|
||||
@@ -0,0 +1,351 @@
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances: dict[type, 'SingletonMeta'] = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cap_list: list[dict] = []
|
||||
self.elements: list[int] = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements: dict[int, list[int]] = {}
|
||||
self.n_used_elements: dict[int, int] = {}
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info) -> int:
|
||||
worker_id = 0 if work_info is None else work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] %
|
||||
len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h: int,
|
||||
w: int,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16) -> bool:
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0) -> None:
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx) -> dict:
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx) -> dict:
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx) -> dict:
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images"
|
||||
in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps: list[str] = (image_data["cap"] if isinstance(
|
||||
image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
single_text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
single_text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
|
||||
# from IPython import embed; embed()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices) -> torch.Tensor:
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data) -> list[dict]:
|
||||
cap_lists = []
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self) -> list:
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -0,0 +1,153 @@
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip) -> bool:
|
||||
if not torch.is_tensor(clip):
|
||||
raise TypeError(f"clip should be Tensor. Got {type(clip)}")
|
||||
|
||||
if not clip.ndimension() == 4:
|
||||
raise ValueError(f"clip should be 4D. Got {clip.dim()}D")
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i:i + h, j:j + w]
|
||||
|
||||
|
||||
def resize(clip, target_size, interpolation_mode) -> torch.Tensor:
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(
|
||||
f"target size should be tuple (height, width), instead got {target_size}"
|
||||
)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
size=target_size,
|
||||
mode=interpolation_mode,
|
||||
align_corners=True,
|
||||
antialias=True,
|
||||
)
|
||||
|
||||
|
||||
def center_crop_th_tw(clip, th, tw, top_crop) -> torch.Tensor:
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
tr = th / tw
|
||||
if h / w > tr:
|
||||
new_h = int(w * tr)
|
||||
new_w = w
|
||||
else:
|
||||
new_h = h
|
||||
new_w = int(h / tr)
|
||||
|
||||
i = 0 if top_crop else int(round((h - new_h) / 2.0))
|
||||
j = int(round((w - new_w) / 2.0))
|
||||
return crop(clip, i, j, new_h, new_w)
|
||||
|
||||
|
||||
def normalize_video(clip) -> torch.Tensor:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError(
|
||||
f"clip tensor should have data type uint8. Got {clip.dtype}")
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
|
||||
class CenterCropResizeVideo:
|
||||
"""
|
||||
First use the short side for cropping length,
|
||||
center crop video, then resize to the specified size
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
) -> None:
|
||||
if len(size) != 2:
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
self.top_crop = top_crop
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_center_crop = center_crop_th_tw(clip,
|
||||
self.size[0],
|
||||
self.size[1],
|
||||
top_crop=self.top_crop)
|
||||
clip_center_crop_resize = resize(
|
||||
clip_center_crop,
|
||||
target_size=self.size,
|
||||
interpolation_mode=self.interpolation_mode,
|
||||
)
|
||||
return clip_center_crop_resize
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class Normalize255:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def __call__(self, clip) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
return normalize_video(clip)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return self.__class__.__name__
|
||||
|
||||
|
||||
class TemporalRandomCrop:
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, size) -> None:
|
||||
self.size = size
|
||||
|
||||
def __call__(self, total_frames) -> tuple[int, int]:
|
||||
rand_end = max(0, total_frames - self.size - 1)
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
@@ -0,0 +1,10 @@
|
||||
from huggingface_hub import HfApi, upload_folder
|
||||
|
||||
api = HfApi()
|
||||
repo_id = "weizhou03/HD-Mixkit-Finetune-Wan" # customize this
|
||||
api.create_repo(repo_id=repo_id, repo_type="dataset")
|
||||
|
||||
upload_folder(repo_id=repo_id,
|
||||
folder_path="/workspace/data/HD-Mixkit-Finetune-Wan",
|
||||
repo_type="dataset",
|
||||
path_in_repo="")
|
||||
@@ -2,19 +2,27 @@
|
||||
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
|
||||
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,
|
||||
get_tensor_model_parallel_world_size, get_world_group,
|
||||
init_distributed_environment, initialize_model_parallel)
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
|
||||
__all__ = [
|
||||
"init_distributed_environment",
|
||||
"initialize_model_parallel",
|
||||
"get_data_parallel_world_size",
|
||||
"get_data_parallel_rank",
|
||||
"get_sequence_model_parallel_rank",
|
||||
"get_sequence_model_parallel_world_size",
|
||||
"get_tensor_model_parallel_rank",
|
||||
"get_tensor_model_parallel_world_size",
|
||||
"cleanup_dist_env_and_memory",
|
||||
"get_world_group",
|
||||
"get_dp_group",
|
||||
"get_sp_group",
|
||||
"model_parallel_is_initialized",
|
||||
]
|
||||
|
||||
@@ -1,16 +1,182 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
|
||||
|
||||
from typing import Optional
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
from torch import Tensor
|
||||
from torch.distributed import ProcessGroup, ReduceOp
|
||||
|
||||
|
||||
class DistributedAutograd:
|
||||
"""Collection of autograd functions for distributed operations.
|
||||
|
||||
This class provides custom autograd functions for distributed operations like all_reduce,
|
||||
all_gather, and all_to_all. Each operation is implemented as a static inner class with
|
||||
proper forward and backward implementations.
|
||||
"""
|
||||
|
||||
class AllReduce(torch.autograd.Function):
|
||||
"""Differentiable all_reduce operation.
|
||||
|
||||
The gradient of all_reduce is another all_reduce operation since the operation
|
||||
combines values from all ranks equally.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any,
|
||||
group: ProcessGroup,
|
||||
input_: Tensor,
|
||||
op: Optional[dist.ReduceOp] = None) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.op = op
|
||||
output = input_.clone()
|
||||
dist.all_reduce(output, group=group, op=op)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None]:
|
||||
grad_output = grad_output.clone()
|
||||
dist.all_reduce(grad_output, group=ctx.group, op=ctx.op)
|
||||
return None, grad_output, None
|
||||
|
||||
class AllGather(torch.autograd.Function):
|
||||
"""Differentiable all_gather operation.
|
||||
|
||||
The operation gathers tensors from all ranks and concatenates them along a specified dimension.
|
||||
The backward pass uses reduce_scatter to efficiently distribute gradients back to source ranks.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
|
||||
world_size: int, dim: int) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.world_size = world_size
|
||||
ctx.dim = dim
|
||||
ctx.input_shape = input_.shape
|
||||
|
||||
input_size = input_.size()
|
||||
output_size = (input_size[0] * world_size, ) + input_size[1:]
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
|
||||
dist.all_gather_into_tensor(output_tensor, input_, group=group)
|
||||
|
||||
output_tensor = output_tensor.reshape((world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
|
||||
# Split the gradient tensor along the gathered dimension
|
||||
dim_size = grad_output.size(ctx.dim) // ctx.world_size
|
||||
grad_chunks = grad_output.reshape(grad_output.shape[:ctx.dim] +
|
||||
(ctx.world_size, dim_size) +
|
||||
grad_output.shape[ctx.dim + 1:])
|
||||
grad_chunks = grad_chunks.movedim(ctx.dim, 0)
|
||||
|
||||
# Each rank only needs its corresponding gradient
|
||||
grad_input = torch.empty(ctx.input_shape,
|
||||
dtype=grad_output.dtype,
|
||||
device=grad_output.device)
|
||||
dist.reduce_scatter_tensor(grad_input,
|
||||
grad_chunks.contiguous(),
|
||||
group=ctx.group)
|
||||
|
||||
return None, grad_input, None, None
|
||||
|
||||
class AllToAll4D(torch.autograd.Function):
|
||||
"""Differentiable all_to_all operation specialized for 4D tensors.
|
||||
|
||||
This operation is particularly useful for attention operations where we need to
|
||||
redistribute data across ranks for efficient parallel processing.
|
||||
|
||||
The operation supports two modes:
|
||||
1. scatter_dim=2, gather_dim=1: Used for redistributing attention heads
|
||||
2. scatter_dim=1, gather_dim=2: Used for redistributing sequence dimensions
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
|
||||
world_size: int, scatter_dim: int,
|
||||
gather_dim: int) -> Tensor:
|
||||
ctx.group = group
|
||||
ctx.world_size = world_size
|
||||
ctx.scatter_dim = scatter_dim
|
||||
ctx.gather_dim = gather_dim
|
||||
|
||||
if world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * world_size
|
||||
shard_hc = hc // world_size
|
||||
|
||||
input_t = input_.reshape(bs, shard_seqlen, world_size, shard_hc,
|
||||
hs).transpose(0, 2).contiguous()
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous()
|
||||
output = output.reshape(bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * world_size
|
||||
shard_seqlen = seqlen // world_size
|
||||
|
||||
input_t = input_.reshape(bs, world_size, shard_seqlen, shard_hc,
|
||||
hs)
|
||||
input_t = input_t.transpose(0, 3).transpose(0, 1).contiguous()
|
||||
input_t = input_t.reshape(world_size, shard_hc, shard_seqlen,
|
||||
bs, hs)
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
|
||||
output = output.reshape(hc, shard_seqlen, bs, hs)
|
||||
output = output.transpose(0, 2).contiguous()
|
||||
output = output.reshape(bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Invalid scatter_dim={scatter_dim}, gather_dim={gather_dim}. "
|
||||
f"Only (scatter_dim=2, gather_dim=1) and (scatter_dim=1, gather_dim=2) are supported."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def backward(
|
||||
ctx: Any,
|
||||
grad_output: Tensor) -> Tuple[None, Tensor, None, None, None]:
|
||||
if ctx.world_size == 1:
|
||||
return None, grad_output, None, None, None
|
||||
|
||||
# For backward pass, we swap scatter_dim and gather_dim
|
||||
output = DistributedAutograd.AllToAll4D.apply(
|
||||
ctx.group, grad_output, ctx.world_size, ctx.gather_dim,
|
||||
ctx.scatter_dim)
|
||||
return None, output, None, None, None
|
||||
|
||||
|
||||
class DeviceCommunicatorBase:
|
||||
"""
|
||||
Base class for device-specific communicator.
|
||||
Base class for device-specific communicator with autograd support.
|
||||
It can use the `cpu_group` to initialize the communicator.
|
||||
If the device has PyTorch integration (PyTorch can recognize its
|
||||
communication backend), the `device_group` will also be given.
|
||||
@@ -33,35 +199,28 @@ class DeviceCommunicatorBase:
|
||||
self.rank_in_group = dist.get_group_rank(self.cpu_group,
|
||||
self.global_rank)
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
dist.all_reduce(input_, group=self.device_group)
|
||||
return input_
|
||||
def all_reduce(self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[dist.ReduceOp] = ReduceOp.SUM) -> torch.Tensor:
|
||||
"""Performs an all_reduce operation with gradient support."""
|
||||
return DistributedAutograd.AllReduce.apply(self.device_group, input_,
|
||||
op)
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
"""Performs an all_gather operation with gradient support."""
|
||||
if dim < 0:
|
||||
# Convert negative dim to positive.
|
||||
dim += input_.dim()
|
||||
input_size = input_.size()
|
||||
# NOTE: we have to use concat-style all-gather here,
|
||||
# stack-style all-gather has compatibility issues with
|
||||
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
|
||||
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
|
||||
# Allocate output tensor.
|
||||
output_tensor = torch.empty(output_size,
|
||||
dtype=input_.dtype,
|
||||
device=input_.device)
|
||||
# All-gather.
|
||||
dist.all_gather_into_tensor(output_tensor,
|
||||
input_,
|
||||
group=self.device_group)
|
||||
# Reshape
|
||||
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
|
||||
output_tensor = output_tensor.movedim(0, dim)
|
||||
output_tensor = output_tensor.reshape(input_size[:dim] +
|
||||
(self.world_size *
|
||||
input_size[dim], ) +
|
||||
input_size[dim + 1:])
|
||||
return output_tensor
|
||||
return DistributedAutograd.AllGather.apply(self.device_group, input_,
|
||||
self.world_size, dim)
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Performs a 4D all-to-all operation with gradient support."""
|
||||
return DistributedAutograd.AllToAll4D.apply(self.device_group, input_,
|
||||
self.world_size,
|
||||
scatter_dim, gather_dim)
|
||||
|
||||
def gather(self,
|
||||
input_: torch.Tensor,
|
||||
@@ -95,81 +254,6 @@ class DeviceCommunicatorBase:
|
||||
output_tensor = None
|
||||
return output_tensor
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
|
||||
|
||||
Args:
|
||||
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
|
||||
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
|
||||
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after all-to-all operation.
|
||||
"""
|
||||
# Bypass the function if we are using only 1 GPU.
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * self.world_size
|
||||
shard_hc = hc // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, shard_seqlen, self.world_size,
|
||||
shard_hc, hs).transpose(0,
|
||||
2).contiguous())
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous().reshape(
|
||||
bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * self.world_size
|
||||
shard_seqlen = seqlen // self.world_size
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, self.world_size, shard_seqlen,
|
||||
shard_hc, hs).transpose(0, 3).transpose(
|
||||
0, 1).contiguous().reshape(
|
||||
self.world_size, shard_hc,
|
||||
shard_seqlen, bs, hs))
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(hc, shard_seqlen, bs,
|
||||
hs).transpose(0, 2).contiguous().reshape(
|
||||
bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
|
||||
@@ -29,17 +29,19 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def all_reduce(self, input_):
|
||||
def all_reduce(self,
|
||||
input_,
|
||||
op: Optional[torch.distributed.ReduceOp] = None):
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
out = pynccl_comm.all_reduce(input_)
|
||||
out = pynccl_comm.all_reduce(input_, op=op)
|
||||
if out is None:
|
||||
# fall back to the default all-reduce using PyTorch.
|
||||
# this usually happens during testing.
|
||||
# when we run the model, allreduce only happens for the TP
|
||||
# group, where we always have either custom allreduce or pynccl.
|
||||
out = input_.clone()
|
||||
torch.distributed.all_reduce(out, group=self.device_group)
|
||||
torch.distributed.all_reduce(out, group=self.device_group, op=op)
|
||||
return out
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
|
||||
@@ -35,7 +35,7 @@ from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
from torch.distributed import Backend, ProcessGroup
|
||||
from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
@@ -260,7 +260,11 @@ class GroupCoordinator:
|
||||
with torch.cuda.stream(stream):
|
||||
yield graph_capture_context
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
def all_reduce(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
User-facing all-reduce function before we actually call the
|
||||
all-reduce operation.
|
||||
@@ -283,10 +287,14 @@ class GroupCoordinator:
|
||||
return torch.ops.vllm.all_reduce(input_,
|
||||
group_name=self.unique_name)
|
||||
else:
|
||||
return self._all_reduce_out_place(input_)
|
||||
return self._all_reduce_out_place(input_, op=op)
|
||||
|
||||
def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
return self.device_communicator.all_reduce(input_)
|
||||
def _all_reduce_out_place(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
|
||||
) -> torch.Tensor:
|
||||
return self.device_communicator.all_reduce(input_, op=op)
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
world_size = self.world_size
|
||||
@@ -647,7 +655,7 @@ class GroupCoordinator:
|
||||
tensor_dict[key] = value
|
||||
return tensor_dict
|
||||
|
||||
def barrier(self):
|
||||
def barrier(self) -> None:
|
||||
"""Barrier synchronization among the group.
|
||||
NOTE: don't use `device_group` here! `barrier` in NCCL is
|
||||
terrible because it is internally a broadcast operation with
|
||||
@@ -696,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=False,
|
||||
use_device_communicator=True,
|
||||
group_name="world",
|
||||
)
|
||||
|
||||
@@ -739,10 +747,10 @@ def set_custom_all_reduce(enable: bool):
|
||||
|
||||
|
||||
def init_distributed_environment(
|
||||
world_size: int = -1,
|
||||
rank: int = -1,
|
||||
world_size: int = 1,
|
||||
rank: int = 0,
|
||||
distributed_init_method: str = "env://",
|
||||
local_rank: int = -1,
|
||||
local_rank: int = 0,
|
||||
backend: str = "nccl",
|
||||
):
|
||||
logger.debug(
|
||||
@@ -786,9 +794,18 @@ 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:
|
||||
"""
|
||||
@@ -844,6 +861,22 @@ 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."""
|
||||
@@ -855,9 +888,20 @@ 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,
|
||||
@@ -868,7 +912,8 @@ 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, backend)
|
||||
sequence_model_parallel_size,
|
||||
data_parallel_size, backend)
|
||||
return
|
||||
|
||||
assert (
|
||||
@@ -887,7 +932,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
|
||||
return _TP is not None and _SP is not None and _DP is not None
|
||||
|
||||
|
||||
_TP_STATE_PATCHED = False
|
||||
@@ -940,6 +985,11 @@ def destroy_model_parallel() -> None:
|
||||
_SP.destroy()
|
||||
_SP = None
|
||||
|
||||
global _DP
|
||||
if _DP:
|
||||
_DP.destroy()
|
||||
_DP = None
|
||||
|
||||
|
||||
def destroy_distributed_environment() -> None:
|
||||
global _WORLD
|
||||
|
||||
@@ -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,7 +118,6 @@ 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,
|
||||
@@ -276,10 +275,10 @@ class VideoGenerator:
|
||||
|
||||
# Save video if requested
|
||||
if batch.save_video:
|
||||
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")
|
||||
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")
|
||||
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", video_path)
|
||||
else:
|
||||
@@ -295,6 +294,9 @@ 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.
|
||||
|
||||
@@ -44,6 +44,8 @@ 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
|
||||
@@ -55,6 +57,8 @@ 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"
|
||||
@@ -70,7 +74,7 @@ class FastVideoArgs:
|
||||
# Text encoder configuration
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = (
|
||||
"fp16",
|
||||
"fp16",
|
||||
# "fp16",
|
||||
)
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
|
||||
@@ -82,10 +86,19 @@ 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
|
||||
@@ -100,6 +113,10 @@ class FastVideoArgs:
|
||||
device_str: Optional[str] = None
|
||||
device = None
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
|
||||
def __post_init__(self):
|
||||
pass
|
||||
|
||||
@@ -132,6 +149,13 @@ class FastVideoArgs:
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--inference-mode",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.inference_mode,
|
||||
help="Whether to use inference mode",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
@@ -168,6 +192,20 @@ 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,
|
||||
@@ -243,6 +281,21 @@ 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,
|
||||
@@ -258,8 +311,16 @@ class FastVideoArgs:
|
||||
parser.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
help="Use CPU offload for the model load",
|
||||
help=
|
||||
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
|
||||
)
|
||||
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,
|
||||
@@ -321,6 +382,10 @@ 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
|
||||
@@ -332,10 +397,20 @@ 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)
|
||||
@@ -423,3 +498,344 @@ def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
# TODO(will): may need to handle this for CI.
|
||||
raise ValueError("Current fastvideo args is not set.")
|
||||
return _current_fastvideo_args
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TrainingArgs(FastVideoArgs):
|
||||
"""
|
||||
Training arguments. Inherits from FastVideoArgs and adds training-specific
|
||||
arguments. If there are any conflicts, the training arguments will take
|
||||
precedence.
|
||||
"""
|
||||
data_path: str = ""
|
||||
dataloader_num_workers: int = 0
|
||||
num_height: int = 0
|
||||
num_width: int = 0
|
||||
num_frames: int = 0
|
||||
|
||||
train_batch_size: int = 0
|
||||
num_latent_t: int = 0
|
||||
group_frame: bool = False
|
||||
group_resolution: bool = False
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
cache_dir: str = ""
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
cfg: float = 0.0
|
||||
precondition_outputs: bool = False
|
||||
|
||||
# validation & logs
|
||||
validation_prompt_dir: str = ""
|
||||
validation_sampling_steps: str = ""
|
||||
validation_guidance_scale: str = ""
|
||||
validation_steps: float = 0.0
|
||||
log_validation: bool = False
|
||||
tracker_project_name: str = ""
|
||||
seed: Optional[int] = None
|
||||
|
||||
# output
|
||||
output_dir: str = ""
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: bool = False
|
||||
logging_dir: str = ""
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
max_train_steps: int = 0
|
||||
gradient_accumulation_steps: int = 0
|
||||
learning_rate: float = 0.0
|
||||
scale_lr: bool = False
|
||||
lr_scheduler: str = ""
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
gradient_checkpointing: bool = False
|
||||
selective_checkpointing: float = 0.0
|
||||
allow_tf32: bool = False
|
||||
mixed_precision: str = ""
|
||||
train_sp_batch_size: int = 0
|
||||
fsdp_sharding_startegy: str = ""
|
||||
|
||||
weighting_scheme: str = ""
|
||||
logit_mean: float = 0.0
|
||||
logit_std: float = 1.0
|
||||
mode_scale: float = 0.0
|
||||
|
||||
num_euler_timesteps: int = 0
|
||||
lr_num_cycles: int = 0
|
||||
lr_power: float = 0.0
|
||||
not_apply_cfg_solver: bool = False
|
||||
distill_cfg: float = 0.0
|
||||
scheduler_type: str = ""
|
||||
linear_quadratic_threshold: float = 0.0
|
||||
linear_range: float = 0.0
|
||||
weight_decay: float = 0.0
|
||||
use_ema: bool = False
|
||||
multi_phased_distill_schedule: str = ""
|
||||
pred_decay_weight: float = 0.0
|
||||
pred_decay_type: str = ""
|
||||
hunyuan_teacher_disable_cfg: bool = False
|
||||
|
||||
# master_weight_type
|
||||
master_weight_type: str = ""
|
||||
|
||||
# For fast checking in LoRA pipeline
|
||||
training_mode: bool = True
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
|
||||
kwargs[attr] = args.data_parallel_size
|
||||
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
|
||||
kwargs[attr] = args.data_parallel_shards
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
parser.add_argument("--data-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to parquet files")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of workers for dataloader")
|
||||
parser.add_argument("--num-height",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of heights")
|
||||
parser.add_argument("--num-width",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of widths")
|
||||
parser.add_argument("--num-frames",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of frames")
|
||||
|
||||
# Training batch and model configuration
|
||||
parser.add_argument("--train-batch-size",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Training batch size")
|
||||
parser.add_argument("--num-latent-t",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of latent time steps")
|
||||
parser.add_argument("--group-frame",
|
||||
action=StoreBoolean,
|
||||
help="Whether to group frames during training")
|
||||
parser.add_argument("--group-resolution",
|
||||
action=StoreBoolean,
|
||||
help="Whether to group resolutions during training")
|
||||
|
||||
# Model paths
|
||||
parser.add_argument("--pretrained-model-name-or-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to pretrained model or model name")
|
||||
parser.add_argument("--dit-model-name-or-path",
|
||||
type=str,
|
||||
required=False,
|
||||
help="Path to DiT model or model name")
|
||||
parser.add_argument("--cache-dir",
|
||||
type=str,
|
||||
help="Directory to cache models")
|
||||
|
||||
# Diffusion settings
|
||||
parser.add_argument("--ema-decay",
|
||||
type=float,
|
||||
default=0.999,
|
||||
help="EMA decay rate")
|
||||
parser.add_argument("--ema-start-step",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Step to start EMA")
|
||||
parser.add_argument("--cfg",
|
||||
type=float,
|
||||
help="Classifier-free guidance scale")
|
||||
parser.add_argument(
|
||||
"--precondition-outputs",
|
||||
action=StoreBoolean,
|
||||
help="Whether to precondition the outputs of the model")
|
||||
|
||||
# Validation and logging
|
||||
parser.add_argument("--validation-prompt-dir",
|
||||
type=str,
|
||||
help="Directory containing validation prompts")
|
||||
parser.add_argument("--validation-sampling-steps",
|
||||
type=str,
|
||||
help="Validation sampling steps")
|
||||
parser.add_argument("--validation-guidance-scale",
|
||||
type=str,
|
||||
help="Validation guidance scale")
|
||||
parser.add_argument("--validation-steps",
|
||||
type=float,
|
||||
help="Number of validation steps")
|
||||
parser.add_argument("--log-validation",
|
||||
action=StoreBoolean,
|
||||
help="Whether to log validation results")
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
parser.add_argument("--seed",
|
||||
type=int,
|
||||
help="Seed for deterministic training")
|
||||
|
||||
# Output configuration
|
||||
parser.add_argument("--output-dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Output directory for checkpoints and logs")
|
||||
parser.add_argument("--checkpoints-total-limit",
|
||||
type=int,
|
||||
help="Maximum number of checkpoints to keep")
|
||||
parser.add_argument("--checkpointing-steps",
|
||||
type=int,
|
||||
help="Steps between checkpoints")
|
||||
parser.add_argument("--resume-from-checkpoint",
|
||||
type=str,
|
||||
help="Path to checkpoint to resume from")
|
||||
parser.add_argument("--logging-dir",
|
||||
type=str,
|
||||
help="Directory for logging")
|
||||
|
||||
# Training configuration
|
||||
parser.add_argument("--num-train-epochs",
|
||||
type=int,
|
||||
help="Number of training epochs")
|
||||
parser.add_argument("--max-train-steps",
|
||||
type=int,
|
||||
help="Maximum number of training steps")
|
||||
parser.add_argument("--gradient-accumulation-steps",
|
||||
type=int,
|
||||
help="Number of steps to accumulate gradients")
|
||||
parser.add_argument("--learning-rate",
|
||||
type=float,
|
||||
required=True,
|
||||
help="Learning rate")
|
||||
parser.add_argument("--scale-lr",
|
||||
action=StoreBoolean,
|
||||
help="Whether to scale learning rate")
|
||||
parser.add_argument("--lr-scheduler",
|
||||
type=str,
|
||||
default="constant",
|
||||
help="Learning rate scheduler type")
|
||||
parser.add_argument("--lr-warmup-steps",
|
||||
type=int,
|
||||
default=10,
|
||||
help="Number of warmup steps for learning rate")
|
||||
parser.add_argument("--max-grad-norm",
|
||||
type=float,
|
||||
help="Maximum gradient norm")
|
||||
parser.add_argument("--gradient-checkpointing",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use gradient checkpointing")
|
||||
parser.add_argument("--selective-checkpointing",
|
||||
type=float,
|
||||
help="Selective checkpointing threshold")
|
||||
parser.add_argument("--allow-tf32",
|
||||
action=StoreBoolean,
|
||||
help="Whether to allow TF32")
|
||||
parser.add_argument("--mixed-precision",
|
||||
type=str,
|
||||
help="Mixed precision training type")
|
||||
parser.add_argument("--train-sp-batch-size",
|
||||
type=int,
|
||||
help="Training spatial parallelism batch size")
|
||||
|
||||
parser.add_argument("--fsdp-sharding-strategy",
|
||||
type=str,
|
||||
help="FSDP sharding strategy")
|
||||
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
default="uniform",
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="mean to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_std",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="std to use when using the `'logit_normal'` weighting scheme.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode_scale",
|
||||
type=float,
|
||||
default=1.29,
|
||||
help=
|
||||
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
)
|
||||
|
||||
# Additional training parameters
|
||||
parser.add_argument("--num-euler-timesteps",
|
||||
type=int,
|
||||
help="Number of Euler timesteps")
|
||||
parser.add_argument("--lr-num-cycles",
|
||||
type=int,
|
||||
help="Number of learning rate cycles")
|
||||
parser.add_argument("--lr-power",
|
||||
type=float,
|
||||
help="Learning rate power")
|
||||
parser.add_argument("--not-apply-cfg-solver",
|
||||
action=StoreBoolean,
|
||||
help="Whether to not apply CFG solver")
|
||||
parser.add_argument("--distill-cfg",
|
||||
type=float,
|
||||
help="Distillation CFG scale")
|
||||
parser.add_argument("--scheduler-type", type=str, help="Scheduler type")
|
||||
parser.add_argument("--linear-quadratic-threshold",
|
||||
type=float,
|
||||
help="Linear quadratic threshold")
|
||||
parser.add_argument("--linear-range", type=float, help="Linear range")
|
||||
parser.add_argument("--weight-decay", type=float, help="Weight decay")
|
||||
parser.add_argument("--use-ema",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use EMA")
|
||||
parser.add_argument("--multi-phased-distill-schedule",
|
||||
type=str,
|
||||
help="Multi-phased distillation schedule")
|
||||
parser.add_argument("--pred-decay-weight",
|
||||
type=float,
|
||||
help="Prediction decay weight")
|
||||
parser.add_argument("--pred-decay-type",
|
||||
type=str,
|
||||
help="Prediction decay type")
|
||||
parser.add_argument("--hunyuan-teacher-disable-cfg",
|
||||
action=StoreBoolean,
|
||||
help="Whether to disable CFG for Hunyuan teacher")
|
||||
parser.add_argument("--master-weight-type",
|
||||
type=str,
|
||||
help="Master weight type")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -1,207 +0,0 @@
|
||||
# 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
|
||||
@@ -9,6 +9,19 @@ import torch.nn as nn
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
|
||||
|
||||
class FP32LayerNorm(nn.LayerNorm):
|
||||
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
origin_dtype = inputs.dtype
|
||||
return torch.nn.functional.layer_norm(
|
||||
inputs.float(),
|
||||
self.normalized_shape,
|
||||
self.weight.float() if self.weight is not None else None,
|
||||
self.bias.float() if self.bias is not None else None,
|
||||
self.eps,
|
||||
).to(origin_dtype)
|
||||
|
||||
|
||||
@CustomOp.register("rms_norm")
|
||||
class RMSNorm(CustomOp):
|
||||
"""Root mean square normalization.
|
||||
@@ -121,10 +134,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -144,9 +156,11 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
# Apply residual connection with gating
|
||||
residual_output = residual + x * gate
|
||||
# Apply normalization
|
||||
normalized = self.norm(residual_output)
|
||||
normalized = self.norm(residual_output.float()).to(
|
||||
residual_output.dtype)
|
||||
# Apply scale and shift
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
modulated = (normalized.float() * (1.0 + scale) + shift).to(
|
||||
residual_output.dtype)
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
@@ -171,15 +185,14 @@ class LayerNormScaleShift(nn.Module):
|
||||
has_weight=elementwise_affine,
|
||||
eps=eps)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
def forward(self, x: torch.Tensor, shift: torch.Tensor,
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
return normalized * (1.0 + scale) + shift
|
||||
normalized = self.norm(x.float()).to(x.dtype)
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
|
||||
@@ -0,0 +1,295 @@
|
||||
# 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
|
||||
@@ -33,9 +33,11 @@ class BaseDiT(nn.Module, ABC):
|
||||
f"Subclasses of BaseDiT must define '{attr}' class variable"
|
||||
)
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any],
|
||||
**kwargs) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.hf_config = hf_config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
@@ -76,6 +78,7 @@ 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,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -441,9 +441,10 @@ 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):
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
self.patch_size = [
|
||||
config.patch_size_t, config.patch_size, config.patch_size
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
# ==============================================================================
|
||||
from typing import Dict, Optional, Tuple
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
@@ -459,11 +459,13 @@ class StepVideoModel(BaseDiT):
|
||||
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
|
||||
]
|
||||
_param_names_mapping = StepVideoConfig()._param_names_mapping
|
||||
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
|
||||
_supported_attention_backends = StepVideoConfig(
|
||||
)._supported_attention_backends
|
||||
|
||||
def __init__(self, config: StepVideoConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.attention_head_dim = config.attention_head_dim
|
||||
self.in_channels = config.in_channels
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -13,8 +13,8 @@ from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
ScaleResidual,
|
||||
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
# from torch.nn import RMSNorm
|
||||
@@ -229,7 +229,8 @@ class WanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
# self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -298,7 +299,7 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
assert orig_dtype != torch.float32
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
@@ -319,6 +320,7 @@ 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,
|
||||
@@ -359,9 +361,11 @@ 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) -> None:
|
||||
super().__init__(config=config)
|
||||
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
|
||||
@@ -6,6 +6,7 @@ import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from copy import deepcopy
|
||||
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
@@ -14,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
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
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 load_fsdp_model
|
||||
from fastvideo.v1.models.loader.fsdp_load import maybe_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,
|
||||
@@ -366,6 +367,7 @@ class TransformerLoader(ComponentLoader):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
hf_config = deepcopy(config)
|
||||
cls_name = config.pop("_class_name")
|
||||
if cls_name is None:
|
||||
raise ValueError(
|
||||
@@ -389,16 +391,36 @@ class TransformerLoader(ComponentLoader):
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
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]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params={"config": dit_config},
|
||||
weight_dir_list=safetensors_list,
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
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)
|
||||
if fastvideo_args.enable_torch_compile:
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
|
||||
@@ -5,22 +5,25 @@
|
||||
# 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, Hashable, List,
|
||||
Optional, Tuple, Type)
|
||||
from typing import (Any, Callable, DefaultDict, Dict, Generator, List, Optional,
|
||||
Tuple, Type, Union)
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.distributed import DeviceMesh, init_device_mesh
|
||||
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
|
||||
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.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.utils import get_param_names_mapping
|
||||
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
|
||||
from fastvideo.v1.utils import set_mixed_precision_policy
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
@@ -51,66 +54,60 @@ 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 load_fsdp_model(
|
||||
def maybe_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,
|
||||
default_dtype: Optional[torch.dtype] = torch.bfloat16,
|
||||
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)
|
||||
|
||||
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",
|
||||
mesh_shape=(get_sequence_model_parallel_world_size(), ),
|
||||
mesh_dim_names=("dp", ),
|
||||
# (Replicate(), Shard(dim=0))
|
||||
mesh_shape=(dp_size, data_parallel_shards),
|
||||
mesh_dim_names=("dp", "sp"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
dp_mesh=device_mesh["dp"])
|
||||
mp_policy=mp_policy,
|
||||
mesh=device_mesh)
|
||||
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
load_fsdp_model_from_full_model_state_dict(
|
||||
load_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,
|
||||
@@ -119,8 +116,9 @@ def load_fsdp_model(
|
||||
if p.is_meta:
|
||||
raise RuntimeError(
|
||||
f"Unexpected param or buffer {n} on meta device.")
|
||||
for p in model.parameters():
|
||||
p.requires_grad = False
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
p.requires_grad = False
|
||||
|
||||
return model
|
||||
|
||||
|
||||
@@ -129,7 +127,9 @@ def shard_model(
|
||||
*,
|
||||
cpu_offload: bool,
|
||||
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,14 +156,17 @@ def shard_model(
|
||||
"""
|
||||
fsdp_kwargs = {
|
||||
"reshard_after_forward": reshard_after_forward,
|
||||
"mesh": dp_mesh
|
||||
"mesh": mesh,
|
||||
"mp_policy": mp_policy,
|
||||
}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
|
||||
|
||||
# Shard the model with FSDP, iterating in reverse to start with
|
||||
# iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
# TODO(will): don't reshard after forward for the last layer to save on the
|
||||
# all-gather that will immediately happen Shard the model with FSDP,
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
@@ -182,25 +185,28 @@ def shard_model(
|
||||
|
||||
|
||||
# TODO(PY): device mesh for cfg parallel
|
||||
def load_fsdp_model_from_full_model_state_dict(
|
||||
model: torch.nn.Module,
|
||||
def load_model_from_full_model_state_dict(
|
||||
model: Union[FSDPModule, 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
|
||||
and loading it into FSDP model (if training) or normal huggingface model
|
||||
Args:
|
||||
model (FSDPModule): Model to generate fully qualified names for cpu_state_dict
|
||||
model (Union[FSDPModule, torch.nn.Module]): 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 offload to CPU is enabled
|
||||
cpu_offload (bool): flag to check if FSDP offload 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
|
||||
@@ -209,19 +215,18 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
Raises:
|
||||
NotImplementedError: If got FSDP with more than 1D.
|
||||
"""
|
||||
meta_sharded_sd = model.state_dict()
|
||||
meta_sd = model.state_dict()
|
||||
|
||||
sharded_sd = {}
|
||||
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
|
||||
to_merge_params: DefaultDict[str, 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 dim=1 according to the merge_index order
|
||||
# cat at output dim according to the merge_index order
|
||||
sorted_tensors = [
|
||||
to_merge_params[target_param_name][i]
|
||||
for i in range(num_params_to_merge)
|
||||
@@ -231,24 +236,25 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
else:
|
||||
continue
|
||||
|
||||
sharded_meta_param = meta_sharded_sd.get(target_param_name)
|
||||
if sharded_meta_param is None:
|
||||
meta_sharded_param = meta_sd.get(target_param_name)
|
||||
if meta_sharded_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(sharded_meta_param, "device_mesh"):
|
||||
if not hasattr(meta_sharded_param, "device_mesh"):
|
||||
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
|
||||
# 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,
|
||||
sharded_meta_param.device_mesh,
|
||||
sharded_meta_param.placements,
|
||||
meta_sharded_param.device_mesh,
|
||||
meta_sharded_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)
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utilities for selecting and loading models."""
|
||||
import contextlib
|
||||
import re
|
||||
from typing import Any, Callable, Dict
|
||||
|
||||
import torch
|
||||
|
||||
@@ -16,3 +18,37 @@ 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
|
||||
@@ -0,0 +1,817 @@
|
||||
# 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
|
||||
@@ -5,9 +5,12 @@ 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,
|
||||
@@ -16,7 +19,12 @@ from fastvideo.v1.utils import (maybe_download_model,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
|
||||
class PipelineWithLoRA(LoRAPipeline, ComposedPipelineBase):
|
||||
"""Type for a pipeline that has both ComposedPipelineBase and LoRAPipeline functionality."""
|
||||
pass
|
||||
|
||||
|
||||
def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
|
||||
"""
|
||||
Only works with valid hf diffusers configs. (model_index.json)
|
||||
We want to build a pipeline based on the inference args mode_path:
|
||||
@@ -45,7 +53,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
|
||||
logger.info("Pipeline instantiated")
|
||||
|
||||
# pipeline is now initialized and ready to use
|
||||
return pipeline
|
||||
return cast(PipelineWithLoRA, pipeline)
|
||||
|
||||
|
||||
__all__ = [
|
||||
@@ -54,4 +62,5 @@ __all__ = [
|
||||
"ComposedPipelineBase",
|
||||
"PipelineRegistry",
|
||||
"ForwardBatch",
|
||||
"LoRAPipeline",
|
||||
]
|
||||
|
||||
@@ -5,19 +5,25 @@ Base class for composed pipelines.
|
||||
This module defines the base class for pipelines that are composed of multiple stages.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from copy import deepcopy
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.pipelines import (PipelineConfig,
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages import PipelineStage
|
||||
from fastvideo.v1.utils import (maybe_download_model,
|
||||
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -34,20 +40,36 @@ 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,
|
||||
model_path: str,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
config: Optional[Dict[str, Any]] = None):
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
required_config_modules: Optional[List[str]] = None):
|
||||
"""
|
||||
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.model_path = model_path
|
||||
self._stages: List[PipelineStage] = []
|
||||
self._stage_name_mapping: Dict[str, PipelineStage] = {}
|
||||
|
||||
if required_config_modules is not None:
|
||||
self._required_config_modules = required_config_modules
|
||||
|
||||
if self._required_config_modules is None:
|
||||
raise NotImplementedError(
|
||||
"Subclass must set _required_config_modules")
|
||||
@@ -59,16 +81,134 @@ class ComposedPipelineBase(ABC):
|
||||
else:
|
||||
self.config = config
|
||||
|
||||
self.maybe_init_distributed_environment(fastvideo_args)
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
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)
|
||||
|
||||
self.initialize_pipeline(fastvideo_args)
|
||||
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
if not fastvideo_args.training_mode:
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
|
||||
def get_module(self, module_name: str) -> Any:
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"if training_mode is True, the pipeline must implement this method")
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"if log_validation is True, the pipeline must implement this method"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls,
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
args: Optional[argparse.Namespace] = None,
|
||||
required_config_modules: Optional[List[str]] = None,
|
||||
**kwargs) -> "ComposedPipelineBase":
|
||||
config = None
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
if isinstance(pipeline_config, str):
|
||||
config.load_from_json(pipeline_config)
|
||||
|
||||
# 2. If users also provide some kwargs, it will override the pipeline config.
|
||||
# The user kwargs shouldn't contain model config parameters!
|
||||
if config is None:
|
||||
logger.warning("No config found for model %s, using default config",
|
||||
model_path)
|
||||
config_args = kwargs
|
||||
else:
|
||||
config_args = shallow_asdict(config)
|
||||
config_args.update(kwargs)
|
||||
|
||||
if args is None or args.inference_mode:
|
||||
fastvideo_args = FastVideoArgs(model_path=model_path,
|
||||
device_str=device or "cuda" if
|
||||
torch.cuda.is_available() else "cpu",
|
||||
**config_args)
|
||||
|
||||
fastvideo_args.model_path = model_path
|
||||
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
|
||||
) else "cpu"
|
||||
for key, value in config_args.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
else:
|
||||
assert args is not None, "args must be provided for training mode"
|
||||
fastvideo_args = TrainingArgs.from_cli_args(args)
|
||||
# TODO(will): fix this so that its not so ugly
|
||||
fastvideo_args.model_path = model_path
|
||||
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
|
||||
) else "cpu"
|
||||
for key, value in config_args.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
|
||||
fastvideo_args.num_gpus = int(os.environ.get("WORLD_SIZE", 1))
|
||||
fastvideo_args.use_cpu_offload = False
|
||||
# make sure we are in training mode
|
||||
fastvideo_args.inference_mode = False
|
||||
# we hijack the precision to be the master weight type so that the
|
||||
# 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.check_fastvideo_args()
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
return cls(model_path,
|
||||
fastvideo_args,
|
||||
required_config_modules=required_config_modules)
|
||||
|
||||
def maybe_init_distributed_environment(self, fastvideo_args: FastVideoArgs):
|
||||
if model_parallel_is_initialized():
|
||||
return
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", -1))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", -1))
|
||||
rank = int(os.environ.get("RANK", -1))
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
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)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
fastvideo_args.device = device
|
||||
|
||||
def get_module(self, module_name: str, default_value: Any = None) -> Any:
|
||||
if module_name not in self.modules:
|
||||
return default_value
|
||||
return self.modules[module_name]
|
||||
|
||||
def add_module(self, module_name: str, module: Any):
|
||||
@@ -110,7 +250,13 @@ class ComposedPipelineBase(ABC):
|
||||
@abstractmethod
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Create the pipeline stages.
|
||||
Create the inference pipeline stages.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
Create the training pipeline stages.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -136,19 +282,21 @@ class ComposedPipelineBase(ABC):
|
||||
modules_config
|
||||
) > 1, "model_index.json must contain at least one pipeline module"
|
||||
|
||||
required_modules = [
|
||||
"vae", "text_encoder", "transformer", "scheduler", "tokenizer"
|
||||
]
|
||||
for module_name in required_modules:
|
||||
for module_name in self.required_config_modules:
|
||||
if module_name not in modules_config:
|
||||
raise ValueError(
|
||||
f"model_index.json must contain a {module_name} module")
|
||||
logger.info("Diffusers config passed sanity checks")
|
||||
|
||||
# all the component models used by the pipeline
|
||||
required_modules = self.required_config_modules
|
||||
logger.info("Loading required modules: %s", required_modules)
|
||||
|
||||
modules = {}
|
||||
for module_name, (transformers_or_diffusers,
|
||||
architecture) in modules_config.items():
|
||||
if module_name not in required_modules:
|
||||
logger.info("Skipping module %s", module_name)
|
||||
continue
|
||||
component_model_path = os.path.join(self.model_path, module_name)
|
||||
module = PipelineComponentLoader.load_module(
|
||||
module_name=module_name,
|
||||
@@ -164,7 +312,6 @@ class ComposedPipelineBase(ABC):
|
||||
logger.warning("Overwriting module %s", module_name)
|
||||
modules[module_name] = module
|
||||
|
||||
required_modules = self.required_config_modules
|
||||
# Check if all required modules were loaded
|
||||
for module_name in required_modules:
|
||||
if module_name not in modules or modules[module_name] is None:
|
||||
@@ -198,7 +345,7 @@ class ComposedPipelineBase(ABC):
|
||||
# Execute each stage
|
||||
logger.info("Running pipeline stages: %s",
|
||||
self._stage_name_mapping.keys())
|
||||
logger.info("Batch: %s", batch)
|
||||
# logger.info("Batch: %s", batch)
|
||||
for stage in self.stages:
|
||||
batch = stage(batch, fastvideo_args)
|
||||
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
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
|
||||
@@ -114,10 +114,17 @@ 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 = []
|
||||
|
||||
@@ -6,10 +6,11 @@ import importlib
|
||||
import pkgutil
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
from typing import AbstractSet, Dict, Optional, Tuple, Type
|
||||
from typing import AbstractSet, Dict, Optional, Tuple, Type, Union
|
||||
|
||||
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__)
|
||||
|
||||
@@ -33,7 +34,7 @@ class _PipelineRegistry:
|
||||
def resolve_pipeline_cls(
|
||||
self,
|
||||
architecture: str,
|
||||
) -> Tuple[Type[ComposedPipelineBase], str]:
|
||||
) -> Tuple[Union[Type[ComposedPipelineBase], Type[LoRAPipeline]], str]:
|
||||
if not architecture:
|
||||
logger.warning("No pipeline architecture is specified")
|
||||
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
# 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
|
||||
@@ -0,0 +1,23 @@
|
||||
# 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,539 @@
|
||||
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
|
||||
@@ -23,7 +23,6 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
st_attn_available = False
|
||||
spec = importlib.util.find_spec("st_attn")
|
||||
@@ -48,6 +47,15 @@ class DenoisingStage(PipelineStage):
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
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,
|
||||
@@ -64,6 +72,8 @@ class DenoisingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
self.transformer.to(fastvideo_args.device)
|
||||
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
@@ -74,7 +84,9 @@ class DenoisingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
# 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
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
@@ -150,6 +162,10 @@ 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
|
||||
@@ -191,15 +207,6 @@ class DenoisingStage(PipelineStage):
|
||||
dtype=target_dtype,
|
||||
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(
|
||||
)
|
||||
@@ -221,6 +228,7 @@ 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,
|
||||
@@ -239,6 +247,7 @@ 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,
|
||||
@@ -288,6 +297,10 @@ 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()
|
||||
@@ -359,3 +372,142 @@ 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]}/'
|
||||
)
|
||||
|
||||
@@ -63,10 +63,15 @@ class TextEncodingStage(PipelineStage):
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
text_encoder = text_encoder.to(fastvideo_args.device)
|
||||
|
||||
assert isinstance(batch.prompt, str)
|
||||
text = preprocess_func(batch.prompt)
|
||||
text_inputs = tokenizer(text, **encoder_config.tokenizer_kwargs).to(
|
||||
fastvideo_args.device)
|
||||
assert isinstance(batch.prompt, (str, list))
|
||||
if isinstance(batch.prompt, str):
|
||||
batch.prompt = [batch.prompt]
|
||||
texts = []
|
||||
for prompt_str in batch.prompt:
|
||||
texts.append(preprocess_func(prompt_str))
|
||||
text_inputs = tokenizer(texts,
|
||||
**encoder_config.tokenizer_kwargs).to(
|
||||
fastvideo_args.device)
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -78,6 +83,8 @@ class TextEncodingStage(PipelineStage):
|
||||
prompt_embeds = postprocess_func(outputs)
|
||||
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
if batch.prompt_attention_mask is not None:
|
||||
batch.prompt_attention_mask.append(attention_mask)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
assert isinstance(batch.negative_prompt, str)
|
||||
@@ -98,6 +105,9 @@ class TextEncodingStage(PipelineStage):
|
||||
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_prompt_embeds)
|
||||
if batch.negative_attention_mask is not None:
|
||||
batch.negative_attention_mask.append(
|
||||
negative_attention_mask)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
text_encoder.to('cpu')
|
||||
|
||||
@@ -20,6 +20,7 @@ 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,
|
||||
@@ -29,7 +30,7 @@ from fastvideo.v1.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class StepVideoPipeline(ComposedPipelineBase):
|
||||
class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = ["transformer", "scheduler", "vae"]
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ 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 (
|
||||
@@ -16,17 +17,23 @@ 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(ComposedPipelineBase):
|
||||
class WanImageToVideoPipeline(LoRAPipeline, 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."""
|
||||
|
||||
|
||||
@@ -8,25 +8,33 @@ 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.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanPipeline(ComposedPipelineBase):
|
||||
class WanPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
Wan video diffusion pipeline with LoRA support.
|
||||
"""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -48,7 +56,37 @@ class WanPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
transformer=self.get_module("transformer", None)))
|
||||
|
||||
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")))
|
||||
|
||||
|
||||
class WanValidationPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
|
||||
"""
|
||||
_required_config_modules = ["vae", "scheduler"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
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", None)))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
|
||||
@@ -1,77 +0,0 @@
|
||||
# type: ignore
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, prepare_fastvideo_args
|
||||
# Fix the import path
|
||||
from fastvideo.v1.inference_engine import InferenceEngine
|
||||
|
||||
|
||||
def initialize_distributed_and_parallelism(fastvideo_args: FastVideoArgs):
|
||||
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(local_rank)
|
||||
init_distributed_environment(world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
device_str = f"cuda:{local_rank}"
|
||||
fastvideo_args.device_str = device_str
|
||||
fastvideo_args.device = torch.device(device_str)
|
||||
assert fastvideo_args.sp_size is not None
|
||||
assert fastvideo_args.tp_size is not None
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=fastvideo_args.sp_size,
|
||||
tensor_model_parallel_size=fastvideo_args.tp_size,
|
||||
)
|
||||
|
||||
|
||||
def main(fastvideo_args: FastVideoArgs):
|
||||
initialize_distributed_and_parallelism(fastvideo_args)
|
||||
engine = InferenceEngine.create_engine(fastvideo_args, )
|
||||
|
||||
if fastvideo_args.prompt_path is not None:
|
||||
with open(fastvideo_args.prompt_path) as f:
|
||||
prompts = [line.strip() for line in f.readlines()]
|
||||
else:
|
||||
if fastvideo_args.prompt is None:
|
||||
raise ValueError("prompt or prompt_path is required")
|
||||
prompts = [fastvideo_args.prompt]
|
||||
|
||||
# Process each prompt
|
||||
for prompt in prompts:
|
||||
outputs = engine.run(
|
||||
prompt=prompt,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
|
||||
# Process outputs
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
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))
|
||||
|
||||
# Save video
|
||||
os.makedirs(os.path.dirname(fastvideo_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(fastvideo_args.output_path,
|
||||
f"{prompt[:100]}.mp4"),
|
||||
frames,
|
||||
fps=fastvideo_args.fps)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
main(fastvideo_args)
|
||||
@@ -1,20 +0,0 @@
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
import os
|
||||
|
||||
|
||||
def main():
|
||||
print(os.environ["RANK"])
|
||||
print(os.environ["WORLD_SIZE"])
|
||||
print(os.environ["LOCAL_RANK"])
|
||||
print(os.environ["MASTER_ADDR"])
|
||||
print(os.environ["MASTER_PORT"])
|
||||
print(os.environ["RANK"])
|
||||
print(os.environ["WORLD_SIZE"])
|
||||
print(os.environ["LOCAL_RANK"])
|
||||
print(os.environ["MASTER_ADDR"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -257,6 +257,7 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
"flow_shift": BASE_PARAMS["flow_shift"],
|
||||
"sp_size": BASE_PARAMS["sp_size"],
|
||||
"tp_size": BASE_PARAMS["tp_size"],
|
||||
"use_cpu_offload": True,
|
||||
}
|
||||
if BASE_PARAMS.get("vae_sp"):
|
||||
init_kwargs["vae_sp"] = True
|
||||
|
||||
@@ -65,6 +65,7 @@ def test_hunyuanvideo_distributed():
|
||||
precision=precision_str)
|
||||
args.device = torch.device(f"cuda:{LOCAL_RANK}")
|
||||
args.dit_config = HunyuanVideoConfig()
|
||||
args.check_fastvideo_args()
|
||||
|
||||
loader = TransformerLoader()
|
||||
model = loader.load(TRANSFORMER_PATH, "", args)
|
||||
|
||||
@@ -37,6 +37,7 @@ def test_wan_transformer():
|
||||
precision=precision_str)
|
||||
args.device = device
|
||||
args.dit_config = WanVideoConfig()
|
||||
args.check_fastvideo_args()
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
from .training_pipeline import TrainingPipeline
|
||||
from .wan_training_pipeline import WanTrainingPipeline
|
||||
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
|
||||
@@ -0,0 +1,107 @@
|
||||
import random
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed.checkpoint.stateful
|
||||
from torch.distributed.checkpoint.state_dict import (StateDictOptions,
|
||||
get_model_state_dict,
|
||||
get_optimizer_state_dict,
|
||||
set_model_state_dict,
|
||||
set_optimizer_state_dict)
|
||||
|
||||
|
||||
class ModelWrapper(torch.distributed.checkpoint.stateful.Stateful):
|
||||
|
||||
def __init__(self, model: torch.nn.Module) -> None:
|
||||
self.model = model
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
return get_model_state_dict(self.model) # type: ignore[no-any-return]
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
set_model_state_dict(
|
||||
self.model,
|
||||
model_state_dict=state_dict,
|
||||
options=StateDictOptions(strict=False),
|
||||
)
|
||||
|
||||
|
||||
class OptimizerWrapper(torch.distributed.checkpoint.stateful.Stateful):
|
||||
|
||||
def __init__(self, model: torch.nn.Module,
|
||||
optimizer: torch.optim.Optimizer) -> None:
|
||||
self.model = model
|
||||
self.optimizer = optimizer
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
return get_optimizer_state_dict( # type: ignore[no-any-return]
|
||||
self.model,
|
||||
self.optimizer,
|
||||
options=StateDictOptions(flatten_optimizer_state_dict=True),
|
||||
)
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
set_optimizer_state_dict(
|
||||
self.model,
|
||||
self.optimizer,
|
||||
optim_state_dict=state_dict,
|
||||
options=StateDictOptions(flatten_optimizer_state_dict=True),
|
||||
)
|
||||
|
||||
|
||||
class SchedulerWrapper(torch.distributed.checkpoint.stateful.Stateful):
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
self.scheduler = scheduler
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
return {"scheduler": self.scheduler.state_dict()}
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
self.scheduler.load_state_dict(state_dict["scheduler"])
|
||||
|
||||
|
||||
class RandomStateWrapper(torch.distributed.checkpoint.stateful.Stateful):
|
||||
|
||||
def __init__(self,
|
||||
noise_generator: Optional[torch.Generator] = None) -> None:
|
||||
self.noise_generator = noise_generator
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
state = {
|
||||
"torch_rng_state": torch.get_rng_state(),
|
||||
"numpy_rng_state": np.random.get_state(),
|
||||
"python_rng_state": random.getstate(),
|
||||
}
|
||||
|
||||
if torch.cuda.is_available():
|
||||
state["cuda_rng_state"] = torch.cuda.get_rng_state()
|
||||
if torch.cuda.device_count() > 1:
|
||||
state["cuda_rng_state_all"] = torch.cuda.get_rng_state_all()
|
||||
|
||||
if self.noise_generator is not None:
|
||||
state["noise_generator_state"] = self.noise_generator.get_state()
|
||||
|
||||
return state
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
if "torch_rng_state" in state_dict:
|
||||
torch.set_rng_state(state_dict["torch_rng_state"])
|
||||
|
||||
if "numpy_rng_state" in state_dict:
|
||||
np.random.set_state(state_dict["numpy_rng_state"])
|
||||
|
||||
if "python_rng_state" in state_dict:
|
||||
random.setstate(state_dict["python_rng_state"])
|
||||
|
||||
# Restore CUDA random state
|
||||
if torch.cuda.is_available():
|
||||
if "cuda_rng_state" in state_dict:
|
||||
torch.cuda.set_rng_state(state_dict["cuda_rng_state"])
|
||||
if "cuda_rng_state_all" in state_dict:
|
||||
torch.cuda.set_rng_state_all(state_dict["cuda_rng_state_all"])
|
||||
|
||||
# Restore noise generator state
|
||||
if "noise_generator_state" in state_dict and self.noise_generator is not None:
|
||||
self.noise_generator.set_state(state_dict["noise_generator_state"])
|
||||
@@ -0,0 +1,574 @@
|
||||
import gc
|
||||
import math
|
||||
import os
|
||||
import traceback
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Iterator
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torchvision
|
||||
from diffusers.optimization import get_scheduler
|
||||
from einops import rearrange
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
|
||||
from fastvideo.v1.distributed import get_sp_group, get_world_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Note: if checking with float32, cannot use flash-attn.
|
||||
GRADIENT_CHECK_DTYPE = torch.bfloat16
|
||||
|
||||
|
||||
class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
"""
|
||||
A pipeline for training a model. All training pipelines should inherit from this class.
|
||||
All reusable components and code should be implemented in this class.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer"]
|
||||
validation_pipeline: ComposedPipelineBase
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
|
||||
Dict[str, Any]]]
|
||||
current_epoch: int = 0
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
raise RuntimeError(
|
||||
"create_pipeline_stages should not be called for training pipeline")
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing training pipeline...")
|
||||
self.device = training_args.device
|
||||
world_group = get_world_group()
|
||||
self.world_size = world_group.world_size
|
||||
self.rank = world_group.rank
|
||||
self.sp_group = get_sp_group()
|
||||
self.rank_in_sp_group = self.sp_group.rank_in_group
|
||||
self.sp_world_size = self.sp_group.world_size
|
||||
self.local_rank = world_group.local_rank
|
||||
self.transformer = self.get_module("transformer")
|
||||
assert self.transformer is not None
|
||||
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
params_to_optimize = self.transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
self.init_steps = 0
|
||||
logger.info("optimizer: %s", self.optimizer)
|
||||
|
||||
self.lr_scheduler = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer,
|
||||
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
|
||||
num_training_steps=training_args.max_train_steps * self.world_size,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset = ParquetVideoTextDataset(
|
||||
training_args.data_path,
|
||||
batch_size=training_args.train_batch_size,
|
||||
rank=self.rank,
|
||||
world_size=self.world_size,
|
||||
cfg_rate=training_args.cfg,
|
||||
num_latent_t=training_args.num_latent_t)
|
||||
|
||||
self.train_dataloader = StatefulDataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=training_args.train_batch_size,
|
||||
num_workers=training_args.
|
||||
dataloader_num_workers, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
|
||||
drop_last=True)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
|
||||
assert training_args.gradient_accumulation_steps is not None
|
||||
assert training_args.sp_size is not None
|
||||
assert training_args.train_sp_batch_size is not None
|
||||
assert training_args.max_train_steps is not None
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
training_args.train_sp_batch_size)
|
||||
self.num_train_epochs = math.ceil(training_args.max_train_steps /
|
||||
self.num_update_steps_per_epoch)
|
||||
|
||||
# TODO(will): is there a cleaner way to track epochs?
|
||||
self.current_epoch = 0
|
||||
|
||||
if self.rank == 0:
|
||||
project = training_args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=training_args)
|
||||
|
||||
@abstractmethod
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"Training pipelines must implement this method")
|
||||
|
||||
@abstractmethod
|
||||
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
|
||||
loader, noise_scheduler, noise_random_generator,
|
||||
gradient_accumulation_steps, sp_size,
|
||||
precondition_outputs, max_grad_norm, weighting_scheme,
|
||||
logit_mean, logit_std, mode_scale):
|
||||
"""
|
||||
Train one step of the model.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"Training pipeline must implement this method")
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
assert training_args is not None
|
||||
training_args.inference_mode = True
|
||||
training_args.use_cpu_offload = False
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
raise ValueError("Validation pipeline is not set")
|
||||
|
||||
logger.info("Starting validation")
|
||||
|
||||
# Create sampling parameters if not provided
|
||||
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
|
||||
|
||||
# Set deterministic seed for validation
|
||||
validation_seed = training_args.seed if training_args.seed is not None else 42
|
||||
torch.manual_seed(validation_seed)
|
||||
torch.cuda.manual_seed_all(validation_seed)
|
||||
|
||||
logger.info("Using validation seed: %s", validation_seed)
|
||||
|
||||
# Prepare validation prompts
|
||||
logger.info('fastvideo_args.validation_prompt_dir: %s',
|
||||
training_args.validation_prompt_dir)
|
||||
validation_dataset = ParquetVideoTextDataset(
|
||||
training_args.validation_prompt_dir,
|
||||
batch_size=1,
|
||||
rank=self.rank,
|
||||
world_size=self.world_size,
|
||||
cfg_rate=training_args.cfg,
|
||||
num_latent_t=training_args.num_latent_t,
|
||||
validation=True)
|
||||
if sampling_param.negative_prompt:
|
||||
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
|
||||
)
|
||||
|
||||
validation_dataloader = StatefulDataLoader(
|
||||
validation_dataset,
|
||||
batch_size=1,
|
||||
num_workers=5, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
|
||||
drop_last=False)
|
||||
|
||||
transformer.eval()
|
||||
|
||||
# Add the transformer to the validation pipeline
|
||||
self.validation_pipeline.add_module("transformer", transformer)
|
||||
# TODO(Peiyuan): those logic should be inside add_module
|
||||
self.validation_pipeline.latent_preparation_stage.transformer = transformer # type: ignore[attr-defined]
|
||||
self.validation_pipeline.denoising_stage.transformer = transformer # type: ignore[attr-defined]
|
||||
|
||||
# Process each validation prompt
|
||||
videos = []
|
||||
captions = []
|
||||
for _, embeddings, masks, infos in validation_dataloader:
|
||||
caption = infos['caption']
|
||||
captions.extend(caption)
|
||||
print(f"rank {self.rank} is running validation")
|
||||
print(f"rank {self.rank} file_name: {infos['file_name']}")
|
||||
prompt_embeds = embeddings.to(training_args.device)
|
||||
prompt_attention_mask = masks.to(training_args.device)
|
||||
|
||||
# Calculate sizes
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8,
|
||||
sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
temporal_compression_factor = training_args.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
|
||||
logger.info(f"rank {self.rank} num_frames: {num_frames}")
|
||||
|
||||
# Prepare batch for validation
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
latents=None,
|
||||
seed=validation_seed, # Use deterministic seed
|
||||
prompt_embeds=[prompt_embeds],
|
||||
prompt_attention_mask=[prompt_attention_mask],
|
||||
negative_prompt_embeds=[negative_prompt_embeds],
|
||||
negative_attention_mask=[negative_prompt_attention_mask],
|
||||
# make sure we use the same height, width, and num_frames as the training pipeline
|
||||
height=training_args.num_height,
|
||||
width=training_args.num_width,
|
||||
num_frames=num_frames,
|
||||
# TODO(will): validation_sampling_steps and
|
||||
# validation_guidance_scale are actually passed in as a list of
|
||||
# values, like "10,20,30". The validation should be run for each
|
||||
# combination of values.
|
||||
# num_inference_steps=fastvideo_args.validation_sampling_steps,
|
||||
num_inference_steps=sampling_param.num_inference_steps,
|
||||
# guidance_scale=fastvideo_args.validation_guidance_scale,
|
||||
guidance_scale=sampling_param.guidance_scale,
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
)
|
||||
|
||||
# Run validation inference
|
||||
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
output_batch = self.validation_pipeline.forward(
|
||||
batch, training_args)
|
||||
samples = output_batch.output
|
||||
|
||||
# Re-enable gradients for training
|
||||
transformer.requires_grad_(True)
|
||||
transformer.train()
|
||||
|
||||
if self.rank_in_sp_group != 0:
|
||||
continue
|
||||
|
||||
# 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
|
||||
world_group = get_world_group()
|
||||
num_sp_groups = world_group.world_size // self.sp_group.world_size
|
||||
|
||||
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
|
||||
# results to global rank 0
|
||||
if self.rank_in_sp_group == 0:
|
||||
if self.rank == 0:
|
||||
# Global rank 0 collects results from all sp_group leaders
|
||||
all_videos = videos # Start with own results
|
||||
all_captions = captions
|
||||
|
||||
# Receive from other sp_group leaders
|
||||
for sp_group_idx in range(1, num_sp_groups):
|
||||
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
|
||||
recv_videos = world_group.recv_object(src=src_rank)
|
||||
recv_captions = world_group.recv_object(src=src_rank)
|
||||
all_videos.extend(recv_videos)
|
||||
all_captions.extend(recv_captions)
|
||||
|
||||
video_filenames = []
|
||||
for i, (video,
|
||||
caption) in enumerate(zip(all_videos, all_captions)):
|
||||
os.makedirs(training_args.output_dir, exist_ok=True)
|
||||
filename = os.path.join(
|
||||
training_args.output_dir,
|
||||
f"validation_step_{global_step}_video_{i}.mp4")
|
||||
imageio.mimsave(filename, video, fps=sampling_param.fps)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
"validation_videos": [
|
||||
wandb.Video(filename, caption=caption) for filename,
|
||||
caption in zip(video_filenames, all_captions)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
else:
|
||||
# Other sp_group leaders send their results to global rank 0
|
||||
world_group.send_object(videos, dst=0)
|
||||
world_group.send_object(captions, dst=0)
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def gradient_check_parameters(self,
|
||||
transformer,
|
||||
latents,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
timesteps,
|
||||
target,
|
||||
eps=5e-2,
|
||||
max_params_to_check=2000) -> float:
|
||||
"""
|
||||
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
|
||||
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
|
||||
"""
|
||||
assert self.training_args is not None
|
||||
# Move all inputs to CPU and clear GPU memory
|
||||
inputs_cpu = {
|
||||
'latents': latents.cpu(),
|
||||
'encoder_hidden_states': encoder_hidden_states.cpu(),
|
||||
'encoder_attention_mask': encoder_attention_mask.cpu(),
|
||||
'timesteps': timesteps.cpu(),
|
||||
'target': target.cpu()
|
||||
}
|
||||
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def compute_loss() -> torch.Tensor:
|
||||
assert self.training_args is not None
|
||||
# Move inputs to GPU, compute loss, cleanup
|
||||
inputs_gpu = {
|
||||
k:
|
||||
v.to(self.training_args.device,
|
||||
dtype=GRADIENT_CHECK_DTYPE
|
||||
if k != 'encoder_attention_mask' else None)
|
||||
for k, v in inputs_cpu.items()
|
||||
}
|
||||
|
||||
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
|
||||
# with torch.autocast(enabled=False, device_type="cuda"):
|
||||
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
|
||||
with set_forward_context(
|
||||
current_timestep=inputs_gpu['timesteps'],
|
||||
attn_metadata=None):
|
||||
model_pred = transformer(
|
||||
hidden_states=inputs_gpu['latents'],
|
||||
encoder_hidden_states=inputs_gpu[
|
||||
'encoder_hidden_states'],
|
||||
timestep=inputs_gpu['timesteps'],
|
||||
encoder_attention_mask=inputs_gpu[
|
||||
'encoder_attention_mask'],
|
||||
return_dict=False)[0]
|
||||
|
||||
if self.training_args.precondition_outputs:
|
||||
sigmas = get_sigmas(self.noise_scheduler,
|
||||
inputs_gpu['latents'].device,
|
||||
inputs_gpu['timesteps'],
|
||||
n_dim=inputs_gpu['latents'].ndim,
|
||||
dtype=inputs_gpu['latents'].dtype)
|
||||
model_pred = inputs_gpu['latents'] - model_pred * sigmas
|
||||
target_adjusted = inputs_gpu['target']
|
||||
else:
|
||||
target_adjusted = inputs_gpu['target']
|
||||
|
||||
loss = torch.mean((model_pred - target_adjusted)**2)
|
||||
|
||||
# Cleanup and return
|
||||
loss_cpu = loss.cpu()
|
||||
del inputs_gpu, model_pred, target_adjusted
|
||||
if 'sigmas' in locals():
|
||||
del sigmas
|
||||
torch.cuda.empty_cache()
|
||||
return loss_cpu.to(self.training_args.device)
|
||||
|
||||
try:
|
||||
# Get analytical gradients
|
||||
transformer.zero_grad()
|
||||
analytical_loss = compute_loss()
|
||||
analytical_loss.backward()
|
||||
|
||||
# Check gradients for selected parameters
|
||||
absolute_errors: list[float] = []
|
||||
param_count = 0
|
||||
|
||||
rank = dist.get_rank()
|
||||
sp_group = get_sp_group()
|
||||
for name, param in transformer.named_parameters():
|
||||
sp_group.barrier()
|
||||
# skip scale_shift_table because it is not sharded
|
||||
if 'scale_shift_table' in name:
|
||||
continue
|
||||
if isinstance(param.grad, torch.distributed.tensor.DTensor):
|
||||
full_grad = param.grad.full_tensor()
|
||||
distributed = True
|
||||
else:
|
||||
full_grad = param.grad
|
||||
distributed = False
|
||||
continue
|
||||
if not (param.requires_grad and param.grad is not None
|
||||
and param_count < max_params_to_check
|
||||
and full_grad.abs().max() > 5e-4):
|
||||
continue
|
||||
if not distributed and rank != 0:
|
||||
continue
|
||||
|
||||
# Get local parameter and gradient tensors
|
||||
local_param = param._local_tensor if hasattr(
|
||||
param, '_local_tensor') else param
|
||||
local_grad = param.grad._local_tensor if hasattr(
|
||||
param.grad, '_local_tensor') else param.grad
|
||||
|
||||
# Find first significant gradient element
|
||||
flat_param = local_param.data.view(-1)
|
||||
flat_grad = local_grad.view(-1)
|
||||
check_idx = next((i for i in range(min(10, flat_param.numel()))
|
||||
if abs(flat_grad[i]) > 1e-4), 0)
|
||||
|
||||
# Store original values
|
||||
orig_value = flat_param[check_idx].item()
|
||||
analytical_grad = flat_grad[check_idx].item()
|
||||
|
||||
# Compute numerical gradient
|
||||
for delta in [eps, -eps]:
|
||||
with torch.no_grad():
|
||||
# only have a single rank modify the parameter
|
||||
# because we are using FSDP
|
||||
if rank == 0:
|
||||
flat_param[check_idx] = orig_value + delta
|
||||
loss = compute_loss()
|
||||
if delta > 0:
|
||||
loss_plus = loss.item()
|
||||
else:
|
||||
loss_minus = loss.item()
|
||||
|
||||
# Restore parameter and compute error
|
||||
with torch.no_grad():
|
||||
flat_param[check_idx] = orig_value
|
||||
|
||||
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
|
||||
abs_error = abs(analytical_grad - numerical_grad)
|
||||
rel_error = abs_error / max(abs(analytical_grad),
|
||||
abs(numerical_grad), 1e-3)
|
||||
absolute_errors.append(abs_error)
|
||||
|
||||
if self.rank == 0:
|
||||
logger.info(
|
||||
"%s[%s]: analytical=%.5f, numerical=%.5f, abs_error=%.2e, rel_error=%.2f%%",
|
||||
name, check_idx, analytical_grad, numerical_grad,
|
||||
abs_error, rel_error * 100)
|
||||
|
||||
# param_count += 1
|
||||
|
||||
# Compute and log statistics
|
||||
if rank == 0 and absolute_errors:
|
||||
min_err, max_err, mean_err = min(absolute_errors), max(
|
||||
absolute_errors
|
||||
), sum(absolute_errors) / len(absolute_errors)
|
||||
logger.info("Gradient check stats: min=%s, max=%s, mean=%s",
|
||||
min_err, max_err, mean_err)
|
||||
|
||||
wandb.log({
|
||||
"grad_check/min_abs_error": min_err,
|
||||
"grad_check/max_abs_error": max_err,
|
||||
"grad_check/mean_abs_error": mean_err,
|
||||
"grad_check/analytical_loss": analytical_loss.item(),
|
||||
})
|
||||
return max_err
|
||||
|
||||
return float('inf')
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Gradient check failed: %s", e)
|
||||
traceback.print_exc()
|
||||
return float('inf')
|
||||
|
||||
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
|
||||
noise_random_generator) -> float | None:
|
||||
"""
|
||||
Setup and perform gradient check on a fresh batch.
|
||||
Args:
|
||||
args: Training arguments
|
||||
loader_iter: Data loader iterator
|
||||
noise_scheduler: Noise scheduler for diffusion
|
||||
noise_random_generator: Random number generator for noise
|
||||
Returns:
|
||||
float or None: Maximum gradient error or None if check is disabled/fails
|
||||
"""
|
||||
assert self.training_args is not None
|
||||
|
||||
try:
|
||||
# Get a fresh batch and process it exactly like train_one_step
|
||||
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
|
||||
loader_iter)
|
||||
|
||||
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
|
||||
check_latents = check_latents.to(self.training_args.device,
|
||||
dtype=GRADIENT_CHECK_DTYPE)
|
||||
check_encoder_hidden_states = check_encoder_hidden_states.to(
|
||||
self.training_args.device, dtype=GRADIENT_CHECK_DTYPE)
|
||||
check_latents = normalize_dit_input("wan", check_latents)
|
||||
batch_size = check_latents.shape[0]
|
||||
check_noise = torch.randn_like(check_latents)
|
||||
|
||||
check_u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=noise_random_generator,
|
||||
logit_mean=args.logit_mean,
|
||||
logit_std=args.logit_std,
|
||||
mode_scale=args.mode_scale,
|
||||
)
|
||||
check_indices = (check_u *
|
||||
noise_scheduler.config.num_train_timesteps).long()
|
||||
check_timesteps = noise_scheduler.timesteps[check_indices].to(
|
||||
device=check_latents.device)
|
||||
|
||||
check_sigmas = get_sigmas(
|
||||
noise_scheduler,
|
||||
check_latents.device,
|
||||
check_timesteps,
|
||||
n_dim=check_latents.ndim,
|
||||
dtype=check_latents.dtype,
|
||||
)
|
||||
check_noisy_model_input = (
|
||||
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
|
||||
|
||||
# Compute target exactly like train_one_step
|
||||
if args.precondition_outputs:
|
||||
check_target = check_latents
|
||||
else:
|
||||
check_target = check_noise - check_latents
|
||||
|
||||
# Perform gradient check with the exact same inputs as training
|
||||
max_grad_error = self.gradient_check_parameters(
|
||||
transformer=self.transformer,
|
||||
latents=
|
||||
check_noisy_model_input, # Use noisy input like in training
|
||||
encoder_hidden_states=check_encoder_hidden_states,
|
||||
encoder_attention_mask=check_encoder_attention_mask,
|
||||
timesteps=check_timesteps,
|
||||
target=check_target,
|
||||
max_params_to_check=100 # Check more parameters
|
||||
)
|
||||
|
||||
if max_grad_error > 5e-2:
|
||||
logger.error("❌ Large gradient error detected: %s",
|
||||
max_grad_error)
|
||||
else:
|
||||
logger.info("✅ Gradient check passed: max error %s",
|
||||
max_grad_error)
|
||||
|
||||
return max_grad_error
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Gradient check setup failed: %s", e)
|
||||
traceback.print_exc()
|
||||
return None
|
||||
@@ -0,0 +1,464 @@
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Dict, 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
|
||||
|
||||
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,
|
||||
):
|
||||
"""
|
||||
Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(
|
||||
mean=logit_mean,
|
||||
std=logit_std,
|
||||
size=(batch_size, ),
|
||||
device="cpu",
|
||||
generator=generator,
|
||||
)
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2)**2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
|
||||
return u
|
||||
|
||||
|
||||
def get_sigmas(noise_scheduler,
|
||||
device,
|
||||
timesteps,
|
||||
n_dim=4,
|
||||
dtype=torch.float32) -> torch.Tensor:
|
||||
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
|
||||
schedule_timesteps = noise_scheduler.timesteps.to(device)
|
||||
timesteps = timesteps.to(device)
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item()
|
||||
for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
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)
|
||||
|
||||
|
||||
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(
|
||||
parameters: Union[torch.Tensor, List[torch.Tensor]],
|
||||
max_norm: float,
|
||||
norm_type: float = 2.0,
|
||||
error_if_nonfinite: bool = False,
|
||||
foreach: Optional[bool] = None,
|
||||
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
|
||||
) -> Optional[torch.Tensor]:
|
||||
global _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES
|
||||
|
||||
if not _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES:
|
||||
try:
|
||||
return clip_grad_norm_(parameters, max_norm, norm_type,
|
||||
error_if_nonfinite, foreach, pp_mesh)
|
||||
except NotImplementedError as e:
|
||||
if "DTensor does not support cross-mesh operation" in str(e):
|
||||
# https://github.com/pytorch/pytorch/issues/134212
|
||||
logger.warning(
|
||||
"DTensor does not support cross-mesh operation. If you haven't fully tensor-parallelized your "
|
||||
"model, while combining other parallelisms such as FSDP, it could be the reason for this error. "
|
||||
"Gradient clipping will be skipped and gradient norm will not be logged."
|
||||
)
|
||||
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)
|
||||
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = True
|
||||
return None
|
||||
|
||||
|
||||
# Copied from https://github.com/pytorch/torchtitan/blob/4a169701555ab9bd6ca3769f9650ae3386b84c6e/torchtitan/utils.py#L362
|
||||
@torch.no_grad()
|
||||
def clip_grad_norm_(
|
||||
parameters: Union[torch.Tensor, List[torch.Tensor]],
|
||||
max_norm: float,
|
||||
norm_type: float = 2.0,
|
||||
error_if_nonfinite: bool = False,
|
||||
foreach: Optional[bool] = None,
|
||||
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Clip the gradient norm of parameters.
|
||||
|
||||
Gradient norm clipping requires computing the gradient norm over the entire model.
|
||||
`torch.nn.utils.clip_grad_norm_` only computes gradient norm along DP/FSDP/TP dimensions.
|
||||
We need to manually reduce the gradient norm across PP stages.
|
||||
See https://github.com/pytorch/torchtitan/issues/596 for details.
|
||||
|
||||
Args:
|
||||
parameters (`torch.Tensor` or `List[torch.Tensor]`):
|
||||
Tensors that will have gradients normalized.
|
||||
max_norm (`float`):
|
||||
Maximum norm of the gradients after clipping.
|
||||
norm_type (`float`, defaults to `2.0`):
|
||||
Type of p-norm to use. Can be `inf` for infinity norm.
|
||||
error_if_nonfinite (`bool`, defaults to `False`):
|
||||
If `True`, an error is thrown if the total norm of the gradients from `parameters` is `nan`, `inf`, or `-inf`.
|
||||
foreach (`bool`, defaults to `None`):
|
||||
Use the faster foreach-based implementation. If `None`, use the foreach implementation for CUDA and CPU native tensors
|
||||
and silently fall back to the slow implementation for other device types.
|
||||
pp_mesh (`torch.distributed.device_mesh.DeviceMesh`, defaults to `None`):
|
||||
Pipeline parallel device mesh. If not `None`, will reduce gradient norm across PP stages.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
Total norm of the gradients
|
||||
"""
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
|
||||
# TODO(aryan): Wait for next Pytorch release to use `torch.nn.utils.get_total_norm`
|
||||
# total_norm = torch.nn.utils.get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
|
||||
total_norm = _get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
|
||||
|
||||
# If total_norm is a DTensor, the placements must be `torch.distributed._tensor.ops.math_ops._NormPartial`.
|
||||
# We can simply reduce the DTensor to get the total norm in this tensor's process group
|
||||
# and then convert it to a local tensor.
|
||||
# It has two purposes:
|
||||
# 1. to make sure the total norm is computed correctly when PP is used (see below)
|
||||
# 2. to return a reduced total_norm tensor whose .item() would return the correct value
|
||||
if isinstance(total_norm, torch.distributed.tensor.DTensor):
|
||||
# Will reach here if any non-PP parallelism is used.
|
||||
# If only using PP, total_norm will be a local tensor.
|
||||
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,
|
||||
group=pp_mesh.get_group())
|
||||
else:
|
||||
total_norm **= norm_type
|
||||
dist.all_reduce(total_norm,
|
||||
op=dist.ReduceOp.SUM,
|
||||
group=pp_mesh.get_group())
|
||||
total_norm **= 1.0 / norm_type
|
||||
|
||||
_clip_grads_with_norm_(parameters, max_norm, total_norm, foreach)
|
||||
return total_norm
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _clip_grads_with_norm_(
|
||||
parameters: Union[torch.Tensor, List[torch.Tensor]],
|
||||
max_norm: float,
|
||||
total_norm: torch.Tensor,
|
||||
foreach: Optional[bool] = None,
|
||||
) -> None:
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
max_norm = float(max_norm)
|
||||
if len(grads) == 0:
|
||||
return
|
||||
grouped_grads: dict[Tuple[torch.device, torch.dtype],
|
||||
Tuple[List[List[torch.Tensor]],
|
||||
List[int]]] = (_group_tensors_by_device_and_dtype(
|
||||
[grads])) # type: ignore[assignment]
|
||||
|
||||
clip_coef = max_norm / (total_norm + 1e-6)
|
||||
|
||||
# Note: multiplying by the clamped coef is redundant when the coef is clamped to 1, but doing so
|
||||
# avoids a `if clip_coef < 1:` conditional which can require a CPU <=> device synchronization
|
||||
# when the gradients do not reside in CPU memory.
|
||||
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||
for (device, _), ([device_grads], _) in grouped_grads.items():
|
||||
if (foreach is None and _has_foreach_support(device_grads, device)) or (
|
||||
foreach and _device_has_foreach_support(device)):
|
||||
torch._foreach_mul_(device_grads, clip_coef_clamped.to(device))
|
||||
elif foreach:
|
||||
raise RuntimeError(
|
||||
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
|
||||
)
|
||||
else:
|
||||
clip_coef_clamped_device = clip_coef_clamped.to(device)
|
||||
for g in device_grads:
|
||||
g.mul_(clip_coef_clamped_device)
|
||||
|
||||
|
||||
def _get_total_norm(
|
||||
tensors: Union[torch.Tensor, List[torch.Tensor]],
|
||||
norm_type: float = 2.0,
|
||||
error_if_nonfinite: bool = False,
|
||||
foreach: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
tensors = [tensors] if isinstance(tensors, torch.Tensor) else list(tensors)
|
||||
norm_type = float(norm_type)
|
||||
if len(tensors) == 0:
|
||||
return torch.tensor(0.0)
|
||||
first_device = tensors[0].device
|
||||
grouped_tensors: dict[tuple[torch.device, torch.dtype],
|
||||
tuple[list[list[torch.Tensor]], list[int]]] = (
|
||||
_group_tensors_by_device_and_dtype(
|
||||
[tensors] # type: ignore[list-item]
|
||||
)) # type: ignore[assignment]
|
||||
|
||||
norms: List[torch.Tensor] = []
|
||||
for (device, _), ([device_tensors], _) in grouped_tensors.items():
|
||||
local_tensors = [
|
||||
t.to_local()
|
||||
if isinstance(t, torch.distributed.tensor.DTensor) else t
|
||||
for t in device_tensors
|
||||
]
|
||||
if (foreach is None and _has_foreach_support(local_tensors, device)
|
||||
) or (foreach and _device_has_foreach_support(device)):
|
||||
norms.extend(torch._foreach_norm(local_tensors, norm_type))
|
||||
elif foreach:
|
||||
raise RuntimeError(
|
||||
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
|
||||
)
|
||||
else:
|
||||
norms.extend(
|
||||
[torch.linalg.vector_norm(g, norm_type) for g in local_tensors])
|
||||
|
||||
total_norm = torch.linalg.vector_norm(
|
||||
torch.stack([norm.to(first_device) for norm in norms]), norm_type)
|
||||
|
||||
if error_if_nonfinite and torch.logical_or(total_norm.isnan(),
|
||||
total_norm.isinf()):
|
||||
raise RuntimeError(
|
||||
f"The total norm of order {norm_type} for gradients from "
|
||||
"`parameters` is non-finite, so it cannot be clipped. To disable "
|
||||
"this error and scale the gradients by the non-finite norm anyway, "
|
||||
"set `error_if_nonfinite=False`")
|
||||
return total_norm
|
||||
|
||||
|
||||
def _get_foreach_kernels_supported_devices() -> list[str]:
|
||||
r"""Return the device type list that supports foreach kernels."""
|
||||
return ["cuda", "xpu", torch._C._get_privateuse1_backend_name()]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _group_tensors_by_device_and_dtype(
|
||||
tensorlistlist: List[List[Optional[torch.Tensor]]],
|
||||
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)
|
||||
|
||||
|
||||
def _device_has_foreach_support(device: torch.device) -> bool:
|
||||
return device.type in (_get_foreach_kernels_supported_devices() +
|
||||
["cpu"]) and not torch.jit.is_scripting()
|
||||
|
||||
|
||||
def _has_foreach_support(tensors: List[torch.Tensor],
|
||||
device: torch.device) -> bool:
|
||||
return _device_has_foreach_support(device) and all(
|
||||
t is None or type(t) in [torch.Tensor] for t in tensors)
|
||||
@@ -0,0 +1,371 @@
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
from collections import deque
|
||||
from copy import deepcopy
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
|
||||
get_world_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.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
|
||||
from fastvideo.v1.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.v1.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
|
||||
normalize_dit_input, save_checkpoint)
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Manual gradient checking flag - set to True to enable gradient verification
|
||||
ENABLE_GRADIENT_CHECK = False
|
||||
|
||||
|
||||
class WanTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Wan.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.vae_config.load_encoder = False
|
||||
validation_pipeline = WanValidationPipeline.from_pretrained(
|
||||
training_args.model_path, args=None, inference_mode=True)
|
||||
|
||||
self.validation_pipeline = validation_pipeline
|
||||
|
||||
def train_one_step(
|
||||
self,
|
||||
transformer,
|
||||
model_type,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader_iter,
|
||||
noise_scheduler,
|
||||
noise_random_generator,
|
||||
gradient_accumulation_steps,
|
||||
sp_size,
|
||||
precondition_outputs,
|
||||
max_grad_norm,
|
||||
weighting_scheme,
|
||||
logit_mean,
|
||||
logit_std,
|
||||
mode_scale,
|
||||
) -> tuple[float, float]:
|
||||
assert self.training_args is not None
|
||||
self.modules["transformer"].requires_grad_(True)
|
||||
self.modules["transformer"].train()
|
||||
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
# Get next batch, handling epoch boundaries gracefully
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
latents, encoder_hidden_states, encoder_attention_mask, infos = batch
|
||||
|
||||
# logger.info("rank: %s, caption: %s",
|
||||
# self.rank,
|
||||
# infos['caption'],
|
||||
# local_main_process_only=False)
|
||||
# TODO(will): don't hardcode bfloat16
|
||||
latents = latents.to(self.training_args.device,
|
||||
dtype=torch.bfloat16)
|
||||
encoder_hidden_states = encoder_hidden_states.to(
|
||||
self.training_args.device, dtype=torch.bfloat16)
|
||||
latents = normalize_dit_input(model_type, latents)
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn_like(latents)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=noise_random_generator,
|
||||
logit_mean=logit_mean,
|
||||
logit_std=logit_std,
|
||||
mode_scale=mode_scale,
|
||||
)
|
||||
indices = (u * noise_scheduler.config.num_train_timesteps).long()
|
||||
timesteps = noise_scheduler.timesteps[indices].to(
|
||||
device=latents.device)
|
||||
if sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
sp_group = get_sp_group()
|
||||
sp_group.broadcast(timesteps, src=0)
|
||||
sigmas = get_sigmas(
|
||||
noise_scheduler,
|
||||
latents.device,
|
||||
timesteps,
|
||||
n_dim=latents.ndim,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
input_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if 'hunyuan' in model_type:
|
||||
input_kwargs["guidance"] = torch.tensor(
|
||||
[1000.0],
|
||||
device=noisy_model_input.device,
|
||||
dtype=torch.bfloat16)
|
||||
with set_forward_context(current_timestep=timesteps,
|
||||
attn_metadata=None):
|
||||
model_pred = transformer(**input_kwargs)
|
||||
|
||||
if precondition_outputs:
|
||||
model_pred = noisy_model_input - model_pred * sigmas
|
||||
target = latents if precondition_outputs else noise - latents
|
||||
|
||||
loss = (torch.mean((model_pred.float() - target.float())**2) /
|
||||
gradient_accumulation_steps)
|
||||
|
||||
loss.backward()
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
|
||||
# local_main_process_only=False)
|
||||
world_group = get_world_group()
|
||||
world_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
# TODO(will): perhaps move this into transformer api so that we can do
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
model_parts = [self.transformer]
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
foreach=None,
|
||||
)
|
||||
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
|
||||
else:
|
||||
grad_norm = 0.0
|
||||
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
return total_loss, grad_norm
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
):
|
||||
assert self.training_args is not None
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
seed = self.training_args.seed if self.training_args.seed is not None else 42
|
||||
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
noise_random_generator = torch.Generator(device="cpu")
|
||||
noise_random_generator.manual_seed(seed)
|
||||
|
||||
logger.info("Initialized random seeds with seed: %s", seed)
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
# Train!
|
||||
assert self.training_args.sp_size is not None
|
||||
assert self.training_args.gradient_accumulation_steps is not None
|
||||
total_batch_size = (self.world_size *
|
||||
self.training_args.gradient_accumulation_steps /
|
||||
self.training_args.sp_size *
|
||||
self.training_args.train_sp_batch_size)
|
||||
logger.info("***** Running training *****")
|
||||
logger.info(" Num examples = %s", len(self.train_dataset))
|
||||
logger.info(" Dataloader size = %s", len(self.train_dataloader))
|
||||
logger.info(" Num Epochs = %s", self.num_train_epochs)
|
||||
logger.info(" Resume training from step %s",
|
||||
self.init_steps) # type: ignore
|
||||
logger.info(" Instantaneous batch size per device = %s",
|
||||
self.training_args.train_batch_size)
|
||||
logger.info(
|
||||
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
|
||||
total_batch_size)
|
||||
logger.info(" Gradient Accumulation steps = %s",
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
logger.info(" Total optimization steps = %s",
|
||||
self.training_args.max_train_steps)
|
||||
logger.info(
|
||||
" Total training parameters per FSDP shard = %s B",
|
||||
sum(p.numel()
|
||||
for p in self.transformer.parameters() if p.requires_grad) /
|
||||
1e9)
|
||||
# print dtype
|
||||
logger.info(" Master weight dtype: %s",
|
||||
self.transformer.parameters().__next__().dtype)
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
logger.info("Loading checkpoint from %s",
|
||||
self.training_args.resume_from_checkpoint)
|
||||
resumed_step = load_checkpoint(
|
||||
self.transformer, self.rank,
|
||||
self.training_args.resume_from_checkpoint, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
noise_random_generator)
|
||||
if resumed_step > 0:
|
||||
self.init_steps = resumed_step
|
||||
logger.info("Successfully resumed from step %s", resumed_step)
|
||||
else:
|
||||
logger.warning(
|
||||
"Failed to load checkpoint, starting from step 0")
|
||||
self.init_steps = 0
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, self.training_args.max_train_steps),
|
||||
initial=self.init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable=self.local_rank > 0,
|
||||
)
|
||||
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
# TODO(will): fix this
|
||||
# for i in range(self.init_steps):
|
||||
# next(loader_iter)
|
||||
# get gpu memory usage
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage before train_one_step: %s MB",
|
||||
gpu_memory_usage)
|
||||
|
||||
# Do validation at the beginning of training
|
||||
# self._log_validation(self.transformer, self.training_args, 0)
|
||||
|
||||
for step in range(self.init_steps + 1,
|
||||
self.training_args.max_train_steps + 1):
|
||||
start_time = time.perf_counter()
|
||||
|
||||
loss, grad_norm = self.train_one_step(
|
||||
self.transformer,
|
||||
# args.model_type,
|
||||
"wan",
|
||||
self.optimizer,
|
||||
self.lr_scheduler,
|
||||
self.train_loader_iter,
|
||||
noise_scheduler,
|
||||
noise_random_generator,
|
||||
self.training_args.gradient_accumulation_steps,
|
||||
self.training_args.sp_size,
|
||||
self.training_args.precondition_outputs,
|
||||
self.training_args.max_grad_norm,
|
||||
self.training_args.weighting_scheme,
|
||||
self.training_args.logit_mean,
|
||||
self.training_args.logit_std,
|
||||
self.training_args.mode_scale,
|
||||
)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
logger.info("GPU memory usage after train_one_step: %s MB",
|
||||
gpu_memory_usage)
|
||||
|
||||
step_time = time.perf_counter() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
# Manual gradient checking - only at first step
|
||||
if step == 1 and ENABLE_GRADIENT_CHECK:
|
||||
logger.info("Performing gradient check at step %s", step)
|
||||
self.setup_gradient_check(args, self.train_loader_iter,
|
||||
noise_scheduler,
|
||||
noise_random_generator)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if self.rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
"learning_rate": self.lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
},
|
||||
step=step,
|
||||
)
|
||||
if step % self.training_args.checkpointing_steps == 0:
|
||||
save_checkpoint(self.transformer, self.rank,
|
||||
self.training_args.output_dir, step,
|
||||
self.optimizer, self.train_dataloader,
|
||||
self.lr_scheduler, noise_random_generator)
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
|
||||
save_checkpoint(self.transformer, self.rank,
|
||||
self.training_args.output_dir,
|
||||
self.training_args.max_train_steps, self.optimizer,
|
||||
self.train_dataloader, self.lr_scheduler,
|
||||
noise_random_generator)
|
||||
|
||||
if get_sp_group():
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = WanTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.forward(None, args)
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.v1.fastvideo_args import TrainingArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.use_cpu_offload = False
|
||||
main(args)
|
||||
+97
-10
@@ -10,10 +10,12 @@ import json
|
||||
import math
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import traceback
|
||||
from dataclasses import fields, is_dataclass
|
||||
from dataclasses import dataclass, fields, is_dataclass
|
||||
from functools import partial, wraps
|
||||
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
|
||||
Union, cast)
|
||||
@@ -22,7 +24,10 @@ import cloudpickle
|
||||
import filelock
|
||||
import torch
|
||||
import yaml
|
||||
from diffusers.loaders.lora_base import (
|
||||
_best_guess_weight_name) # watch out for potetential removal from diffusers
|
||||
from huggingface_hub import snapshot_download
|
||||
from remote_pdb import RemotePdb
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -448,14 +453,14 @@ def import_pynvml():
|
||||
return pynvml
|
||||
|
||||
|
||||
def maybe_download_model(model_path: str,
|
||||
def maybe_download_model(model_name_or_path: str,
|
||||
local_dir: Optional[str] = None,
|
||||
download: bool = True) -> str:
|
||||
"""
|
||||
Check if the model path is a Hugging Face Hub model ID and download it if needed.
|
||||
|
||||
Args:
|
||||
model_path: Local path or Hugging Face Hub model ID
|
||||
model_name_or_path: Local path or Hugging Face Hub model ID
|
||||
local_dir: Local directory to save the model
|
||||
download: Whether to download the model from Hugging Face Hub
|
||||
|
||||
@@ -464,27 +469,47 @@ def maybe_download_model(model_path: str,
|
||||
"""
|
||||
|
||||
# If the path exists locally, return it
|
||||
if os.path.exists(model_path):
|
||||
logger.info("Model already exists locally at %s", model_path)
|
||||
return model_path
|
||||
if os.path.exists(model_name_or_path):
|
||||
logger.info("Model already exists locally at %s", model_name_or_path)
|
||||
return model_name_or_path
|
||||
|
||||
# Otherwise, assume it's a HF Hub model ID and try to download it
|
||||
try:
|
||||
logger.info("Downloading model snapshot from HF Hub for %s...",
|
||||
model_path)
|
||||
with get_lock(model_path):
|
||||
model_name_or_path)
|
||||
with get_lock(model_name_or_path):
|
||||
local_path = snapshot_download(
|
||||
repo_id=model_path,
|
||||
repo_id=model_name_or_path,
|
||||
ignore_patterns=["*.onnx", "*.msgpack"],
|
||||
local_dir=local_dir)
|
||||
logger.info("Downloaded model to %s", local_path)
|
||||
return str(local_path)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Could not find model at {model_path} and failed to download from HF Hub: {e}"
|
||||
f"Could not find model at {model_name_or_path} and failed to download from HF Hub: {e}"
|
||||
) from e
|
||||
|
||||
|
||||
def maybe_download_lora(model_name_or_path: str,
|
||||
local_dir: Optional[str] = None,
|
||||
download: bool = True) -> str:
|
||||
"""
|
||||
Check if the model path is a Hugging Face Hub model ID and download it if needed.
|
||||
Args:
|
||||
model_name_or_path: Local path or Hugging Face Hub model ID
|
||||
local_dir: Local directory to save the model
|
||||
download: Whether to download the model from Hugging Face Hub
|
||||
|
||||
Returns:
|
||||
Local path to the model
|
||||
"""
|
||||
|
||||
local_path = maybe_download_model(model_name_or_path, local_dir, download)
|
||||
weight_name = _best_guess_weight_name(model_name_or_path,
|
||||
file_extension=".safetensors")
|
||||
return os.path.join(local_path, weight_name)
|
||||
|
||||
|
||||
def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Verify that the model directory contains a valid diffusers configuration.
|
||||
@@ -645,3 +670,65 @@ class TypeBasedDispatcher:
|
||||
if isinstance(obj, ty):
|
||||
return fn(obj)
|
||||
raise ValueError(f"Invalid object: {obj}")
|
||||
|
||||
|
||||
# For non-torch.distributed debugging
|
||||
def remote_breakpoint() -> None:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
s.bind(("localhost", 0)) # Let the OS pick an ephemeral port.
|
||||
port = s.getsockname()[1]
|
||||
RemotePdb(host="localhost", port=port).set_trace()
|
||||
|
||||
|
||||
@dataclass
|
||||
class MixedPrecisionState:
|
||||
master_dtype: Optional[torch.dtype] = None
|
||||
param_dtype: Optional[torch.dtype] = None
|
||||
reduce_dtype: Optional[torch.dtype] = None
|
||||
output_dtype: Optional[torch.dtype] = None
|
||||
compute_dtype: Optional[torch.dtype] = None
|
||||
|
||||
|
||||
# Thread-local storage for mixed precision state
|
||||
_mixed_precision_state = threading.local()
|
||||
|
||||
|
||||
def get_mixed_precision_state() -> MixedPrecisionState:
|
||||
"""Get the current mixed precision state."""
|
||||
if not hasattr(_mixed_precision_state, 'state'):
|
||||
raise ValueError("Mixed precision state not set")
|
||||
return cast(MixedPrecisionState, _mixed_precision_state.state)
|
||||
|
||||
|
||||
def set_mixed_precision_policy(master_dtype: torch.dtype,
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
output_dtype: Optional[torch.dtype] = None):
|
||||
"""Set mixed precision policy globally.
|
||||
|
||||
Args:
|
||||
param_dtype: Parameter dtype used for training
|
||||
reduce_dtype: Reduction dtype used for gradients
|
||||
output_dtype: Optional output dtype
|
||||
"""
|
||||
state = MixedPrecisionState(
|
||||
master_dtype=master_dtype,
|
||||
param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
output_dtype=output_dtype,
|
||||
)
|
||||
_mixed_precision_state.state = state
|
||||
|
||||
|
||||
def get_compute_dtype() -> torch.dtype:
|
||||
"""Get the current compute dtype from mixed precision policy.
|
||||
|
||||
Returns:
|
||||
torch.dtype: The compute dtype to use, defaults to get_default_dtype() if no policy set
|
||||
"""
|
||||
if not hasattr(_mixed_precision_state, 'state'):
|
||||
return torch.get_default_dtype()
|
||||
else:
|
||||
state = get_mixed_precision_state()
|
||||
return state.param_dtype
|
||||
|
||||
@@ -47,6 +47,13 @@ class Executor(ABC):
|
||||
})
|
||||
return cast(ForwardBatch, outputs[0]["output_batch"])
|
||||
|
||||
@abstractmethod
|
||||
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
|
||||
"""
|
||||
Set the LoRA adapter for the workers.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def collective_rpc(self,
|
||||
method: Union[str, Callable[..., _R]],
|
||||
|
||||
@@ -5,6 +5,7 @@ import multiprocessing as mp
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
from multiprocessing.connection import Connection
|
||||
from typing import Any, Dict, Optional, TextIO, cast
|
||||
|
||||
import psutil
|
||||
@@ -29,14 +30,14 @@ RESET = '\033[0;0m'
|
||||
class Worker:
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int,
|
||||
rank: int, pipe):
|
||||
rank: int, pipe: Connection, master_port: int):
|
||||
self.fastvideo_args = fastvideo_args
|
||||
self.local_rank = local_rank
|
||||
self.rank = rank
|
||||
# TODO(will): don't hardcode this
|
||||
self.distributed_init_method = "env://"
|
||||
self.pipe = pipe
|
||||
|
||||
self.master_port = master_port
|
||||
self.init_device()
|
||||
|
||||
# Init request dispatcher
|
||||
@@ -76,7 +77,7 @@ class Worker:
|
||||
f"Unsupported device: {self.fastvideo_args.device_str}")
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
os.environ["MASTER_PORT"] = str(self.master_port)
|
||||
os.environ["LOCAL_RANK"] = str(self.local_rank)
|
||||
os.environ["RANK"] = str(self.rank)
|
||||
|
||||
@@ -92,6 +93,9 @@ class Worker:
|
||||
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
|
||||
return cast(ForwardBatch, output_batch)
|
||||
|
||||
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
|
||||
self.pipeline.set_lora_adapter(lora_nickname, lora_path)
|
||||
|
||||
def shutdown(self) -> Dict[str, Any]:
|
||||
"""Gracefully shut down the worker process"""
|
||||
logger.info("Worker %d shutting down...",
|
||||
@@ -191,7 +195,7 @@ def init_worker_distributed_environment(
|
||||
|
||||
|
||||
def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
|
||||
rank: int, pipe):
|
||||
rank: int, pipe: Connection, master_port: int):
|
||||
# Add process-specific prefix to stdout and stderr
|
||||
process_name = mp.current_process().name
|
||||
pid = os.getpid()
|
||||
@@ -206,8 +210,9 @@ def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
|
||||
logger.info("Worker %d initializing...",
|
||||
rank,
|
||||
local_main_process_only=False)
|
||||
|
||||
try:
|
||||
worker = Worker(fastvideo_args, local_rank, rank, pipe)
|
||||
worker = Worker(fastvideo_args, local_rank, rank, pipe, master_port)
|
||||
logger.info("Worker %d sending ready", rank)
|
||||
pipe.send({
|
||||
"status": "ready",
|
||||
|
||||
@@ -3,6 +3,7 @@ import contextlib
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import time
|
||||
from multiprocessing.process import BaseProcess
|
||||
from typing import Any, Callable, List, Optional, Union, cast
|
||||
@@ -28,6 +29,15 @@ class MultiprocExecutor(Executor):
|
||||
|
||||
self.workers: List[BaseProcess] = []
|
||||
self.worker_pipes = []
|
||||
self.master_port = None
|
||||
|
||||
for port in range(29503, 65535):
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
if s.connect_ex(('localhost', port)) != 0:
|
||||
self.master_port = port
|
||||
break
|
||||
if self.master_port is None:
|
||||
raise ValueError("No unused port found to use as master port")
|
||||
|
||||
# Create pipes and start workers
|
||||
for rank in range(self.world_size):
|
||||
@@ -39,7 +49,8 @@ class MultiprocExecutor(Executor):
|
||||
kwargs=dict(fastvideo_args=self.fastvideo_args,
|
||||
local_rank=rank,
|
||||
rank=rank,
|
||||
pipe=worker_pipe))
|
||||
pipe=worker_pipe,
|
||||
master_port=self.master_port))
|
||||
worker.start()
|
||||
self.workers.append(worker)
|
||||
|
||||
@@ -62,6 +73,13 @@ class MultiprocExecutor(Executor):
|
||||
})
|
||||
return cast(ForwardBatch, responses[0]["output_batch"])
|
||||
|
||||
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
|
||||
self.collective_rpc("set_lora_adapter",
|
||||
kwargs={
|
||||
"lora_nickname": lora_nickname,
|
||||
"lora_path": lora_path
|
||||
})
|
||||
|
||||
def collective_rpc(self,
|
||||
method: Union[str, Callable],
|
||||
timeout: Optional[float] = None,
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
DATA_DIR=data/cats_480_2_latents_parq_neg/combined_parquet_dataset
|
||||
VALIDATION_DIR=data/cats_480_2_latents_parq_neg/validation_parquet_dataset
|
||||
NUM_GPUS=4
|
||||
CUDA_VISIBLE_DEVICES=4,5,6,7
|
||||
# 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
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
fastvideo/v1/training/wan_training_pipeline.py\
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_prompt_dir "$VALIDATION_DIR"\
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 20 \
|
||||
--num_gpus 4 \
|
||||
--sp_size 4 \
|
||||
--tp_size 4 \
|
||||
--dp_size 1 \
|
||||
--dp_shards 4 \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 1\
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=5000 \
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=6000 \
|
||||
--validation_steps 10\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="data/wan_finetune_crush"\
|
||||
--tracker_project_name wan_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 81 \
|
||||
--shift 3 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
@@ -0,0 +1,53 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
DATA_DIR=data/crush-smol_parq/combined_parquet_dataset
|
||||
VALIDATION_DIR=data/crush-smol_parq/validation_parquet_dataset
|
||||
NUM_GPUS=1
|
||||
CUDA_VISIBLE_DEVICES=4,5,6,7
|
||||
# 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
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
fastvideo/v1/training/wan_training_pipeline.py\
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_prompt_dir "$VALIDATION_DIR"\
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 14 \
|
||||
--num_gpus 1 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--dp_size 1 \
|
||||
--dp_shards 1 \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 1\
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=5000 \
|
||||
--learning_rate=1e-5\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=6000 \
|
||||
--validation_steps 100\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="data/wan_finetune_crush"\
|
||||
--tracker_project_name finetrainers-wan \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 77 \
|
||||
--shift 3 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
+9
-2
@@ -19,7 +19,7 @@ dependencies = [
|
||||
|
||||
# Machine Learning & Transformers
|
||||
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
|
||||
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.0", "bitsandbytes",
|
||||
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
|
||||
"torch==2.6.0", "torchvision",
|
||||
|
||||
# Acceleration & Optimization
|
||||
@@ -37,16 +37,23 @@ dependencies = [
|
||||
"gradio>=5.22.0", "moviepy==1.0.3", "flask",
|
||||
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
|
||||
# System & Monitoring Tools
|
||||
"gpustat", "watch",
|
||||
"gpustat", "watch", "remote-pdb",
|
||||
|
||||
# Kernel & Packaging
|
||||
"wheel",
|
||||
|
||||
# Training Dependencies
|
||||
"torchdata",
|
||||
"pyarrow",
|
||||
"datasets",
|
||||
"av",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
# flash-attn: pip install flash-attn==2.7.4.post1 --no-cache-dir --no-build-isolation
|
||||
|
||||
|
||||
lint = [
|
||||
"pre-commit==4.0.1",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
#https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tree/main
|
||||
DATA_DIR=./data
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/wan\
|
||||
--model_type "wan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Image-Vid-Finetune-Wan/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Image-Vid-Finetune-Wan/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 32 \
|
||||
--sp_size 1 \
|
||||
--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 "8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 5.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/wan"\
|
||||
--tracker_project_name Wan_PCM \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 81 \
|
||||
--shift 3 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,54 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
DATA_DIR=data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
|
||||
VALIDATION_DIR=data/HD-Mixkit-Finetune-Wan/validation_parquet_dataset
|
||||
NUM_GPUS=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_finetune/checkpoint-5"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
fastvideo/v1/training/wan_training_pipeline.py\
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_prompt_dir "$VALIDATION_DIR"\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 4 \
|
||||
--sp_size $NUM_GPUS \
|
||||
--tp_size $NUM_GPUS \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 5\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=120 \
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=50 \
|
||||
--validation_steps 20\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="$DATA_DIR/outputs/wan_finetune"\
|
||||
--tracker_project_name wan_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 81 \
|
||||
--shift 3 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "fp32" \
|
||||
--max_grad_norm 1.0 \
|
||||
|
||||
# --resume_from_checkpoint "$CHECKPOINT_PATH"
|
||||
@@ -6,19 +6,17 @@ export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--sp_size $num_gpus \
|
||||
--tp_size $num_gpus \
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_inference_steps 6 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 17 \
|
||||
--prompt_path ./assets/prompt.txt \
|
||||
--num-frames 125 \
|
||||
--num-inference-steps 6 \
|
||||
--guidance-scale 1 \
|
||||
--embedded-cfg-scale 6 \
|
||||
--flow-shift 17 \
|
||||
--prompt "A beautiful woman in a red dress walking down a street" \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--vae-sp
|
||||
--output-path outputs_video/
|
||||
@@ -7,19 +7,17 @@ export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--sp_size $num_gpus \
|
||||
--tp_size $num_gpus \
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size $num_gpus \
|
||||
--tp-size $num_gpus \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 7 \
|
||||
--prompt_path ./assets/prompt.txt \
|
||||
--num-frames 125 \
|
||||
--num-inference-steps 50 \
|
||||
--guidance-scale 1 \
|
||||
--embedded-cfg-scale 6 \
|
||||
--flow-shift 7 \
|
||||
--prompt "A beautiful woman in a red dress walking down a street" \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--vae-sp
|
||||
--output-path outputs_video/
|
||||
|
||||
@@ -8,19 +8,17 @@ export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
|
||||
# Note that the tp_size and sp_size should be the same and equal to the number
|
||||
# of GPUs. They are used for different parallel groups. sp_size is used for
|
||||
# dit model and tp_size is used for encoder models.
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--sp_size ${num_gpus} \
|
||||
--tp_size ${num_gpus} \
|
||||
fastvideo generate \
|
||||
--model-path $MODEL_BASE \
|
||||
--sp-size ${num_gpus} \
|
||||
--tp-size ${num_gpus} \
|
||||
--height 768 \
|
||||
--width 1280 \
|
||||
--num_frames 117 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 7 \
|
||||
--prompt_path ./assets/prompt.txt \
|
||||
--num-frames 117 \
|
||||
--num-inference-steps 50 \
|
||||
--guidance-scale 1 \
|
||||
--embedded-cfg-scale 6 \
|
||||
--flow-shift 7 \
|
||||
--prompt "A beautiful woman in a red dress walking down a street" \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--vae-sp
|
||||
--output-path outputs_video/
|
||||
|
||||
@@ -5,19 +5,18 @@ export FASTVIDEO_ATTENTION_BACKEND=
|
||||
num_gpus=2
|
||||
url='127.0.0.1'
|
||||
model_dir=data/stepvideo-t2v
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/v1/sample/v1_fastvideo_inference.py \
|
||||
--sp_size ${num_gpus} \
|
||||
--tp_size ${num_gpus} \
|
||||
fastvideo generate \
|
||||
--model-path $model_dir \
|
||||
--sp-size ${num_gpus} \
|
||||
--tp-size ${num_gpus} \
|
||||
--height 256 \
|
||||
--width 256 \
|
||||
--num_frames 29 \
|
||||
--num_inference_steps 50 \
|
||||
--embedded_cfg_scale 9.0 \
|
||||
--guidance_scale 9.0 \
|
||||
--prompt_path ./assets/prompt.txt \
|
||||
--num-frames 29 \
|
||||
--num-inference-steps 50 \
|
||||
--embedded-cfg-scale 9.0 \
|
||||
--guidance-scale 9.0 \
|
||||
--prompt "A beautiful woman in a red dress walking down a street" \
|
||||
--seed 1024 \
|
||||
--output_path outputs_stepvideo/ \
|
||||
--model_path $model_dir \
|
||||
--flow_shift 13.0 \
|
||||
--vae_precision bf16
|
||||
--output-path outputs_stepvideo/ \
|
||||
--flow-shift 13.0 \
|
||||
--vae-precision bf16
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user