Compare commits

..
Author SHA1 Message Date
SolitaryThinker 32b5439b2c checkpoint 2025-07-29 20:46:04 +00:00
SolitaryThinker 9f842d4c63 revert other stages 2025-07-29 01:15:50 +00:00
SolitaryThinker 02a0b1548a revert example 2025-07-29 01:15:50 +00:00
SolitaryThinker 5bec789377 revert example 2025-07-29 01:15:50 +00:00
SolitaryThinker 510577e8d0 ti2v wan2.2 2025-07-29 01:15:49 +00:00
SolitaryThinker cd0ed23b82 lint 2025-07-29 01:13:54 +00:00
SolitaryThinker 1bce2933d8 move wan i2v dmd pipeline 2025-07-29 01:13:00 +00:00
SolitaryThinker 48607164f8 Fastwan 14B and inp sampling param and optional image_path 2025-07-29 01:12:58 +00:00
SolitaryThinker 5ec46ba818 checkpoint 2025-07-29 01:11:01 +00:00
JerryZhou54 a403e66a47 Add wan2.2 5B T2V 2025-07-29 00:30:04 +00:00
Jinzhe Pan 7b6c8aee99 [2/3][Preprocess] refactor pipeline registry & file structure (#639) 2025-07-27 23:30:17 -07:00
Yongqi Chen 6284eaa363 [Feature][Distill]Add DMD+VSA joint training example (#654) 2025-07-27 18:18:01 -04:00
Yongqi Chen 636524e87f [Feature] Add Wan-14B-T2V-VSA CLI inference; add master port args (#653) 2025-07-27 07:13:44 -04:00
Yongqi Chen 202b2f3972 [Feature] Ignore [union-attr] and [override] mypy check and remove from training (#652) 2025-07-27 04:34:38 -04:00
Yongqi Chen 247fe273d8 [Feature] Add DMD T2V training pipeline (#651) 2025-07-27 03:35:51 -04:00
William Lin cb320dfa3a [bugfix] VideoGenerator improperly extracts output_video_name (#649) 2025-07-26 19:46:26 -07:00
Kevin Lin d8bb5abc46 [CI] Fix ComfyUI publisher ID (#648) 2025-07-25 19:02:52 -07:00
Kevin Lin cc703eca51 [CI] Add publish workflow for ComfyUI (#647) 2025-07-25 18:39:05 -07:00
William Lin 81c9df629c [core] Add offloading for vae and image encoder and rename offloading args (#643) 2025-07-25 17:55:03 -07:00
Yongqi Chen d3c0c52208 [Feature] Add prompt_txt support for CLI inference; Add DMD CLI inference (#646) 2025-07-25 19:44:10 -04:00
William Lin 744e0555c0 [misc] Use FASTVIDEO_STAGE_LOGGING for perf timing of stage (#644) 2025-07-25 16:20:53 -07:00
Jinzhe Pan 3a38f7dfdc [1/3][Preprocess] refactor preprocessing configs (#638) 2025-07-25 14:27:12 -07:00
William Lin f572319bd9 [Feature] Remove V1 folder (#642) 2025-07-24 22:43:12 -07:00
84 changed files with 3993 additions and 620 deletions
+28
View File
@@ -0,0 +1,28 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'hao-ai-lab' }}
steps:
- name: Check out code
uses: actions/checkout@v4
with:
submodules: true
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+1 -1
View File
@@ -60,7 +60,7 @@ repos:
rev: v1.15.0
hooks:
- id: mypy
args: [--python-version, '3.10', --follow-imports, "skip" ]
args: [--python-version, '3.10', --follow-imports, "skip", "--disable-error-code", "union-attr", "--disable-error-code", "override" ]
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
- repo: local
hooks:
Binary file not shown.

After

Width:  |  Height:  |  Size: 31 KiB

+1 -1
View File
@@ -310,7 +310,7 @@
"value": -99999,
"cachedValue": "fp16"
},
"use_cpu_offload": {
"dit_cpu_offload": {
"isAuto": true,
"value": -99999,
"cachedValue": true
@@ -517,7 +517,7 @@
"value": "fp16",
"cachedValue": "fp16"
},
"use_cpu_offload": {
"dit_cpu_offload": {
"isAuto": true,
"value": true,
"cachedValue": true
+4 -4
View File
@@ -105,7 +105,7 @@ class VideoGenerator:
"precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"use_cpu_offload": ([True, False], {
"dit_cpu_offload": ([True, False], {
"default": False
}),
}
@@ -204,7 +204,7 @@ class VideoGenerator:
vae_config=None,
text_encoder_config=None,
dit_config=None,
use_cpu_offload=None,
dit_cpu_offload=None,
):
print('Running FastVideo inference')
@@ -259,8 +259,8 @@ class VideoGenerator:
raw_generation_args['tp_size'] = tp_size
if sp_size is not None:
raw_generation_args['sp_size'] = sp_size
if use_cpu_offload is not None:
raw_generation_args['use_cpu_offload'] = use_cpu_offload
if dit_cpu_offload is not None:
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
generation_args = {
k: v
+1 -1
View File
@@ -552,7 +552,7 @@ app.registerExtension({
]
const floatWidgetNames = ["embedded_cfg_scale", "guidance_scale"]
const comboWidgetNames = ["vae_tiling", "vae_precision", "vae_sp", "text_encoder_precision", "precision",
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "use_cpu_offload", "enable_teacache"
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload", "enable_teacache"
]
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
+1 -1
View File
@@ -27,7 +27,7 @@ def main():
model_name = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
config = PipelineConfig.from_pretrained(model_name)
config.vae_precision = "fp16"
config.use_cpu_offload = True
config.dit_cpu_offload = True
# Create the generator
generator = VideoGenerator.from_pretrained(
+19
View File
@@ -0,0 +1,19 @@
# Wan2.1-T2V-1.3B Distill Example
These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using DMD-only and DMD+VSA methods.
### 1. Download dataset:
```bash
bash examples/distill/Wan-Syn-480P/download_dataset.sh
```
### 2. Configure and run distillation:
#### For DMD-only distillation:
```bash
sbatch examples/distill/Wan-Syn-480P/distill_dmd_t2v.slurm
```
#### For DMD+VSA distillation:
```bash
sbatch examples/distill/Wan-Syn-480P/distill_dmd_VSA_t2v.slurm
```
@@ -0,0 +1,137 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=8
#SBATCH --ntasks=8
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate your_env
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir="checkpoints/wan_t2v_finetune"
--max_train_steps=4000
--train_batch_size=1
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--num_latent_t 16
--num_height 448
--num_width 832
--num_frames 61
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 64
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-5
--mixed_precision="bf16"
--checkpointing_steps=500
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "bf16"
--ema_start_step 0
--flow_shift 8
--seed 1000
)
# DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,757,522'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--real_score_guidance_scale 3.5
--VSA_sparsity 0.8
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -0,0 +1,136 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=8
#SBATCH --ntasks=8
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate your_env
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd
--output_dir="checkpoints/wan_t2v_finetune"
--max_train_steps=4000
--train_batch_size=1
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--num_latent_t 16
--num_height 448
--num_width 832
--num_frames 61
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 64
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "1.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-5
--mixed_precision="bf16"
--checkpointing_steps=500
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "bf16"
--ema_start_step 0
--flow_shift 8
--seed 1000
)
# DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,757,522'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--real_score_guidance_scale 3.5
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
@@ -0,0 +1,516 @@
{
"data": [
{
"caption": "In the video, a woman is elegantly showcasing her earrings, bringing attention to their intricate design with a gentle touch of her fingers. She is bathed in ambient purple and pink lighting, which casts a soft glow on her delicate features and enhances the vivid tones of her lipstick and eye makeup. Her hair is styled to frame her face smoothly, emphasizing the contours of her jawline and cheekbones. The background features a blurred neon light, adding an artistic and modern touch to the overall aesthetic.",
"video_path": "Fashion/mixkit-face-of-an-elegant-and-captivating-woman-41914_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
"video_path": "Man/mixkit-a-rancher-riding-a-horse-at-sunset-1143_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a dimly lit, eerie setting, a mysterious pink bottle labeled \"Authentic 100% organic POISON\" sits prominently in the foreground, casting a menacing aura. The bottle is accentuated by green fog, which swirls lightly around it, enhancing its sinister allure. Behind it, a shadowy golden bottle adorned with a spider emblem subtly emerges, adding an extra layer of mystery to the scene. Dim candles provide faint, flickering light, which complements the dark atmosphere, making the setting ideal for an illusion of hidden dangers.",
"video_path": "smoke/mixkit-poison-in-halloween-ritual-33879_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video opens with a tranquil scene in the heart of a dense forest, emphasizing two large, textured tree trunks in the foreground framing the view. Sunlight filters through the canopy above, casting intricate patterns of light and shadow on the trees and the ground. Between the tree trunks, a clear view of a calm, muddy river unfolds, its surface shimmering under the gentle sunlight. The riverbank is decorated with a variety of small bushes and vibrant foliage, subtly transitioning into the deep greens of tall, leafy plants. In the background, the dense forest looms, filled with dark, towering trees, their branches intertwining to form an intricate canopy. The scene is bathed in the soft glow of the sun, creating a serene and picturesque setting. Occasional sunbeams pierce through the foliage, adding a magical aura to the landscape. The vibrant reds and oranges of the smaller plants add contrast, bringing warmth to the earthy tones of the scenery. Overall, this harmonious blend of natural elements creates a peaceful and idyllic forest setting.",
"video_path": "forest/mixkit-view-of-a-river-between-two-old-trees-560_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a martial artist dressed in a traditional white uniform with a black belt demonstrates a series of precise movements against a stark black background. The individual gracefully transitions between stances, embodying a sense of focused discipline and control. Each motion is executed with a deliberate pace, showcasing the fluidity of martial arts techniques. The soft lighting creates subtle highlights on the uniform, adding depth to the figure as it moves. The practitioner begins with an open-hand pose, feet firmly grounded, gradually shifting to a powerful forward punch. The fluidity of the sequence displays a mastery of balance and poise. Every trajectory of the limbs is precise and deliberate, capturing the elegance and strength of martial arts. The serene, isolated setting enhances the intensity and concentration of the practitioner. This visual presentation is an elegant interplay of motion and stillness, displaying the art form's discipline and grace.",
"video_path": "Man/mixkit-a-young-man-practicing-his-karate-moves-49635_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A tranquil coastal scene unfolds with a drone's aerial view capturing a serene beach landscape. The camera glides over a quiet stretch of sandy shoreline, where gentle waves kiss the shore under a clear blue sky. Nestled amidst lush palm trees are a series of traditional thatched-roof huts, their earthy tones blending harmoniously with the natural surroundings. The sandy beach stretches endlessly, bordered by the rhythmic dance of ocean waves on one side and verdant greenery on the other. A pair of white umbrellas is set up on the sand, suggesting a place to relax and enjoy the sun. In the distance, two small human figures can be seen walking leisurely along the water's edge, leaving faint footprints behind them. The scene exudes a calm and inviting atmosphere, with the soft rustle of palm leaves and the whisper of the ocean breeze almost audible. The overall composition is a captivating blend of nature's tranquility and architectural simplicity. This picturesque setting invites viewers to imagine themselves steps away from this idyllic coastal escape.",
"video_path": "beach/mixkit-sunny-beach-in-a-dynamic-shot-from-a-drone-44383_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A lone figure stands on a large, moss-covered rock, surrounded by the soft rush of a nearby stream. The figure is wearing white sneakers and shorts, with a plaid shirt that hangs loosely in the breeze. The lighting creates dramatic shadows, enhancing the textures of the rock and the subtle movement of the water below. In the background, a waterfall cascades into the stream, completing this tranquil and serene nature scene.",
"video_path": "forest/mixkit-woman-standing-in-front-of-waterfall-559_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In an industrial setting, a person leans casually against a railing, exuding a sense of confidence and composure. They are wearing a striking outfit, consisting of a vibrant, patterned jacket over a simple white crop top, creating a bold contrast. The atmosphere is infused with warm, ambient lighting that casts soft shadows on the concrete walls and metallic surfaces. Intricate wiring and pipes form an intricate backdrop, enhancing the urban aesthetic. Their relaxed posture and direct, engaging gaze suggest a sense of ease in this industrial environment. This scene encapsulates a blend of modern fashion and gritty, urban architecture, creating a visually compelling narrative.",
"video_path": "Fashion/mixkit-portrait-of-a-hipster-woman-walking-down-a-stairs-1297_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A man is energetically stretching in an open-air setting, surrounded by rows of vibrant red seats that suggest an amphitheater or outdoor venue. He wears a sleeveless black shirt layered with a hooded vest, emphasizing his athletic build as he engages in a warm-up routine. Behind him, the striking modern architecture of the building features geometric panels, with large sections of glass and overlapping metallic beams creating a dynamic backdrop. The scene captures the contrast between his focused movements and the static, bold design of the structure, while the surrounding greenery adds a touch of nature to the environment. The overall atmosphere is one of preparation and anticipation, with the man appearing determined and ready for an upcoming event or performance.",
"video_path": "Sport/mixkit-man-doing-arm-stretches-595_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A young woman is seated on the floor in front of a plush, beige tufted couch, fully engrossed in sorting through a stack of papers. Her dark hair falls loosely past her shoulders, and she wears a green plaid shirt, contributing to the casual yet focused atmosphere. She gently places the papers onto a small round white table, occasionally lifting individual sheets to examine them more closely. Her expression shifts subtly, reflecting concentration and contemplation as she processes the information on the pages. Two small, round nested tables hold her documents, along with a small plant in a gray pot, adding a touch of greenery to the scene. The background features a dark paneled wall, creating a contrasting backdrop for the light-colored furniture. The setting is tranquil and organized, the couch and tables arranged symmetrically, conveying a sense of harmony. A calculator rests on the smaller table, hinting at a task involving calculations or budgeting.",
"video_path": "Woman/mixkit-frustrated-woman-throws-paperwork-on-the-floor-4526_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A heavily rusted metal gate stands firmly locked, with two vertical bars joined by a thick, old chain that loops elegantly around them. The chain's texture is coarse and rugged, its surface reflecting varying shades of orange and brown, indicative of years exposed to the elements. At the heart of the chain, a black iron padlock, slightly worn yet imposing, secures the gate, its curves and edges smooth against the aged links. The gate's metalwork is outlined by a backdrop of soft, blurred greenery, suggesting a serene and isolated location beyond the barrier. Tall trees rise in the distance, their trunks and leaves creating a lush, forest-like setting that contrasts with the gate's severe rust. A pathway leads away from the gate, its surface uneven with patches of moss and weathered stone visible in the soft focus, inviting yet inaccessible. The ambiance is quiet and mysterious, with a sense of abandonment hanging subtly in the air, evoking curiosity about what lies beyond. Shadows play across the gate, cast by branches swaying gently in the breeze, adding to the dynamic interaction of light and texture. This scene, rich in detail and atmosphere, captures the viewer's imagination, evoking both the allure of the forbidden and the beauty of decay.",
"video_path": "forest/mixkit-rusty-fence-with-a-chain-of-a-property-in-nature-5294_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
"video_path": "People/mixkit-small-group-of-people-doing-yoga-together-43730_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
"video_path": "sea/mixkit-dolphins-underwater-4133_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a young woman stands against a vibrant graffiti-covered wall, deeply engrossed in her smartphone. Her expression reflects a mix of focus and subtle satisfaction as she interacts with the screen. She wears a black floral-patterned top, which contrasts with the bright, abstract shapes and bold colors of the mural behind her. As she continues to engage with her phone, a series of like count notifications appear on the screen, indicating a growing online appreciation. The wall behind her features a striking mix of geometric and organic shapes, including swirls of teal, orange, and black, with large humanoid figures in a pop-art style. Her long, light-brown hair frames her face, adding a calm, composed aura amidst the lively backdrop. The video captures a blend of contemporary digital interaction and expressive urban art, creating a dynamic yet harmonious scene.",
"video_path": "Girl/mixkit-girl-looking-at-the-likes-in-her-post-4914_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A young mother and her baby sit comfortably on a bed, surrounded by an inviting, cozy atmosphere. The woman, wearing a sleeveless top and jeans, is gently engaging with the baby, who is dressed in an adorable animal-print onesie. The child is seated on the bed with colorful toys scattered around, including a plush toy and a board book. The warm glow from a hanging lamp casts a soft light on them, enhancing the serene environment. Pillows are propped up against the headboard, providing a cushioned backdrop as the mother leans slightly over to interact with the baby. A small bottle is visible beside her, suggesting a nurturing setting. Her hand gestures animatedly as she holds up a soft, white cushion with red and blue accents, likely stimulating the baby\u2019s curiosity. Their shared moment is filled with affection and joy, a perfect snapshot of familial bonding.",
"video_path": "Baby/mixkit-loving-mother-and-her-baby-playing-with-soft-toys-49966_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A young girl with long brown hair sits at a round wooden table, engrossed in working on her laptop. The laptop screen is a vivid green, suggesting a green screen effect is in use. To her left, a doll dressed in a yellow and white outfit is casually laid on top of some books, adding a playful and innocent touch to the scene. The setting is cozy, with sheer curtains in the background allowing soft natural light to spill into the room. The girl's posture and focused attention on the laptop suggest she is either playing a game or learning something new. This serene and domestic atmosphere is complemented by the slight blur of a dark couch in the foreground, framing the focused activity of the child.",
"video_path": "Girl/mixkit-little-girl-doing-homework-on-a-laptop-4757_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "An expansive view of a calm bay reveals a fleet of sailboats, each anchored in a regimented line stretching toward the horizon. The water is a serene blue, reflecting the soft hues of the early morning sky. A gentle breeze is indicated by the subtle ripples trailing behind the boats, while a single, larger vessel cuts a distinct path, leaving a graceful wake in its journey to the open sea. On one side, a cluster of modern high-rise buildings stands, contrasting against the natural simplicity of the water, suggesting a blend of urban and marine life. The distant shoreline is barely visible, softened by the atmospheric perspective, giving a sense of endless waters meeting the sky. The overall mood is peaceful and orderly, with the boats appearing almost as sentinels guarding the expanse of the tranquil bay.",
"video_path": "beach/mixkit-flying-backwards-over-the-sea-near-a-coast-50187_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a person is standing in the center of a dark, featureless space, illuminated by a spotlight that emphasizes their presence. The individual is dressed in a traditional martial arts uniform, known as a gi, which is predominantly white with a black belt tied around the waist, indicating a high level of expertise. The background remains pitch black, creating a stark contrast with the brightly lit figure, ensuring complete focus on them. The person's expression is serious and focused, reflecting a deep sense of discipline and concentration. Their hands move gracefully, transitioning through various martial arts stances, demonstrating practiced skill and fluidity. The uniform's crisp fabric folds and subtly reflects the light, further highlighting each precise movement. Despite the simplicity of the environment, the scene is dynamic, with each motion capturing the essence of martial arts practice. The video effectively conveys a sense of calm strength and mastery, making it ideal for an AI to recreate with attention to posture, lighting, and attire.",
"video_path": "Sport/mixkit-karate-fighter-bowing-to-the-front-49706_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a dimly lit room bathed in a mix of neon purple and blue lights, a focused individual is seated in a gaming chair. She wears a white hoodie and large headphones with cat ears that glow softly, creating a striking silhouette. Her hands rest on a keyboard, typing swiftly as she concentrates intently on the screen in front of her. The atmosphere exudes a sense of intensity and immersion, with the soft-colored lighting enhancing the futuristic vibe. Her long hair cascades down her shoulders, adding a touch of elegance to the otherwise tech-centric setting. The overall scene captures the essence of a dedicated gamer deeply engaged in her virtual world.",
"video_path": "earth/mixkit-a-young-woman-wearing-headphones-with-rgb-lights-suddenly-gets-51621_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "Inside a dimly-lit bus, five individuals are seated along the rows of worn seats, each subtly illuminated by the colorful lights emanating from overhead. On the left, a woman sits with a relaxed posture, her curly hair accented by a patterned scarf, wearing a plaid outfit paired with bright neon socks. Next to her, a person clad in a denim jacket appears deep in thought, resting their head on a hand. Further back, another figure in a bucket hat and oversized yellow attire gazes across the aisle, evoking a sense of introspection. The atmosphere is enriched by the soft glow of red and green lights, bathing the bus interior in an almost surreal ambiance, creating a compelling tableau of urban life.",
"video_path": "Music/mixkit-conceptual-urban-fashion-42581_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "An aerial view captures two tennis players on a court, with one dressed in white on the left and another in red on the right. They are mid-game, each poised for action with rackets in hand, accentuated by their strategic positioning at opposite baselines. The court itself is a stark, deep blue, bordered by the vibrant green of the surrounding area, with a dark central net dividing the space. Long shadows stretch dramatically across the ground, suggesting a late afternoon setting. The subtly textured surface of the court contrasts with the crisp, white lines marking its boundaries and sections. This scene creates a vivid, balanced composition, highlighting both the competitive tension and serene atmosphere of the game.",
"video_path": "People/mixkit-two-people-playing-tennis-aerial-view-880_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a vibrant, dreamlike setting, a lone figure moves energetically against a backdrop of deep blue and purple hues, casting emotive shadows that ripple with dynamic motion. The figure, almost obscured by a smeared effect, suggests a rhythmic dance or a passionate performance, arms blurred as they sweep through colorful, streaked lighting. A neon glow accentuates their form, particularly highlighting the face which is abstractly illuminated in bursts of orange and red, suggesting intense emotional expression. The scene is dominated by two primary elements \u2013 the figure\u2019s motion and the dramatic lighting, creating a synergy of human emotion and visual spectacle. Swirling trails of light seem to intertwine with the figure, like a visual symphony of movement and color that floods the space. The lighting changes, casting intricate patterns on the figure and the surrounding space, giving the impression of a kaleidoscope in motion. Despite the blurred and abstract portrayal, there is a sense of focus conveyed through the figure\u2019s intent movements, akin to a conductor orchestrating a visual and auditory performance. The environment resonates with an electric energy, suggesting a seamless fusion of art and technology. As the visual drama unfolds, the scene invites viewers to lose themselves in the abstract dance and the play of vivid luminance.",
"video_path": "Music/mixkit-dancer-dancing-with-a-light-bar-in-his-hands-42221_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a brightly lit studio, a photographer wearing a denim jacket focuses intently, capturing shots with a professional camera. Facing him, a model stands gracefully, adjusting her long, flowing hair with delicate movements. The scene is characterized by strong contrasts; the model's soft pink attire and gentle gestures complement the rugged, precise demeanor of the photographer. Positioned against a minimalist backdrop, the pair work seamlessly, with the camera\u2019s lens pointed directly at the model, capturing her elegance. The soft, diffused lighting casts a gentle glow on both subjects, creating an airy and ethereal atmosphere perfect for a high-fashion photo shoot.",
"video_path": "Fashion/mixkit-professional-photo-session-with-a-young-female-model-41621_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video showcases a serene, expansive landscape covered with a variety of trees dotting the hills. The hills gently slope across the frame, with patches of dry grass contrasting against the lush green foliage. Tall trees with dense canopies stand elegantly, casting soft shadows on the ground below. The sunlight bathes the entire scene, highlighting the varied textures of the leaves and terrain. Gaps between the trees reveal a narrow dirt path meandering through the hills, suggesting a sense of quiet solitude. The undulating hills extend into the distance, creating depth and a calming sense of vast space. The verdant hues of the leaves contrast with the earthy tones of the hills, enhancing the visual richness. In the background, a faint outline of distant hills can be seen, blurred softly by the atmospheric perspective. This tranquil setting could be efficiently recreated in a virtual environment by focusing on its layered composition, color palette, and natural textures.",
"video_path": "forest/mixkit-aerial-panorama-of-a-sunny-mountain-landscape-40846_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
"video_path": "Car/mixkit-skiers-on-a-snowy-slope-3327_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The scene unfolds within a dimly lit bus, where three young individuals are seated, each absorbed in their unique world. To the left, a person with tied-back hair rests their head on their hand, dressed casually in a jacket and jeans, projecting a relaxed demeanor. Central to the frame is another individual, sitting upright with intense focus, donning a plaid blazer and oversize hoops, enhancing their confident presence. The muted green and red lighting casts an atmospheric glow, adding depth and intrigue to the setting. On the right, a person in a bucket hat and striped shirt leans back, appearing contemplative as they adjust their hat with a nonchalant gesture. The interplay of light and shadow highlights their expressions, creating an intimate and cinematic ambiance. Together, these figures form a cohesive tableau, capturing a moment of introspection amid a bustling yet serene urban environment.",
"video_path": "City/mixkit-three-models-posing-to-the-lens-while-on-board-a-42575_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
"video_path": "Sport/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A woman stands confidently in front of a large array of solar panels, her navy blue jumpsuit contrasting against the lush green grass beneath her feet. Her expression is calm and focused, eyes facing directly ahead, suggesting a deep connection to the subject matter\u2014renewable energy. The sunlight bathes the scene in warm hues, casting gentle shadows and highlighting the geometric precision of the solar panels' grid-like structure. The background reveals a blend of nature and technology, as the panels are anchored on a grassy slope with foliage on the left side of the frame. This composition captures a harmonious blend of human innovation and environmental consciousness, accentuated by the serene outdoor setting.",
"video_path": "Business/mixkit-woman-standing-in-front-of-a-solar-panel-4880_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, two people are working at a wooden desk, using an iMac computer. One person, wearing a white knit sweater, is using the apple wireless mouse with their right hand, while their left hand rests on the sleek white keyboard. Their movements are smooth yet intentional, suggesting they are focused on a task on the computer screen. The monitor displays a well-organized array of files and folders, hinting at a task that involves detailed organization or detailed data navigation. The second person, only subtly visible, sits closely by and appears to observe or assist, creating a collaborative atmosphere. Their presence adds a quiet dynamic to the scene, as if they are ready to provide input or guidance. Sticky notes with handwritten notes are attached to the monitor\u2019s stand, adding a touch of personal organization amidst the digital workspace. The focus on the keyboard and mouse emphasizes a streamlined workflow, indicative of a productive work environment. The overall ambiance is calm and focuses on teamwork, technology, and efficient workspace management.",
"video_path": "People/mixkit-person-with-glasses-working-on-a-desktop-computer-3248_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A man stands in front of a modern glass facade, taking off a dark hoodie to reveal his gray tank top underneath. His arms are lifted high as he maneuvers the hoodie over his head, showcasing a fluid motion that conveys a sense of calm and routine. The lighting highlights the contours of his muscles, emphasizing a combination of strength and quiet determination. Behind him, the reflective surface of the glass panels provides a subtle backdrop, enhancing the focus on his focused and serene demeanor.",
"video_path": "Sport/mixkit-man-puts-on-sleeveless-hoodie-603_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video displays a captivating dance of fiery orange flames against a stark black background, creating an intense visual contrast. The flames twist and intertwine, forming symmetrical, swirling patterns that expand and contract rhythmically across the frame. Each fiery tendril seems to be alive, moving with an almost hypnotic fluidity that captures the viewer's attention. The illumination from the flames casts subtle shadows, enhancing the depth and texture of the scene. Overall, the dynamic movement and vibrant color palette create an atmosphere of both beauty and power.",
"video_path": "fire/mixkit-two-orange-flames-on-black-background-685_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In this scene, a person is seated in a dimly lit room, possibly a recording studio, holding several drumsticks in their hands. The individual's face is partially obscured by sunglasses, adding a touch of mystery to their demeanor. They are wearing a colorful, patterned shirt with a mix of orange and blue tones that stands out against the darker background. The person appears focused and engaged with the drumsticks, their hands prominently displayed. The ambient light casts warm, soft shadows, emphasizing the texture and colors of their shirt and the wooden drumsticks. The room features wooden paneling, which complements the overall cozy, music-centric setting of the scene. The use of perspective centers on the drumsticks, highlighting the importance of rhythm and music in the captured moment.",
"video_path": "Music/mixkit-drummer-stretching-before-playing-42783_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A man is casually sitting on a sofa, engrossed in his meal and entertainment. He is holding a TV remote in one hand while reaching for food with the other, indicating a laid-back, comfortable evening. The table before him is filled with takeout containers, revealing a variety of appetizers and dishes, suggestive of a casual dining experience at home. The background is defined by colorful patterned cushions, adding a cozy, homey feel to the scene. Warm, ambient lighting highlights the relaxed atmosphere, casting soft shadows that contribute to the intimate setting. In this moment, he takes a bite of a sandwich, comfortably balancing his attention between food and whatever is playing on the screen.",
"video_path": "Man/mixkit-man-watching-tv-and-eating-fast-food-26089_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The scene opens to a breathtaking view of a tranquil ocean horizon at dusk, displaying a vibrant tapestry of oranges, pinks, and purples as the sun sets. In the foreground, tall, swaying palm trees frame the scene, their silhouettes stark against the colorful sky. The ocean itself shimmers with reflections of the sunset, creating a peaceful, almost ethereal atmosphere. A small boat can be seen in the distance, centered on the horizon, adding a sense of scale and solitude to the scene. The waves gently lap the shore, creating faint patterns on the sandy beach, which stretches across the foreground. Above, the sky is dotted with scattered clouds that catch the last light of the day, enhancing the drama and beauty of the scene. The overall mood is serene and contemplative, capturing a perfect moment of nature\u2019s grandeur.",
"video_path": "beach/mixkit-sunset-with-sailing-boats-2166_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A man sits hunched on a couch, the weight of emotions clearly visible on his posture. He wears a simple, gray t-shirt, and his head is bowed, resting in his hands, which cover most of his face, obscuring his features. The gentle light filtering through sheer curtains in the background casts a soft glow upon him, emphasizing the contrast between his static form and the hazy brightness behind. His elbows rest upon his knees, suggesting a posture of deep contemplation or distress. The simplicity of the room, with its muted colors, highlights the focus on the man's internal struggle. Delicate detailing on the fabric of his shirt adds texture, enhancing the scene's realism. Subtle changes in the natural light indicate the passage of time, as the man remains unmoving, absorbed in thought. This intimate moment captures a profound vulnerability, making the scene universally relatable and poignant.",
"video_path": "Man/mixkit-worried-and-sad-man-with-his-head-down-4701_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A pair of hands, belonging to an unseen figure, carefully unrolls a large sheet of crisp, white paper on a dark wooden table. The lighting is warm, casting a gentle glow that highlights the textures of the paper and the wood grain of the table. As the paper unfurls, the edges reveal the faint beginnings of a colorful map printed on its surface. The arms, clad in a casual gray T-shirt, suggest a relaxed and focused task at hand. Each motion is deliberate, with fingers deftly guiding the paper, ensuring it lays flat without creases. In the background, a hint of a red curtain can be seen, adding a touch of color and depth to the setting. The composition of the scene emphasizes the contrast between the bright paper and the rich tones of the surroundings. This serene and methodical action evokes a sense of exploration and preparation.",
"video_path": "Man/mixkit-unrolling-a-world-map-on-a-table-21626_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A young woman sits on a vibrant green seat inside a bus, illuminated by the soft glow of pink and blue lights. Her outfit is a striking mix of colors: a neon pink top paired with a jacket featuring dark sleeves, and jeans that provide a neutral contrast. She wears large, hoop earrings that catch the light as she moves slightly, exuding an air of cool confidence. Her gaze is directed thoughtfully to the side, suggesting contemplation or daydreaming during her commute. The metallic pole beside her adds a geometric element to the composition, reflecting the kaleidoscope of neon hues. The background is a clean, futuristic white, serving as a blank canvas that amplifies the neon atmosphere. Her relaxed posture and the modern bus setting create a scene that captures a blend of urban life and personal introspection.",
"video_path": "City/mixkit-fashion-model-posing-on-a-bus-42578_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.",
"video_path": "Car/mixkit-curve-on-a-snowy-forest-road-3317_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video showcases a vibrant urban skyline during twilight, with towering buildings reflecting the warm hues of the setting sun. A series of tall, cylindrical structures dominate the foreground, adjacent to a complex of industrial equipment and grids. The scene includes modern high-rise buildings with glass exteriors, capturing the evolving architecture of a bustling cityscape. A prominent structure labeled \"CITY OF AUSTIN POWER PLANT\" stands out, highlighting the industrial theme amidst the urban backdrop. The soft glow of city lights begins to pierce the approaching dusk, creating an inviting yet dynamic atmosphere. Shadows cast by the buildings add depth and contrast, emphasizing their massive scale and intricate designs. The overall composition is balanced between the natural light of the sunset and the artificial illumination of the city, offering a compelling visual narrative.",
"video_path": "Car/mixkit-slow-air-travel-in-reverse-over-a-big-city-49841_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the scene, a striking architectural structure dominates the view, bathed in a soft, ambient light. The enormous yellow arches serve as the centerpiece, drawing the eye upwards with their majestic curves and towering presence. The smooth, clean surfaces of the structure reflect the light, highlighting the texture and depth of the architecture. In the foreground, blurred streaks of headlights and taillights suggest the motion of vehicles passing by, adding dynamic energy to the otherwise still scene. The contrast between the fast-moving lights and the static arches creates a balanced composition. To the left, a lone streetlamp and a small tree provide a touch of nature and urban elements against the monumental backdrop. The night sky subtly peeks through the gaps in the structure, hinting at a clear, calm evening. Shadows from the arches create patterns on the ground, adding an intricate detail to the scene. Overall, the combination of light, shadow, and movement makes for a dramatic and visually captivating moment.",
"video_path": "Car/mixkit-a-fast-timelapse-of-the-street-with-a-monumental-yellow-50993_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A tranquil marina comes into full view under the golden hues of a setting sun. A collection of gleaming yachts and boats are neatly moored, their reflections shimmering softly on the gentle water. The sun's low position casts elongated shadows over the bustling harbor scene, while rolling hillsides surround the distant cityscape. The skyline is interspersed with modern buildings and clusters of residences, adding layers to the vibrant community. At the center, a broad wooden pier juts confidently into the harbor, extending an invitation for leisurely strolls. To the left, various shops and colorful structures line the waterfront, indicating a vibrant coastal economy. The entire atmosphere exudes a serene yet lively charm, balancing the hustle of maritime activity with the peacefulness of the encroaching dusk. It's a scene of calm anticipation, as if the whole place holds its breath before the night's events unfold.",
"video_path": "beach/mixkit-harbor-on-a-tourist-coast-with-many-boats-and-yachts-40077_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video features a confident individual standing atop a structure against a clear blue sky, exuding a sense of freedom and style. The person is clad in a striking yellow button-up shirt tied at the waist, and beneath it, they wear a simple white top that adds to their relaxed yet stylish appearance. Completing the ensemble are high-waisted white jeans paired with a black belt, adding a touch of contrast. Around their neck is a bold red scarf, providing a splash of color and an air of vintage flair. The person's sunglasses, tinted in yellow, reflect the sunlight and contribute to the overall cool and composed demeanor. Their hair is styled elegantly, pulled back with headphones resting over the ears, suggesting they are immersed in music. One hand casually grazes the headphones, while the other rests gently on the railing, grounding the individual in the moment. The scene is an effortless blend of fashion and tranquility, capturing the spirit of sunny, carefree days.",
"video_path": "Music/mixkit-standing-woman-listening-to-music-460_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A ballerina gracefully spins and moves across a pink-hued studio, her poised figure accentuated by a shimmering white tutu and bodice. The background, a continuous wash of soft pink, provides a serene and ethereal atmosphere, emphasizing her fluid movements. Her arms extend with elegance, highlighting the delicacy and precision of her ballet pose, while her focused expression adds intensity to the scene. The subtle details of her costume, combined with the pink monochromatic ambiance, create a dreamlike spectacle, ideal for an AI to envision a oneiric dance setting.",
"video_path": "Dance/mixkit-portrait-of-a-ballerina-spinning-with-pink-background-40163_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The scene unfolds with two human figures in the distance, making their way through a serene meadow, thick with tall golden grass swaying gently in the breeze. The sun hangs low in the sky, casting a soft, diffused glow that illuminates the landscape with a warm, ethereal light. These figures, clad in hiking gear, move deliberately, suggesting they're either embarking on or concluding a journey. Their silhouettes contrast against the lush greenery of the surrounding trees, whose branches reach out, framing the horizon. The play of light and shadow among the trees creates a quilt of textures, with each leaf catching a hint of the sun's dying rays. This tranquil setting evokes a sense of calm and adventure, capturing the quintessential beauty of nature\u2019s landscape.",
"video_path": "People/mixkit-landscape-in-nature-while-two-people-are-jogging-44348_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A large cargo ship is docked at an industrial port, its white superstructure contrasting with the deep green and yellow of its deck. The foreground is dominated by the calm, deep blue waters of the harbor, which reflect the vessel\u2019s imposing presence. Surrounding the ship, a series of industrial buildings and storage facilities are visible, hinting at the bustling activity of the port. The deck is intricately detailed, featuring an array of pipes, equipment, and railings, showcasing the ship's functionality and purpose. In the background, a paved area with green patches and a few parked vehicles adds to the busy, industrious atmosphere of the scene.",
"video_path": "sea/mixkit-empty-cargo-ship-waiting-at-the-port-4209_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A lone climber ascends a towering rock face, clad in a pink shirt and gray pants, displaying a determined and focused expression. The climber navigates the rugged surface, where the texture of the rock is peppered with natural pockets and crevices that offer handholds and footholds. Sunlight casts soft shadows across the cliff, highlighting the intricate patterns and the climber\u2019s strategic movements. The cliff looms high, with sparse vegetation breaking the monotony of the stone, while distant rocky formations form a dramatic backdrop against the clear blue sky. The climber\u2019s gear, including a harness and chalk bag, underscores the adventure and challenge woven into this majestic, vertical journey.",
"video_path": "Sport/mixkit-mountaineer-girl-climbing-a-steep-rocky-mountain-41089_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A person is seen in a close-up shot, skillfully adjusting the tuning pegs of a guitar, showcasing a focused and practiced hand. The image is in black and white, highlighting the contrast between the textures of the instrument and the clothing. The individual's shirt, visible in the background, adds a soft, subtle texture, while the dark tones of the guitar neck create depth in the scene. This composition captures a moment of concentration and finesse, perfect for recreating an intimate musical setting.",
"video_path": "Music/mixkit-guitarist-playing-so-inspired-black-and-white-shot-44178_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A musician is playing a large brass instrument with the words \"Brass Band\" clearly visible on its bell. The scene is set against a vibrant yellow backdrop, casting a warm glow on the subject. The musician wears a dark cap and a matching suit, adding a formal touch to his attire. He is deeply focused on his performance, with the instrument's intricate tubing adding complexity to the visual composition. The lighting creates dramatic shadows and highlights, emphasizing the musician's expression and the instrument's metallic sheen. This harmonious blend of color and form captures the essence of a live brass band performance.",
"video_path": "Music/mixkit-musician-playing-the-trombone-while-dancing-43752_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a lone musician stands gracefully in front of a grand cathedral, playing an accordion while surrounded by the lively water display of a central fountain. Dressed in a casual ensemble, he wears a light-colored shirt, dark pants, and a flat cap that gives him a vintage charm. His posture is relaxed, yet engaged, as he sways gently in rhythm with the music, casting soft shadows on the cobblestone steps beneath him. The backdrop features the cathedral's towering twin spires, with intricate stonework that casts a rich, historical aura around the scene. Sunlight bathes the entire setting, enhancing the golden hues of the cathedral facade and creating a halo-like effect around the musician. The fountain's water jets splash playfully, catching glimmers of light and adding a dynamic element to the tranquil atmosphere. The scene captures a harmonious blend of architectural majesty and human creativity, framed by the clear, azure sky that extends infinitely above. It's a vivid depiction of solitude and artistry, set against a timeless urban landscape.",
"video_path": "Music/mixkit-man-plays-an-accordion-in-front-of-a-fountain-630_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the tranquil video, a person sits in a meditative pose on a gentle hillside, silhouetted against the dawning sky. The person is facing the breathtaking sunrise, with their back slightly turned to the viewer, wearing a simple, light-colored shirt. Their right hand rests on their knee, fingers relaxed in a common meditation mudra, symbolizing calmness and peace. The sky, a stunning blend of soft oranges and deep purples, gradually brightens, casting a warm glow over the lush, green landscape. To the left, the outlines of distant urban buildings can be seen against the horizon, adding a contrast between nature and city life. A river reflecting the sky's colors meanders through the scene, lending a serene, flowing dynamic to the landscape. Trees rise and fall gently across the terrain, their leaves rustling only faintly in the morning breeze. The person remains still and focused, embodying a moment of mindfulness and connection with nature. This visual captures a harmonious balance, evoking a sense of tranquility and introspection.",
"video_path": "City/mixkit-girl-meditating-in-yoga-pose-at-sunset-4803_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A serene landscape video captures a breathtaking panoramic view of a vast valley covered in a gentle mist. The undulating hills are lush with dense greenery, their rich foliage creating a vibrant border on the left side of the frame. The mist weaves through the landscape like a soft, ethereal blanket, lending a dream-like quality to the scene. In the distance, several mountain peaks emerge, their dark outlines contrasting against the pale blue sky. A few faint, wispy clouds drift lazily across the horizon, complementing the tranquil atmosphere. The sunlight filters through the haze, casting a warm glow and highlighting different textures of the flora. The overall mood is calm and contemplative, inviting the viewer to pause and appreciate nature's untouched beauty. The composition emphasizes depth and expansiveness, drawing attention to the harmony between earth and sky. This captivating scene embodies tranquility, offering a perfect backdrop for meditation or relaxation.",
"video_path": "forest/mixkit-flying-over-a-hill-with-a-view-of-the-surrounding-49743_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In this scene, a bearded individual is intently focused on their smartphone, with the sun setting in the background, casting a warm glow across the cityscape. The person, partially visible, is wearing a dark, buttoned shirt that contrasts with the golden hue of the sunset. Their hands are holding the smartphone delicately but purposefully, reflecting a sense of engagement and focus on the screen. The sunlight creates a striking lens flare effect, enhancing the dramatic atmosphere of the moment as it glimmers off the phone\u2019s surface. The surrounding environment hints at an elevated vantage point, providing a panoramic view of the urban landscape below.",
"video_path": "City/mixkit-guy-texting-at-sunset-265_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In an expansive, industrial space defined by towering columns and high ceilings, a solitary figure takes center stage. The person, dressed in dark, fitted clothing, assumes a powerful, dynamic stance with one leg bent forward and both arms outstretched in a horizontal arc. Framing this pose are intense flames that engulf their arms, creating a striking visual contrast against the muted tones of the room. The fire forms a brilliant halo of orange and yellow, casting flickering shadows on the weathered walls and worn, tiled floor. This interplay between light and dark showcases the dancer's poise and agility, as they maintain balance amidst the intense heat. Windows line the background, their panes dimly illuminated by the daylight filtering in, adding depth and perspective to the scene. The entire performance evokes a sense of raw energy and elemental mastery, as the figure continues to manipulate the fire in a seamless, mesmerizing display.",
"video_path": "fire/mixkit-expert-juggler-doing-tricks-with-a-stick-with-fire-43663_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A man is playing the violin, focused intently on his music. His fingers gracefully dance along the strings, flawlessly executing each note. He holds the violin close to his chin with a sense of familiarity and expertise. The rich, warm tones of the violin reflect in the soft lighting of the room. He wears a dark shirt, and a subtle necklace rests against his chest, adding a personal touch to his attire. The bow moves smoothly across the strings, producing a melody that seems to fill the space with emotion. His expression is one of concentration and passion, immersing himself fully in the performance. The background is softly blurred, bringing the violin's intricate craftsmanship and his precise movements into sharp focus. This serene and intimate moment captures the essence of his musical artistry.",
"video_path": "Music/mixkit-fiddler-playing-a-song-639_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the dimly lit parking garage, two figures engage in an impromptu game of soccer. The first person, wearing a light grey shirt and black pants with three white stripes, skillfully maneuvers the ball with precise footwork. The ground is slick with patches of water, reflecting the vibrant neon lights above. A second figure, clad in dark clothing, stands poised in the background, ready to intercept. The space is defined by stark yellow lines and orange safety bollards, adding structure to the chaotic energy of the scene. The soccer ball glides smoothly across the wet floor, kicking up droplets as it passes. Despite the muted colors of the environment, the players' movements are dynamic and full of life. Their shadowy silhouettes dance with the reflecting light, creating a mesmerizing visual interplay. The atmosphere is charged with focus and camaraderie, encapsulating the essence of a late-night urban soccer experience.",
"video_path": "Sport/mixkit-player-making-skillful-play-in-a-street-soccer-game-43504_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A lone climber is seen scaling a towering vertical rock face, demonstrating remarkable strength and focus. Dressed in a light-colored shirt and jeans, the climber grips the stone tightly, navigating the rough textures and crevices with precision. The sheer cliff is massive, exhibiting a range of natural hues from light tan to deep gray, accentuating the climber's figure against the vast rocky backdrop. Surrounding the cliff, scattered greenery and rugged terrain provide a sense of wilderness and isolation. The scene portrays a daring ascension requiring concentration and skill, capturing the essence of human endeavor against nature's formidable beauty.",
"video_path": "Sport/mixkit-skilled-mountaineer-climbing-a-gigantic-mountain-41083_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In this serene landscape, a lush meadow stretches across the foreground, dotted with vibrant yellow wildflowers swaying gently in the breeze. A towering tree stands majestically on the right side, its branches reaching wide under the bright blue sky filled with fluffy white clouds. On the left, dense trees form a natural corridor leading to the horizon, suggesting a sense of journey and possibility. The richness of the green grass contrasts beautifully with the golden hue of the distant fields, creating a harmonious palette of nature\u2019s colors. The play of light and shadow adds depth and dimension, evoking a tranquil, inviting atmosphere. It's a scene where nature\u2019s beauty simply commands attention, offering a perfect escape into tranquility.",
"video_path": "sky/mixkit-countryside-meadow-4075_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
"video_path": "mountain/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a cozy, dimly lit caf\u00e9, a woman sits alone at a rustic wooden table, fully engrossed in her reading. Her dark, wavy hair frames her face as she leans forward over an open book, suggesting deep focus and contemplation. The caf\u00e9\u2019s ambiance is warm, with hanging pendant lights casting a soft glow over the wooden shelves lined with jars and coffee paraphernalia in the background. A small cup of coffee rests just within her reach, alongside a glass dome encasing a solitary pastry, adding a touch of tranquility to the scene. Her casual attire, a denim jacket over a simple shirt, complements the laid-back, comfortable setting of the caf\u00e9. The contrast between her concentrated expression and the bustling, yet subdued caf\u00e9 atmosphere creates a harmonious, serene visual. The overall composition captures a quiet moment of introspection amidst the gentle hum of caf\u00e9 life.",
"video_path": "Woman/mixkit-woman-drinking-coffee-in-a-cafe-223_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a vast, deserted landscape under the night sky, a solitary figure stands at a small music setup, illuminated by strategically placed lights. The person is engrossed in playing a keyboard, with various electronic equipment surrounding them, casting soft glows of orange and blue hues across the scene. To the left, a large circular light adds a dramatic focal point, highlighting the intense contrast between the darkness and the lit performance area. This setup, with its minimalistic design and strategic lighting, creates a captivating and easily recognizable scene that merges the serene, expansive backdrop with an intimate, focused music performance.",
"video_path": "Music/mixkit-talented-dj-playing-in-a-lonely-desert-42414_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a bustling urban scene, cars zoom past a weathered building, their blurred motion a testament to the city\u2019s lively pace. The building, with its faded yellow and brown facade, boasts graffiti that speaks of both art and decay, framing the scene with an air of urban grit. A solitary figure stands slightly to the side, clad casually in a gray top and mustard trousers, gazing into the street, seemingly detached from the surrounding flurry. The motion of the traffic creates a dynamic contrast against the static backdrop, emphasizing the relentless movement of the city. As the video progresses, a bright yellow taxi appears, slowing down as it approaches the figure, adding a pop of color to the desaturated hues of the environment. The interaction suggests a routine, a possibly daily exchange between the driver and the pedestrian, hinting at the rhythms of city life. Overhead, a soft, overcast sky casts a diffused light, lending the scene a subdued, timeless quality. Small elements, like the vertical pole cutting through the frame and the distant chatter of urban sounds, complete this vivid tableau of urban existence.",
"video_path": "Car/mixkit-morning-in-the-street-time-lapse-1648_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "A young woman sits on a curb in a tranquil park, basking in the golden hue of the setting sun. Beside her, a collie dog rests calmly, its fur illuminated by the warm sunlight, creating a serene glow. The woman's hand gently strokes the dog's back, highlighting the bond and affection between them. Tall trees surround the pair, casting elongated shadows on the leaf-laden ground, adding to the peaceful and intimate ambiance of the scene.",
"video_path": "Pets/mixkit-a-woman-pets-a-dog-in-a-park-1562_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a grand, majestic elephant stands in an open, sunlit field, its massive form dominating the scene. The elephant's skin is a tapestry of earthy tones, with rough, textured wrinkles that add character to its already imposing presence. Its trunk, a powerful and flexible appendage, moves gently, swaying as the elephant possibly enjoys the warmth of the day. The background is a blur of greenery, suggesting a lively environment filled with trees and shrubs that provide a natural habitat. Light plays on the elephant's skin, highlighting patches of dust and dirt that give it an authentic wilderness look. The scene captures the tranquility and majesty of this gentle giant in its natural surroundings.",
"video_path": "Zoo/mixkit-wet-elephant-in-the-savanna-3663_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a fluffy dog with brown patches is intently engaged with a bright red toy shaped like a fire hydrant, which has a yellow and orange rope attached. The dog's body is relaxed as it lies on a plain white background, concentrating on nudging and playfully biting the toy. Its ears perk up slightly with curiosity, and its eyes are fixated on the toy, suggesting a scene of focused playfulness. The neutral tones of the dog's fur contrast starkly against the vivid red of the toy, creating a visually striking moment.",
"video_path": "Pets/mixkit-a-cute-border-collie-dog-play-with-a-fire-street-50662_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
}
]
}
+15 -16
View File
@@ -1,42 +1,41 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
OUTPUT_PATH = "video_samples_5b"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
model_id = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
model_id,
# FastVideo will automatically handle distributed setup
num_gpus=2,
num_gpus=1,
use_fsdp_inference=True,
use_cpu_offload=False
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
# image_encoder_cpu_offload=False,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.2-TI2V-5B-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
"A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic."
)
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
+18 -4
View File
@@ -16,8 +16,9 @@ def main():
num_gpus=1,
use_fsdp_inference=True,
# Adjust these offload parameters if you have < 32GB of VRAM
text_encoder_offload=False,
use_cpu_offload=False,
text_encoder_cpu_offload=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
VSA_sparsity=0.8,
)
load_end_time = time.perf_counter()
@@ -35,11 +36,24 @@ def main():
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
end_time = time.perf_counter()
gen_time = end_time - start_time
e2e_gen_time = end_time - start_time
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"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.")
start_time = time.perf_counter()
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
end_time = time.perf_counter()
gen_time2 = end_time - start_time
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate video: {gen_time} seconds")
print(f"Time taken for e2e generation: {e2e_gen_time} seconds")
print(f"Time taken to generate video2: {gen_time2} seconds")
if __name__ == "__main__":
+42
View File
@@ -0,0 +1,42 @@
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_5b"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
model_id = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
generator = VideoGenerator.from_pretrained(
model_id,
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
# image_encoder_cpu_offload=False,
)
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.2-TI2V-5B-Diffusers")
# sampling_param.num_frames = 45
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
# prompt2 = (
# "A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic."
# )
# video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
main()
+2 -2
View File
@@ -9,8 +9,8 @@ def main():
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
pipeline_config=config,
use_fsdp_inference=False, # Disable FSDP for MPS
use_cpu_offload=True,
text_encoder_offload=True,
dit_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
disable_autocast=False,
num_gpus=1,
+2 -2
View File
@@ -16,8 +16,8 @@ def main():
# if num_gpus > 1, FastVideo will automatically handle distributed setup
pipeline_config=pipeline_config,
use_fsdp_inference=False, # Disable FSDP for MPS
use_cpu_offload=True,
text_encoder_offload=True,
dit_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
disable_autocast=False,
num_gpus=1,
@@ -9,7 +9,7 @@ def main():
gen = VideoGenerator.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
use_cpu_offload=False,
dit_cpu_offload=False,
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
+145
View File
@@ -0,0 +1,145 @@
import dataclasses
from typing import Any, Optional
from fastvideo.configs.utils import update_config_from_args
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
@dataclasses.dataclass
class PreprocessConfig:
"""Configuration for preprocessing operations."""
# Model and dataset configuration
model_path: str = ""
dataset_path: str = ""
dataset_output_dir: str = "./output"
# Dataloader configuration
dataloader_num_workers: int = 1
preprocess_video_batch_size: int = 2
# Saver configuration
samples_per_file: int = 64
flush_frequency: int = 256
# Video processing parameters
max_height: int = 480
max_width: int = 848
num_frames: int = 163
video_length_tolerance_range: float = 2.0
train_fps: int = 30
speed_factor: float = 1.0
drop_short_ratio: float = 1.0
do_temporal_sample: bool = False
# Model configuration
training_cfg_rate: float = 0.0
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "preprocess") -> FlexibleArgumentParser:
"""Add preprocessing configuration arguments to the parser."""
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
preprocess_args = parser.add_argument_group("Preprocessing Arguments")
# Model & Dataset
preprocess_args.add_argument(f"--{prefix_with_dot}model-path",
type=str,
default=PreprocessConfig.model_path,
help="Path to the model for preprocessing")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-path",
type=str,
default=PreprocessConfig.dataset_path,
help="Path to the dataset directory for preprocessing")
preprocess_args.add_argument(
f"--{prefix_with_dot}dataset-output-dir",
type=str,
default=PreprocessConfig.dataset_output_dir,
help="The output directory where the dataset will be written.")
# Dataloader
preprocess_args.add_argument(
f"--{prefix_with_dot}dataloader-num-workers",
type=int,
default=PreprocessConfig.dataloader_num_workers,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process."
)
preprocess_args.add_argument(
f"--{prefix_with_dot}preprocess-video-batch-size",
type=int,
default=PreprocessConfig.preprocess_video_batch_size,
help="Batch size (per device) for the training dataloader.")
# Saver
preprocess_args.add_argument(f"--{prefix_with_dot}samples-per-file",
type=int,
default=PreprocessConfig.samples_per_file,
help="Number of samples per output file")
preprocess_args.add_argument(f"--{prefix_with_dot}flush-frequency",
type=int,
default=PreprocessConfig.flush_frequency,
help="How often to save to parquet files")
# Video processing parameters
preprocess_args.add_argument(f"--{prefix_with_dot}max-height",
type=int,
default=PreprocessConfig.max_height,
help="Maximum height for video processing")
preprocess_args.add_argument(f"--{prefix_with_dot}max-width",
type=int,
default=PreprocessConfig.max_width,
help="Maximum width for video processing")
preprocess_args.add_argument(f"--{prefix_with_dot}num-frames",
type=int,
default=PreprocessConfig.num_frames,
help="Number of frames to process")
preprocess_args.add_argument(
f"--{prefix_with_dot}video-length-tolerance-range",
type=float,
default=PreprocessConfig.video_length_tolerance_range,
help="Video length tolerance range")
preprocess_args.add_argument(f"--{prefix_with_dot}train-fps",
type=int,
default=PreprocessConfig.train_fps,
help="Training FPS")
preprocess_args.add_argument(f"--{prefix_with_dot}speed-factor",
type=float,
default=PreprocessConfig.speed_factor,
help="Speed factor for video processing")
preprocess_args.add_argument(f"--{prefix_with_dot}drop-short-ratio",
type=float,
default=PreprocessConfig.drop_short_ratio,
help="Ratio for dropping short videos")
preprocess_args.add_argument(
f"--{prefix_with_dot}do-temporal-sample",
action=StoreBoolean,
default=PreprocessConfig.do_temporal_sample,
help="Whether to do temporal sampling")
# Model Training configuration
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
return parser
@classmethod
def from_kwargs(cls, kwargs: dict[str,
Any]) -> Optional["PreprocessConfig"]:
"""Create PreprocessConfig from keyword arguments."""
preprocess_config = cls()
if not update_config_from_args(
preprocess_config, kwargs, prefix="preprocess", pop_args=True):
return None
return preprocess_config
def check_preprocess_config(self) -> None:
if self.dataset_path == "":
raise ValueError("dataset_path must be set for preprocess mode")
if self.samples_per_file <= 0:
raise ValueError("samples_per_file must be greater than 0")
if self.flush_frequency <= 0:
raise ValueError("flush_frequency must be greater than 0")
+1 -1
View File
@@ -1,7 +1,7 @@
{
"embedded_cfg_scale": 6,
"flow_shift": 17,
"use_cpu_offload": false,
"dit_cpu_offload": false,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -89,6 +90,7 @@ class WanVideoArchConfig(DiTArchConfig):
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: Optional[int] = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
+11 -2
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
from dataclasses import dataclass, field
import torch
@@ -9,6 +10,7 @@ from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96
decoder_base_dim: Optional[int] = None
z_dim: int = 16
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
@@ -51,14 +53,21 @@ class WanVAEArchConfig(VAEArchConfig):
2.8251,
1.9160,
)
temporal_compression_ratio = 4
spatial_compression_ratio = 8
is_residual: bool = False
in_channels: int = 3
out_channels: int = 3
patch_size: Optional[int] = None
scale_factor_temporal: int = 4
scale_factor_spatial: int = 8
clip_output: bool = True
def __post_init__(self):
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(
self.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
self.temporal_compression_ratio = self.scale_factor_temporal
self.spatial_compression_ratio = self.scale_factor_spatial
@dataclass
+3
View File
@@ -85,6 +85,9 @@ class PipelineConfig:
# DMD parameters
dmd_denoising_steps: list[int] | None = field(default=None)
# Wan2.2 TI2V parameters
ti2v_task: bool = False
# Compilation
# enable_torch_compile: bool = False
+6
View File
@@ -8,6 +8,7 @@ from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (FastWanT2V480PConfig,
Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
from fastvideo.logger import init_logger
@@ -26,7 +27,12 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers": FastWanT2V480PConfig,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
# "Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
# "Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
# Add other specific weight variants
}
+23
View File
@@ -111,3 +111,26 @@ class FastWanT2V480PConfig(WanT2V480PConfig):
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
"""Base configuration for FastWan T2V 1.3B 480P pipeline architecture with DMD"""
# Denoising stage
flow_shift: int = 5
ti2v_task: bool = True
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
pass
@dataclass
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
pass
+7
View File
@@ -23,6 +23,7 @@ class SamplingParam:
negative_prompt: str | None = None
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
# Batch info
num_videos_per_prompt: int = 1
@@ -106,6 +107,12 @@ class SamplingParam:
default=SamplingParam.output_path,
help="Path to save the generated video",
)
parser.add_argument(
"--output-video-name",
type=str,
default=SamplingParam.output_video_name,
help="Name of the output video",
)
parser.add_argument(
"--num-videos-per-prompt",
type=int,
+16 -3
View File
@@ -6,11 +6,14 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.wan import (FastWanT2V480PConfig,
from fastvideo.configs.sample.wan import (Wan2_1_Fun_1_3B_InP_SamplingParam,
Wan2_2_TI2V_5B_SamplingParam,
WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam)
WanT2V_14B_SamplingParam,
Wan2_2_TI2V_5B_SamplingParam,
Wan2_1_Fun_1_3B_InP_SamplingParam)
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -24,8 +27,18 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
# "Wan-AI/Wan2.2-T2V-A14B-Diffusers":
# Wan2_2_T2V_A14B_SamplingParam,
# "Wan-AI/Wan2.2-I2V-A14B-Diffusers":
# Wan2_2_I2V_A14B_SamplingParam,
# Add other specific weight variants
}
+48
View File
@@ -105,3 +105,51 @@ class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
height: int = 448
width: int = 832
fps: int = 16
# =============================================
# ============= Wan2.1 Fun Models =============
# =============================================
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
# =============================================
# ============= Wan2.2 TI2V Models =============
# =============================================
@dataclass
class Wan2_2_Base_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.2 TI2V 5B model."""
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
@dataclass
class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
"""Sampling parameters for Wan2.2 TI2V 5B model."""
# Video parameters
height: int = 704
width: int = 1280
num_frames: int = 121
fps: int = 24
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 50
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
pass
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
pass
+17 -1
View File
@@ -1,10 +1,11 @@
import argparse
from typing import Any
def update_config_from_args(config: Any,
args_dict: dict[str, Any],
prefix: str = "",
pop_args: bool = False) -> None:
pop_args: bool = False) -> bool:
"""
Update configuration object from arguments dictionary.
@@ -43,3 +44,18 @@ def update_config_from_args(config: Any,
for key in args_to_remove:
if key not in args_not_to_remove:
args_dict.pop(key)
return len(args_to_remove) > 0
def clean_cli_args(args: argparse.Namespace) -> dict[str, Any]:
"""
Clean the arguments by removing the ones that not explicitly provided by the user.
"""
provided_args = {}
for k, v in vars(args).items():
if (v is not None and hasattr(args, '_provided')
and k in args._provided):
provided_args[k] = v
return provided_args
+1 -1
View File
@@ -1,7 +1,7 @@
{
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"dit_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
@@ -1,7 +1,7 @@
{
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"dit_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
+13 -3
View File
@@ -58,9 +58,18 @@ class GenerateSubcommand(CLISubcommand):
"model_path must be provided either in config file or via --model-path"
)
if 'prompt' not in merged_args or not merged_args['prompt']:
# Check if either prompt or prompt_txt is provided
has_prompt = 'prompt' in merged_args and merged_args['prompt']
has_prompt_txt = 'prompt_txt' in merged_args and merged_args[
'prompt_txt']
if not (has_prompt or has_prompt_txt):
raise ValueError("Either prompt or prompt_txt must be provided")
if has_prompt and has_prompt_txt:
raise ValueError(
"prompt must be provided either in config file or via --prompt")
"Cannot provide both 'prompt' and 'prompt_txt'. Use only one of them."
)
init_args = {
k: v
@@ -73,11 +82,12 @@ class GenerateSubcommand(CLISubcommand):
}
model_path = init_args.pop('model_path')
prompt = generation_args.pop('prompt')
prompt = generation_args.pop('prompt', None)
generator = VideoGenerator.from_pretrained(model_path=model_path,
**init_args)
# Call generate_video - it handles both single and batch modes
generator.generate_video(prompt=prompt, **generation_args)
def validate(self, args: argparse.Namespace) -> None:
+77 -5
View File
@@ -9,6 +9,7 @@ diffusion models.
import math
import os
import time
from copy import deepcopy
from typing import Any
import imageio
@@ -98,15 +99,15 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
prompt: str | None = None,
sampling_param: SamplingParam | None = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray]:
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]]:
"""
Generate a video based on the given prompt.
Args:
prompt: The prompt to use for generation
prompt: The prompt to use for generation (optional if prompt_txt is provided)
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
output_path: Path to save the video (overrides the one in fastvideo_args)
output_video_name: Name of the video file to save. Default is the first 100 characters of the prompt.
@@ -123,8 +124,73 @@ class VideoGenerator:
callback_steps: Number of steps between each callback
Returns:
Either the output dictionary or the list of frames depending on return_frames
Either the output dictionary, list of frames, or list of results for batch processing
"""
# Handle batch processing from text file
if self.fastvideo_args.prompt_txt is not None:
prompt_txt_path = self.fastvideo_args.prompt_txt
if not os.path.exists(prompt_txt_path):
raise FileNotFoundError(
f"Prompt text file not found: {prompt_txt_path}")
# Read prompts from file
with open(prompt_txt_path, encoding='utf-8') as f:
prompts = [line.strip() for line in f if line.strip()]
if not prompts:
raise ValueError(f"No prompts found in file: {prompt_txt_path}")
logger.info("Found %d prompts in %s", len(prompts), prompt_txt_path)
if sampling_param is not None:
original_output_video_name = sampling_param.output_video_name
else:
original_output_video_name = None
results = []
for i, batch_prompt in enumerate(prompts):
logger.info("Processing prompt %d/%d: %s...", i + 1,
len(prompts), batch_prompt[:100])
try:
# Generate video for this prompt using the same logic below
if sampling_param is not None and original_output_video_name is not None:
sampling_param.output_video_name = original_output_video_name + f"_{i}"
result = self._generate_single_video(
batch_prompt, sampling_param, **kwargs)
# Add prompt info to result
if isinstance(result, dict):
result["prompt_index"] = i
result["prompt"] = batch_prompt
results.append(result)
logger.info("Successfully generated video for prompt %d",
i + 1)
except Exception as e:
logger.error("Failed to generate video for prompt %d: %s",
i + 1, e)
continue
logger.info(
"Completed batch processing. Generated %d videos successfully.",
len(results))
return results
# Single prompt generation (original behavior)
if prompt is None:
raise ValueError("Either prompt or prompt_txt must be provided")
return self._generate_single_video(prompt, sampling_param, **kwargs)
def _generate_single_video(
self,
prompt: str,
sampling_param: SamplingParam | None = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray]:
"""Internal method for single video generation"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
@@ -137,6 +203,8 @@ class VideoGenerator:
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
else:
sampling_param = deepcopy(sampling_param)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
@@ -222,6 +290,11 @@ class VideoGenerator:
output_path: {sampling_param.output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
# Use prompt[:100] for video name
if sampling_param.output_video_name is None:
sampling_param.output_video_name = prompt[:100]
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
@@ -229,7 +302,6 @@ class VideoGenerator:
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
output_video_name=kwargs.get("output_video_name", prompt[:100]),
)
# Run inference
+7
View File
@@ -27,6 +27,7 @@ if TYPE_CHECKING:
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_SERVER_DEV_MODE: bool = False
FASTVIDEO_STAGE_LOGGING: bool = False
def get_default_cache_root() -> str:
@@ -172,6 +173,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
# - "FLASH_ATTN": use FlashAttention
# - "SLIDING_TILE_ATTN" : use Sliding Tile Attention
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
# - "SAGE_ATTN": use Sage Attention
"FASTVIDEO_ATTENTION_BACKEND":
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
@@ -199,6 +201,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
# e.g. `/reset_prefix_cache`
"FASTVIDEO_SERVER_DEV_MODE":
lambda: bool(int(os.getenv("FASTVIDEO_SERVER_DEV_MODE", "0"))),
# If set, fastvideo will enable stage logging, which will print the time
# taken for each stage
"FASTVIDEO_STAGE_LOGGING":
lambda: bool(int(os.getenv("FASTVIDEO_STAGE_LOGGING", "0"))),
}
# end-env-vars-definition
+230 -13
View File
@@ -6,9 +6,12 @@ import argparse
import dataclasses
from contextlib import contextmanager
from dataclasses import field
from enum import Enum
from typing import Any
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.configs.utils import clean_cli_args
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
@@ -16,17 +19,58 @@ from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
def clean_cli_args(args: argparse.Namespace) -> dict[str, Any]:
class ExecutionMode(str, Enum):
"""
Clean the arguments by removing the ones that not explicitly provided by the user.
Enumeration for different pipeline modes.
Inherits from str to allow string comparison for backward compatibility.
"""
provided_args = {}
for k, v in vars(args).items():
if (v is not None and hasattr(args, '_provided')
and k in args._provided):
provided_args[k] = v
INFERENCE = "inference"
PREPROCESS = "preprocess"
FINETUNING = "finetuning"
DISTILLATION = "distillation"
return provided_args
@classmethod
def from_string(cls, value: str) -> "ExecutionMode":
"""Convert string to ExecutionMode enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid mode: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [mode.value for mode in cls]
class WorkloadType(str, Enum):
"""
Enumeration for different workload types.
Inherits from str to allow string comparison for backward compatibility.
"""
I2V = "i2v" # Image to Video
T2V = "t2v" # Text to Video
T2I = "t2i" # Text to Image
I2I = "i2i" # Image to Image
@classmethod
def from_string(cls, value: str) -> "WorkloadType":
"""Convert string to WorkloadType enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid workload type: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [workload.value for workload in cls]
# args for fastvideo framework
@@ -35,6 +79,12 @@ class FastVideoArgs:
# Model and path configuration (for convenience)
model_path: str
# Running mode
mode: ExecutionMode = ExecutionMode.INFERENCE
# Workload type
workload_type: WorkloadType = WorkloadType.T2V
# Cache strategy
cache_strategy: str = "none"
@@ -56,6 +106,7 @@ class FastVideoArgs:
dist_timeout: int | None = None # timeout for torch.distributed
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
preprocess_config: PreprocessConfig | None = None
# LoRA parameters
# (Wenxuan) prefer to keep it here instead of in pipeline config to not make it complicated.
@@ -67,9 +118,12 @@ class FastVideoArgs:
output_type: str = "pil"
use_cpu_offload: bool = True # For DiT
# CPU offload parameters
dit_cpu_offload: bool = True
use_fsdp_inference: bool = True
text_encoder_offload: bool = True
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
# STA (Sliding Tile Attention) parameters
@@ -85,9 +139,15 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
# Master port for distributed training/inference
master_port: int | None = None
# Stage verification
enable_stage_verification: bool = True
# Prompt text file for batch processing
prompt_txt: str | None = None
# model paths for correct deallocation
model_paths: dict[str, str] = field(default_factory=dict)
model_loaded: dict[str, bool] = field(default_factory=lambda: {
@@ -120,6 +180,24 @@ class FastVideoArgs:
help="Directory containing StepVideo model",
)
# Running mode
parser.add_argument(
"--mode",
type=str,
choices=ExecutionMode.choices(),
default=FastVideoArgs.mode.value,
help="The mode to run FastVideo",
)
# Workload type
parser.add_argument(
"--workload-type",
type=str,
choices=WorkloadType.choices(),
default=FastVideoArgs.workload_type.value,
help="The workload type",
)
# distributed_executor_backend
parser.add_argument(
"--distributed-executor-backend",
@@ -198,6 +276,15 @@ class FastVideoArgs:
help="Output type for the generated video",
)
# Prompt text file for batch processing
parser.add_argument(
"--prompt-txt",
type=str,
default=FastVideoArgs.prompt_txt,
help=
"Path to a text file containing prompts (one per line) for batch processing",
)
# STA (Sliding Tile Attention) parameters
parser.add_argument(
"--STA-mode",
@@ -226,7 +313,7 @@ class FastVideoArgs:
)
parser.add_argument(
"--use-cpu-offload",
"--dit-cpu-offload",
action=StoreBoolean,
help=
"Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
@@ -243,6 +330,17 @@ class FastVideoArgs:
help=
"Use CPU offload for text encoder. Enable if run out of memory.",
)
parser.add_argument(
"--image-encoder-cpu-offload",
action=StoreBoolean,
help=
"Use CPU offload for image encoder. Enable if run out of memory.",
)
parser.add_argument(
"--vae-cpu-offload",
action=StoreBoolean,
help="Use CPU offload for VAE. Enable if run out of memory.",
)
parser.add_argument(
"--pin-cpu-memory",
action=StoreBoolean,
@@ -265,6 +363,14 @@ class FastVideoArgs:
help="Validation sparsity for VSA",
)
# Master port for distributed training/inference
parser.add_argument(
"--master-port",
type=int,
default=FastVideoArgs.master_port,
help="Master port for distributed training/inference",
)
# Stage verification
parser.add_argument(
"--enable-stage-verification",
@@ -276,6 +382,9 @@ class FastVideoArgs:
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
# Add preprocessing configuration arguments
PreprocessConfig.add_cli_args(parser)
return parser
@classmethod
@@ -285,11 +394,26 @@ class FastVideoArgs:
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
kwargs: dict[str, Any] = {}
for attr in attrs:
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
kwargs['pipeline_config'] = pipeline_config
elif attr == 'preprocess_config':
preprocess_config = PreprocessConfig.from_kwargs(provided_args)
kwargs['preprocess_config'] = preprocess_config
elif attr == 'mode':
# Convert string to ExecutionMode enum
mode_value = getattr(args, attr, FastVideoArgs.mode.value)
kwargs['mode'] = ExecutionMode.from_string(
mode_value) if isinstance(mode_value, str) else mode_value
elif attr == 'workload_type':
# Convert string to WorkloadType enum
workload_type_value = getattr(args, 'workload_type',
FastVideoArgs.workload_type.value)
kwargs['workload_type'] = WorkloadType.from_string(
workload_type_value) if isinstance(
workload_type_value, str) else workload_type_value
# Use getattr with default value from the dataclass for potentially missing attributes
else:
# Get the field to check if it has a default_factory
@@ -308,7 +432,18 @@ class FastVideoArgs:
@classmethod
def from_kwargs(cls, **kwargs: Any) -> "FastVideoArgs":
# Convert mode string to enum if necessary
if 'mode' in kwargs and isinstance(kwargs['mode'], str):
kwargs['mode'] = ExecutionMode.from_string(kwargs['mode'])
# Convert workload_type string to enum if necessary
if 'workload_type' in kwargs and isinstance(kwargs['workload_type'],
str):
kwargs['workload_type'] = WorkloadType.from_string(
kwargs['workload_type'])
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
kwargs['preprocess_config'] = PreprocessConfig.from_kwargs(kwargs)
return cls(**kwargs)
def check_fastvideo_args(self) -> None:
@@ -316,6 +451,33 @@ class FastVideoArgs:
if current_platform.is_mps():
self.use_fsdp_inference = False
# Validate mode and inference_mode consistency
assert isinstance(
self.mode, ExecutionMode
), f"Mode must be an ExecutionMode enum, got {type(self.mode)}"
assert self.mode in ExecutionMode.choices(
), f"Invalid execution mode: {self.mode}"
# Validate workload type
assert isinstance(
self.workload_type, WorkloadType
), f"Workload type must be a WorkloadType enum, got {type(self.workload_type)}"
assert self.workload_type in WorkloadType.choices(
), f"Invalid workload type: {self.workload_type}"
if self.mode in [ExecutionMode.DISTILLATION, ExecutionMode.FINETUNING
] and self.inference_mode:
logger.warning(
"Mode is 'training' but inference_mode is True. Setting inference_mode to False."
)
self.inference_mode = False
elif self.mode in [ExecutionMode.INFERENCE, ExecutionMode.PREPROCESS
] and not self.inference_mode:
logger.warning(
"Mode is '%s' but inference_mode is False. Setting inference_mode to True.",
self.mode)
self.inference_mode = True
if not self.inference_mode:
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
@@ -346,6 +508,18 @@ class FastVideoArgs:
self.pipeline_config.check_pipeline_config()
# Add preprocessing config validation if needed
if self.mode == ExecutionMode.PREPROCESS:
if self.preprocess_config is None:
raise ValueError(
"preprocess_config is not set in FastVideoArgs when mode is PREPROCESS"
)
if self.preprocess_config.model_path == "":
self.preprocess_config.model_path = self.model_path
if not self.pipeline_config.vae_config.load_encoder:
self.pipeline_config.vae_config.load_encoder = True
self.preprocess_config.check_preprocess_config()
_current_fastvideo_args = None
@@ -491,6 +665,12 @@ class TrainingArgs(FastVideoArgs):
lora_alpha: int | None = None
lora_training: bool = False
# distillation args
generator_update_interval: int = 5
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
@@ -503,6 +683,19 @@ class TrainingArgs(FastVideoArgs):
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
elif attr == 'mode':
# Convert string to ExecutionMode enum
mode_value = getattr(args, attr, ExecutionMode.FINETUNING.value)
kwargs[attr] = ExecutionMode.from_string(
mode_value) if isinstance(mode_value, str) else mode_value
elif attr == 'workload_type':
# Convert string to WorkloadType enum
workload_type_value = getattr(args, 'workload_type',
WorkloadType.T2V.value)
kwargs[attr] = WorkloadType.from_string(
workload_type_value) if isinstance(
workload_type_value, str) else workload_type_value
# Use getattr with default value from the dataclass for potentially missing attributes
else:
# Get the field to check its default value
field = dataclasses.fields(cls)[next(
@@ -777,4 +970,28 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
# Distillation arguments
parser.add_argument("--generator-update-interval",
type=int,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
help="Minimum step ratio")
parser.add_argument("--max-timestep-ratio",
type=float,
default=TrainingArgs.max_timestep_ratio,
help="Maximum step ratio")
parser.add_argument("--real-score-guidance-scale",
type=float,
default=TrainingArgs.real_score_guidance_scale,
help="Teacher guidance scale")
return parser
def parse_int_list(value: str) -> list[int]:
if not value:
return []
return [int(x.strip()) for x in value.split(",")]
+15 -5
View File
@@ -72,6 +72,8 @@ class ComponentLoader(ABC):
module_loaders = {
"scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"),
"real_score_transformer": (TransformerLoader, "diffusers"),
"fake_score_transformer": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
@@ -247,10 +249,10 @@ class TextEncoderLoader(ComponentLoader):
target_device: torch.device,
fastvideo_args: FastVideoArgs,
dtype: str = "fp16"):
use_cpu_offload = fastvideo_args.text_encoder_offload and len(
use_cpu_offload = fastvideo_args.text_encoder_cpu_offload and len(
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
if fastvideo_args.text_encoder_offload:
if fastvideo_args.text_encoder_cpu_offload:
target_device = torch.device(
"mps") if current_platform.is_mps() else torch.device("cpu")
@@ -324,7 +326,10 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config)
target_device = get_local_torch_device()
if fastvideo_args.image_encoder_cpu_offload:
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
target_device = get_local_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(
model_path, encoder_config, target_device, fastvideo_args,
@@ -375,10 +380,15 @@ class VAELoader(ComponentLoader):
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
if fastvideo_args.vae_cpu_offload:
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
target_device = get_local_torch_device()
with set_default_torch_dtype(PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(get_local_torch_device())
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -441,7 +451,7 @@ class TransformerLoader(ComponentLoader):
device=get_local_torch_device(),
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
cpu_offload=fastvideo_args.use_cpu_offload,
cpu_offload=fastvideo_args.dit_cpu_offload,
fsdp_inference=fastvideo_args.use_fsdp_inference,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
+3 -1
View File
@@ -16,7 +16,9 @@ from typing import NoReturn, TypeVar, cast
import cloudpickle
from torch import nn
from fastvideo.logger import logger
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# huggingface class name: (component_name, fastvideo module name, fastvideo class name)
_TEXT_TO_VIDEO_DIT_MODELS = {
@@ -1,6 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2024 TSAIL Team and The HuggingFace Team. All rights reserved.
# Copyright 2025 TSAIL 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.
@@ -14,34 +12,36 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# DISCLAIMER: check https://arxiv.org/abs/2302.04867 and https://github.com/wl-zhao/UniPC for more info
# DISCLAIMER: check https://huggingface.co/papers/2302.04867 and https://github.com/wl-zhao/UniPC for more info
# The codebase is modified based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/schedulers/scheduling_dpmsolver_multistep.py
# ==============================================================================
#
# Modified from diffusers==0.33.0.dev0
# Modified from diffusers==0.35.0.dev0
#
# ==============================================================================
import math
from typing import List, Optional, Tuple, Union
import numpy as np
import scipy.stats
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
from diffusers.utils import deprecate, is_scipy_available
from diffusers.schedulers.scheduling_utils import KarrasDiffusionSchedulers, SchedulerMixin, SchedulerOutput
from fastvideo.models.schedulers.base import BaseScheduler
if is_scipy_available():
import scipy.stats
# Copied from diffusers.schedulers.scheduling_ddpm.betas_for_alpha_bar
def betas_for_alpha_bar(
num_diffusion_timesteps,
max_beta=0.999,
alpha_transform_type="cosine",
) -> torch.Tensor:
):
"""
Create a beta schedule that discretizes the given alpha_t_bar function, which defines the cumulative product of
(1-beta) over time from t = [0,1].
@@ -62,17 +62,16 @@ def betas_for_alpha_bar(
"""
if alpha_transform_type == "cosine":
def alpha_bar_fn(t: float) -> float:
return math.cos((t + 0.008) / 1.008 * math.pi / 2)**2
def alpha_bar_fn(t):
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
elif alpha_transform_type == "exp":
def alpha_bar_fn(t: float) -> float:
def alpha_bar_fn(t):
return math.exp(t * -12.0)
else:
raise ValueError(
f"Unsupported alpha_transform_type: {alpha_transform_type}")
raise ValueError(f"Unsupported alpha_transform_type: {alpha_transform_type}")
betas = []
for i in range(num_diffusion_timesteps):
@@ -83,9 +82,9 @@ def betas_for_alpha_bar(
# Copied from diffusers.schedulers.scheduling_ddim.rescale_zero_terminal_snr
def rescale_zero_terminal_snr(betas: torch.Tensor) -> torch.Tensor:
def rescale_zero_terminal_snr(betas):
"""
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
Rescales betas to have zero terminal SNR Based on https://huggingface.co/papers/2305.08891 (Algorithm 1)
Args:
@@ -108,8 +107,7 @@ def rescale_zero_terminal_snr(betas: torch.Tensor) -> torch.Tensor:
alphas_bar_sqrt -= alphas_bar_sqrt_T
# Scale so the first timestep is back to the old value.
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 -
alphas_bar_sqrt_T)
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
# Convert alphas_bar_sqrt to betas
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
@@ -176,6 +174,8 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
use_beta_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use beta sigmas for step sizes in the noise schedule during the sampling process. Refer to [Beta
Sampling is All You Need](https://huggingface.co/papers/2407.12173) for more information.
use_flow_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use flow 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.
@@ -200,7 +200,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
beta_start: float = 0.0001,
beta_end: float = 0.02,
beta_schedule: str = "linear",
trained_betas: np.ndarray | list[float] | None = None,
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
solver_order: int = 2,
prediction_type: str = "epsilon",
thresholding: bool = False,
@@ -209,44 +209,38 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
predict_x0: bool = True,
solver_type: str = "bh2",
lower_order_final: bool = True,
disable_corrector: tuple[int, ...] = (),
disable_corrector: List[int] = [],
solver_p: SchedulerMixin = None,
use_karras_sigmas: bool | None = False,
use_exponential_sigmas: bool | None = False,
use_beta_sigmas: bool | None = False,
use_flow_sigmas: bool | None = False,
flow_shift: float | None = 1.0,
use_karras_sigmas: Optional[bool] = False,
use_exponential_sigmas: Optional[bool] = False,
use_beta_sigmas: Optional[bool] = False,
use_flow_sigmas: Optional[bool] = False,
flow_shift: Optional[float] = 1.0,
timestep_spacing: str = "linspace",
steps_offset: int = 0,
final_sigmas_type: str | None = "zero", # "zero", "sigma_min"
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
rescale_betas_zero_snr: bool = False,
use_dynamic_shifting: bool = False,
time_shift_type: str = "exponential",
):
if sum([
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
self.config.use_karras_sigmas
]) > 1:
if self.config.use_beta_sigmas and not is_scipy_available():
raise ImportError("Make sure to install scipy if you want to use beta sigmas.")
if sum([self.config.use_beta_sigmas, self.config.use_exponential_sigmas, self.config.use_karras_sigmas]) > 1:
raise ValueError(
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
)
if trained_betas is not None:
self.betas = torch.tensor(trained_betas, dtype=torch.float32)
elif beta_schedule == "linear":
self.betas = torch.linspace(beta_start,
beta_end,
num_train_timesteps,
dtype=torch.float32)
self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32)
elif beta_schedule == "scaled_linear":
# this schedule is very specific to the latent diffusion model.
self.betas = torch.linspace(beta_start**0.5,
beta_end**0.5,
num_train_timesteps,
dtype=torch.float32)**2
self.betas = torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2
elif beta_schedule == "squaredcos_cap_v2":
# Glide cosine schedule
self.betas = betas_for_alpha_bar(num_train_timesteps)
else:
raise NotImplementedError(
f"{beta_schedule} is not implemented for {self.__class__}")
raise NotImplementedError(f"{beta_schedule} is not implemented for {self.__class__}")
if rescale_betas_zero_snr:
self.betas = rescale_zero_terminal_snr(self.betas)
@@ -263,7 +257,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.alpha_t = torch.sqrt(self.alphas_cumprod)
self.sigma_t = torch.sqrt(1 - self.alphas_cumprod)
self.lambda_t = torch.log(self.alpha_t) - torch.log(self.sigma_t)
self.sigmas = ((1 - self.alphas_cumprod) / self.alphas_cumprod)**0.5
self.sigmas = ((1 - self.alphas_cumprod) / self.alphas_cumprod) ** 0.5
# standard deviation of the initial noise distribution
self.init_noise_sigma = 1.0
@@ -272,27 +266,22 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
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__}")
raise NotImplementedError(f"{solver_type} is not implemented for {self.__class__}")
self.predict_x0 = predict_x0
# setable values
self.num_inference_steps: int | None = None
timesteps = np.linspace(0,
num_train_timesteps - 1,
num_train_timesteps,
dtype=np.float32)[::-1].copy()
self.num_inference_steps = None
timesteps = np.linspace(0, num_train_timesteps - 1, num_train_timesteps, dtype=np.float32)[::-1].copy()
self.timesteps = torch.from_numpy(timesteps)
self.model_outputs = [None] * solver_order
self.timestep_list: list[int | torch.Tensor] = [None] * solver_order
self.timestep_list = [None] * solver_order
self.lower_order_nums = 0
self.disable_corrector = list(disable_corrector)
self.disable_corrector = disable_corrector
self.solver_p = solver_p
self.last_sample = None
self._step_index: int | None = None
self._begin_index: int | None = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication
BaseScheduler.__init__(self)
@@ -324,9 +313,9 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
self._begin_index = begin_index
def set_timesteps(self,
num_inference_steps: int,
device: str | torch.device = None):
def set_timesteps(
self, num_inference_steps: int, device: Union[str, torch.device] = None, mu: Optional[float] = None
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
@@ -336,41 +325,40 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
"""
# "linspace", "leading", "trailing" corresponds to annotation of Table 2. of https://arxiv.org/abs/2305.08891
# "linspace", "leading", "trailing" corresponds to annotation of Table 2. of https://huggingface.co/papers/2305.08891
if mu is not None:
assert self.config.use_dynamic_shifting and self.config.time_shift_type == "exponential"
self.config.flow_shift = np.exp(mu)
if self.config.timestep_spacing == "linspace":
timesteps = (np.linspace(
0, self.config.num_train_timesteps - 1, num_inference_steps +
1).round()[::-1][:-1].copy().astype(np.int64))
timesteps = (
np.linspace(0, self.config.num_train_timesteps - 1, num_inference_steps + 1)
.round()[::-1][:-1]
.copy()
.astype(np.int64)
)
elif self.config.timestep_spacing == "leading":
step_ratio = self.config.num_train_timesteps // (
num_inference_steps + 1)
step_ratio = self.config.num_train_timesteps // (num_inference_steps + 1)
# creates integer timesteps by multiplying by ratio
# casting to int to avoid issues when num_inference_step is power of 3
timesteps = (np.arange(0, num_inference_steps + 1) *
step_ratio).round()[::-1][:-1].copy().astype(np.int64)
timesteps = (np.arange(0, num_inference_steps + 1) * step_ratio).round()[::-1][:-1].copy().astype(np.int64)
timesteps += self.config.steps_offset
elif self.config.timestep_spacing == "trailing":
step_ratio = self.config.num_train_timesteps / num_inference_steps
# creates integer timesteps by multiplying by ratio
# casting to int to avoid issues when num_inference_step is power of 3
timesteps = np.arange(self.config.num_train_timesteps, 0,
-step_ratio).round().copy().astype(np.int64)
timesteps = np.arange(self.config.num_train_timesteps, 0, -step_ratio).round().copy().astype(np.int64)
timesteps -= 1
else:
raise ValueError(
f"{self.config.timestep_spacing} is not supported. Please make sure to choose one of 'linspace', 'leading' or 'trailing'."
)
sigmas = np.array(
((1 - self.alphas_cumprod) / self.alphas_cumprod)**0.5)
sigmas = np.array(((1 - self.alphas_cumprod) / self.alphas_cumprod) ** 0.5)
if self.config.use_karras_sigmas:
log_sigmas = np.log(sigmas)
sigmas = np.flip(sigmas).copy()
sigmas = self._convert_to_karras(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array([
self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas
]).round()
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas]).round()
if self.config.final_sigmas_type == "sigma_min":
sigma_last = sigmas[-1]
elif self.config.final_sigmas_type == "zero":
@@ -383,10 +371,8 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
elif self.config.use_exponential_sigmas:
log_sigmas = np.log(sigmas)
sigmas = np.flip(sigmas).copy()
sigmas = self._convert_to_exponential(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array(
[self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
sigmas = self._convert_to_exponential(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
if self.config.final_sigmas_type == "sigma_min":
sigma_last = sigmas[-1]
elif self.config.final_sigmas_type == "zero":
@@ -399,10 +385,8 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
elif self.config.use_beta_sigmas:
log_sigmas = np.log(sigmas)
sigmas = np.flip(sigmas).copy()
sigmas = self._convert_to_beta(
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array(
[self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
sigmas = self._convert_to_beta(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
if self.config.final_sigmas_type == "sigma_min":
sigma_last = sigmas[-1]
elif self.config.final_sigmas_type == "zero":
@@ -413,12 +397,9 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
)
sigmas = np.concatenate([sigmas, [sigma_last]]).astype(np.float32)
elif self.config.use_flow_sigmas:
alphas = np.linspace(1, 1 / self.config.num_train_timesteps,
num_inference_steps + 1)
alphas = np.linspace(1, 1 / self.config.num_train_timesteps, num_inference_steps + 1)
sigmas = 1.0 - alphas
sigmas = np.flip(
self.config.flow_shift * sigmas /
(1 + (self.config.flow_shift - 1) * sigmas))[:-1].copy()
sigmas = np.flip(self.config.flow_shift * sigmas / (1 + (self.config.flow_shift - 1) * sigmas))[:-1].copy()
timesteps = (sigmas * self.config.num_train_timesteps).copy()
if self.config.final_sigmas_type == "sigma_min":
sigma_last = sigmas[-1]
@@ -432,8 +413,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
else:
sigmas = np.interp(timesteps, np.arange(0, len(sigmas)), sigmas)
if self.config.final_sigmas_type == "sigma_min":
sigma_last = ((1 - self.alphas_cumprod[0]) /
self.alphas_cumprod[0])**0.5
sigma_last = ((1 - self.alphas_cumprod[0]) / self.alphas_cumprod[0]) ** 0.5
elif self.config.final_sigmas_type == "zero":
sigma_last = 0
else:
@@ -443,8 +423,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
sigmas = np.concatenate([sigmas, [sigma_last]]).astype(np.float32)
self.sigmas = torch.from_numpy(sigmas)
self.timesteps = torch.from_numpy(timesteps).to(device=device,
dtype=torch.int64)
self.timesteps = torch.from_numpy(timesteps).to(device=device, dtype=torch.int64)
self.num_inference_steps = len(timesteps)
@@ -459,8 +438,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
# 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
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:
@@ -471,31 +449,25 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
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
https://huggingface.co/papers/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
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.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"
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)
@@ -503,7 +475,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
return sample
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._sigma_to_t
def _sigma_to_t(self, sigma, log_sigmas) -> torch.Tensor:
def _sigma_to_t(self, sigma, log_sigmas):
# get log sigma
log_sigma = np.log(np.maximum(sigma, 1e-10))
@@ -511,9 +483,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
dists = log_sigma - log_sigmas[:, np.newaxis]
# get sigmas range
low_idx = np.cumsum(
(dists >= 0),
axis=0).argmax(axis=0).clip(max=log_sigmas.shape[0] - 2)
low_idx = np.cumsum((dists >= 0), axis=0).argmax(axis=0).clip(max=log_sigmas.shape[0] - 2)
high_idx = low_idx + 1
low = log_sigmas[low_idx]
@@ -529,20 +499,18 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
return t
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._sigma_to_alpha_sigma_t
def _sigma_to_alpha_sigma_t(
self, sigma: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
def _sigma_to_alpha_sigma_t(self, sigma):
if self.config.use_flow_sigmas:
alpha_t = 1 - sigma
sigma_t = sigma
else:
alpha_t = 1 / ((sigma**2 + 1)**0.5)
alpha_t = 1 / ((sigma**2 + 1) ** 0.5)
sigma_t = sigma * alpha_t
return alpha_t, sigma_t
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
def _convert_to_karras(self, in_sigmas: torch.Tensor,
num_inference_steps) -> torch.Tensor:
def _convert_to_karras(self, in_sigmas: torch.Tensor, num_inference_steps) -> torch.Tensor:
"""Constructs the noise schedule of Karras et al. (2022)."""
# Hack to make sure that other schedulers which copy this function don't break
@@ -562,14 +530,13 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
rho = 7.0 # 7.0 is the value used in the paper
ramp = np.linspace(0, 1, num_inference_steps)
min_inv_rho = sigma_min**(1 / rho)
max_inv_rho = sigma_max**(1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
min_inv_rho = sigma_min ** (1 / rho)
max_inv_rho = sigma_max ** (1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
def _convert_to_exponential(self, in_sigmas: torch.Tensor,
num_inference_steps: int) -> torch.Tensor:
def _convert_to_exponential(self, in_sigmas: torch.Tensor, num_inference_steps: int) -> torch.Tensor:
"""Constructs an exponential noise schedule."""
# Hack to make sure that other schedulers which copy this function don't break
@@ -587,17 +554,13 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.exp(
np.linspace(math.log(sigma_max), math.log(sigma_min),
num_inference_steps))
sigmas = np.exp(np.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps))
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
def _convert_to_beta(self,
in_sigmas: torch.Tensor,
num_inference_steps: int,
alpha: float = 0.6,
beta: float = 0.6) -> torch.Tensor:
def _convert_to_beta(
self, in_sigmas: torch.Tensor, num_inference_steps: int, alpha: float = 0.6, beta: float = 0.6
) -> torch.Tensor:
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
# Hack to make sure that other schedulers which copy this function don't break
@@ -615,12 +578,15 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.array([
sigma_min + (ppf * (sigma_max - sigma_min)) for ppf in [
scipy.stats.beta.ppf(timestep, alpha, beta)
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
sigmas = np.array(
[
sigma_min + (ppf * (sigma_max - sigma_min))
for ppf in [
scipy.stats.beta.ppf(timestep, alpha, beta)
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
]
]
])
)
return sigmas
def convert_model_output(
@@ -650,8 +616,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
"missing `sample` as a required keyword argument")
raise ValueError("missing `sample` as a required keyword argument")
if timestep is not None:
deprecate(
"timesteps",
@@ -694,14 +659,15 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`, or"
" `v_prediction` for the UniPCMultistepScheduler.")
" `v_prediction` for the UniPCMultistepScheduler."
)
def multistep_uni_p_bh_update(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
order: int | None = None,
order: int = None,
**kwargs,
) -> torch.Tensor:
"""
@@ -721,20 +687,17 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
`torch.Tensor`:
The sample tensor at the previous timestep.
"""
prev_timestep = args[0] if len(args) > 0 else kwargs.pop(
"prev_timestep", None)
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")
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")
raise ValueError("missing `order` as a required keyword argument")
if prev_timestep is not None:
deprecate(
"prev_timestep",
@@ -751,8 +714,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
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]
sigma_t, sigma_s0 = self.sigmas[self.step_index + 1], self.sigmas[self.step_index]
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
@@ -766,149 +728,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
D1s = []
for i in range(1, order):
si = self.step_index - i
mi: torch.Tensor = 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)
D1s.append((mi - m0) / rk)
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_tensor: torch.Tensor = torch.stack(R)
b = torch.tensor(b, device=device)
D1s_tensor: torch.Tensor | None = None
if len(D1s) > 0:
D1s_tensor = 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:
rhos_p = torch.linalg.solve(R_tensor[:-1, :-1],
b[:-1]).to(device).to(x.dtype)
else:
D1s_tensor = None
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s_tensor is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p, D1s_tensor)
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_tensor is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p, D1s_tensor)
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 = None,
this_sample: torch.Tensor | None = None,
order: int | None = None,
**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]
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 = []
for i in range(1, order):
si = self.step_index - (i + 1)
mi: torch.Tensor = model_output_list[-(i + 1)]
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
@@ -943,8 +763,145 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
R = torch.stack(R)
b = torch.tensor(b, device=device)
D1s_tensor: torch.Tensor | None = torch.stack(
D1s, dim=1) if len(D1s) > 0 else None
if 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:
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)
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)
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: int = None,
**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]
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 = []
for i in range(1, order):
si = self.step_index - (i + 1)
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)
D1s.append((mi - m0) / rk)
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 len(D1s) > 0:
D1s = torch.stack(D1s, dim=1)
else:
D1s = None
# for order 1, we use a simplified version
if order == 1:
@@ -954,18 +911,16 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s_tensor is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1],
D1s_tensor)
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_tensor is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1],
D1s_tensor)
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
@@ -974,7 +929,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
return x_t
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.index_for_timestep
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
@@ -994,7 +949,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
return step_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._init_step_index
def _init_step_index(self, timestep) -> None:
def _init_step_index(self, timestep):
"""
Initialize the step_index counter for the scheduler.
"""
@@ -1009,10 +964,10 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def step(
self,
model_output: torch.Tensor,
timestep: int | torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
) -> SchedulerOutput | tuple:
) -> Union[SchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC.
@@ -1041,13 +996,11 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
if self.step_index is None:
self._init_step_index(timestep)
assert self.step_index is not None
use_corrector = (self.step_index > 0
and self.step_index - 1 not in self.disable_corrector
and self.last_sample is not None)
use_corrector = (
self.step_index > 0 and self.step_index - 1 not in self.disable_corrector and self.last_sample is not None
)
model_output_convert = self.convert_model_output(model_output,
sample=sample)
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,
@@ -1064,19 +1017,16 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.timestep_list[-1] = timestep
if self.config.lower_order_final:
this_order = min(self.config.solver_order,
len(self.timesteps) - self.step_index)
this_order = min(self.config.solver_order, len(self.timesteps) - self.step_index)
else:
this_order = self.config.solver_order
self.this_order: int = min(this_order, self.lower_order_nums +
1) # warmup for multistep
self.this_order = 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
model_output=model_output, # pass the original non-converted model output, in case solver-p is used
sample=sample,
order=self.this_order,
)
@@ -1085,16 +1035,14 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
self.lower_order_nums += 1
# upon completion increase step index by one
assert self._step_index is not None
self._step_index += 1
if not return_dict:
return (prev_sample, )
return (prev_sample,)
return SchedulerOutput(prev_sample=prev_sample)
def scale_model_input(self, sample: torch.Tensor, *args,
**kwargs) -> torch.Tensor:
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.
@@ -1117,25 +1065,18 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
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):
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)
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
]
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]
@@ -1152,4 +1093,4 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
return self.config.num_train_timesteps
+381 -45
View File
@@ -32,22 +32,123 @@ CACHE_T = 2
is_first_frame = contextvars.ContextVar("is_first_frame", default=False)
feat_cache = contextvars.ContextVar("feat_cache", default=None)
feat_idx = contextvars.ContextVar("feat_idx", default=0)
first_chunk = contextvars.ContextVar("first_chunk", default=None)
@contextmanager
def forward_context(first_frame_arg=False,
feat_cache_arg=None,
feat_idx_arg=None):
feat_idx_arg=None,
first_chunk_arg=None):
is_first_frame_token = is_first_frame.set(first_frame_arg)
feat_cache_token = feat_cache.set(feat_cache_arg)
feat_idx_token = feat_idx.set(feat_idx_arg)
first_chunk_token = first_chunk.set(first_chunk_arg)
try:
yield
finally:
is_first_frame.reset(is_first_frame_token)
feat_cache.reset(feat_cache_token)
feat_idx.reset(feat_idx_token)
first_chunk.reset(first_chunk_token)
class AvgDown3D(nn.Module):
def __init__(
self,
in_channels,
out_channels,
factor_t,
factor_s=1,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.factor_t = factor_t
self.factor_s = factor_s
self.factor = self.factor_t * self.factor_s * self.factor_s
assert in_channels * self.factor % out_channels == 0
self.group_size = in_channels * self.factor // out_channels
def forward(self, x: torch.Tensor) -> torch.Tensor:
pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
pad = (0, 0, 0, 0, pad_t, 0)
x = F.pad(x, pad)
B, C, T, H, W = x.shape
x = x.view(
B,
C,
T // self.factor_t,
self.factor_t,
H // self.factor_s,
self.factor_s,
W // self.factor_s,
self.factor_s,
)
x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
x = x.view(
B,
C * self.factor,
T // self.factor_t,
H // self.factor_s,
W // self.factor_s,
)
x = x.view(
B,
self.out_channels,
self.group_size,
T // self.factor_t,
H // self.factor_s,
W // self.factor_s,
)
x = x.mean(dim=2)
return x
class DupUp3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
factor_t,
factor_s=1,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.factor_t = factor_t
self.factor_s = factor_s
self.factor = self.factor_t * self.factor_s * self.factor_s
assert out_channels * self.factor % in_channels == 0
self.repeats = out_channels * self.factor // in_channels
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x.repeat_interleave(self.repeats, dim=1)
x = x.view(
x.size(0),
self.out_channels,
self.factor_t,
self.factor_s,
self.factor_s,
x.size(2),
x.size(3),
x.size(4),
)
x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
x = x.view(
x.size(0),
self.out_channels,
x.size(2) * self.factor_t,
x.size(4) * self.factor_s,
x.size(6) * self.factor_s,
)
_first_chunk = first_chunk.get()
if _first_chunk:
x = x[:, :, self.factor_t - 1 :, :, :]
return x
class WanCausalConv3d(nn.Conv3d):
r"""
@@ -158,20 +259,26 @@ class WanResample(nn.Module):
- 'downsample3d': 3D downsampling with zero-padding, convolution, and causal 3D convolution.
"""
def __init__(self, dim: int, mode: str) -> None:
def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None:
super().__init__()
self.dim = dim
self.mode = mode
# default to dim //2
if upsample_out_dim is None:
upsample_out_dim = dim // 2
# layers
if mode == "upsample2d":
self.resample = nn.Sequential(
WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
nn.Conv2d(dim, dim // 2, 3, padding=1))
nn.Conv2d(dim, upsample_out_dim, 3, padding=1),
)
elif mode == "upsample3d":
self.resample = nn.Sequential(
WanUpsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
nn.Conv2d(dim, dim // 2, 3, padding=1))
nn.Conv2d(dim, upsample_out_dim, 3, padding=1),
)
self.time_conv = WanCausalConv3d(dim,
dim * 2, (3, 1, 1),
padding=(1, 0, 0))
@@ -444,6 +551,40 @@ class WanMidBlock(nn.Module):
return x
class WanResidualDownBlock(nn.Module):
def __init__(self, in_dim, out_dim, dropout, num_res_blocks, temperal_downsample=False, down_flag=False):
super().__init__()
# Shortcut path with downsample
self.avg_shortcut = AvgDown3D(
in_dim,
out_dim,
factor_t=2 if temperal_downsample else 1,
factor_s=2 if down_flag else 1,
)
# Main path with residual blocks and downsample
resnets = []
for _ in range(num_res_blocks):
resnets.append(WanResidualBlock(in_dim, out_dim, dropout))
in_dim = out_dim
self.resnets = nn.ModuleList(resnets)
# Add the final downsample block
if down_flag:
mode = "downsample3d" if temperal_downsample else "downsample2d"
self.downsampler = WanResample(out_dim, mode=mode)
else:
self.downsampler = None
def forward(self, x):
x_copy = x.clone()
for resnet in self.resnets:
x = resnet(x)
if self.downsampler is not None:
x = self.downsampler(x)
return x + self.avg_shortcut(x_copy)
class WanEncoder3d(nn.Module):
r"""
@@ -462,6 +603,7 @@ class WanEncoder3d(nn.Module):
def __init__(
self,
in_channels: int = 3,
dim=128,
z_dim=4,
dim_mult=(1, 2, 4, 4),
@@ -470,6 +612,7 @@ class WanEncoder3d(nn.Module):
temperal_downsample=(True, True, False),
dropout=0.0,
non_linearity: str = "silu",
is_residual: bool = False, # wan 2.2 vae use a residual downblock
):
super().__init__()
self.dim = dim
@@ -486,26 +629,36 @@ class WanEncoder3d(nn.Module):
scale = 1.0
# init block
self.conv_in = WanCausalConv3d(3, dims[0], 3, padding=1)
self.conv_in = WanCausalConv3d(in_channels, dims[0], 3, padding=1)
# downsample blocks
self.down_blocks = nn.ModuleList([])
for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=True)):
# residual (+attention) blocks
for _ in range(num_res_blocks):
if is_residual:
self.down_blocks.append(
WanResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
self.down_blocks.append(WanAttentionBlock(out_dim))
in_dim = out_dim
WanResidualDownBlock(
in_dim,
out_dim,
dropout,
num_res_blocks,
temperal_downsample=temperal_downsample[i] if i != len(dim_mult) - 1 else False,
down_flag=i != len(dim_mult) - 1,
)
)
else:
for _ in range(num_res_blocks):
self.down_blocks.append(WanResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
self.down_blocks.append(WanAttentionBlock(out_dim))
in_dim = out_dim
# downsample block
if i != len(dim_mult) - 1:
mode = "downsample3d" if temperal_downsample[
i] else "downsample2d"
self.down_blocks.append(WanResample(out_dim, mode=mode))
scale /= 2.0
# downsample block
if i != len(dim_mult) - 1:
mode = "downsample3d" if temperal_downsample[i] else "downsample2d"
self.down_blocks.append(WanResample(out_dim, mode=mode))
scale /= 2.0
# middle blocks
self.mid_block = WanMidBlock(out_dim,
@@ -572,6 +725,83 @@ class WanEncoder3d(nn.Module):
x = self.conv_out(x)
return x
class WanResidualUpBlock(nn.Module):
"""
A block that handles upsampling for the WanVAE decoder.
Args:
in_dim (int): Input dimension
out_dim (int): Output dimension
num_res_blocks (int): Number of residual blocks
dropout (float): Dropout rate
temperal_upsample (bool): Whether to upsample on temporal dimension
up_flag (bool): Whether to upsample or not
non_linearity (str): Type of non-linearity to use
"""
def __init__(
self,
in_dim: int,
out_dim: int,
num_res_blocks: int,
dropout: float = 0.0,
temperal_upsample: bool = False,
up_flag: bool = False,
non_linearity: str = "silu",
):
super().__init__()
self.in_dim = in_dim
self.out_dim = out_dim
if up_flag:
self.avg_shortcut = DupUp3D(
in_dim,
out_dim,
factor_t=2 if temperal_upsample else 1,
factor_s=2,
)
else:
self.avg_shortcut = None
# create residual blocks
resnets = []
current_dim = in_dim
for _ in range(num_res_blocks + 1):
resnets.append(WanResidualBlock(current_dim, out_dim, dropout, non_linearity))
current_dim = out_dim
self.resnets = nn.ModuleList(resnets)
# Add upsampling layer if needed
if up_flag:
upsample_mode = "upsample3d" if temperal_upsample else "upsample2d"
self.upsampler = WanResample(out_dim, mode=upsample_mode, upsample_out_dim=out_dim)
else:
self.upsampler = None
self.gradient_checkpointing = False
def forward(self, x):
"""
Forward pass through the upsampling block.
Args:
x (torch.Tensor): Input tensor
feat_cache (list, optional): Feature cache for causal convolutions
feat_idx (list, optional): Feature index for cache management
Returns:
torch.Tensor: Output tensor
"""
x_copy = x.clone()
for resnet in self.resnets:
x = resnet(x)
if self.upsampler is not None:
x = self.upsampler(x)
if self.avg_shortcut is not None:
x = x + self.avg_shortcut(x_copy)
return x
class WanUpBlock(nn.Module):
"""
@@ -663,6 +893,8 @@ class WanDecoder3d(nn.Module):
temperal_upsample=(False, True, True),
dropout=0.0,
non_linearity: str = "silu",
out_channels: int = 3,
is_residual: bool = False,
):
super().__init__()
self.dim = dim
@@ -677,7 +909,6 @@ class WanDecoder3d(nn.Module):
# dimensions
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
scale = 1.0 / 2**(len(dim_mult) - 2)
# init block
self.conv_in = WanCausalConv3d(z_dim, dims[0], 3, padding=1)
@@ -693,33 +924,44 @@ class WanDecoder3d(nn.Module):
for i, (in_dim,
out_dim) in enumerate(zip(dims[:-1], dims[1:], strict=True)):
# residual (+attention) blocks
if i > 0:
if i > 0 and not is_residual:
# wan vae 2.1
in_dim = in_dim // 2
# Determine if we need upsampling
# determine if we need upsampling
up_flag = i != len(dim_mult) - 1
# determine upsampling mode, if not upsampling, set to None
upsample_mode = None
if i != len(dim_mult) - 1:
upsample_mode = "upsample3d" if temperal_upsample[
i] else "upsample2d"
if up_flag and temperal_upsample[i]:
upsample_mode = "upsample3d"
elif up_flag:
upsample_mode = "upsample2d"
# Create and add the upsampling block
up_block = WanUpBlock(
in_dim=in_dim,
out_dim=out_dim,
num_res_blocks=num_res_blocks,
dropout=dropout,
upsample_mode=upsample_mode,
non_linearity=non_linearity,
)
if is_residual:
up_block = WanResidualUpBlock(
in_dim=in_dim,
out_dim=out_dim,
num_res_blocks=num_res_blocks,
dropout=dropout,
temperal_upsample=temperal_upsample[i] if up_flag else False,
up_flag=up_flag,
non_linearity=non_linearity,
)
else:
up_block = WanUpBlock(
in_dim=in_dim,
out_dim=out_dim,
num_res_blocks=num_res_blocks,
dropout=dropout,
upsample_mode=upsample_mode,
non_linearity=non_linearity,
)
self.up_blocks.append(up_block)
# Update scale for next iteration
if upsample_mode is not None:
scale *= 2.0
# output blocks
self.norm_out = WanRMS_norm(out_dim, images=False)
self.conv_out = WanCausalConv3d(out_dim, 3, 3, padding=1)
self.conv_out = WanCausalConv3d(out_dim, out_channels, 3, padding=1)
self.gradient_checkpointing = False
@@ -776,6 +1018,74 @@ class WanDecoder3d(nn.Module):
x = self.conv_out(x)
return x
def patchify(x, patch_size):
if patch_size == 1:
return x
if x.dim() == 4:
# x shape: [batch_size, channels, height, width]
batch_size, channels, height, width = x.shape
# Ensure height and width are divisible by patch_size
if height % patch_size != 0 or width % patch_size != 0:
raise ValueError(f"Height ({height}) and width ({width}) must be divisible by patch_size ({patch_size})")
# Reshape to [batch_size, channels, height//patch_size, patch_size, width//patch_size, patch_size]
x = x.view(batch_size, channels, height // patch_size, patch_size, width // patch_size, patch_size)
# Rearrange to [batch_size, channels * patch_size * patch_size, height//patch_size, width//patch_size]
x = x.permute(0, 1, 3, 5, 2, 4).contiguous()
x = x.view(batch_size, channels * patch_size * patch_size, height // patch_size, width // patch_size)
elif x.dim() == 5:
# x shape: [batch_size, channels, frames, height, width]
batch_size, channels, frames, height, width = x.shape
# Ensure height and width are divisible by patch_size
if height % patch_size != 0 or width % patch_size != 0:
raise ValueError(f"Height ({height}) and width ({width}) must be divisible by patch_size ({patch_size})")
# Reshape to [batch_size, channels, frames, height//patch_size, patch_size, width//patch_size, patch_size]
x = x.view(batch_size, channels, frames, height // patch_size, patch_size, width // patch_size, patch_size)
# Rearrange to [batch_size, channels * patch_size * patch_size, frames, height//patch_size, width//patch_size]
x = x.permute(0, 1, 4, 6, 2, 3, 5).contiguous()
x = x.view(batch_size, channels * patch_size * patch_size, frames, height // patch_size, width // patch_size)
else:
raise ValueError(f"Invalid input shape: {x.shape}")
return x
def unpatchify(x, patch_size):
if patch_size == 1:
return x
if x.dim() == 4:
# x shape: [b, (c * patch_size * patch_size), h, w]
batch_size, c_patches, height, width = x.shape
channels = c_patches // (patch_size * patch_size)
# Reshape to [b, c, patch_size, patch_size, h, w]
x = x.view(batch_size, channels, patch_size, patch_size, height, width)
# Rearrange to [b, c, h * patch_size, w * patch_size]
x = x.permute(0, 1, 4, 2, 5, 3).contiguous()
x = x.view(batch_size, channels, height * patch_size, width * patch_size)
elif x.dim() == 5:
# x shape: [batch_size, (channels * patch_size * patch_size), frame, height, width]
batch_size, c_patches, frames, height, width = x.shape
channels = c_patches // (patch_size * patch_size)
# Reshape to [b, c, patch_size, patch_size, f, h, w]
x = x.view(batch_size, channels, patch_size, patch_size, frames, height, width)
# Rearrange to [b, c, f, h * patch_size, w * patch_size]
x = x.permute(0, 1, 4, 5, 2, 6, 3).contiguous()
x = x.view(batch_size, channels, frames, height * patch_size, width * patch_size)
return x
class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
r"""
@@ -795,24 +1105,43 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
self.z_dim = config.z_dim
self.temperal_downsample = list(config.temperal_downsample)
self.temperal_upsample = list(config.temperal_downsample)[::-1]
if config.decoder_base_dim is None:
decoder_base_dim = config.base_dim
else:
decoder_base_dim = config.decoder_base_dim
self.latents_mean = list(config.latents_mean)
self.latents_std = list(config.latents_std)
self.shift_factor = config.shift_factor
if config.load_encoder:
self.encoder = WanEncoder3d(config.base_dim, self.z_dim * 2,
config.dim_mult, config.num_res_blocks,
config.attn_scales,
self.temperal_downsample,
config.dropout)
self.encoder = WanEncoder3d(
in_channels=config.in_channels,
dim=config.base_dim,
z_dim=self.z_dim * 2,
dim_mult=config.dim_mult,
num_res_blocks=config.num_res_blocks,
attn_scales=config.attn_scales,
temperal_downsample=self.temperal_downsample,
dropout=config.dropout,
is_residual=config.is_residual,
)
self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
if config.load_decoder:
self.decoder = WanDecoder3d(config.base_dim, self.z_dim,
config.dim_mult, config.num_res_blocks,
config.attn_scales,
self.temperal_upsample, config.dropout)
self.decoder = WanDecoder3d(
dim=decoder_base_dim,
z_dim=self.z_dim,
dim_mult=config.dim_mult,
num_res_blocks=config.num_res_blocks,
attn_scales=config.attn_scales,
temperal_upsample=self.temperal_upsample,
dropout=config.dropout,
out_channels=config.out_channels,
is_residual=config.is_residual,
)
self.use_feature_cache = config.use_feature_cache
@@ -838,6 +1167,8 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
def encode(self, x: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
self.clear_cache()
if self.config.patch_size is not None:
x = patchify(x, patch_size=self.config.patch_size)
with forward_context(feat_cache_arg=self._enc_feat_map,
feat_idx_arg=self._enc_conv_idx):
t = x.shape[2]
@@ -903,12 +1234,17 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
for i in range(iter_):
feat_idx.set(0)
if i == 0:
first_chunk.set(True)
out = self.decoder(x[:, :, i:i + 1, :, :])
else:
first_chunk.set(False)
out_ = self.decoder(x[:, :, i:i + 1, :, :])
out = torch.cat([out, out_], 2)
out = torch.clamp(out, min=-1.0, max=1.0)
if self.config.clip_output:
out = torch.clamp(out, min=-1.0, max=1.0)
if self.config.patch_size is not None:
out = unpatchify(out, patch_size=self.config.patch_size)
self.clear_cache()
else:
out = ParallelTiledVAE.decode(self, z)
+23 -12
View File
@@ -12,7 +12,8 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.pipelines.pipeline_registry import PipelineRegistry
from fastvideo.pipelines.pipeline_registry import (PipelineType,
get_pipeline_registry)
from fastvideo.utils import (maybe_download_model,
verify_model_config_and_directory)
@@ -24,7 +25,10 @@ class PipelineWithLoRA(LoRAPipeline, ComposedPipelineBase):
pass
def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
def build_pipeline(
fastvideo_args: FastVideoArgs,
pipeline_type: PipelineType | str = PipelineType.BASIC
) -> PipelineWithLoRA:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
@@ -37,30 +41,37 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
model_path = maybe_download_model(model_path)
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
pipeline_architecture = config.get("_class_name")
if pipeline_architecture is None:
config = verify_model_config_and_directory(model_path)
pipeline_name = config.get("_class_name")
if pipeline_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
pipeline_cls, pipeline_architecture = PipelineRegistry.resolve_pipeline_cls(
pipeline_architecture)
# Get the appropriate pipeline registry based on pipeline_type
logger.info(
"Building pipeline of type: %s", pipeline_type.value if isinstance(
pipeline_type, PipelineType) else pipeline_type)
pipeline_registry = get_pipeline_registry(pipeline_type)
# instantiate the pipeline
if isinstance(pipeline_type, str):
pipeline_type = PipelineType.from_string(pipeline_type)
pipeline_cls = pipeline_registry.resolve_pipeline_cls(
pipeline_name, pipeline_type, fastvideo_args.workload_type)
# instantiate the pipelines
pipeline = pipeline_cls(model_path, fastvideo_args)
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
logger.info("Pipelines instantiated")
return cast(PipelineWithLoRA, pipeline)
__all__ = [
"build_pipeline",
"list_available_pipelines",
"ComposedPipelineBase",
"PipelineRegistry",
"ForwardBatch",
"LoRAPipeline",
"TrainingBatch",
+6
View File
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""
Basic inference pipelines for fastvideo.
This package contains basic pipelines for video and image generation.
"""
@@ -0,0 +1,81 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (ImageEncodingStage, ConditioningStage,
DecodingStage, DmdDenoisingStage,
EncodingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
logger = init_logger(__name__)
class WanImageToVideoDmdPipeline(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"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanImageToVideoDmdPipeline
@@ -62,6 +62,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self))
self.add_stage(stage_name="decoding_stage",
+17 -4
View File
@@ -139,7 +139,7 @@ class ComposedPipelineBase(ABC):
for key, value in kwargs.items():
setattr(fastvideo_args, key, value)
fastvideo_args.use_cpu_offload = False
fastvideo_args.dit_cpu_offload = 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
@@ -234,6 +234,9 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# @TODO(Wei): Temporary hack
model_index.pop("boundary_ratio")
model_index.pop("expand_timesteps")
# some sanity checks
assert len(
@@ -242,8 +245,11 @@ class ComposedPipelineBase(ABC):
for module_name in self.required_config_modules:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.warning(
"model_index.json does not contain a %s module, adding %s to model_index",
module_name, module_name)
if 'transformer' in module_name:
model_index[module_name] = model_index['transformer']
# all the component models used by the pipeline
required_modules = self.required_config_modules
@@ -252,6 +258,8 @@ class ComposedPipelineBase(ABC):
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in model_index.items():
if transformers_or_diffusers is None:
continue
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
@@ -259,7 +267,12 @@ class ComposedPipelineBase(ABC):
logger.info("Using module %s already provided", module_name)
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
if 'transformer' in module_name:
loading_module_name = module_name.split("_")[-1]
else:
loading_module_name = module_name
component_model_path = os.path.join(self.model_path,
loading_module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
@@ -148,6 +148,8 @@ class TrainingBatch:
# Dataloader batch outputs
latents: torch.Tensor | None = None
raw_latent_shape: torch.Tensor | None = None
noise_latents: torch.Tensor | None = None
encoder_hidden_states: torch.Tensor | None = None
encoder_attention_mask: torch.Tensor | None = None
# i2v
@@ -155,6 +157,7 @@ class TrainingBatch:
image_embeds: torch.Tensor | None = None
image_latents: torch.Tensor | None = None
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
@@ -162,6 +165,7 @@ class TrainingBatch:
sigmas: torch.Tensor | None = None
noise: torch.Tensor | None = None
attn_metadata_vsa: AttentionMetadata | None = None
attn_metadata: AttentionMetadata | None = None
# input kwargs
@@ -173,3 +177,13 @@ class TrainingBatch:
# Training outputs
total_loss: float | None = None
grad_norm: float | None = None
# Distillation-specific attributes
encoder_hidden_states_neg: torch.Tensor | None = None
encoder_attention_mask_neg: torch.Tensor | None = None
conditional_dict: dict[str, Any] | None = None
unconditional_dict: dict[str, Any] | None = None
# Distillation losses
generator_loss: float = 0.0
fake_score_loss: float = 0.0
+199 -52
View File
@@ -6,84 +6,231 @@ import importlib
import pkgutil
from collections.abc import Set
from dataclasses import dataclass, field
from enum import Enum
from functools import lru_cache
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
logger = init_logger(__name__)
# map pipeline name to folder name
_PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
}
class PipelineType(str, Enum):
"""
Enumeration for different pipeline types.
Inherits from str to allow string comparison for backward compatibility.
"""
BASIC = "basic"
PREPROCESS = "preprocess"
TRAINING = "training"
@classmethod
def from_string(cls, value: str) -> "PipelineType":
"""Convert string to PipelineType enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid pipeline type: {value}. Must be one of: {', '.join([t.value for t in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings."""
return [pipeline_type.value for pipeline_type in cls]
@dataclass
class _PipelineRegistry:
# Keyed by pipeline_arch
pipelines: dict[str, type[ComposedPipelineBase]
| None] = field(default_factory=dict)
# Keyed by pipeline_type -> architecture -> pipeline_name
# pipelines[pipeline_type][architecture][pipeline_name] = pipeline_cls
pipelines: dict[str, dict[str, dict[str, type[ComposedPipelineBase]
| None]]] = field(default_factory=dict)
def get_supported_archs(self) -> Set[str]:
return self.pipelines.keys()
def get_supported_archs(self, pipeline_name_in_config: str,
pipeline_type: PipelineType) -> Set[str]:
"""Get supported architectures, optionally filtered by pipeline type and workload type."""
arch = _PIPELINE_NAME_TO_ARCHITECTURE_NAME[pipeline_name_in_config]
return set(self.pipelines[pipeline_type.value][arch].keys())
def _load_preprocessing_pipeline_cls(
self, workload_type: WorkloadType,
arch: str) -> type[ComposedPipelineBase] | None:
if workload_type == WorkloadType.I2V:
pipeline_name = "I2VPreprocessPipeline"
elif workload_type == WorkloadType.T2V:
pipeline_name = "T2VPreprocessPipeline"
else:
raise ValueError(f"Invalid workload type: {workload_type.value}")
return self.pipelines[
PipelineType.PREPROCESS.value][arch][pipeline_name]
def _try_load_pipeline_cls(
self, pipeline_arch: str) -> type[ComposedPipelineBase] | None:
if pipeline_arch not in self.pipelines:
self, pipeline_name_in_config: str, pipeline_type: PipelineType,
workload_type: WorkloadType
) -> type[ComposedPipelineBase] | type[LoRAPipeline] | None:
"""Try to load a pipeline class for the given architecture, pipeline type, and workload type."""
arch = _PIPELINE_NAME_TO_ARCHITECTURE_NAME[pipeline_name_in_config]
if (pipeline_type.value not in self.pipelines
or arch not in self.pipelines[pipeline_type.value]):
return None
return self.pipelines[pipeline_arch]
if pipeline_type == PipelineType.PREPROCESS:
return self._load_preprocessing_pipeline_cls(workload_type, arch)
elif pipeline_type == PipelineType.BASIC:
return self.pipelines[
pipeline_type.value][arch][pipeline_name_in_config]
elif pipeline_type == PipelineType.TRAINING:
pass
else:
raise ValueError(f"Invalid pipeline type: {pipeline_type.value}")
return None
def resolve_pipeline_cls(
self,
architecture: str,
) -> tuple[type[ComposedPipelineBase] | type[LoRAPipeline], str]:
if not architecture:
pipeline_name_in_config: str,
pipeline_type: PipelineType,
workload_type: WorkloadType,
) -> type[ComposedPipelineBase] | type[LoRAPipeline]:
"""Resolve pipeline class based on pipeline name in the config, pipeline type, and workload type."""
if not pipeline_name_in_config:
logger.warning("No pipeline architecture is specified")
pipeline_cls = self._try_load_pipeline_cls(architecture)
pipeline_cls = self._try_load_pipeline_cls(pipeline_name_in_config,
pipeline_type, workload_type)
if pipeline_cls is not None:
return (pipeline_cls, architecture)
supported_archs = self.get_supported_archs()
return pipeline_cls
supported_archs = self.get_supported_archs(pipeline_name_in_config,
pipeline_type)
raise ValueError(
f"Pipeline architectures {architecture} are not supported for now. "
f"Pipeline architecture '{pipeline_name_in_config}' is not supported for pipeline type '{pipeline_type.value}' "
f"and workload type '{workload_type.value}'. "
f"Supported architectures: {supported_archs}")
@lru_cache
def import_pipeline_classes():
pipeline_arch_name_to_cls = {}
package_name = "fastvideo.pipelines"
package = importlib.import_module(package_name)
for _, name, ispkg in pkgutil.iter_modules(package.__path__,
package_name + "."):
if ispkg:
if name.split(".")[-1] == "stages":
continue
sub_package_name = name
sub_package = importlib.import_module(sub_package_name)
for _, name, ispkg in pkgutil.iter_modules(sub_package.__path__,
sub_package_name + "."):
try:
module = importlib.import_module(name)
except Exception as e:
logger.warning("Ignore import error when loading %s. %s",
name, e)
continue
if hasattr(module, "EntryClass"):
entry = module.EntryClass
if isinstance(
entry, list
): # To support multiple pipeline classes in one module
for tmp in entry:
assert (
tmp.__name__ not in pipeline_arch_name_to_cls
), f"Duplicated pipeline implementation for {tmp.__name__}"
pipeline_arch_name_to_cls[tmp.__name__] = tmp
else:
assert (
entry.__name__ not in pipeline_arch_name_to_cls
), f"Duplicated pipeline implementation for {entry.__name__}"
pipeline_arch_name_to_cls[entry.__name__] = entry
return pipeline_arch_name_to_cls
def import_pipeline_classes(
pipeline_types: list[PipelineType] | PipelineType | None = None
) -> dict[str, dict[str, dict[str, type[ComposedPipelineBase] | None]]]:
"""
Import pipeline classes based on the pipeline type and workload type.
Args:
pipeline_types: The pipeline types to load (basic, preprocessing, training).
If None, loads all types.
Returns:
A three-level nested dictionary:
{pipeline_type: {architecture_name: {pipeline_name: pipeline_cls}}}
e.g., {"basic": {"wan": {"WanPipeline": WanPipeline}}}
"""
type_to_arch_to_pipeline_dict: dict[str,
dict[str,
dict[str,
type[ComposedPipelineBase]
| None]]] = {}
package_name: str = "fastvideo.pipelines"
# Determine which pipeline types to scan
if isinstance(pipeline_types, list):
pipeline_types_to_scan = [
pipeline_type.value for pipeline_type in pipeline_types
]
elif isinstance(pipeline_types, PipelineType):
pipeline_types_to_scan = [pipeline_types.value]
else:
pipeline_types_to_scan = [pt.value for pt in PipelineType]
logger.info("Loading pipelines for types: %s", pipeline_types_to_scan)
for pipeline_type_str in pipeline_types_to_scan:
arch_to_pipeline_dict: dict[str, dict[str, type[ComposedPipelineBase]
| None]] = {}
# Try to load from pipeline-type-specific directory first
pipeline_type_package_name = f"{package_name}.{pipeline_type_str}"
try:
pipeline_type_package = importlib.import_module(
pipeline_type_package_name)
logger.debug("Successfully imported %s", pipeline_type_package_name)
for _, arch, ispkg in pkgutil.iter_modules(
pipeline_type_package.__path__):
pipeline_dict: dict[str, type[ComposedPipelineBase] | None] = {}
arch_package_name = f"{pipeline_type_package_name}.{arch}"
if ispkg:
arch_package = importlib.import_module(arch_package_name)
for _, module_name, ispkg in pkgutil.walk_packages(
arch_package.__path__, arch_package_name + "."):
if not ispkg:
pipeline_module = importlib.import_module(
module_name)
if hasattr(pipeline_module, "EntryClass"):
if isinstance(pipeline_module.EntryClass, list):
for pipeline in pipeline_module.EntryClass:
pipeline_name = pipeline.__name__
assert (
pipeline_name not in pipeline_dict
), f"Duplicated pipeline implementation for {pipeline_name} in {pipeline_type_str}.{arch_package_name}"
pipeline_dict[pipeline_name] = pipeline
else:
pipeline_name = pipeline_module.EntryClass.__name__
assert (
pipeline_name not in pipeline_dict
), f"Duplicated pipeline implementation for {pipeline_name} in {pipeline_type_str}.{arch_package_name}"
pipeline_dict[
pipeline_name] = pipeline_module.EntryClass
arch_to_pipeline_dict[arch] = pipeline_dict
except ImportError as e:
raise ImportError(
f"Could not import {pipeline_type_package_name} when importing pipeline classes: {e}"
) from None
type_to_arch_to_pipeline_dict[pipeline_type_str] = arch_to_pipeline_dict
# Log summary
total_pipelines = sum(
len(pipeline_dict)
for arch_to_pipeline_dict in type_to_arch_to_pipeline_dict.values()
for pipeline_dict in arch_to_pipeline_dict.values())
logger.info("Loaded %d pipeline classes across %d types", total_pipelines,
len(pipeline_types_to_scan))
return type_to_arch_to_pipeline_dict
PipelineRegistry = _PipelineRegistry(import_pipeline_classes())
def get_pipeline_registry(
pipeline_type: PipelineType | str | None = None) -> _PipelineRegistry:
"""
Get a pipeline registry for the specified mode, pipeline type, and workload type.
Args:
pipeline_type: Pipeline type to load. If None and mode is provided, will be derived from mode.
Returns:
A pipeline registry instance.
"""
if isinstance(pipeline_type, str):
pipeline_type = PipelineType.from_string(pipeline_type)
pipeline_classes = import_pipeline_classes(pipeline_type)
return _PipelineRegistry(pipeline_classes)
@@ -23,7 +23,7 @@ def main(args) -> None:
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"use_cpu_offload": False,
"dit_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
+2 -2
View File
@@ -12,6 +12,7 @@ from abc import ABC, abstractmethod
import torch
import fastvideo.envs as envs
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
@@ -145,8 +146,7 @@ class PipelineStage(ABC):
raise
# Execute the actual stage logic
# envs.ENABLE_STAGE_LOGGING
if False:
if envs.FASTVIDEO_STAGE_LOGGING:
logger.info("[%s] Starting execution", stage_name)
start_time = time.perf_counter()
+2 -1
View File
@@ -133,7 +133,8 @@ class DecodingStage(PipelineStage):
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
self.vae.to("cpu")
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
if torch.backends.mps.is_available():
del self.vae
+49 -9
View File
@@ -10,7 +10,9 @@ from collections.abc import Iterable
from typing import Any
import torch
import torchvision.transforms.functional as TF
from einops import rearrange
from PIL import Image
from tqdm.auto import tqdm
from fastvideo.attention import get_attn_backend
@@ -30,7 +32,7 @@ from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import dict_to_3d_list
from fastvideo.utils import best_output_size, dict_to_3d_list, masks_like
try:
from fastvideo.attention.backends.sliding_tile_attn import (
@@ -57,10 +59,11 @@ class DenoisingStage(PipelineStage):
the initial noise into the final output.
"""
def __init__(self, transformer, scheduler, pipeline=None) -> None:
def __init__(self, transformer, scheduler, vae, pipeline=None) -> None:
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
self.vae = vae
self.pipeline = weakref.ref(pipeline) if pipeline else None
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
@@ -119,13 +122,13 @@ class DenoisingStage(PipelineStage):
sp_group = sp_world_size > 1
if sp_group:
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
latents = latents[:, :, rank_in_sp_group, :, :, :]
batch.latents = latents
if batch.image_latent is not None:
image_latent = rearrange(batch.image_latent,
"b t (n s) h w -> b t n s h w",
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
batch.image_latent = image_latent
@@ -194,9 +197,49 @@ class DenoisingStage(PipelineStage):
# Expand latents for I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
elif fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
# TI2V directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
# preprocess
img = batch.pil_image
assert img is not None
ih, iw = img.height, img.width
dh, dw = self.patch_size[1] * self.vae_stride[
1], self.patch_size[2] * self.vae_stride[2]
max_area = 704 * 1280
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
assert img.width == ow and img.height == oh
# to tensor
img = TF.to_tensor(img).sub_(0.5).div_(0.5).to(
self.device).unsqueeze(1)
# return [
# self.model.encode(u.unsqueeze(0),
# self.scale).float().squeeze(0)
# for u in videos
# ]
z = self.vae.encode(img.unsqueeze(0))
mask1, mask2 = masks_like([latent_model_input],
zero=True,
generator=batch.generator)
latent_model_input = (
1. - mask2[0]) * z[0] + mask2[0] * latent_model_input
assert torch.isnan(latent_model_input).sum() == 0
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
@@ -758,9 +801,7 @@ class DmdDenoisingStage(DenoisingStage):
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
pred_video.shape[:2],
dtype=torch.long,
device=pred_video.device)
[1], dtype=torch.long, device=pred_video.device)
noise = torch.randn(video_raw_latent_shape,
device=self.device,
dtype=pred_video.dtype)
@@ -771,8 +812,7 @@ class DmdDenoisingStage(DenoisingStage):
noise = noise[:, rank_in_sp_group, :, :, :, :]
latents = self.scheduler.add_noise(
pred_video.flatten(0, 1), noise.flatten(0, 1),
next_timestep.flatten(0, 1)).unflatten(
0, pred_video.shape[:2])
next_timestep).unflatten(0, pred_video.shape[:2])
else:
latents = pred_video
+2 -1
View File
@@ -135,7 +135,8 @@ class EncodingStage(PipelineStage):
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
self.vae.to("cpu")
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
return batch
+1 -1
View File
@@ -67,7 +67,7 @@ class ImageEncodingStage(PipelineStage):
batch.image_embeds.append(image_embeds)
if fastvideo_args.use_cpu_offload:
if fastvideo_args.image_encoder_cpu_offload:
self.image_encoder.to('cpu')
return batch
+6
View File
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""
Training pipelines for fastvideo.v1.
This package contains pipelines for training diffusion models.
"""
@@ -33,7 +33,7 @@ WAN_LORA_PARAMS = {
"fps": 24,
"neg_prompt": "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
"text-encoder-precision": ("fp32",),
"use_cpu_offload": True,
"dit_cpu_offload": True,
}
# LoRA configurations for testing
@@ -68,7 +68,7 @@ def test_merge_lora_weights(model_id):
lora_path = lora_config["lora_path"]
args = FastVideoArgs.from_kwargs(
model_path=model_id,
use_cpu_offload=True,
dit_cpu_offload=True,
dit_precision="bf16",
)
pipe = build_pipeline(args)
@@ -113,7 +113,7 @@ def test_lora_inference_similarity(ATTENTION_BACKEND, model_id):
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"flow_shift": BASE_PARAMS["flow_shift"],
"use_cpu_offload": BASE_PARAMS["use_cpu_offload"],
"dit_cpu_offload": BASE_PARAMS["dit_cpu_offload"],
}
if "text-encoder-precision" in BASE_PARAMS:
init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
@@ -228,7 +228,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,
"dit_cpu_offload": True,
}
if BASE_PARAMS.get("vae_sp"):
init_kwargs["vae_sp"] = True
@@ -62,7 +62,7 @@ def test_hunyuanvideo_distributed():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=True,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=HunyuanVideoConfig(), dit_precision=precision_str))
args.device = torch.device(f"cuda:{LOCAL_RANK}")
@@ -34,7 +34,7 @@ def test_wan_transformer():
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=True,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
+1
View File
@@ -43,6 +43,7 @@ def test_hunyuan_vae():
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=HunyuanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_cpu_offload = False
loader = VAELoader()
model = loader.load(VAE_PATH, args)
+1
View File
@@ -32,6 +32,7 @@ def test_wan_vae():
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=WanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_cpu_offload = False
loader = VAELoader()
model2 = loader.load(VAE_PATH, args)
+2 -1
View File
@@ -1,4 +1,5 @@
from .distillation_pipeline import DistillationPipeline
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
+794
View File
@@ -0,0 +1,794 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import os
import time
from abc import abstractmethod
from collections import deque
from collections.abc import Iterator
from typing import Any
import imageio
import numpy as np
import torch
import torch.nn.functional as F
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_local_torch_device, get_sp_group,
get_world_group)
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
TrainingBatch)
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases, load_checkpoint,
pred_noise_to_pred_video, save_checkpoint, shift_timestep)
from fastvideo.utils import is_vsa_available, set_random_seed
import wandb # isort: skip
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class DistillationPipeline(TrainingPipeline):
"""
A distillation pipeline for training a 3 step model.
Inherits from TrainingPipeline to reuse training infrastructure.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
video_latent_shape: tuple[int, ...]
video_latent_shape_sp: tuple[int, ...]
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):
"""Initialize the distillation training pipeline with multiple models."""
logger.info("Initializing distillation pipeline...")
super().initialize_training_pipeline(training_args)
self.noise_scheduler = self.get_module("scheduler")
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
# self.transformer is the generator model
self.real_score_transformer = self.get_module("real_score_transformer")
self.fake_score_transformer = self.get_module("fake_score_transformer")
self.real_score_transformer.requires_grad_(False)
self.real_score_transformer.eval()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.fake_score_transformer = apply_activation_checkpointing(
self.fake_score_transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
self.real_score_transformer = apply_activation_checkpointing(
self.real_score_transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
# Initialize optimizers
fake_score_params = list(
filter(lambda p: p.requires_grad,
self.fake_score_transformer.parameters()))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.fake_score_lr_scheduler = get_scheduler(
training_args.lr_scheduler,
optimizer=self.fake_score_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,
)
logger.info(
"Distillation optimizers initialized: generator and fake_score")
self.generator_update_interval = self.training_args.generator_update_interval
logger.info(
"Distillation pipeline initialized with generator_update_interval=%s",
self.generator_update_interval)
self.denoising_step_list = torch.tensor(
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps",
len(self.denoising_step_list))
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.min_timestep = int(self.training_args.min_timestep_ratio *
self.num_train_timestep)
self.max_timestep = int(self.training_args.max_timestep_ratio *
self.num_train_timestep)
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
"""Initialize validation pipeline - must be implemented by subclasses."""
raise NotImplementedError(
"Distillation pipelines must implement this method")
def _prepare_distillation(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Prepare training environment for distillation."""
self.transformer.requires_grad_(True)
self.transformer.train()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
return training_batch
def _build_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
training_batch: TrainingBatch) -> TrainingBatch:
if text_dict is None:
raise ValueError(
"text_dict cannot be None for distillation pipeline")
training_batch.input_kwargs = {
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep,
"return_dict": False,
}
training_batch.noise_latents = noise_input
# After setting noise_latents, it's guaranteed to be not None
return training_batch
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
latents = training_batch.latents
dtype = latents.dtype
index = torch.randint(0,
len(self.denoising_step_list), [1],
device=self.device,
dtype=torch.long)
timestep = self.denoising_step_list[index]
noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=dtype)
if self.sp_world_size > 1:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
noisy_latent = self.noise_scheduler.add_noise(latents.flatten(0, 1),
noise.flatten(0, 1),
timestep).unflatten(
0,
(1, latents.shape[1]))
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
return pred_video
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
with torch.no_grad():
timestep = torch.randint(0,
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
timestep = shift_timestep(
timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
timestep = timestep.clamp(self.min_timestep, self.max_timestep)
noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_pred_video.dtype)
if self.sp_world_size > 1:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
noisy_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
fake_score_pred_noise = self.fake_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
pred_fake_video = pred_noise_to_pred_video(
pred_noise=fake_score_pred_noise.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(
0, fake_score_pred_noise.shape[:2])
# real_score_transformer cond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
real_score_pred_noise_cond = self.real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
pred_real_video_cond = pred_noise_to_pred_video(
pred_noise=real_score_pred_noise_cond.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(
0, real_score_pred_noise_cond.shape[:2])
# real_score_transformer uncond forward
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.unconditional_dict,
training_batch)
real_score_pred_noise_uncond = self.real_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
pred_real_video_uncond = pred_noise_to_pred_video(
pred_noise=real_score_pred_noise_uncond.flatten(0, 1),
noise_input_latent=training_batch.noise_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(
0, real_score_pred_noise_uncond.shape[:2])
real_score_pred_video = pred_real_video_cond + (
pred_real_video_cond -
pred_real_video_uncond) * self.real_score_guidance_scale
grad = (pred_fake_video - real_score_pred_video) / torch.abs(
generator_pred_video - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
generator_pred_video.float(),
(generator_pred_video.float() - grad.float()).detach())
return dmd_loss
def faker_score_forward(
self, training_batch: TrainingBatch
) -> tuple[TrainingBatch, torch.Tensor]:
with torch.no_grad(), set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
generator_pred_video = self._generator_forward(training_batch)
fake_score_timestep = torch.randint(0,
self.num_train_timestep, [1],
device=self.device,
dtype=torch.long)
fake_score_timestep = shift_timestep(
fake_score_timestep,
self.timestep_shift, # type: ignore
self.num_train_timestep)
fake_score_timestep = fake_score_timestep.clamp(self.min_timestep,
self.max_timestep)
fake_score_noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=generator_pred_video.dtype)
if self.sp_world_size > 1:
fake_score_noise = rearrange(fake_score_noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
fake_score_noise = fake_score_noise[:, self.
rank_in_sp_group, :, :, :, :]
noisy_generator_pred_video = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), fake_score_noise.flatten(0, 1),
fake_score_timestep).unflatten(0,
(1, generator_pred_video.shape[1]))
with set_forward_context(current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
training_batch = self._build_distill_input_kwargs(
noisy_generator_pred_video, fake_score_timestep,
training_batch.conditional_dict, training_batch)
fake_score_pred_noise = self.fake_score_transformer(
**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
target = fake_score_noise - generator_pred_video
denoising_loss = torch.mean((fake_score_pred_noise - target)**2)
return training_batch, denoising_loss
def _clip_model_grad_norm_(self, training_batch: TrainingBatch,
transformer) -> TrainingBatch:
max_grad_norm = self.training_args.max_grad_norm
if max_grad_norm is not None:
model_parts = [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,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
super()._prepare_dit_inputs(training_batch)
conditional_dict = {
"encoder_hidden_states": training_batch.encoder_hidden_states,
"encoder_attention_mask": training_batch.encoder_attention_mask,
}
unconditional_dict = {
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
training_batch.conditional_dict = conditional_dict
training_batch.unconditional_dict = unconditional_dict
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
self.video_latent_shape = training_batch.latents.shape
training_batch.raw_latent_shape = training_batch.latents.shape
if self.sp_world_size > 1:
training_batch.latents = rearrange(
training_batch.latents,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
training_batch.latents = training_batch.latents[:, self.
rank_in_sp_group, :, :, :, :]
self.video_latent_shape_sp = training_batch.latents.shape
return training_batch
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
gradient_accumulation_steps = getattr(self.training_args,
'gradient_accumulation_steps', 1)
batches = []
# Collect N batches for gradient accumulation
for _ in range(gradient_accumulation_steps):
batch = self._prepare_distillation(training_batch)
batch = self._get_next_batch(batch)
batch = self._normalize_dit_input(batch)
batch = self._prepare_dit_inputs(batch)
batch = self._build_attention_metadata(batch)
batch.attn_metadata_vsa = copy.deepcopy(batch.attn_metadata)
if batch.attn_metadata is not None:
batch.attn_metadata.VSA_sparsity = 0.0 # type: ignore
batches.append(batch)
self.optimizer.zero_grad()
total_dmd_loss = 0.0
if (self.current_trainstep % self.generator_update_interval == 0):
for batch in batches:
batch_stu = copy.deepcopy(batch)
with set_forward_context(
current_timestep=batch_stu.timesteps,
attn_metadata=batch_stu.attn_metadata_vsa):
generator_pred_video = self._generator_forward(batch_stu)
with set_forward_context(current_timestep=batch_stu.timesteps,
attn_metadata=batch_stu.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=batch_stu)
with set_forward_context(
current_timestep=batch_stu.timesteps,
attn_metadata=batch_stu.attn_metadata_vsa):
(dmd_loss / gradient_accumulation_steps).backward()
total_dmd_loss += dmd_loss.detach().item()
self._clip_model_grad_norm_(batch_stu, self.transformer)
self.optimizer.step()
self.lr_scheduler.step()
self.optimizer.zero_grad(set_to_none=True)
avg_dmd_loss = torch.tensor(total_dmd_loss /
gradient_accumulation_steps,
device=self.device)
world_group = get_world_group()
world_group.all_reduce(avg_dmd_loss,
op=torch.distributed.ReduceOp.AVG)
training_batch.generator_loss = avg_dmd_loss.item()
else:
training_batch.generator_loss = 0.0
self.fake_score_optimizer.zero_grad()
total_fake_score_loss = 0.0
for batch in batches:
batch_critic = copy.deepcopy(batch)
batch_critic, fake_score_loss = self.faker_score_forward(
batch_critic)
with set_forward_context(current_timestep=batch_critic.timesteps,
attn_metadata=batch_critic.attn_metadata):
(fake_score_loss / gradient_accumulation_steps).backward()
total_fake_score_loss += fake_score_loss.detach().item()
self._clip_model_grad_norm_(batch_critic, self.fake_score_transformer)
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
self.fake_score_optimizer.zero_grad(set_to_none=True)
avg_fake_score_loss = torch.tensor(total_fake_score_loss /
gradient_accumulation_steps,
device=self.device)
world_group = get_world_group()
world_group.all_reduce(avg_fake_score_loss,
op=torch.distributed.ReduceOp.AVG)
training_batch.fake_score_loss = avg_fake_score_loss.item()
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_loss
return training_batch
def _resume_from_checkpoint(self) -> None: #TODO(yongqi)
"""Resume training from checkpoint with distillation models."""
logger.info("Loading distillation checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
self.transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
# TODO: Add checkpoint loading for critic and teacher models
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 = -1
def _log_training_info(self) -> None:
"""Log distillation-specific training information."""
# First call parent class method to get basic training info
super()._log_training_info()
# Then add distillation-specific information
logger.info("Distillation-specific settings:")
logger.info(" Generator update ratio: %s",
self.generator_update_interval)
assert isinstance(self.training_args, TrainingArgs)
logger.info(" Max gradient norm: %s", self.training_args.max_grad_norm)
logger.info(
" Real score transformer parameters: %s B",
sum(p.numel()
for p in self.real_score_transformer.parameters()) / 1e9)
logger.info(
" Fake score transformer parameters: %s B",
sum(p.numel()
for p in self.fake_score_transformer.parameters()) / 1e9)
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
training_args.inference_mode = True
training_args.dit_cpu_offload = True
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
logger.info("Using validation seed: %s", self.seed)
# Prepare validation prompts
logger.info('rank: %s: fastvideo_args.validation_dataset_file: %s',
self.global_rank,
training_args.validation_dataset_file,
local_main_process_only=False)
validation_dataset = ValidationDataset(
training_args.validation_dataset_file)
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
# Log validation results for this step
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
logger.info("rank: %s: num_inference_steps: %s",
self.global_rank,
num_inference_steps,
local_main_process_only=False)
step_videos: list[np.ndarray] = []
step_captions: list[str] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
validation_batch,
num_inference_steps)
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
batch_negative, training_args)
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
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))
step_videos.append(frames)
# Log validation results for this step
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.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_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, strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions, strict=True)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
# Re-enable gradients for training
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def train(self) -> None:
"""Main training loop with distillation-specific logging."""
assert self.training_args.seed is not None, "seed must be set"
seed = self.training_args.seed
# Set the same seed within each SP group to ensure reproducibility
if self.sp_world_size > 1:
# Use the same seed for all processes within the same SP group
sp_group_seed = seed + (self.global_rank // self.sp_world_size)
set_random_seed(sp_group_seed)
logger.info("Rank %s: Using SP group seed %s", self.global_rank,
sp_group_seed)
else:
set_random_seed(seed + self.global_rank)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
self.train_loader_iter = iter(self.train_dataloader)
step_times: deque[float] = deque(maxlen=100)
self._log_training_info()
self._log_validation(self.transformer, self.training_args, 0)
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
disable=self.local_rank > 0,
)
use_vsa = vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if use_vsa:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
if vsa_decay_interval_steps > 1:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = vsa_sparsity
else:
current_vsa_sparsity = 0.0
training_batch = TrainingBatch()
self.current_trainstep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
total_loss = training_batch.total_loss
generator_loss = training_batch.generator_loss
fake_score_loss = training_batch.fake_score_loss
grad_norm = training_batch.grad_norm
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"total_loss": f"{total_loss:.4f}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.global_rank == 0:
# Prepare logging data
log_data = {
"train_total_loss": total_loss,
"train_fake_score_loss": fake_score_loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
}
# Only log generator loss when generator is actually trained
if (step % self.generator_update_interval == 0):
log_data["train_generator_loss"] = generator_loss
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
wandb.log(log_data, step=step)
if step % self.training_args.checkpointing_steps == 0 and step > 0:
print("rank", self.global_rank, "save checkpoint at step", step)
save_checkpoint(
self.transformer,
self.global_rank, #TODO(yongqi)
self.training_args.output_dir,
step,
self.optimizer,
self.train_dataloader,
self.lr_scheduler,
self.noise_random_generator)
if self.transformer:
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)
wandb.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
+15 -50
View File
@@ -97,12 +97,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
assert training_args.seed is not None
self.seed = training_args.seed
assert self.transformer is not None
self.set_schemas()
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
@@ -151,10 +150,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
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 /
@@ -183,10 +178,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
@@ -222,11 +213,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert self.noise_random_generator is not None
latents = training_batch.latents
batch_size = latents.shape[0]
noise = torch.randn(latents.shape,
@@ -262,23 +248,22 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.raw_latent_shape = training_batch.latents.shape
return training_batch
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
latents = training_batch.latents
assert latents is not None
assert training_batch.timesteps is not None
latents_shape = training_batch.raw_latent_shape
patch_size = self.training_args.pipeline_config.dit_config.patch_size
current_vsa_sparsity = training_batch.current_vsa_sparsity
assert latents_shape is not None
assert training_batch.timesteps is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
dit_seq_shape = [
latents.shape[2] * self.sp_world_size // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
latents_shape[2] // patch_size[0],
latents_shape[3] // patch_size[1],
latents_shape[4] // patch_size[2]
]
training_batch.attn_metadata = VideoSparseAttentionMetadata(
current_timestep=training_batch.timesteps,
@@ -291,12 +276,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
@@ -314,19 +293,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.latents is not None
assert training_batch.noise is not None
assert training_batch.sigmas is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
assert training_batch.input_kwargs is not None
input_kwargs = training_batch.input_kwargs
# if 'hunyuan' in self.training_args.model_type:
@@ -340,7 +310,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
attn_metadata=training_batch.attn_metadata):
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
assert training_batch.sigmas is not None
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
assert training_batch.latents is not None
assert training_batch.noise is not None
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
# make sure no implicit broadcasting happens
@@ -360,7 +333,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
# TODO(will): perhaps move this into transformer api so that we can do
@@ -382,8 +354,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
training_batch = self._prepare_training(training_batch)
for _ in range(self.training_args.gradient_accumulation_steps):
@@ -422,7 +392,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _resume_from_checkpoint(self) -> None:
assert self.training_args is not None
logger.info("Loading checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
@@ -438,12 +407,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.init_steps = 0
def train(self) -> None:
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
logger.info('rank: %s: start training',
self.global_rank,
local_main_process_only=False)
assert self.training_args is not None
if not self.post_init_called:
self.post_init()
num_trainable_params = _get_trainable_params(self.transformer)
@@ -469,7 +437,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
step_times: deque[float] = deque(maxlen=100)
self._log_training_info()
self._log_validation(self.transformer, self.training_args, 1)
self._log_validation(self.transformer, self.training_args, 0)
# Train!
progress_bar = tqdm(
@@ -549,9 +517,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
cleanup_dist_env_and_memory()
def _log_training_info(self) -> None:
assert self.training_args is not None
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 *
@@ -592,6 +557,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
sampling_param.width = training_args.num_width
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
@@ -617,9 +583,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
"""
Generate a validation video and log it to wandb to check the quality during training.
"""
assert training_args is not None
training_args.inference_mode = True
training_args.use_cpu_offload = True
training_args.dit_cpu_offload = True
if not training_args.log_validation:
return
if self.validation_pipeline is None:
+9
View File
@@ -576,3 +576,12 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
return timestep
t = timestep / num_train_timestep
denominator = 1 + (shift - 1) * t
return num_train_timestep * (shift * t / denominator)
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.wan.wan_dmd_pipeline import WanDMDPipeline
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanDistillationPipeline(DistillationPipeline):
"""
A distillation pipeline for Wan that uses a single transformer model.
The main transformer serves as the student model, and copies are made for teacher and critic.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.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.dit_cpu_offload = True
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus)
self.validation_pipeline = validation_pipeline
def main(args) -> None:
logger.info("Starting Wan distillation pipeline...")
# Create pipeline with original args
pipeline = WanDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
# Start training
pipeline.train()
logger.info("Wan distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
main(args)
@@ -0,0 +1,233 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from einops import rearrange
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.dataloader.schema import (pyarrow_schema_i2v,
pyarrow_schema_i2v_validation)
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.pipelines.wan.wan_i2v_dmd_pipeline import (
WanImageToVideoDmdPipeline)
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanI2VDistillationPipeline(DistillationPipeline):
"""
A distillation pipeline for Wan that uses a single transformer model.
The main transformer serves as the student model, and copies are made for teacher and critic.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_i2v
self.validation_dataset_schema = pyarrow_schema_i2v_validation
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.dit_cpu_offload = False
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
validation_pipeline = WanImageToVideoDmdPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
dit_cpu_offload=True)
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
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 = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.preprocessed_image = pil_image.to(
get_local_torch_device())
training_batch.image_embeds = clip_features.to(get_local_torch_device())
training_batch.image_latents = image_latents.to(
get_local_torch_device())
training_batch.infos = infos
return training_batch
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.image_path = validation_batch['video_path']
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
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.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(
get_local_torch_device(), dtype=torch.bfloat16)
temporal_compression_ratio = 4
num_frames = (self.training_args.num_latent_t -
1) * temporal_compression_ratio + 1
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, 1:] = 0
first_frame_mask = mask_lat_size[:, :, :1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
dim=2)
mask_lat_size = mask_lat_size.view(batch_size, -1,
temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(
image_latents.device).to(dtype=torch.bfloat16)
image_latents = torch.cat([mask_lat_size, image_latents], dim=1)
training_batch.image_latents = image_latents
if self.sp_world_size > 1:
image_latents = rearrange(image_latents,
"b c (n t) h w -> b c n t h w",
n=self.sp_world_size).contiguous()
image_latents = image_latents[:, :, self.rank_in_sp_group, :, :, :]
training_batch.image_latents = image_latents
return training_batch
def _build_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor],
training_batch: TrainingBatch) -> TrainingBatch:
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_local_torch_device(),
dtype=torch.bfloat16)
noisy_model_input = torch.cat(
[noise_input,
training_batch.image_latents.permute(0, 2, 1, 3, 4)],
dim=2)
training_batch.input_kwargs = {
"hidden_states": noisy_model_input.permute(0, 2, 1, 3,
4), # bs, c, t, h, w
"encoder_hidden_states": text_dict["encoder_hidden_states"],
"encoder_attention_mask": text_dict["encoder_attention_mask"],
"timestep": timestep,
"encoder_hidden_states_image": image_embeds,
"return_dict": False,
}
training_batch.noise_latents = noise_input
return training_batch
def main(args) -> None:
logger.info("Starting Wan distillation pipeline...")
# Create pipeline with original args
pipeline = WanI2VDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
# Start training
pipeline.train()
logger.info("Wan distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
@@ -12,8 +12,9 @@ from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.basic.wan.wan_i2v_pipeline import (
WanImageToVideoPipeline)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.pipelines.wan.wan_i2v_pipeline import WanImageToVideoPipeline
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
@@ -46,7 +47,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.use_cpu_offload = True
args_copy.dit_cpu_offload = True
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
self.validation_pipeline = WanImageToVideoPipeline.from_pretrained(
@@ -59,12 +60,9 @@ class WanI2VTrainingPipeline(TrainingPipeline):
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
use_cpu_offload=True)
dit_cpu_offload=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
@@ -102,12 +100,6 @@ class WanI2VTrainingPipeline(TrainingPipeline):
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert self.noise_random_generator is not None
assert training_batch.image_latents is not None
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
@@ -144,12 +136,6 @@ class WanI2VTrainingPipeline(TrainingPipeline):
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
assert training_batch.image_embeds is not None
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
@@ -186,6 +172,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
sampling_param.image_path = validation_batch['video_path']
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
@@ -225,5 +212,5 @@ if __name__ == "__main__":
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
args.dit_cpu_offload = False
main(args)
+4 -4
View File
@@ -6,7 +6,7 @@ from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.wan.wan_pipeline import WanPipeline
from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available
@@ -36,11 +36,11 @@ class WanTrainingPipeline(TrainingPipeline):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.use_cpu_offload = True
args_copy.dit_cpu_offload = True
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanPipeline.from_pretrained(
training_args.model_path,
args=None,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
@@ -70,5 +70,5 @@ if __name__ == "__main__":
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
args.dit_cpu_offload = False
main(args)
+60
View File
@@ -812,3 +812,63 @@ def set_random_seed(seed: int) -> None:
@lru_cache(maxsize=1)
def is_vsa_available() -> bool:
return importlib.util.find_spec("vsa") is not None
def masks_like(tensor,
zero=False,
generator=None,
p=0.2) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
assert isinstance(tensor, list)
out1 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
out2 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensor]
if zero:
if generator is not None:
for u, v in zip(out1, out2, strict=False):
random_num = torch.rand(1,
generator=generator,
device=generator.device).item()
if random_num < p:
u[:, 0] = torch.normal(mean=-3.5,
std=0.5,
size=(1, ),
device=u.device,
generator=generator).expand_as(
u[:, 0]).exp()
v[:, 0] = torch.zeros_like(v[:, 0])
else:
u[:, 0] = u[:, 0]
v[:, 0] = v[:, 0]
else:
for u, v in zip(out1, out2, strict=False):
u[:, 0] = torch.zeros_like(u[:, 0])
v[:, 0] = torch.zeros_like(v[:, 0])
return out1, out2
def best_output_size(w, h, dw, dh, expected_area):
# float output size
ratio = w / h
ow = (expected_area * ratio)**0.5
oh = expected_area / ow
# process width first
ow1 = int(ow // dw * dw)
oh1 = int(expected_area / ow1 // dh * dh)
assert ow1 % dw == 0 and oh1 % dh == 0 and ow1 * oh1 <= expected_area
ratio1 = ow1 / oh1
# process height first
oh2 = int(oh // dh * dh)
ow2 = int(expected_area / oh2 // dw * dw)
assert oh2 % dh == 0 and ow2 % dw == 0 and ow2 * oh2 <= expected_area
ratio2 = ow2 / oh2
# compare ratios
if max(ratio / ratio1, ratio1 / ratio) < max(ratio / ratio2,
ratio2 / ratio):
return ow1, oh1
else:
return ow2, oh2
+16 -8
View File
@@ -31,15 +31,23 @@ 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")
# Check if master_port is provided in fastvideo_args
if hasattr(
self.fastvideo_args,
'master_port') and self.fastvideo_args.master_port is not None:
self.master_port = self.fastvideo_args.master_port
logger.info("Using provided master port: %s", self.master_port)
else:
# Auto-find available port
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
else:
raise ValueError("No unused port found to use as master port")
logger.info("Auto-selected master port: %s", self.master_port)
# Create pipes and start workers
for rank in range(self.world_size):
+10
View File
@@ -139,3 +139,13 @@ column_limit = 80
line_length = 80
use_parentheses = true
skip_gitignore = true
[project.urls]
Repository = "https://github.com/hao-ai-lab/FastVideo"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "fastvideo"
DisplayName = "FastVideo"
Icon = "https://raw.githubusercontent.com/hao-ai-lab/FastVideo/main/comfyui/assets/logo.png"
+58
View File
@@ -0,0 +1,58 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY=
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=mini_i2v_dataset/crush-smol_preprocessed/combined_parquet_dataset/
VALIDATION_DIR=mini_i2v_dataset/crush-smol_raw/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
export TOKENIZERS_PARALLELISM=false
# make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_distillation_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_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 20 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 500 \
--validation_steps 100 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune" \
--tracker_project_name Wan_distillation \
--num_height 448 \
--num_width 832 \
--num_frames 77 \
--flow_shift 8 \
--validation_guidance_scale "1.0" \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--vae_precision "bf16" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--generator_update_interval 5 \
--dmd_denoising_steps '1000,757,522' \
--min_timestep_ratio 0.02 \
--max_timestep_ratio 0.98 \
--real_score_guidance_scale 3.5 \
--seed 1024
+60
View File
@@ -0,0 +1,60 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY=
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=mini_i2v_dataset/crush-smol_preprocessed/combined_parquet_dataset/
VALIDATION_DIR=mini_i2v_dataset/crush-smol_raw/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
export TOKENIZERS_PARALLELISM=false
# Train generator with VSA
# Make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_distillation_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_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 0 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 500 \
--validation_steps 100 \
--validation_sampling_steps "3" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--output_dir "outputs_dmd/wan_finetune" \
--tracker_project_name Wan_distillation \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 8 \
--validation_guidance_scale "1.0" \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--vae_precision "bf16" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--generator_update_interval 5 \
--dmd_denoising_steps '1000,757,522' \
--min_timestep_ratio 0.02 \
--max_timestep_ratio 0.98 \
--real_score_guidance_scale 3.5 \
--seed 1024 \
--VSA_sparsity 0.8
+1 -1
View File
@@ -9,7 +9,7 @@ NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# 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
# Make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/training/wan_training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
+1 -1
View File
@@ -14,7 +14,7 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
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
# Make sure that num_latent_t is a multiple of sp_size
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/training/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
+5 -3
View File
@@ -1,9 +1,10 @@
#!/bin/bash
num_gpus=2
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# You can either use --prompt or --prompt-txt, but not both.
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
@@ -14,8 +15,9 @@ fastvideo generate \
--num-frames 77 \
--num-inference-steps 50 \
--fps 16 \
--guidance-scale 3.0 \
--prompt "A beautiful woman in a red dress walking down a street" \
--guidance-scale 6.0 \
--flow-shift 8.0 \
--prompt-txt assets/prompt.txt \
--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" \
--seed 1024 \
--output-path outputs_video/
+23
View File
@@ -0,0 +1,23 @@
#!/bin/bash
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
export MODEL_BASE=FastVideo/FastWan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# You can either use --prompt or --prompt-txt, but not both.
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--height 448 \
--width 832 \
--num-frames 61 \
--num-inference-steps 3 \
--fps 16 \
--prompt-txt assets/prompt.txt \
--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" \
--seed 1024 \
--output-path outputs_video_dmd/ \
--VSA-sparsity 0.8 \
--dmd-denoising-steps "1000,757,522"