Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a8459deeb | ||
|
|
c5f9ea53b2 | ||
|
|
d32a7184da | ||
|
|
2930abe456 | ||
|
|
b93ef4289d | ||
|
|
401bdbd316 | ||
|
|
1048d79cf8 | ||
|
|
1e8406162d | ||
|
|
03edd35c83 |
+12
-1
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -20,5 +20,7 @@ setup(
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[]
|
||||
install_requires=[
|
||||
"flash-attn >= 2.7.1",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,93 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="examples/dataset/mixkit/validation_64.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[@]}"
|
||||
@@ -0,0 +1,28 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
# MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="examples/dataset/crush_smol/crush_smol_prompts.txt"
|
||||
DATA_MERGE_PATH="test.txt"
|
||||
DATA_MERGE_PATH="/mnt/weka/home/hao.zhang/wl/Self-Forcing/prompts/vidprom_1.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_a14b_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 \
|
||||
--flow_shift 12.0 \
|
||||
--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
@@ -15,6 +15,7 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--flow_shift 5.0 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
@@ -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.
|
||||
@@ -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)
|
||||
|
||||
@@ -27,6 +27,7 @@ class DiTArchConfig(ArchConfig):
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -87,6 +87,7 @@ class PipelineConfig:
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -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
|
||||
@@ -169,6 +170,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",
|
||||
|
||||
@@ -144,18 +144,26 @@ 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
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
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
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
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
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
+68
@@ -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
|
||||
@@ -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()),
|
||||
])
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -375,7 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.num_frame_per_block = 1
|
||||
self.num_frame_per_block = 3
|
||||
self.independent_first_frame = False
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@@ -45,11 +45,17 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
sigma_start, self.sigma_min, num_inference_steps)
|
||||
if self.inverse_timesteps:
|
||||
self.sigmas = torch.flip(self.sigmas, dims=[0])
|
||||
logger.info("before shift sigmas: %s", self.sigmas)
|
||||
logger.info("sigma length: %s", len(self.sigmas))
|
||||
logger.info("shift: %s", self.shift)
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
logger.info("after shift sigmas: %s", self.sigmas)
|
||||
logger.info("after shift sigmas length: %s", len(self.sigmas))
|
||||
if self.reverse_sigmas:
|
||||
self.sigmas = 1 - self.sigmas
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
logger.info("final timesteps: %s", self.timesteps)
|
||||
if training:
|
||||
x = self.timesteps
|
||||
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
|
||||
@@ -62,8 +68,15 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
elif timestep.ndim == 0:
|
||||
# handles the case where timestep is a scalar, this occurs when we
|
||||
# use this scheduler for ODE trajectory
|
||||
timestep = timestep.unsqueeze(0)
|
||||
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep = timestep.to(model_output.device)
|
||||
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -12,7 +12,6 @@ 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
|
||||
@@ -21,6 +20,10 @@ from tqdm import tqdm
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
ode_text_only_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -36,8 +39,6 @@ from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
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__)
|
||||
|
||||
@@ -59,26 +60,22 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
|
||||
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)
|
||||
assert fastvideo_args.pipeline_config.flow_shift == 12
|
||||
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,
|
||||
logger.info("before 38 steps sigmas:")
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=38,
|
||||
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
|
||||
logger.info("before 39 steps sigmas:")
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=39,
|
||||
denoising_strength=1.0)
|
||||
logger.info("before 40 steps sigmas:")
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=40,
|
||||
denoising_strength=1.0)
|
||||
self.modules["scheduler"].timesteps = self.modules["scheduler"].timesteps.to(torch.int64)
|
||||
logger.info("after casting timesteps: %s", self.modules["scheduler"].timesteps)
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
@@ -97,6 +94,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
pipeline=self,
|
||||
))
|
||||
@@ -184,8 +182,10 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
batch.negative_attention_mask = [
|
||||
negative_prompt_attention_mask
|
||||
]
|
||||
batch.num_inference_steps = 48
|
||||
batch.num_inference_steps = 40
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
batch.return_trajectory_decoded = False
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
@@ -223,9 +223,14 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
for i, latent in enumerate(trajectory_latents[0]):
|
||||
logger.info(f"sum for timestep %s is %s", trajectory_timesteps[i], latent.float().sum())
|
||||
|
||||
for i in [0, 12, 24, 36]:
|
||||
logger.info(f"sum for timestep %s is %s", i, trajectory_latents[0][i].float().sum())
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
batch_data: list[dict[str, Any]] = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
@@ -254,14 +259,16 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
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,
|
||||
# Create record for Parquet dataset (text-only ODE schema)
|
||||
record: dict[str, Any] = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=sample_extra_features)
|
||||
caption=valid_data["text"][idx],
|
||||
trajectory_latents=sample_extra_features[
|
||||
"trajectory_latents"],
|
||||
trajectory_timesteps=sample_extra_features[
|
||||
"trajectory_timesteps"],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
@@ -289,78 +296,10 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush()
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
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()
|
||||
|
||||
@@ -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
|
||||
@@ -1,5 +1,6 @@
|
||||
import argparse
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
@@ -13,6 +14,8 @@ 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 +26,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(),
|
||||
@@ -42,13 +54,19 @@ def main(args) -> None:
|
||||
PreprocessPipeline = PreprocessPipeline_T2V
|
||||
elif args.preprocess_task == "i2v":
|
||||
PreprocessPipeline = PreprocessPipeline_I2V
|
||||
elif args.preprocess_task == "text_only":
|
||||
PreprocessPipeline = PreprocessPipeline_Text
|
||||
elif args.preprocess_task == "ode_trajectory":
|
||||
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
|
||||
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
|
||||
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
|
||||
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 +105,12 @@ 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("--flow_shift", type=float, default=None)
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
|
||||
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)
|
||||
|
||||
@@ -204,10 +204,20 @@ 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
|
||||
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
|
||||
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
|
||||
logger.info("boundary_timestep: %s", boundary_timestep)
|
||||
logger.info("boundary_ratio: %s", boundary_ratio)
|
||||
logger.info("self.scheduler.num_train_timesteps: %s", self.scheduler.num_train_timesteps)
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
|
||||
@@ -253,6 +263,13 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
logger.info("guidance_scale: %s", batch.guidance_scale)
|
||||
logger.info("guidance_scale_2: %s", batch.guidance_scale_2)
|
||||
logger.info("num_inference_steps: %s", num_inference_steps)
|
||||
logger.info("timesteps: %s", timesteps)
|
||||
timesteps[0] = timesteps[0] - 1.0
|
||||
# timesteps = timesteps.to(torch.int64)
|
||||
logger.info("new timesteps: %s", timesteps)
|
||||
for i, t in enumerate(timesteps):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
@@ -264,6 +281,7 @@ class DenoisingStage(PipelineStage):
|
||||
self.transformer_2.parameters()).device.type
|
||||
== 'cuda'):
|
||||
self.transformer_2.to('cpu')
|
||||
logger.info("using high-noise stage for timestep: %s", t)
|
||||
current_model = self.transformer
|
||||
current_guidance_scale = batch.guidance_scale
|
||||
else:
|
||||
@@ -272,6 +290,7 @@ class DenoisingStage(PipelineStage):
|
||||
self.transformer.parameters(
|
||||
)).device.type == 'cuda':
|
||||
self.transformer.to('cpu')
|
||||
logger.info("using low-noise stage for timestep: %s", t)
|
||||
current_model = self.transformer_2
|
||||
current_guidance_scale = batch.guidance_scale_2
|
||||
assert current_model is not None, "current_model is None"
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
|
||||
@@ -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,16 +53,14 @@ 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)
|
||||
assert total == 5
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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!"
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user