Compare commits

..
Author SHA1 Message Date
SolitaryThinker eb56c7f8c5 add wan22 test 2025-09-15 10:17:02 +00:00
SolitaryThinker e63fe531d8 fix 2025-09-15 06:16:20 +00:00
Wenxuan Tanandgemini-code-assist[bot] 2930abe456 [Bugfix] Fix VMoba requirements (#802)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-09-14 18:28:52 -07:00
William Lin b93ef4289d [bugfix] Fix empty PipelineConfigs for Wan2.2 A14B (#800) 2025-09-13 17:31:38 -07:00
401bdbd316 [self-forcing] [3/n] Text embed only preprocessing (#797)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-09-13 14:03:53 -07:00
William Lin 1048d79cf8 [bugfix] pin gradio version and set current_vsa_sparsity in TrainingPipeline (#798) 2025-09-11 17:04:47 -07:00
1e8406162d [bugfix] Fix delta calculation (#796)
Co-authored-by: zbchu2 <zbchu2@iflytek.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-09-11 16:31:23 -07:00
William Lin 03edd35c83 [preprocessing] [self-forcing] [2/n] Improve preprocessing and add ode trajectory dataset schema (#794) 2025-09-10 17:33:57 -07:00
48 changed files with 779 additions and 891 deletions
+12 -1
View File
@@ -198,4 +198,15 @@ steps:
env:
- TEST_TYPE=inference_vmoba
agents:
queue: "default"
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Unit Tests"
env:
- TEST_TYPE=unit_test
agents:
queue: "default"
+4
View File
@@ -118,6 +118,10 @@ case "$TEST_TYPE" in
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
+34 -9
View File
@@ -62,8 +62,8 @@ on:
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
run_unit_test:
description: "Run unit-test"
required: false
default: false
type: boolean
@@ -93,6 +93,7 @@ jobs:
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
unit-test: ${{ steps.filter.outputs.unit-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -102,6 +103,8 @@ jobs:
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.10'
- 'docker/Dockerfile.python3.11'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/sliding_tile_attn/**'
@@ -155,6 +158,9 @@ jobs:
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
unit-test:
- 'fastvideo/**'
- *common-paths
encoder-test:
needs: change-filter
@@ -333,23 +339,42 @@ jobs:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
nightly-test:
unit-test:
needs: change-filter
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
job_id: "unit-test"
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
# nightly-test:
# if: >-
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
# uses: ./.github/workflows/runpod-test.yml
# with:
# job_id: "nightly-test"
# gpu_type: "NVIDIA A40"
# gpu_count: 4
# volume_size: 100
# disk_size: 100
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
# timeout_minutes: 30
# secrets:
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
# Add other jobs to this list as you create them
+3
View File
@@ -64,3 +64,6 @@ docs/source/distillation/examples/
!docs/source/_static/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
preprocess_output_text/
+3 -1
View File
@@ -20,5 +20,7 @@ setup(
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[]
install_requires=[
"flash-attn >= 2.7.1",
]
)
+10 -2
View File
@@ -6,8 +6,16 @@ import time
import os
import torch
from typing import Tuple
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
try:
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
from functools import lru_cache
from einops import rearrange
+9
View File
@@ -0,0 +1,9 @@
# VidProm Dataset
From [Self-Forcing](https://github.com/gdhe17/Self-Forcing) repository.
## Download the dataset
```bash
./download_dataset.sh
```
@@ -0,0 +1,3 @@
#! /bin/bash
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
@@ -1,47 +0,0 @@
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.
The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
The video shows a large orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
@@ -1,93 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 77
--warp_denoising_step
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--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 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,24 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory"
@@ -1,40 +0,0 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -28,4 +28,4 @@
"num_frames": 77
}
]
}
}
+4 -4
View File
@@ -5,7 +5,6 @@ from dataclasses import dataclass
import torch
from einops import rearrange
from flash_attn.bert_padding import pad_input
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
@@ -134,6 +133,8 @@ class VMOBAAttentionImpl(AttentionImpl):
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
from flash_attn.bert_padding import pad_input
self.pad_input = pad_input
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
@@ -169,7 +170,6 @@ class VMOBAAttentionImpl(AttentionImpl):
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
# torch.distributed.breakpoint()
query, chunk_size = process_moba_input(query,
attn_metadata.patch_resolution,
moba_chunk_size)
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = process_moba_output(hidden_states,
attn_metadata.patch_resolution,
moba_chunk_size)
@@ -92,6 +92,9 @@ class WanVideoArchConfig(DiTArchConfig):
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
+2 -3
View File
@@ -4,13 +4,12 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
"get_pipeline_config_cls_from_name"
]
+3 -5
View File
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
# isort: off
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig)
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
WanT2V720PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -48,7 +48,6 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -61,7 +60,6 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
+15 -8
View File
@@ -82,7 +82,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
default_factory=CLIPVisionConfig)
image_encoder_precision: str = "fp32"
def __post_init__(self):
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -108,19 +108,17 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
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):
flow_shift: float | None = 5.0
ti2v_task: bool = True
expand_timesteps: bool = True
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self.dit_config.expand_timesteps = self.expand_timesteps
@dataclass
@@ -132,12 +130,21 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
@dataclass
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
pass
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio
@dataclass
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
pass
class Wan2_2_I2V_A14B_Config(WanI2V480PConfig):
flow_shift: float | None = 5.0
boundary_ratio: float | None = 0.900
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.boundary_ratio = self.boundary_ratio
# =============================================
+7 -14
View File
@@ -40,6 +40,7 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
# TeaCache parameters
enable_teacache: bool = False
@@ -47,8 +48,6 @@ class SamplingParam:
# Misc
save_video: bool = True
return_frames: bool = False
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
@@ -169,6 +168,12 @@ class SamplingParam:
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--boundary-ratio",
type=float,
default=SamplingParam.boundary_ratio,
help="Boundary timestep ratio",
)
parser.add_argument(
"--save-video",
action="store_true",
@@ -200,18 +205,6 @@ class SamplingParam:
help=
"Path to a JSON file containing V-MoBA specific configurations.",
)
parser.add_argument(
"--return-trajectory-latents",
action="store_true",
default=SamplingParam.return_trajectory_latents,
help="Whether to return the trajectory",
)
parser.add_argument(
"--return-trajectory-decoded",
action="store_true",
default=SamplingParam.return_trajectory_decoded,
help="Whether to return the decoded trajectory",
)
return parser
+8 -4
View File
@@ -144,18 +144,22 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 4.0
guidance_scale_2: float = 3.0
guidance_scale: float = 4.0 # high_noise
guidance_scale_2: float = 3.0 # low_noise
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 3.5
guidance_scale_2: float = 3.5
guidance_scale: float = 3.5 # high_noise
guidance_scale_2: float = 3.5 # low_noise
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
# =============================================
+1 -1
View File
@@ -37,7 +37,7 @@ def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
seed=args.seed)
def gettextdataset(args) -> TextDataset:
return TextDataset(data_merge_path=args.data_merge_path,
@@ -1,5 +1,7 @@
from typing import Any
import numpy as np
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
@@ -120,3 +122,69 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
})
return records
def ode_text_only_record_creator(
video_name: str, text_embedding: np.ndarray, caption: str,
trajectory_latents: np.ndarray,
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
Args:
video_name: Base name/id for the sample (without extension).
text_embedding: Text encoder output array [SeqLen, Dim].
caption: Original text prompt.
trajectory_latents: Collected trajectory latents array.
trajectory_timesteps: Collected timesteps array.
Returns:
dict suitable for records_to_table(…, pyarrow_schema_ode_trajectory_text_only)
"""
assert trajectory_latents is not None, "trajectory_latents is required"
assert trajectory_timesteps is not None, "trajectory_timesteps is required"
record = {
"id": f"text_{video_name}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"file_name": video_name,
"caption": caption,
"media_type": "text",
}
record.update({
"trajectory_latents_bytes": trajectory_latents.tobytes(),
"trajectory_latents_shape": list(trajectory_latents.shape),
"trajectory_latents_dtype": str(trajectory_latents.dtype),
})
record.update({
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
return record
def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
caption: str) -> dict[str, Any]:
"""Create a text-only record matching pyarrow_schema_text_only.
Args:
text_name: Base id/name for the text sample.
text_embedding: Text encoder output array [SeqLen, Dim].
caption: Original text prompt.
Returns:
dict suitable for records_to_table(…, pyarrow_schema_text_only)
"""
record = {
"id": f"text_{text_name}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"caption": caption,
}
return record
+14
View File
@@ -102,3 +102,17 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # Always 'text' for text-only
])
pyarrow_schema_text_only = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("caption", pa.string()),
])
-3
View File
@@ -344,9 +344,6 @@ class VideoGenerator:
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
def set_lora_adapter(self,
+5 -3
View File
@@ -77,9 +77,11 @@ class BaseLayerWithLoRA(nn.Module):
lora_A = self.lora_A.to_local()
if not self.merged and not self.disable_lora:
delta = x @ (
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(x, non_blocking=True))
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(x, non_blocking=True))
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank:
delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore
@@ -40,7 +40,7 @@ class ComposedPipelineBase(ABC):
_extra_config_module_map: dict[str, str] = {}
training_args: TrainingArgs | None = None
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
modules: dict[str, Any] = {}
modules: dict[str, torch.nn.Module] = {}
post_init_called: bool = False
# TODO(will): args should support both inference args and training args
@@ -237,20 +237,19 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# @TODO(Wei): Temporary hack
if "boundary_ratio" in model_index and model_index[
"boundary_ratio"] is not None:
logger.info(
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
)
self.required_config_modules.append("transformer_2")
if fastvideo_args.boundary_ratio is None:
logger.info(
"MoE pipeline detected. Setting boundary ratio to %s",
model_index["boundary_ratio"])
fastvideo_args.boundary_ratio = model_index["boundary_ratio"]
logger.info("MoE pipeline detected. Setting boundary ratio to %s",
model_index["boundary_ratio"])
fastvideo_args.pipeline_config.dit_config.boundary_ratio = model_index[
"boundary_ratio"]
model_index.pop("boundary_ratio", None)
# used by Wan2.2 ti2v
model_index.pop("expand_timesteps", None)
# some sanity checks
@@ -283,8 +282,8 @@ class ComposedPipelineBase(ABC):
architecture) in model_index.items():
if transformers_or_diffusers is None:
logger.warning(
"Module in model_index.json has null value, removing from required_config_modules"
)
"Module %s in model_index.json has null value, removing from required_config_modules",
module_name)
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
+2 -10
View File
@@ -129,6 +129,7 @@ class ForwardBatch:
timesteps: torch.Tensor | None = None
timestep: torch.Tensor | float | int | None = None
step_index: int | None = None
boundary_ratio: float | None = None
# Scheduler parameters
num_inference_steps: int = 50
@@ -147,12 +148,7 @@ class ForwardBatch:
modules: dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion)
output: torch.Tensor | None = None
return_trajectory_latents: bool = False
return_trajectory_decoded: bool = False
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None
output: Any = None
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
@@ -211,10 +207,6 @@ class TrainingBatch:
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
# ODE trajectory supervision
trajectory_latents: torch.Tensor | None = None
trajectory_timesteps: torch.Tensor | None = None
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
timesteps: torch.Tensor | None = None
@@ -10,6 +10,8 @@ from torch.utils.data import DataLoader
from tqdm import tqdm
from fastvideo.dataset import getdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.preprocessing_datasets import PreprocessBatch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -17,8 +19,6 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import TextEncodingStage
from fastvideo.workflow.preprocess.parquet_io import (ParquetDatasetWriter,
records_to_table)
logger = init_logger(__name__)
@@ -423,7 +423,3 @@ class BasePreprocessPipeline(ComposedPipelineBase):
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
num_processed_samples = 0
def _final_flush_if_any(self):
if hasattr(self, 'dataset_writer'):
self.dataset_writer.flush()
@@ -1,399 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
ODE Trajectory Data Preprocessing pipeline implementation.
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
using the modular pipeline architecture.
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
"""
import os
from collections.abc import Iterator
from typing import Any
import numpy as np
import pyarrow as pa
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
from fastvideo.workflow.preprocess.parquet_io import (ParquetDatasetWriter,
records_to_table)
logger = init_logger(__name__)
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int
def get_pyarrow_schema(self) -> pa.Schema:
"""Return the PyArrow schema for ODE Trajectory pipeline."""
return pyarrow_schema_ode_trajectory_text_only
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
logger.info('WTF flow_shift: %s',
fastvideo_args.pipeline_config.flow_shift)
assert fastvideo_args.pipeline_config.flow_shift == 5
# self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
# shift=fastvideo_args.pipeline_config.flow_shift)
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
denoising_strength=1.0)
# logger.info('WTF scheduler timesteps: %s',
# self.modules["scheduler"].timesteps)
# scheduler = FlowMatchScheduler(
# shift=8.0, sigma_min=0.0, extra_one_step=True)
# device = get_local_torch_device()
# # scheduler.num_train_timesteps = 100
# scheduler.set_timesteps(num_inference_steps=50, denoising_strength=1.0)
# scheduler.sigmas = scheduler.sigmas.to(device)
# self.modules["scheduler"] = scheduler
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="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data and generate trajectory information."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
# Add fps and duration if available in data
if "fps" in data:
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
if "duration" in data:
valid_data["duration"] = [
data["duration"][i] for i in valid_indices
]
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.num_inference_steps = 48
batch.return_trajectory_latents = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
extra_features = {
"trajectory_latents": trajectory_latents,
"trajectory_timesteps": trajectory_timesteps
}
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).numpy()
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (without VAE latents for text-only)
record = self.create_text_only_record(
args,
video_name=video_name,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
self.get_pyarrow_schema())
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self.num_processed_samples = 0
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush()
if written:
logger.info("Final flush wrote %s samples", written)
def create_text_only_record(
self,
args,
video_name: str,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{video_name}_{idx}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "text",
}
assert extra_features is not None, "extra_features is required"
assert "trajectory_latents" in extra_features, "trajectory_latents is required"
assert "trajectory_timesteps" in extra_features, "trajectory_timesteps is required"
# Add trajectory data if available
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features[
"trajectory_latents"][idx] if isinstance(
extra_features["trajectory_latents"],
list) else extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes":
trajectory_latents.tobytes(),
"trajectory_latents_shape":
list(trajectory_latents.shape),
"trajectory_latents_dtype":
str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features[
"trajectory_timesteps"][idx] if isinstance(
extra_features["trajectory_timesteps"],
list) else extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes":
trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape":
list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype":
str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for video preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing videos",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.video_data: dict[str, Any] = {} # Store video metadata and paths
self.latent_data: dict[str, Any] = {} # Store latent tensors
self.preprocess_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_ODE_Trajectory
@@ -0,0 +1,184 @@
# SPDX-License-Identifier: Apache-2.0
"""
Text-only Data Preprocessing pipeline implementation.
This module contains an implementation of the Text-only Data Preprocessing pipeline
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
"""
import os
from collections.abc import Iterator
from typing import Any
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import text_only_record_creator
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import TextEncodingStage
logger = init_logger(__name__)
class PreprocessPipeline_Text(BasePreprocessPipeline):
"""Text-only preprocessing pipeline implementation."""
_required_config_modules = ["text_encoder", "tokenizer"]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int = 0
def get_pyarrow_schema(self):
"""Return the PyArrow schema for text-only pipeline."""
return pyarrow_schema_text_only
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
def preprocess_text_only(self, fastvideo_args: FastVideoArgs, args):
"""Preprocess text-only data."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
logger.info("===== prompt_embeds: %s", prompt_embeds.shape)
logger.info("===== prompt_attention_masks: %s",
prompt_attention_masks.shape)
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, text_path in save_pbar:
text_name = os.path.basename(text_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Create record for Parquet dataset (text-only schema)
record = text_only_record_creator(
text_name=text_name,
text_embedding=text_embedding,
caption=valid_data["text"][idx],
)
batch_data.append(record)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
pyarrow_schema_text_only)
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self.num_processed_samples = 0
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush(write_remainder=True)
if written:
logger.info("Final flush wrote %s samples", written)
# Text-only record creation moved to fastvideo.dataset.dataloader.record_schema
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading text dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for text preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing text",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.text_data: dict[str, Any] = {} # Store text metadata and paths
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text
+28 -11
View File
@@ -1,5 +1,6 @@
import argparse
import os
from typing import Any
from fastvideo import PipelineConfig
from fastvideo.configs.models.vaes import WanVAEConfig
@@ -9,10 +10,10 @@ from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -23,13 +24,22 @@ def main(args) -> None:
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
"flow_shift": 5,
}
kwargs: dict[str, Any] = {}
if args.preprocess_task == "text_only":
kwargs = {
"text_encoder_cpu_offload": False,
}
else:
# Full config for video/image processing
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
@@ -38,17 +48,20 @@ def main(args) -> None:
text_encoder_cpu_offload=False,
pipeline_config=pipeline_config,
)
if args.preprocess_task == "t2v":
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "ode_trajectory":
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
elif args.preprocess_task == "text_only":
PreprocessPipeline = PreprocessPipeline_Text
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}")
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
logger.info("Preprocess task: %s using %s", args.preprocess_task,
PreprocessPipeline.__name__)
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -87,7 +100,11 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
+50 -93
View File
@@ -50,63 +50,6 @@ class DecodingStage(PipelineStage):
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
return result
@torch.no_grad()
def decode(self, latents: torch.Tensor,
fastvideo_args: FastVideoArgs) -> torch.Tensor:
"""
Decode latent representations into pixel space using VAE.
Args:
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
fastvideo_args: Configuration containing:
- disable_autocast: Whether to disable automatic mixed precision (default: False)
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
Returns:
Decoded video tensor with shape (batch, channels, frames, height, width),
normalized to [0, 1] range and moved to CPU as float32
"""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device,
latents.dtype)
else:
latents += self.vae.shift_factor
# Decode latents
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
return image
@torch.no_grad()
def forward(
self,
@@ -116,28 +59,13 @@ class DecodingStage(PipelineStage):
"""
Decode latent representations into pixel space.
This method processes the batch through the VAE decoder, converting latent
representations to pixel-space video/images. It also optionally decodes
trajectory latents for visualization purposes.
Args:
batch: The current batch containing:
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
- return_trajectory_decoded (optional): Flag to decode trajectory latents
- trajectory_latents (optional): Latents at different timesteps
- trajectory_timesteps (optional): Corresponding timesteps
fastvideo_args: Configuration containing:
- output_type: "latent" to skip decoding, otherwise decode to pixels
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
- model_loaded: Track VAE loading state
- model_paths: Path to VAE model if loading needed
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
Modified batch with:
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
- trajectory_decoded (if requested): List of decoded frames per timestep
The batch with decoded outputs.
"""
# load vae if not already loaded (used for memory constrained devices)
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["vae"]:
loader = VAELoader()
@@ -147,29 +75,58 @@ class DecodingStage(PipelineStage):
pipeline.add_module("vae", self.vae)
fastvideo_args.model_loaded["vae"] = True
if fastvideo_args.output_type == "latent":
frames = batch.latents
else:
frames = self.decode(batch.latents, fastvideo_args)
self.vae = self.vae.to(get_local_torch_device())
# decode trajectory latents if needed
if batch.return_trajectory_decoded:
batch.trajectory_decoded = []
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
for idx in range(batch.trajectory_latents.shape[1]):
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
cur_timestep = batch.trajectory_timesteps[idx]
logger.info("decoding trajectory latent for timestep: %s",
cur_timestep)
decoded_frames = self.decode(cur_latent, fastvideo_args)
batch.trajectory_decoded.append(decoded_frames.cpu().float())
latents = batch.latents
# TODO(will): remove this once we add input/output validation for stages
if latents is None:
raise ValueError("Latents must be provided")
# Skip decoding if output type is latent
if fastvideo_args.output_type == "latent":
image = latents
else:
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32
) and not fastvideo_args.disable_autocast
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device,
latents.dtype)
else:
latents += self.vae.shift_factor
# Decode latents
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
# Convert to CPU float32 for compatibility
frames = frames.cpu().float()
image = image.cpu().float()
# Update batch with decoded image
batch.output = frames
batch.output = image
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
+8 -30
View File
@@ -204,8 +204,14 @@ class DenoisingStage(PipelineStage):
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
if fastvideo_args.boundary_ratio is not None:
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_ratio is not None:
logger.info("Overriding boundary ratio from %s to %s",
boundary_ratio, batch.boundary_ratio)
boundary_ratio = batch.boundary_ratio
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
@@ -247,10 +253,6 @@ class DenoisingStage(PipelineStage):
patch_size[2])
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -431,11 +433,6 @@ class DenoisingStage(PipelineStage):
latents = (1. - mask2[0]) * z + mask2[0] * latents
# latents = latents.unsqueeze(0)
# save trajectory latents if needed
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
trajectory_latents.append(latents)
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
@@ -443,28 +440,9 @@ class DenoisingStage(PipelineStage):
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
else:
trajectory_tensor = None
trajectory_timesteps_tensor = None
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
if batch.return_trajectory_latents:
trajectory_tensor = trajectory_tensor.to(
get_local_torch_device())
trajectory_tensor = sequence_model_parallel_all_gather(
trajectory_tensor, dim=3)
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
# Update batch with final latents
batch.latents = latents
@@ -0,0 +1,123 @@
import numpy as np
from fastvideo.dataset.dataloader.record_schema import (
basic_t2v_record_creator,
i2v_record_creator,
ode_text_only_record_creator,
text_only_record_creator,
)
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
def _mk_basic_batch(N: int) -> PreprocessBatch:
batch = PreprocessBatch(data_type="video")
batch.video_file_name = [f"vid_{i}" for i in range(N)]
batch.prompt = [f"caption_{i}" for i in range(N)]
batch.width = [640 for _ in range(N)]
batch.height = [360 for _ in range(N)]
batch.fps = [4 for _ in range(N)]
batch.num_frames = [2 for _ in range(N)]
# Latents: shape (N, C, T, H, W); per-record use latents[idx]
batch.latents = np.zeros((N, 4, 2, 8, 8), dtype=np.float32)
# Prompt embeds: list of per-record arrays [Seq, Dim]
batch.prompt_embeds = [np.ones((6, 16), dtype=np.float32) for _ in range(N)]
return batch
def test_basic_t2v_record_creator_fields():
N = 2
batch = _mk_basic_batch(N)
records = basic_t2v_record_creator(batch)
assert isinstance(records, list) and len(records) == N
for i, rec in enumerate(records):
assert rec["id"] == batch.video_file_name[i]
# Latents bytes/shape/dtype
assert isinstance(rec["vae_latent_bytes"], (bytes, bytearray))
assert rec["vae_latent_shape"] == list(batch.latents[i].shape)
assert rec["vae_latent_dtype"] == str(batch.latents[i].dtype)
# Text embedding
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(batch.prompt_embeds[i].shape)
assert rec["text_embedding_dtype"] == str(batch.prompt_embeds[i].dtype)
# Meta
assert rec["caption"] == batch.prompt[i]
assert rec["media_type"] == "video"
assert rec["width"] == int(batch.width[i])
assert rec["height"] == int(batch.height[i])
assert rec["num_frames"] == batch.latents[i].shape[1]
def test_i2v_record_creator_additional_fields():
N = 3
batch = _mk_basic_batch(N)
# image_embeds is a list of length 1, with an array of shape [N, D]
batch.image_embeds = [np.ones((N, 32), dtype=np.float32)]
# first frame latent per record
batch.image_latent = np.zeros((N, 4, 1, 8, 8), dtype=np.float32)
# pil image per record
batch.pil_image = np.zeros((N, 8, 8, 3), dtype=np.uint8)
records = i2v_record_creator(batch)
assert isinstance(records, list) and len(records) == N
for i, rec in enumerate(records):
# clip feature
assert isinstance(rec["clip_feature_bytes"], (bytes, bytearray))
assert rec["clip_feature_shape"] == list(batch.image_embeds[0][i].shape)
assert rec["clip_feature_dtype"] == str(batch.image_embeds[0][i].dtype)
# first frame latent
assert isinstance(rec["first_frame_latent_bytes"], (bytes, bytearray))
assert rec["first_frame_latent_shape"] == list(batch.image_latent[i].shape)
assert rec["first_frame_latent_dtype"] == str(batch.image_latent[i].dtype)
# pil image
assert isinstance(rec["pil_image_bytes"], (bytes, bytearray))
assert rec["pil_image_shape"] == list(batch.pil_image[i].shape)
assert rec["pil_image_dtype"] == str(batch.pil_image[i].dtype)
def test_ode_text_only_record_creator():
video_name = "ex"
caption = "a prompt"
text_embedding = np.ones((6, 16), dtype=np.float32)
traj = np.ones((5, 4, 2, 2), dtype=np.float32)
tsteps = np.arange(5, dtype=np.float32)
rec = ode_text_only_record_creator(
video_name=video_name,
text_embedding=text_embedding,
caption=caption,
trajectory_latents=traj,
trajectory_timesteps=tsteps,
)
assert rec["id"] == f"text_{video_name}"
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(text_embedding.shape)
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
assert rec["file_name"] == video_name
assert rec["caption"] == caption
assert rec["media_type"] == "text"
# Trajectory fields
assert isinstance(rec["trajectory_latents_bytes"], (bytes, bytearray))
assert rec["trajectory_latents_shape"] == list(traj.shape)
assert rec["trajectory_latents_dtype"] == str(traj.dtype)
assert isinstance(rec["trajectory_timesteps_bytes"], (bytes, bytearray))
assert rec["trajectory_timesteps_shape"] == list(tsteps.shape)
assert rec["trajectory_timesteps_dtype"] == str(tsteps.dtype)
def test_text_only_record_creator():
text_name = "note1"
caption = "a prompt"
text_embedding = np.ones((7, 16), dtype=np.float32)
rec = text_only_record_creator(
text_name=text_name,
text_embedding=text_embedding,
caption=caption,
)
assert rec["id"] == f"text_{text_name}"
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(text_embedding.shape)
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
assert rec["caption"] == caption
+4
View File
@@ -117,3 +117,7 @@ def run_inference_lora_tests():
@app.function(gpu="L40S:2", image=image, timeout=900)
def run_distill_dmd_tests():
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
@@ -28,14 +28,14 @@ HUNYUAN_PARAMS = {
"width": 1280,
"num_frames": 45,
"num_inference_steps": 6,
"guidance_scale": 1,
"embedded_cfg_scale": 6,
"flow_shift": 17,
# "guidance_scale": 1,
# "embedded_cfg_scale": 6,
# "flow_shift": 17,
"seed": 1024,
"sp_size": 2,
"tp_size": 1,
"vae_sp": True,
"fps": 24,
# "vae_sp": True,
# "fps": 24,
}
WAN_T2V_PARAMS = {
@@ -45,14 +45,14 @@ WAN_T2V_PARAMS = {
"width": 832,
"num_frames": 45,
"num_inference_steps": 20,
"guidance_scale": 3,
"embedded_cfg_scale": 6,
"flow_shift": 7.0,
# "guidance_scale": 3,
# "embedded_cfg_scale": 6,
# "flow_shift": 7.0,
"seed": 1024,
"sp_size": 2,
"tp_size": 1,
"vae_sp": True,
"fps": 24,
# "vae_sp": True,
# "fps": 24,
"neg_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",
"text-encoder-precision": ("fp32",)
}
@@ -64,18 +64,33 @@ WAN_I2V_PARAMS = {
"width": 832,
"num_frames": 45,
"num_inference_steps": 6,
"guidance_scale": 5.0,
"embedded_cfg_scale": 6,
"flow_shift": 7.0,
# "guidance_scale": 5.0,
# "embedded_cfg_scale": 6,
# "flow_shift": 7.0,
"seed": 1024,
"sp_size": 2,
"tp_size": 1,
"vae_sp": True,
"fps": 24,
"neg_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",
# "vae_sp": True,
# "fps": 24,
# "neg_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",
"text-encoder-precision": ("fp32",)
}
WAN2_2_I2V_PARAMS = {
"num_gpus": 2,
"model_path": "Wan-AI/Wan2.2-I2V-A14B-Diffusers",
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 20,
# "guidance_scale": 5.0,
# "embedded_cfg_scale": 6,
# "flow_shift": 7.0,
"seed": 1024,
"sp_size": 2,
"tp_size": 1,
}
MODEL_TO_PARAMS = {
"FastHunyuan-diffusers": HUNYUAN_PARAMS,
"Wan2.1-T2V-1.3B-Diffusers": WAN_T2V_PARAMS,
@@ -83,6 +98,7 @@ MODEL_TO_PARAMS = {
I2V_MODEL_TO_PARAMS = {
"Wan2.1-I2V-14B-480P-Diffusers": WAN_I2V_PARAMS,
"Wan2.2-I2V-A14B-Diffusers": WAN2_2_I2V_PARAMS,
}
TEST_PROMPTS = [
@@ -125,7 +141,6 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"flow_shift": BASE_PARAMS["flow_shift"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
}
@@ -142,10 +157,7 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"guidance_scale": BASE_PARAMS["guidance_scale"],
"embedded_cfg_scale": BASE_PARAMS["embedded_cfg_scale"],
"seed": BASE_PARAMS["seed"],
"fps": BASE_PARAMS["fps"],
}
if "neg_prompt" in BASE_PARAMS:
generation_kwargs["neg_prompt"] = BASE_PARAMS["neg_prompt"]
@@ -225,7 +237,6 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"flow_shift": BASE_PARAMS["flow_shift"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
"dit_cpu_offload": True,
@@ -242,10 +253,7 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"guidance_scale": BASE_PARAMS["guidance_scale"],
"embedded_cfg_scale": BASE_PARAMS["embedded_cfg_scale"],
"seed": BASE_PARAMS["seed"],
"fps": BASE_PARAMS["fps"],
}
if "neg_prompt" in BASE_PARAMS:
generation_kwargs["neg_prompt"] = BASE_PARAMS["neg_prompt"]
@@ -38,7 +38,8 @@ def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
data_type="video",
latents=torch.randn(B, 2),
prompt_embeds=[torch.randn(B, 1, 1)],
prompt_attention_mask=[torch.ones(B, 1)],
# Attention mask should be integer dtype in real pipelines
prompt_attention_mask=[torch.ones(B, 1, dtype=torch.int64)],
)
batch.video_file_name = [f"vid_{i}" for i in range(B)]
@@ -52,13 +53,13 @@ def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
out_dir = tmp_path / "saver_out"
saver.save_and_write_parquet_batch(batch, str(out_dir))
# First flush: should write one full file (3 rows), keep 2 in buffer
saver.flush_tables(str(out_dir))
saver.flush_tables()
files = sorted(out_dir.rglob("*.parquet"))
assert len(files) == 1
assert pq.read_table(str(files[0])).num_rows == 3
# Final flush: write remainder 2 rows
saver.flush_last(str(out_dir))
saver.flush_tables(write_remainder=True)
files2 = sorted(out_dir.rglob("*.parquet"))
assert len(files2) == 2
total = sum(pq.read_table(str(f)).num_rows for f in files2)
+1 -1
View File
@@ -478,7 +478,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
current_vsa_sparsity = current_decay_times * vsa_decay_rate
elif vmoba_available:
# TODO: add vmoba sparsity scheduling here
pass
current_vsa_sparsity = 0.0
else:
current_vsa_sparsity = 0.0
-18
View File
@@ -23,14 +23,10 @@ from typing import Any, TypeVar, cast
import cloudpickle
import filelock
import imageio
import numpy as np
import torch
import torchvision
import yaml
from diffusers.loaders.lora_base import (
_best_guess_weight_name) # watch out for potetential removal from diffusers
from einops import rearrange
from huggingface_hub import snapshot_download
from remote_pdb import RemotePdb
from torch.distributed.fsdp import MixedPrecisionPolicy
@@ -890,17 +886,3 @@ def best_output_size(w, h, dw, dh, expected_area):
return ow1, oh1
else:
return ow2, oh2
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor],
output_path: str, fps: int):
# Process outputs
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(output_path), exist_ok=True)
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
+6 -11
View File
@@ -222,7 +222,7 @@ class ParquetDatasetSaver:
# If flush is needed
if self.num_processed_samples >= self.flush_frequency:
self.flush_tables(output_dir)
self.flush_tables()
def _process_non_padded_embeddings(
self, prompt_embeds: torch.Tensor,
@@ -245,12 +245,7 @@ class ParquetDatasetSaver:
return non_padded_embeds
def _convert_batch_to_pyarrow_table(self,
batch_data: list[dict]) -> pa.Table:
# Deprecated path, kept for backward compatibility if needed.
return records_to_table(batch_data, self.schema)
def flush_tables(self, output_dir: str, write_remainder: bool = False):
def flush_tables(self, write_remainder: bool = False):
"""Flush buffered records to disk.
Args:
@@ -266,16 +261,16 @@ class ParquetDatasetSaver:
remainder = self.num_processed_samples % self.samples_per_file
self.num_processed_samples = 0 if write_remainder else remainder
def flush_last(self, output_dir: str):
"""Flush and write any remaining rows (final flush)."""
self.flush_tables(output_dir, write_remainder=True)
def clean_up(self) -> None:
"""Clean up all tables"""
self.flush_tables(write_remainder=True)
self._writer = None
self.num_processed_samples = 0
gc.collect()
def __del__(self):
self.clean_up()
def build_dataset(preprocess_config: PreprocessConfig, split: str,
validator: Callable[[dict[str, Any]], bool]) -> Dataset:
@@ -4,6 +4,8 @@ from typing import cast
from torch.utils.data import DataLoader
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.dataset.dataloader.record_schema import (
basic_t2v_record_creator, i2v_record_creator)
from fastvideo.dataset.dataloader.schema import (pyarrow_schema_i2v,
pyarrow_schema_t2v)
from fastvideo.distributed.parallel_state import get_world_rank
@@ -13,8 +15,6 @@ from fastvideo.pipelines.pipeline_registry import PipelineType
from fastvideo.workflow.preprocess.components import (
ParquetDatasetSaver, PreprocessingDataValidator, VideoForwardBatchBuilder,
build_dataset)
from fastvideo.workflow.preprocess.record_schema import (
basic_t2v_record_creator, i2v_record_creator)
from fastvideo.workflow.workflow_base import WorkflowBase
logger = init_logger(__name__)
@@ -35,8 +35,7 @@ class PreprocessWorkflowI2V(PreprocessWorkflow):
self.processed_dataset_saver.save_and_write_parquet_batch(
forward_batch, self.training_dataset_output_dir)
self.processed_dataset_saver.flush_tables(
self.training_dataset_output_dir)
self.processed_dataset_saver.flush_tables()
self.processed_dataset_saver.clean_up()
# Validation dataset preprocessing
@@ -51,6 +50,5 @@ class PreprocessWorkflowI2V(PreprocessWorkflow):
self.processed_dataset_saver.save_and_write_parquet_batch(
forward_batch, self.validation_dataset_output_dir)
self.processed_dataset_saver.flush_tables(
self.validation_dataset_output_dir)
self.processed_dataset_saver.flush_tables()
self.processed_dataset_saver.clean_up()
@@ -35,8 +35,7 @@ class PreprocessWorkflowT2V(PreprocessWorkflow):
self.processed_dataset_saver.save_and_write_parquet_batch(
forward_batch, self.training_dataset_output_dir)
self.processed_dataset_saver.flush_tables(
self.training_dataset_output_dir)
self.processed_dataset_saver.flush_tables()
self.processed_dataset_saver.clean_up()
# Validation dataset preprocessing
@@ -51,6 +50,5 @@ class PreprocessWorkflowT2V(PreprocessWorkflow):
self.processed_dataset_saver.save_and_write_parquet_batch(
forward_batch, self.validation_dataset_output_dir)
self.processed_dataset_saver.flush_tables(
self.validation_dataset_output_dir)
self.processed_dataset_saver.flush_tables()
self.processed_dataset_saver.clean_up()
+1 -1
View File
@@ -34,7 +34,7 @@ dependencies = [
# Miscellaneous Utilities
"tqdm", "pytest", "PyYAML==6.0.1", "protobuf>=5.28.3",
"gradio>=5.22.0", "moviepy>=2.0.0", "flask",
"gradio==5.41.0", "moviepy>=2.0.0", "flask",
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
# System & Monitoring Tools
"gpustat", "watch", "remote-pdb",
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output_text
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..7}; do
# Calculate the starting file number for this node
start_file=$((node_id * 8 + 1))
echo "Launching text-only node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
echo "sbatch --job-name=text-${node_id} --output=preprocess_output_text/preprocess-text-node-${node_id}.out --error=preprocess_output_text/preprocess-text-node-${node_id}.err scripts/preprocess/syn_text.slurm $start_file $node_id"
sbatch --job-name=text-${node_id} \
--output=preprocess_output_text/preprocess-text-node-${node_id}.out \
--error=preprocess_output_text/preprocess-text-node-${node_id}.err \
scripts/preprocess/syn_text.slurm $start_file $node_id
done
echo "All 8 text-only nodes launched successfully!"
+88
View File
@@ -0,0 +1,88 @@
#!/bin/bash
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=8
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=16
#SBATCH --mem=960G
#SBATCH --exclusive
#SBATCH --time=72:00:00
# conda init
# source ~/conda/miniconda/bin/activate
# PYTHON_VIRTUAL_ENVIRONMENT=fastvideo-train-yq
# conda activate $PYTHON_VIRTUAL_ENVIRONMENT
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices for TEXT-ONLY preprocessing"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
num_gpus=1
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3 4 5 6 7)
GPU_NUM=1
MODEL_TYPE="wan"
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
echo "Processing TEXT-ONLY data"
# Run 8 parallel preprocessing jobs on this node
for i in {1..8}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
DATA_MERGE_PATH="prompts/v2m_${file_num}.txt"
# Create unique output directory based on node and GPU for text-only processing
OUTPUT_DIR="data/test-text-preprocessing/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
end_cpu=$(( start_cpu+1 ))
echo "Starting GPU $gpu processing text-only file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the text-only preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 2 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "text_only" &
done
# Wait for all jobs on this node to complete
wait
echo "All text-only processing blocks completed!"
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
--preprocess_task "i2v"
@@ -22,4 +22,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--samples_per_file 1 \
--flush_frequency 1 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
--preprocess_task "t2v"