Compare commits

...
1 Commits
Author SHA1 Message Date
SolitaryThinker 74e58b7f64 debug dark video 2025-07-12 01:50:35 +00:00
5 changed files with 353 additions and 3 deletions
+4 -1
View File
@@ -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
}
+17
View File
@@ -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
+2 -2
View File
@@ -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(
+232
View File
@@ -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)
+98
View File
@@ -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