Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
74e58b7f64 |
@@ -6,7 +6,8 @@ from typing import Any
|
||||
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
|
||||
from fastvideo.v1.configs.sample.wan import (Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam)
|
||||
@@ -23,6 +24,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -94,3 +94,20 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
|
||||
-13.02252404
|
||||
]))
|
||||
|
||||
|
||||
# =============================================
|
||||
# ============= Wan2.1 Fun Models =============
|
||||
# =============================================
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
|
||||
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
guidance_scale: float = 6.0
|
||||
num_inference_steps: int = 50
|
||||
|
||||
@@ -429,7 +429,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
self.seed)
|
||||
logger.info("Initialized random seeds with seed: %s", self.seed)
|
||||
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=3)
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
@@ -439,7 +439,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
|
||||
step_times: deque[float] = deque(maxlen=100)
|
||||
|
||||
self._log_training_info()
|
||||
self._log_validation(self.transformer, self.training_args, 1)
|
||||
# self._log_validation(self.transformer, self.training_args, 1)
|
||||
|
||||
# Train!
|
||||
progress_bar = tqdm(
|
||||
|
||||
@@ -0,0 +1,232 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
scan_parquet_mt.py
|
||||
|
||||
Recursively scans all Parquet files under the given root directory.
|
||||
If any row in a file contains a black frame (all pixels below threshold),
|
||||
writes a new parquet file with "filtered_" prefix and deletes the original.
|
||||
|
||||
Features
|
||||
--------
|
||||
• ThreadPoolExecutor for parallel I/O
|
||||
• tqdm progress bar with per-file updates
|
||||
• --dry-run flag for a safe preview
|
||||
• --workers flag to control thread count
|
||||
• Writes filtered files in same location with "filtered_" prefix
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import os
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
from PIL import Image
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
|
||||
def process_file(path: Path,
|
||||
black_threshold: float = 5.0,
|
||||
dry_run: bool = False,
|
||||
images_output_dir: Path | None = None) -> int:
|
||||
"""Process a parquet file and write a filtered version with prefix. Returns number of rows removed."""
|
||||
|
||||
# Skip if already filtered
|
||||
if path.stem.startswith("filtered_"):
|
||||
# tqdm.write(f"[SKIP] Already filtered: {path}")
|
||||
# return 0
|
||||
truncate_prefix = True
|
||||
else:
|
||||
truncate_prefix = False
|
||||
|
||||
# Read the entire table
|
||||
table = pq.read_table(path)
|
||||
total_rows = len(table)
|
||||
|
||||
# Track which rows to keep
|
||||
rows_to_keep = []
|
||||
rows_removed = 0
|
||||
|
||||
# Check each row
|
||||
for row_idx in range(total_rows):
|
||||
row = table.slice(row_idx, 1).to_pylist()[0]
|
||||
|
||||
# Skip if any field is None
|
||||
if row["pil_image_bytes"] is None or row[
|
||||
"pil_image_shape"] is None or row["pil_image_dtype"] is None:
|
||||
tqdm.write(
|
||||
f"[WARN] Row {row_idx} in {path} has None values, keeping it")
|
||||
rows_to_keep.append(row_idx)
|
||||
continue
|
||||
|
||||
# Convert bytes to numpy array
|
||||
image_bytes = row["pil_image_bytes"]
|
||||
shape = row["pil_image_shape"]
|
||||
dtype = row["pil_image_dtype"]
|
||||
|
||||
# Convert bytes to numpy array with proper shape and dtype
|
||||
image_array = np.frombuffer(image_bytes,
|
||||
dtype=np.float32).reshape(shape)
|
||||
image_array = image_array.squeeze(
|
||||
0) # Remove single-dimensional entries if any
|
||||
|
||||
# Convert to uint8 for checking black frames
|
||||
if image_array.dtype != np.uint8:
|
||||
# Normalize to 0-255 range
|
||||
img_min = image_array.min()
|
||||
img_max = image_array.max()
|
||||
if img_max > img_min:
|
||||
image_uint8 = ((image_array - img_min) / (img_max - img_min) *
|
||||
255).astype(np.uint8)
|
||||
else:
|
||||
image_uint8 = np.zeros_like(image_array, dtype=np.uint8)
|
||||
else:
|
||||
image_uint8 = image_array
|
||||
|
||||
mean_value = np.mean(image_uint8)
|
||||
|
||||
# Check if the frame is black
|
||||
if mean_value < black_threshold:
|
||||
tqdm.write(
|
||||
f"[INFO] Found black frame in {path} row {row_idx} (mean={mean_value:.2f})"
|
||||
)
|
||||
rows_removed += 1
|
||||
|
||||
# Save black frame for inspection with unique ID
|
||||
if images_output_dir:
|
||||
# Create unique filename using file path, row index, and UUID
|
||||
# Hash the full path to handle duplicate filenames from different directories
|
||||
path_hash = hashlib.md5(str(
|
||||
path.absolute()).encode()).hexdigest()[:8]
|
||||
unique_id = f"{path.stem}_{path_hash}_row{row_idx}_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
image_uint8 = np.transpose(image_uint8, (1, 2, 0))
|
||||
img = Image.fromarray(image_uint8)
|
||||
output_path = images_output_dir / f"{unique_id}.png"
|
||||
img.save(output_path)
|
||||
tqdm.write(f"[INFO] Saved black frame to {output_path}")
|
||||
else:
|
||||
# Keep this row
|
||||
rows_to_keep.append(row_idx)
|
||||
|
||||
# Process based on what we found
|
||||
if rows_removed > 0:
|
||||
if dry_run:
|
||||
tqdm.write(
|
||||
f"[DRY-RUN] Would remove {rows_removed} rows from {path}")
|
||||
tqdm.write(
|
||||
f"[DRY-RUN] Would create filtered_{path.name} and delete original"
|
||||
)
|
||||
else:
|
||||
# Create new table with only the rows to keep
|
||||
if rows_to_keep:
|
||||
# Filter the table to keep only non-black frames
|
||||
new_table = table.take(rows_to_keep)
|
||||
|
||||
if truncate_prefix:
|
||||
name = path.name.replace("filtered_", "")
|
||||
print(f"[INFO] Truncating prefix for {path.name} to {name}")
|
||||
else:
|
||||
name = path.name
|
||||
|
||||
# Create output path with "filtered_" prefix in same directory
|
||||
output_path = path.parent / f"filtered_{name}"
|
||||
|
||||
# Write the filtered table
|
||||
pq.write_table(new_table, output_path)
|
||||
tqdm.write(
|
||||
f"[INFO] Wrote filtered parquet ({rows_removed} rows removed) to {output_path}"
|
||||
)
|
||||
|
||||
# Delete original file
|
||||
path.unlink()
|
||||
|
||||
tqdm.write(f"[INFO] Deleted original file: {path}")
|
||||
else:
|
||||
# All rows were black, just delete the file
|
||||
path.unlink()
|
||||
tqdm.write(
|
||||
f"[INFO] Deleted {path} (all {rows_removed} rows were black frames)"
|
||||
)
|
||||
else:
|
||||
# No black frames found
|
||||
if dry_run:
|
||||
tqdm.write(f"[DRY-RUN] No black frames in {path}")
|
||||
else:
|
||||
# Create output path with "filtered_" prefix
|
||||
output_path = path.parent / f"filtered_{path.name}"
|
||||
|
||||
# Just copy the table as-is
|
||||
pq.write_table(table, output_path)
|
||||
tqdm.write(
|
||||
f"[INFO] No black frames in {path}, created {output_path}")
|
||||
|
||||
# Delete original file
|
||||
path.unlink()
|
||||
tqdm.write(f"[INFO] Deleted original file: {path}")
|
||||
|
||||
return rows_removed
|
||||
|
||||
|
||||
def handle_file(path: Path, dry_run: bool, black_threshold: float,
|
||||
images_output_dir: Path | None) -> None:
|
||||
"""Process a single parquet file."""
|
||||
try:
|
||||
process_file(path,
|
||||
black_threshold=black_threshold,
|
||||
dry_run=dry_run,
|
||||
images_output_dir=images_output_dir)
|
||||
except Exception as exc:
|
||||
tqdm.write(f"[ERROR] Failed to process {path}: {exc}")
|
||||
|
||||
|
||||
def main(root: Path, dry_run: bool, workers: int,
|
||||
black_threshold: float) -> None:
|
||||
# Create output directory for black frame images
|
||||
images_output_dir = Path.cwd() / f"filtered_{int(black_threshold)}"
|
||||
images_output_dir.mkdir(exist_ok=True)
|
||||
tqdm.write(f"[INFO] Saving black frames to {images_output_dir}")
|
||||
|
||||
parquet_files = list(root.rglob("*.parquet"))
|
||||
if not parquet_files:
|
||||
print(f"[INFO] No Parquet files found in {root}")
|
||||
return
|
||||
|
||||
workers = max(1, workers)
|
||||
with tqdm(total=len(parquet_files), desc="Scanning", unit="file") as bar:
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
futures = [
|
||||
pool.submit(handle_file, fp, dry_run, black_threshold,
|
||||
images_output_dir) for fp in parquet_files
|
||||
]
|
||||
for _ in as_completed(futures):
|
||||
bar.update()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description=
|
||||
"Filter black frames from Parquet files, write with 'filtered_' prefix, and delete originals."
|
||||
)
|
||||
parser.add_argument("--folder", type=Path, help="Root directory to scan")
|
||||
parser.add_argument("--dry-run",
|
||||
action="store_true",
|
||||
help="Preview changes only")
|
||||
parser.add_argument(
|
||||
"--workers",
|
||||
type=int,
|
||||
default=os.cpu_count() or 32,
|
||||
help="Number of worker threads (default: CPU count)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--threshold",
|
||||
type=float,
|
||||
default=5.0,
|
||||
help="Black frame threshold (default: 5.0)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args.folder, args.dry_run, args.workers, args.threshold)
|
||||
@@ -0,0 +1,98 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=v-i-1
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=vsa-i2v/1.3B-1e5.out
|
||||
#SBATCH --error=vsa-i2v/1.3B-1e5.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY='73190d8c0de18a14eb3444e222f9432d247d1e30'
|
||||
# will key
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/train/
|
||||
DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents_i2v/test_filter/
|
||||
VALIDATION_DIR=/mnt/weka/home/hao.zhang/wl/FastVideo/data/mixkit/validation.json
|
||||
NUM_GPUS=8
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# OUTPUT_PATH="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn_77x768x1280/VSA_I2V_1.3B_1e5_bs64"
|
||||
OUTPUT_PATH="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn_77x768x1280/VSA_I2V_1.3B_1e5_bs32"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path $MODEL_PATH \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 1 \
|
||||
--tp_size 1 \
|
||||
--num_gpus 8 \
|
||||
--hsdp_replicate_dim 8 \
|
||||
--hsdp-shard-dim 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--max_train_steps 4500 \
|
||||
--learning_rate 2e-5 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 1000 \
|
||||
--validation_steps 300 \
|
||||
--validation_sampling_steps "50" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.1 \
|
||||
--seed 1024 \
|
||||
--output_dir $OUTPUT_PATH \
|
||||
--tracker_project_name VSA_finetune \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 61 \
|
||||
--flow_shift 3 \
|
||||
--validation_guidance_scale "6.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--weight_decay 1e-4 \
|
||||
--max_grad_norm 1.0 \
|
||||
--VSA_decay_rate 0.03 \
|
||||
--VSA_decay_interval_steps 50 \
|
||||
--VSA_sparsity 0.9
|
||||
Reference in New Issue
Block a user