Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
63a02d55f4 | ||
|
|
3f7d495fbd | ||
|
|
808bb8e697 | ||
|
|
661fa51d9e | ||
|
|
59843ee2d1 | ||
|
|
cb80364aaa | ||
|
|
92cc893cfc | ||
|
|
6ef5b6b28f | ||
|
|
948e9d7610 | ||
|
|
ec5a7c7e73 | ||
|
|
c31dd6853e | ||
|
|
ac919b1c87 | ||
|
|
d4eab03809 | ||
|
|
c1332cb4fb | ||
|
|
8bc4e538d0 | ||
|
|
495dd27d33 | ||
|
|
4d0571eabf | ||
|
|
7f0c06e7f2 | ||
|
|
aea5cadb5f | ||
|
|
c5444621f9 | ||
|
|
c6d0233eab | ||
|
|
de8e9e20b4 | ||
|
|
4fe7638101 | ||
|
|
6c23b4552a | ||
|
|
7c6f8f3f23 | ||
|
|
21f25791e2 | ||
|
|
d203b5f557 | ||
|
|
b588be05d1 | ||
|
|
42d4f524a0 | ||
|
|
7349e07cec | ||
|
|
8af0c43a65 | ||
|
|
34a0a81be1 |
@@ -135,3 +135,6 @@ fastvideo/tests/ssim/reference_videos/**
|
||||
*.nvimlog
|
||||
.nvimlog
|
||||
.python-version
|
||||
|
||||
# Gradio runtime/flagging artifacts
|
||||
.gradio/
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
#!/usr/bin/env bash
|
||||
# FastVideo (trackwan_bidir) env — CORRECTED build. CUDA 12.9 container, aarch64 + GB200 (sm_100a).
|
||||
# Ordering fix: install repo-pinned torch (2.11.0 cu128) FIRST, build the local Blackwell kernel
|
||||
# against it, then core deps (relaxed kernel pin so the local 0.3.0 satisfies), then FA4.
|
||||
set +e
|
||||
say(){ echo; echo "==================== $* ===================="; }
|
||||
|
||||
export DEBIAN_FRONTEND=noninteractive
|
||||
export CUDA_HOME=${CUDA_HOME:-/usr/local/cuda}
|
||||
export PATH=$CUDA_HOME/bin:/mnt/lustre/vlm-s4duan/bin:$PATH
|
||||
export LD_LIBRARY_PATH=$CUDA_HOME/lib64:$LD_LIBRARY_PATH
|
||||
export UV_CACHE_DIR=/mnt/lustre/vlm-s4duan/.uv/cache
|
||||
export UV_PYTHON_INSTALL_DIR=/mnt/lustre/vlm-s4duan/.uv/python
|
||||
export TORCH_CUDA_ARCH_LIST=10.0a
|
||||
export MAX_JOBS=${MAX_JOBS:-64}
|
||||
cd /mnt/lustre/vlm-s4duan/FastVideo || { echo NO_REPO; exit 1; }
|
||||
|
||||
say "ENSURE toolchain"
|
||||
if ! command -v g++ >/dev/null || ! command -v cmake >/dev/null || ! command -v ninja >/dev/null; then
|
||||
apt-get update -y >/tmp/apt.log 2>&1 && apt-get install -y --no-install-recommends g++ gcc make cmake ninja-build git ca-certificates >>/tmp/apt.log 2>&1 && echo "apt ok" || { echo APT_FAIL; tail -20 /tmp/apt.log; }
|
||||
fi
|
||||
nvcc --version | tail -2
|
||||
|
||||
say "FRESH venv (python 3.12 on Lustre)"
|
||||
rm -rf .venv
|
||||
uv venv --python 3.12 --seed .venv || exit 1
|
||||
source .venv/bin/activate
|
||||
python --version
|
||||
|
||||
say "STAGE 1: torch 2.11.0 cu128 (repo pin, Blackwell-capable)"
|
||||
uv pip install --torch-backend=cu128 torch==2.11.0 torchvision torchaudio 2>&1 | tail -15
|
||||
python -c "import torch;print('TORCH',torch.__version__,torch.version.cuda,'avail',torch.cuda.is_available())" || echo TORCH_FAIL
|
||||
|
||||
say "STAGE 2: submodules + build local fastvideo-kernel (sm_100a) against torch 2.11"
|
||||
git submodule update --init --recursive fastvideo-kernel/include/cutlass fastvideo-kernel/include/tk 2>&1 | tail -5
|
||||
( cd fastvideo-kernel && TORCH_CUDA_ARCH_LIST=10.0a bash ./build.sh ) 2>&1 | tail -20
|
||||
echo "KERNEL_EXIT=$?"
|
||||
python -c "import fastvideo_kernel;print('kernel import OK')" || echo KERNEL_IMPORT_FAIL
|
||||
|
||||
say "STAGE 3: core deps uv pip install -e .[dev] (local kernel already satisfies relaxed pin)"
|
||||
uv pip install -e ".[dev]" 2>&1 | tail -40
|
||||
echo "CORE_EXIT=$?"
|
||||
python -c "import fastvideo;print('fastvideo import OK')" || echo FASTVIDEO_IMPORT_FAIL
|
||||
|
||||
say "STAGE 4: FA4 — OFFICIAL Dao-AILab flash_attn/cute (NOT the XOR-op fork; fork is stale/broken)"
|
||||
# Package name is flash-attn-4; pass the bare git URL (name-qualified spec is rejected).
|
||||
# --prerelease allow lets it pull the matching nvidia-cutlass-dsl==4.6.0.dev0 + quack>=0.5.3.
|
||||
uv pip install --prerelease allow \
|
||||
"git+https://github.com/Dao-AILab/flash-attention.git#subdirectory=flash_attn/cute" 2>&1 | tail -25
|
||||
python -c "from flash_attn.cute.interface import _flash_attn_fwd,_flash_attn_bwd; print('FA4_CUTE_OK')" 2>&1
|
||||
|
||||
say "VERIFY"
|
||||
python - <<'PY'
|
||||
import importlib, torch
|
||||
print("torch", torch.__version__, torch.version.cuda, "cuda_avail", torch.cuda.is_available())
|
||||
if torch.cuda.is_available():
|
||||
print("device0", torch.cuda.get_device_name(0), "cc", torch.cuda.get_device_capability(0))
|
||||
x = torch.randn(2048, 2048, device="cuda", dtype=torch.bfloat16)
|
||||
y = (x @ x); torch.cuda.synchronize()
|
||||
print("bf16 matmul on GPU OK", tuple(y.shape))
|
||||
for m in ["fastvideo", "fastvideo_kernel", "flash_attn"]:
|
||||
try: importlib.import_module(m); print("import OK:", m)
|
||||
except Exception as e: print("import FAIL:", m, repr(e)[:180])
|
||||
try:
|
||||
from flash_attn.cute.interface import _flash_attn_fwd; print("FA4 flash_attn.cute.interface OK")
|
||||
except Exception as e:
|
||||
print("FA4 cute interface:", repr(e)[:180])
|
||||
PY
|
||||
|
||||
say "DONE build_env.sh (v2)"
|
||||
@@ -0,0 +1,38 @@
|
||||
#!/usr/bin/env bash
|
||||
set +e
|
||||
export HF_HOME=/mnt/lustre/vlm-s4duan/.hf
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
cd /mnt/lustre/vlm-s4duan/FastVideo || exit 1
|
||||
source .venv/bin/activate
|
||||
|
||||
BASE=/mnt/lustre/vlm-s4duan/models/Wan2.1-Fun-1.3B-InP-Diffusers
|
||||
OUT=/mnt/lustre/vlm-s4duan/models/trackwan_1.3b_init
|
||||
mkdir -p /mnt/lustre/vlm-s4duan/models
|
||||
|
||||
echo "==================== download Fun-InP diffusers base ===================="
|
||||
python - <<PY
|
||||
from huggingface_hub import snapshot_download
|
||||
p = snapshot_download("weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers", local_dir="$BASE")
|
||||
print("downloaded to", p)
|
||||
PY
|
||||
echo "DL_EXIT=$?"
|
||||
|
||||
echo "==================== base transformer sanity (I2V => in_channels 36) ===================="
|
||||
python - <<PY
|
||||
import json
|
||||
c = json.load(open("$BASE/transformer/config.json"))
|
||||
print("in_channels:", c.get("in_channels"), "image_dim:", c.get("image_dim"), "num_layers:", c.get("num_layers"), "hidden:", c.get("hidden_size") or c.get("dim"))
|
||||
PY
|
||||
|
||||
echo "==================== convert_trackwan_init.py ===================="
|
||||
python data_pipeline/convert_trackwan_init.py --base "$BASE" --out "$OUT"
|
||||
echo "CONVERT_EXIT=$?"
|
||||
|
||||
echo "==================== trackwan init result ===================="
|
||||
ls -la "$OUT" 2>&1 | head -25
|
||||
python - <<PY
|
||||
import json
|
||||
c = json.load(open("$OUT/transformer/config.json"))
|
||||
print("trackwan in_channels:", c.get("in_channels"), "| has track_config:", "track_config" in c)
|
||||
PY
|
||||
echo "==================== DONE build_trackwan_init ===================="
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
#!/bin/bash
|
||||
# Phase 0, step 0 — carve a SMALL overfit dataset out of the 720p openvid parquets.
|
||||
# Symlinks a few data_chunk parquets into a dedicated dir; the map-style loader walks it for
|
||||
# *.parquet, so a couple of chunks (~32 clips each) is a good overfit set. Non-destructive
|
||||
# (symlinks only). Re-run to rebuild; it clears the stale map_style_cache.
|
||||
set -uo pipefail
|
||||
|
||||
SRC_ROOT=${SRC_ROOT:-/home/hal-shared/motionstream/data/openvid-wantrack-parquets}
|
||||
SRC_SHARD=${SRC_SHARD:-shard000}
|
||||
N_CHUNKS=${N_CHUNKS:-2} # ~32 clips/chunk -> ~64 clips
|
||||
OUT=${OUT:-/home/hal-kevin/data/motion-stream-test/overfit_subset_720p/combined_parquet_dataset}
|
||||
|
||||
src_worker="$SRC_ROOT/$SRC_SHARD/combined_parquet_dataset/worker_0"
|
||||
[ -d "$src_worker" ] || { echo "[subset] source not found: $src_worker" >&2; exit 1; }
|
||||
|
||||
dst_worker="$OUT/worker_0"
|
||||
rm -rf "$OUT" # drop old subset + its map_style_cache
|
||||
mkdir -p "$dst_worker"
|
||||
|
||||
n=0
|
||||
for f in $(ls "$src_worker"/data_chunk_*.parquet | sort -V | head -n "$N_CHUNKS"); do
|
||||
ln -s "$(readlink -f "$f")" "$dst_worker/$(basename "$f")"
|
||||
n=$((n + 1))
|
||||
done
|
||||
|
||||
echo "[subset] linked $n parquet chunk(s) from $SRC_SHARD -> $OUT"
|
||||
python - "$OUT" <<'PY'
|
||||
import glob, sys, pyarrow.parquet as pq
|
||||
fs = glob.glob(f"{sys.argv[1]}/**/*.parquet", recursive=True)
|
||||
rows = sum(pq.ParquetFile(f).metadata.num_rows for f in fs)
|
||||
print(f"[subset] {len(fs)} file(s), {rows} clips total")
|
||||
PY
|
||||
echo "[subset] point the overfit config data_path at: $OUT"
|
||||
Executable
+30
@@ -0,0 +1,30 @@
|
||||
#!/bin/bash
|
||||
# Phase 0, step 1 — build the d64 + bias 14B WanTrack init for the overfit.
|
||||
# CPU only, needs ~62GB RAM (loads the 14B base) -> run on a COMPUTE node, not the login node.
|
||||
# pretrained channels are preserved; the added track slot is ZERO-init (--pe-init zero, matching
|
||||
# upstream's trackwan_14b_i2v_d64_zero_init_bias) and the track_encoder gets default init WITH bias
|
||||
# (--use-bias-defaults), matching TRACKWAN_TRACK_BIAS=1 at train time.
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
BASE=${BASE:-/home/hal-kevin/models/Wan2.1-I2V-14B-720P-Diffusers}
|
||||
OUT=${OUT:-/home/hal-kevin/models/trackwan_14b_i2v_d64_bias_init}
|
||||
ID_DIM=${ID_DIM:-64}
|
||||
PE_INIT=${PE_INIT:-zero}
|
||||
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base "$BASE" \
|
||||
--out "$OUT" \
|
||||
--id-dim "$ID_DIM" \
|
||||
--pe-init "$PE_INIT" \
|
||||
--use-bias-defaults
|
||||
|
||||
echo "[init] built $OUT (id_dim=$ID_DIM, pe-init=$PE_INIT, bias=on)"
|
||||
echo "[init] expect: 'added 4 track_encoder tensors' (2 weights + 2 bias)"
|
||||
python -c "
|
||||
from safetensors import safe_open
|
||||
ks=[k for k in safe_open('$OUT/transformer/diffusion_pytorch_model.safetensors','pt').keys() if 'track_encoder' in k]
|
||||
import json; c=json.load(open('$OUT/transformer/config.json'))
|
||||
print('[init] in_channels', c['in_channels'], '| track_config.id_dim', c['track_config']['id_dim'])
|
||||
print('[init] track_encoder keys:', sorted(ks))
|
||||
"
|
||||
Executable
+56
@@ -0,0 +1,56 @@
|
||||
#!/bin/bash
|
||||
# Phase 0, step 2 — overfit the track pathway. 14B/720p, d64+bias, sparse, HEAD TRAINABLE.
|
||||
# Run the SAME command on BOTH racks (2x4 GB200). MASTER_ADDR = rack0 host on both; NODE_RANK
|
||||
# differs (0 on rack0, 1 on rack1). STAGE=A uses fixed track IDs (lock onto the pattern), STAGE=B
|
||||
# uses random IDs; run A first, then B (B resumes A's checkpoint via resume_from_checkpoint: latest).
|
||||
#
|
||||
# rack0: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 STAGE=A bash data_pipeline/720_stage_1/02_run_overfit.sh
|
||||
# rack1: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 STAGE=A bash data_pipeline/720_stage_1/02_run_overfit.sh
|
||||
# Extra args pass through to the trainer, e.g. ... STAGE=B bash 02_run_overfit.sh --training.loop.max_train_steps 1600
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
: "${MASTER_ADDR:?set MASTER_ADDR to rack-0 hostname (reachable from both racks)}"
|
||||
: "${NODE_RANK:?set NODE_RANK: 0 on the master rack, 1 on the other}"
|
||||
STAGE=${STAGE:-A}
|
||||
MASTER_PORT=${MASTER_PORT:-29502}
|
||||
CFG=data_pipeline/720_stage_1/finetune_wantrack_overfit_14b_720p_d64_bias.yaml
|
||||
|
||||
# Stage A = deterministic track IDs (fixed sampling); Stage B = random IDs.
|
||||
# B writes to its OWN output dir (keeps A's dir pristine, A/B checkpoints separated) — mirrors
|
||||
# upstream chain_stepAB.sh + run_stepB_seed.sh. Seed B's dir once with A's final checkpoint via
|
||||
# 02b_seed_stageB.sh BEFORE launching B, so resume_from_checkpoint: latest picks up A's weights.
|
||||
OUT_A=/home/hal-kevin/data/motion-stream-test/overfit_14b_720p_d64_bias_out
|
||||
OUT_B=/home/hal-kevin/data/motion-stream-test/overfit_14b_720p_d64_bias_stageB_out
|
||||
case "$STAGE" in
|
||||
A) FIXED=1; OUT_DIR=$OUT_A ;;
|
||||
B) FIXED=0; OUT_DIR=$OUT_B ;;
|
||||
*) echo "STAGE must be A or B (got '$STAGE')" >&2; exit 1 ;;
|
||||
esac
|
||||
if [ "$STAGE" = B ] && [ ! -d "$OUT_B" ]; then
|
||||
echo "[overfit] ERROR: STAGE=B but $OUT_B does not exist." >&2
|
||||
echo " Seed it first (run ONCE): bash data_pipeline/720_stage_1/02b_seed_stageB.sh" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Pin NCCL to the four active 400G InfiniBand HCAs (keep off the 200G Ethernet port).
|
||||
export NCCL_IB_HCA=mlx5_0,mlx5_1,mlx5_3,mlx5_4
|
||||
# export NCCL_DEBUG=INFO # uncomment on the FIRST launch to confirm NET/IB, then re-comment.
|
||||
|
||||
FASTVIDEO_FA4=1 \
|
||||
TRACKWAN_TRACK_BIAS=1 WANTRACK_FREEZE_HEAD=0 \
|
||||
WANTRACK_AUG=1 WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 \
|
||||
WANTRACK_PMASK=0 WANTRACK_MASK_CHUNK=0 \
|
||||
WANTRACK_IMAGE_COND=0 \
|
||||
WANTRACK_FIXED_SAMPLE=${FIXED} \
|
||||
torchrun \
|
||||
--nnodes=2 --nproc_per_node=4 --node_rank=${NODE_RANK} \
|
||||
--rdzv_id=overfit_14b_720p --rdzv_backend=c10d --rdzv_endpoint=${MASTER_ADDR}:${MASTER_PORT} \
|
||||
--log-dir ~/FastVideo/torchrun_logs \
|
||||
-m fastvideo.train.entrypoint.train \
|
||||
--config ${CFG} \
|
||||
--training.checkpoint.output_dir ${OUT_DIR} \
|
||||
--training.tracker.run_name overfit-14b-720p-d64-bias-stage${STAGE} \
|
||||
--training.checkpoint.resume_from_checkpoint latest \
|
||||
"$@" \
|
||||
2>&1 | tee -a data_pipeline/720_stage_1/overfit_node${NODE_RANK}_stage${STAGE}.log
|
||||
Executable
+41
@@ -0,0 +1,41 @@
|
||||
#!/bin/bash
|
||||
# Seed Stage B's output dir with Stage A's FINAL checkpoint, so B runs in its own dir while its
|
||||
# config still uses resume_from_checkpoint: latest (which is what makes crash-restarts safe).
|
||||
# Mirrors upstream examples/train/run_stepB_seed.sh.
|
||||
#
|
||||
# Run this ONCE (not per-rack) on the shared filesystem, AFTER Stage A finishes, BEFORE launching
|
||||
# Stage B. Hardlinked (cp -al), not copied: a 14B training-state checkpoint is 100s of GB, and both
|
||||
# dirs are on the same filesystem (/home), so hardlinks are instant and use no extra space.
|
||||
#
|
||||
# bash data_pipeline/720_stage_1/02b_seed_stageB.sh
|
||||
# # then, on BOTH racks:
|
||||
# MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 STAGE=B bash data_pipeline/720_stage_1/02_run_overfit.sh \
|
||||
# --training.loop.max_train_steps 3000
|
||||
set -uo pipefail
|
||||
|
||||
OUT_A=${OUT_A:-/home/hal-kevin/data/motion-stream-test/overfit_14b_720p_d64_bias_out}
|
||||
OUT_B=${OUT_B:-/home/hal-kevin/data/motion-stream-test/overfit_14b_720p_d64_bias_stageB_out}
|
||||
|
||||
# Pick A's latest COMPLETE checkpoint (dcp/.metadata present) — refuse to seed off a partial save.
|
||||
STEP=${STEP:-}
|
||||
if [ -z "$STEP" ]; then
|
||||
for d in $(for c in "$OUT_A"/checkpoint-*; do n=${c##*checkpoint-}; echo "$n"; done | sort -rn); do
|
||||
[ -f "$OUT_A/checkpoint-$d/dcp/.metadata" ] && { STEP=$d; break; }
|
||||
done
|
||||
fi
|
||||
[ -n "$STEP" ] || { echo "[B-seed] no complete checkpoint found under $OUT_A" >&2; exit 1; }
|
||||
|
||||
SRC="$OUT_A/checkpoint-$STEP"
|
||||
DST="$OUT_B/checkpoint-$STEP"
|
||||
[ -f "$SRC/dcp/.metadata" ] || { echo "[B-seed] $SRC incomplete (no dcp/.metadata) — refusing" >&2; exit 1; }
|
||||
|
||||
mkdir -p "$OUT_B"
|
||||
if [ -e "$DST" ]; then
|
||||
echo "[B-seed] $DST already exists — leaving it alone"
|
||||
else
|
||||
echo "[B-seed] hardlinking $SRC -> $DST"
|
||||
cp -al "$SRC" "$DST"
|
||||
fi
|
||||
echo "[B-seed] Stage B dir now seeded at step $STEP:"
|
||||
for c in "$OUT_B"/checkpoint-*; do echo " $c"; done
|
||||
echo "[B-seed] Launch B with --training.loop.max_train_steps > $STEP (e.g. $((STEP + 1000)))."
|
||||
Executable
+38
@@ -0,0 +1,38 @@
|
||||
#!/bin/bash
|
||||
# Phase 0, step 3 — export the overfit DCP checkpoint to a diffusers model dir.
|
||||
# The Phase-1 merge (convert_trackwan_init_v2.py --track-src/--pe-src) reads this diffusers dir.
|
||||
# Only 1 GPU needed (DCP reshards automatically). Run on a compute node.
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
# --checkpoint accepts an output_dir (auto-picks the latest checkpoint-<step>), a specific
|
||||
# checkpoint-<step> dir, or its dcp/ subdir.
|
||||
# Default = the Stage-B dir (the FINAL overfit; B refines A and is what the merge lifts). Override
|
||||
# CKPT to the Stage-A dir only if you deliberately want to export A.
|
||||
CKPT=${CKPT:-/home/hal-kevin/data/motion-stream-test/overfit_14b_720p_d64_bias_stageB_out}
|
||||
OUT=${OUT:-/home/hal-kevin/models/overfit14b_export}
|
||||
CFG=${CFG:-data_pipeline/720_stage_1/finetune_wantrack_overfit_14b_720p_d64_bias.yaml}
|
||||
|
||||
# TRACKWAN_TRACK_BIAS=1 MUST match training: it toggles bias on track_encoder.{proj,temporal_conv}
|
||||
# at build time (track_encoder.py:73). The overfit trained with bias=1, so the checkpoint carries
|
||||
# those bias params; building bias-less here fails to load them ("track_encoder.proj.bias not found").
|
||||
TRACKWAN_TRACK_BIAS=1 \
|
||||
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
--checkpoint "$CKPT" \
|
||||
--output-dir "$OUT" \
|
||||
--config "$CFG" \
|
||||
--role student \
|
||||
--overwrite \
|
||||
|| { echo "[export] FAILED (see traceback above) -- nothing written" >&2; exit 1; }
|
||||
|
||||
# Guard against a silent partial write (the entrypoint can exit 0 yet write no weights).
|
||||
ls "$OUT"/transformer/*.safetensors >/dev/null 2>&1 \
|
||||
|| { echo "[export] ERROR: no transformer/*.safetensors under $OUT -- export did not complete" >&2; exit 1; }
|
||||
|
||||
echo "[export] wrote diffusers model -> $OUT"
|
||||
echo "[export] next (Phase 1 merge):"
|
||||
echo " python data_pipeline/convert_trackwan_init_v2.py \\"
|
||||
echo " --base /home/hal-kevin/models/Wan2.1-I2V-14B-720P-Diffusers \\"
|
||||
echo " --out /home/hal-kevin/models/trackwan_14b_i2v_d64_merged_from_overfit_bias \\"
|
||||
echo " --id-dim 64 --pe-init random \\"
|
||||
echo " --track-src $OUT/transformer --pe-src $OUT/transformer"
|
||||
Executable
+29
@@ -0,0 +1,29 @@
|
||||
#!/bin/bash
|
||||
# Step 04 — MERGE. Graft the overfit's co-adapted track pathway (track_encoder + patch-embed track
|
||||
# slot [36:52]) onto a PRISTINE 14B base, discarding the overfit's base degradation. CPU only,
|
||||
# ~62GB RAM -> run on a compute node. Needs 03_export.sh to have produced the overfit diffusers dir.
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
BASE=${BASE:-/home/hal-kevin/models/Wan2.1-I2V-14B-720P-Diffusers}
|
||||
SRC=${SRC:-/home/hal-kevin/models/overfit14b_export} # produced by 03_export.sh
|
||||
OUT=${OUT:-/home/hal-kevin/models/trackwan_14b_i2v_d64_merged_from_overfit_bias}
|
||||
ID_DIM=${ID_DIM:-64}
|
||||
|
||||
[ -d "$SRC/transformer" ] || { echo "[merge] $SRC/transformer not found — run 03_export.sh first" >&2; exit 1; }
|
||||
|
||||
# --track-src + --pe-src from the SAME export = the merge: encoder AND its co-adapted track slot
|
||||
# lifted together (bias copied through); everything else comes from the pristine --base.
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base "$BASE" \
|
||||
--out "$OUT" \
|
||||
--id-dim "$ID_DIM" --pe-init random \
|
||||
--track-src "$SRC/transformer" \
|
||||
--pe-src "$SRC/transformer"
|
||||
|
||||
echo "[merge] built merged init -> $OUT"
|
||||
python -c "
|
||||
import json; c=json.load(open('$OUT/transformer/config.json'))
|
||||
print('[merge] in_channels', c['in_channels'], '| id_dim', c['track_config']['id_dim'])
|
||||
"
|
||||
echo "[merge] next: 05_run_openvid_stage1.sh (init_from defaults to this dir)"
|
||||
+41
@@ -0,0 +1,41 @@
|
||||
#!/bin/bash
|
||||
# Step 05 — OpenVid Stage 1 (the big run, ~10 days). 14B/720p, merged init, HEAD FROZEN, sparse.
|
||||
# Run the SAME command on BOTH racks. MASTER_ADDR = rack0 host on both; NODE_RANK 0 on rack0, 1 on rack1.
|
||||
# rack0: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 bash data_pipeline/720_stage_1/05_run_openvid_stage1.sh
|
||||
# rack1: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 bash data_pipeline/720_stage_1/05_run_openvid_stage1.sh
|
||||
# resume_from_checkpoint: latest -> relaunch both to continue after any interruption.
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
: "${MASTER_ADDR:?set MASTER_ADDR to rack-0 hostname (reachable from both racks)}"
|
||||
: "${NODE_RANK:?set NODE_RANK: 0 on the master rack, 1 on the other}"
|
||||
MASTER_PORT=${MASTER_PORT:-29503}
|
||||
NNODES=${NNODES:-8} # 8 nodes x 4 GB200 = 32 GPUs (matches num_gpus:32 / replicate 8 in the config)
|
||||
GPUS_PER_NODE=${GPUS_PER_NODE:-4}
|
||||
CFG=data_pipeline/720_stage_1/finetune_wantrack_openvid_stage1_14b_720p_d64_bias.yaml
|
||||
|
||||
export NCCL_IB_HCA=mlx5_0,mlx5_1,mlx5_3,mlx5_4
|
||||
# export NCCL_DEBUG=INFO # uncomment on the FIRST launch to confirm NET/IB, then re-comment.
|
||||
|
||||
# Stage 1: sparse conditioning, HEAD FROZEN (train the DiT to use the merged track pathway), no masking.
|
||||
# Every WANTRACK_/TRACKWAN_ knob is set EXPLICITLY (no reliance on code defaults) to match upstream
|
||||
# stage-1 (D) and to avoid inheriting the overfit launcher's opposite settings (IMAGE_COND=0, FIXED_SAMPLE=1).
|
||||
FASTVIDEO_FA4=1 \
|
||||
TRACKWAN_TRACK_BIAS=1 \
|
||||
WANTRACK_FREEZE_HEAD=1 \
|
||||
WANTRACK_IMAGE_COND=1 \
|
||||
WANTRACK_AUG=1 \
|
||||
WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 WANTRACK_EXTRA_MODE=random \
|
||||
WANTRACK_FIXED_SAMPLE=0 \
|
||||
WANTRACK_PMASK=0 WANTRACK_MASK_CHUNK=0 \
|
||||
WANTRACK_TRACK_DROP=0 WANTRACK_MOTION_DROP=0 WANTRACK_TEXT_DROP=0 \
|
||||
torchrun \
|
||||
--nnodes=${NNODES} --nproc_per_node=${GPUS_PER_NODE} --node_rank=${NODE_RANK} \
|
||||
--rdzv_id=openvid_stage1_14b_720p --rdzv_backend=c10d --rdzv_endpoint=${MASTER_ADDR}:${MASTER_PORT} \
|
||||
--log-dir ~/FastVideo/torchrun_logs \
|
||||
-m fastvideo.train.entrypoint.train \
|
||||
--config ${CFG} \
|
||||
--training.distributed.num_gpus $((NNODES * GPUS_PER_NODE)) \
|
||||
--training.checkpoint.resume_from_checkpoint latest \
|
||||
"$@" \
|
||||
2>&1 | tee -a data_pipeline/720_stage_1/openvid_stage1_node${NODE_RANK}.log
|
||||
@@ -0,0 +1,66 @@
|
||||
#!/bin/bash
|
||||
# Step 05 (2-RACK variant) — OpenVid Stage 1 on 2 racks (2x4 = 8 GB200) instead of 8 nodes.
|
||||
#
|
||||
# SAME effective global batch as the 8-node 05: grad_accum is bumped 4x (4 -> 16) to compensate for
|
||||
# 4x fewer GPUs, so the optimization is equivalent (see note below) — only the wall-clock is ~4x
|
||||
# (~48 days vs ~12). Run the SAME command on BOTH racks; MASTER_ADDR = rack0 host on both,
|
||||
# NODE_RANK 0 on rack0, 1 on rack1.
|
||||
# rack0: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 bash data_pipeline/720_stage_1/05b_run_openvid_stage1_2rack.sh
|
||||
# rack1: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 bash data_pipeline/720_stage_1/05b_run_openvid_stage1_2rack.sh
|
||||
# resume_from_checkpoint: latest -> relaunch both to continue after any interruption.
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
: "${MASTER_ADDR:?set MASTER_ADDR to rack-0 hostname (reachable from both racks)}"
|
||||
: "${NODE_RANK:?set NODE_RANK: 0 on the master rack, 1 on the other}"
|
||||
MASTER_PORT=${MASTER_PORT:-29503}
|
||||
NNODES=${NNODES:-2} # 2 racks x 4 GB200 = 8 GPUs
|
||||
GPUS_PER_NODE=${GPUS_PER_NODE:-4}
|
||||
CFG=data_pipeline/720_stage_1/finetune_wantrack_openvid_stage1_14b_720p_d64_bias.yaml
|
||||
# bf16 parquets (complete 259k set, ~4.7TB). Loader honors the per-field _dtype; same clips/order
|
||||
# as fp32 (so val_sample_indices are unchanged). Override to the fp32 set if ever needed:
|
||||
# DATA_PATH=/home/hal-shared/motionstream/data/openvid-wantrack-parquets
|
||||
DATA_PATH=${DATA_PATH:-/home/hal-shared/motionstream/data/openvid-wantrack-parquets-bf16}
|
||||
|
||||
# --- batch math: hold the effective global batch EQUAL to the 8-node 05 ------------------
|
||||
# 8-node config: replicate 8 x shard 4 = 32 GPUs, grad_accum 4.
|
||||
# 2 racks: replicate 2 x shard 4 = 8 GPUs, grad_accum 16.
|
||||
# 4x fewer GPUs x 4x grad_accum = same effective batch under EITHER counting convention
|
||||
# (num_gpus- or replicate_dim-based both scale by 4). Model still shards across 4 GPUs
|
||||
# (shard_dim 4) as in the 8-node run, so per-GPU memory is unchanged (no OOM risk from this).
|
||||
HSDP_REPLICATE=${HSDP_REPLICATE:-2}
|
||||
HSDP_SHARD=${HSDP_SHARD:-4}
|
||||
GRAD_ACCUM=${GRAD_ACCUM:-16}
|
||||
|
||||
export NCCL_IB_HCA=mlx5_0,mlx5_1,mlx5_3,mlx5_4
|
||||
# export NCCL_DEBUG=INFO # uncomment on the FIRST launch to confirm NET/IB, then re-comment.
|
||||
|
||||
# Stage 1 env — identical to 05 (every knob explicit): sparse, HEAD FROZEN, CLIP on, no masking.
|
||||
FASTVIDEO_FA4=1 \
|
||||
TRACKWAN_TRACK_BIAS=1 \
|
||||
WANTRACK_FREEZE_HEAD=1 \
|
||||
WANTRACK_IMAGE_COND=1 \
|
||||
WANTRACK_AUG=1 \
|
||||
WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 WANTRACK_EXTRA_MODE=random \
|
||||
WANTRACK_FIXED_SAMPLE=0 \
|
||||
WANTRACK_PMASK=0 WANTRACK_MASK_CHUNK=0 \
|
||||
WANTRACK_TRACK_DROP=0 WANTRACK_MOTION_DROP=0 WANTRACK_TEXT_DROP=0 \
|
||||
torchrun \
|
||||
--nnodes=${NNODES} --nproc_per_node=${GPUS_PER_NODE} --node_rank=${NODE_RANK} \
|
||||
--rdzv_id=openvid_stage1_14b_720p_2rack --rdzv_backend=c10d --rdzv_endpoint=${MASTER_ADDR}:${MASTER_PORT} \
|
||||
--log-dir ~/FastVideo/torchrun_logs \
|
||||
-m fastvideo.train.entrypoint.train \
|
||||
--config ${CFG} \
|
||||
--training.data.data_path ${DATA_PATH} \
|
||||
--training.distributed.num_gpus $((NNODES * GPUS_PER_NODE)) \
|
||||
--training.distributed.hsdp_replicate_dim ${HSDP_REPLICATE} \
|
||||
--training.distributed.hsdp_shard_dim ${HSDP_SHARD} \
|
||||
--training.loop.gradient_accumulation_steps ${GRAD_ACCUM} \
|
||||
--training.checkpoint.resume_from_checkpoint latest \
|
||||
--training.checkpoint.training_state_checkpointing_steps 20 \
|
||||
--training.checkpoint.checkpoints_total_limit 50 \
|
||||
--callbacks.track_validation.validate_at_start false \
|
||||
--callbacks.track_validation.every_steps 100 \
|
||||
--callbacks.track_validation.val_sample_indices "[1660, 1888]" \
|
||||
"$@" \
|
||||
2>&1 | tee -a data_pipeline/720_stage_1/openvid_stage1_2rack_node${NODE_RANK}.log
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
#!/bin/bash
|
||||
# Step 06 — export the OpenVid Stage-1 DCP checkpoint to a diffusers dir. Stage 2 (07) inits its
|
||||
# WEIGHTS from this export (init_from, fresh optimizer/step), so it must be a diffusers model dir,
|
||||
# not a DCP resume. Only 1 GPU needed. Run on a compute node.
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
CKPT=${CKPT:-/home/hal-kevin/data/motion-stream-test/openvid_stage1_14b_720p_out} # auto-picks latest
|
||||
OUT=${OUT:-/home/hal-kevin/models/openvid_stage1_14b_export}
|
||||
CFG=${CFG:-data_pipeline/720_stage_1/finetune_wantrack_openvid_stage1_14b_720p_d64_bias.yaml}
|
||||
|
||||
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
--checkpoint "$CKPT" \
|
||||
--output-dir "$OUT" \
|
||||
--config "$CFG" \
|
||||
--role student \
|
||||
--overwrite
|
||||
|
||||
echo "[export] stage-1 teacher -> $OUT"
|
||||
echo "[export] next: 07_run_synth_stage2.sh (init_from defaults to this dir)"
|
||||
+41
@@ -0,0 +1,41 @@
|
||||
#!/bin/bash
|
||||
# Step 07 — Stage 2 (robustness, ~1 day). 14B/720p, inits from exported Stage-1 teacher, HEAD FROZEN,
|
||||
# sparse + heavy masking/dropout. Produces the FINAL bidir teacher. Run on BOTH racks (same command,
|
||||
# NODE_RANK 0/1, MASTER_ADDR = rack0).
|
||||
# rack0: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 bash data_pipeline/720_stage_1/07_run_synth_stage2.sh
|
||||
# rack1: MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 bash data_pipeline/720_stage_1/07_run_synth_stage2.sh
|
||||
set -uo pipefail
|
||||
cd ~/FastVideo
|
||||
|
||||
: "${MASTER_ADDR:?set MASTER_ADDR to rack-0 hostname (reachable from both racks)}"
|
||||
: "${NODE_RANK:?set NODE_RANK: 0 on the master rack, 1 on the other}"
|
||||
MASTER_PORT=${MASTER_PORT:-29504}
|
||||
NNODES=${NNODES:-8} # 8 nodes x 4 GB200 = 32 GPUs (matches num_gpus:32 / replicate 8 in the config)
|
||||
GPUS_PER_NODE=${GPUS_PER_NODE:-4}
|
||||
CFG=data_pipeline/720_stage_1/finetune_wantrack_synth_stage2_14b_720p_d64_bias.yaml
|
||||
|
||||
export NCCL_IB_HCA=mlx5_0,mlx5_1,mlx5_3,mlx5_4
|
||||
# export NCCL_DEBUG=INFO # first launch only.
|
||||
|
||||
# Stage 2: sparse + robustness masking/dropout, HEAD still FROZEN.
|
||||
# Every WANTRACK_/TRACKWAN_ knob is set EXPLICITLY (no code defaults) to match upstream stage-2 (E):
|
||||
# same as stage-1 but with masking (PMASK=0.2/CHUNK=8) and dropout (TRACK=0.5, MOTION=0.3) ON.
|
||||
FASTVIDEO_FA4=1 \
|
||||
TRACKWAN_TRACK_BIAS=1 \
|
||||
WANTRACK_FREEZE_HEAD=1 \
|
||||
WANTRACK_IMAGE_COND=1 \
|
||||
WANTRACK_AUG=1 \
|
||||
WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 WANTRACK_EXTRA_MODE=random \
|
||||
WANTRACK_FIXED_SAMPLE=0 \
|
||||
WANTRACK_PMASK=0.2 WANTRACK_MASK_CHUNK=8 \
|
||||
WANTRACK_TRACK_DROP=0.5 WANTRACK_MOTION_DROP=0.3 WANTRACK_TEXT_DROP=0 \
|
||||
torchrun \
|
||||
--nnodes=${NNODES} --nproc_per_node=${GPUS_PER_NODE} --node_rank=${NODE_RANK} \
|
||||
--rdzv_id=synth_stage2_14b_720p --rdzv_backend=c10d --rdzv_endpoint=${MASTER_ADDR}:${MASTER_PORT} \
|
||||
--log-dir ~/FastVideo/torchrun_logs \
|
||||
-m fastvideo.train.entrypoint.train \
|
||||
--config ${CFG} \
|
||||
--training.distributed.num_gpus $((NNODES * GPUS_PER_NODE)) \
|
||||
--training.checkpoint.resume_from_checkpoint latest \
|
||||
"$@" \
|
||||
2>&1 | tee -a data_pipeline/720_stage_1/synth_stage2_node${NODE_RANK}.log
|
||||
@@ -0,0 +1,83 @@
|
||||
# Stage 1 — Bidirectional 14B/720p track teacher (full pipeline)
|
||||
|
||||
The complete bidirectional teacher: overfit the track pathway → merge it into a pristine 14B
|
||||
base → OpenVid stage-1 (frozen head) → synth stage-2 (robustness). The final teacher then seeds
|
||||
the causal (Self-Forcing) student. Full rationale: `../notes/trackwan_14b_720p_teacher_merge_plan.md`.
|
||||
|
||||
Recipe: **d64 + bias, sparse conditioning, flow_shift 6, 720p.** Hardware: 2×4 GB200, 400G IB,
|
||||
manual 2-node `torchrun` (rack0 = `hpc-rack-1-6` = NODE_RANK 0; rack1 = `hpc-rack-1-8` = NODE_RANK 1;
|
||||
`MASTER_ADDR` = rack0 on both).
|
||||
|
||||
## Files / run order
|
||||
|
||||
| Step | Script | What | Where |
|
||||
|---|---|---|---|
|
||||
| 0 | `00_make_overfit_subset.sh` | carve a small overfit set (symlinks ~2 parquet chunks) | login/compute |
|
||||
| 1 | `01_build_init.sh` | build `trackwan_14b_i2v_d64_bias_init` | compute (CPU, ~62GB RAM) |
|
||||
| 2 | `02_run_overfit.sh` | overfit track pathway, head trainable (Stage A→B) | **both racks** |
|
||||
| 3 | `03_export.sh` | overfit DCP → diffusers | compute (1 GPU) |
|
||||
| 4 | `04_merge.sh` | graft pathway into pristine 14B base → merged init | compute (CPU, ~62GB) |
|
||||
| 5 | `05_run_openvid_stage1.sh` | OpenVid stage-1, **head frozen** (~10 days) | **both racks** |
|
||||
| 6 | `06_export_stage1.sh` | stage-1 DCP → diffusers | compute (1 GPU) |
|
||||
| 7 | `07_run_synth_stage2.sh` | stage-2 robustness, head frozen (~1 day) → **final teacher** | **both racks** |
|
||||
|
||||
Configs (referenced by the scripts): `finetune_wantrack_overfit_14b_720p_d64_bias.yaml`,
|
||||
`finetune_wantrack_openvid_stage1_14b_720p_d64_bias.yaml`,
|
||||
`finetune_wantrack_synth_stage2_14b_720p_d64_bias.yaml`.
|
||||
|
||||
## Commands
|
||||
|
||||
```bash
|
||||
# 0-1: data + init
|
||||
bash data_pipeline/720_stage_1/00_make_overfit_subset.sh
|
||||
bash data_pipeline/720_stage_1/01_build_init.sh # compute node
|
||||
|
||||
# 2: overfit — both racks, Stage A (fixed IDs) then Stage B (random IDs, resumes A)
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 STAGE=A bash data_pipeline/720_stage_1/02_run_overfit.sh # rack0
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 STAGE=A bash data_pipeline/720_stage_1/02_run_overfit.sh # rack1
|
||||
# ...when track-following is clear in validation, stop and run Stage B (bump steps as needed):
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 STAGE=B bash data_pipeline/720_stage_1/02_run_overfit.sh --training.loop.max_train_steps 1600
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 STAGE=B bash data_pipeline/720_stage_1/02_run_overfit.sh --training.loop.max_train_steps 1600
|
||||
|
||||
# 3-4: export + merge
|
||||
bash data_pipeline/720_stage_1/03_export.sh # compute node
|
||||
bash data_pipeline/720_stage_1/04_merge.sh # compute node
|
||||
|
||||
# 5: OpenVid stage-1 — both racks (~10 days; resume by relaunching)
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 bash data_pipeline/720_stage_1/05_run_openvid_stage1.sh # rack0
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 bash data_pipeline/720_stage_1/05_run_openvid_stage1.sh # rack1
|
||||
|
||||
# 6-7: export stage-1 + synth stage-2 -> final teacher
|
||||
bash data_pipeline/720_stage_1/06_export_stage1.sh # compute node
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=0 bash data_pipeline/720_stage_1/07_run_synth_stage2.sh # rack0
|
||||
MASTER_ADDR=hpc-rack-1-6 NODE_RANK=1 bash data_pipeline/720_stage_1/07_run_synth_stage2.sh # rack1
|
||||
```
|
||||
|
||||
## Notes & decisions baked in
|
||||
|
||||
- **Topology / batch:** stage-1 and stage-2 target the **upstream recipe on 8 nodes × 4 GB200 = 32
|
||||
GPUs** — `num_gpus 32, hsdp_replicate_dim 8, hsdp_shard_dim 4, grad_accum 4 → global bs 128`.
|
||||
Stage-1 = 4,800 steps = 2.4 epochs over 259k ≈ **~12 days**; stage-2 = 600 steps ≈ **~1.5 days**.
|
||||
(On only 8 GPUs / 2 nodes this recipe is ~49 days — if you drop back, set `NNODES=2`,
|
||||
`hsdp_replicate_dim 2`, `grad_accum 2` → global bs 16, and use ~8000 stage-1 steps for a ~10-day
|
||||
half-epoch.) The overfit (steps 0–3) stays on **2 nodes** — it's ~64 clips, so 32 GPUs / bs 128
|
||||
would exceed the dataset; leave it at `NNODES=2` there.
|
||||
- **Launching the 8-node stages (05, 07):** the scripts take `NNODES` (default 8) and pass a matching
|
||||
`--training.distributed.num_gpus`. But manually running one command on each of 8 nodes (NODE_RANK
|
||||
0…7) is impractical — **use a SLURM launcher** (`srun` spans all nodes with `--node-rank=$SLURM_PROCID`,
|
||||
as in `examples/train/run_slurm.sh` / `run_wan14b_held.sh`). The overfit (02) is fine to launch
|
||||
manually on 2 nodes.
|
||||
- **Env knobs** live in the launch scripts: overfit = `FREEZE_HEAD=0`; stage-1/2 = `FREEZE_HEAD=1`;
|
||||
stage-2 adds `TRACK_DROP=0.5 MOTION_DROP=0.3 PMASK=0.2 MASK_CHUNK=8`. `TRACKWAN_TRACK_BIAS=1` and
|
||||
`WANTRACK_SPARSE=1 EXTRA_RANDOM=20` throughout.
|
||||
- **Stage-2 data** defaults to the openvid parquets + masking (you have no synth generated). Swap
|
||||
`data_path` to a synth set if you build one; the robustness comes from the masking either way.
|
||||
- **`dit_precision`:** fp32 master (default) for the real teacher stages; upstream uses bf16 — flip
|
||||
only if memory-blocked.
|
||||
- **Monitor track-following, not loss** — the `track_validation` with-track vs no-track/adversarial
|
||||
deltas. Gate the overfit (step 2) on these before merging, and watch stage-1's first validation
|
||||
(step 2000) before committing the full ~10 days. Check efficiency with `python scripts/mfu_estimate.py`.
|
||||
- **Not runnable end-to-end yet:** steps 4–7 chain via default paths (04 reads 03's export, 05 reads
|
||||
04's merge, etc.), so they're correct now but only *run* once the prior step's output exists.
|
||||
- **Two things to smoke-test once:** `dcp_to_diffusers` on a real checkpoint (steps 3 & 6), and that
|
||||
the overfit actually learns track-following on ~64 clips (if weak: `N_CHUNKS=4` in step 0, or more steps).
|
||||
@@ -0,0 +1,86 @@
|
||||
# Step 05 config — OpenVid Stage 1. 14B/720p, merged d64+bias init, HEAD FROZEN, sparse.
|
||||
# The big run: the DiT learns to USE the frozen (merged) track pathway on the full 720p openvid set.
|
||||
#
|
||||
# BATCH: upstream recipe = global bs 128 on 32 GPUs (8 nodes x 4). This config targets that: 4800
|
||||
# steps x 128 = 614k samples = 2.4 epochs over 259k, ~12 days on 32 GPUs. (On only 8 GPUs, bs 128
|
||||
# would be grad_accum 16 = ~49 days; if you drop back to 8 nodes/GPUs, use grad_accum 2 = global
|
||||
# bs 16 and ~8000 steps for a ~10-day half-epoch instead.)
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /home/hal-kevin/models/trackwan_14b_i2v_d64_merged_from_overfit_bias # from 04_merge.sh
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /home/hal-kevin/models/trackwan_14b_i2v_d64_merged_from_overfit_bias
|
||||
# dit_precision: fp32 master (default) for the real teacher (standard MFU / stable updates at lr 1e-5).
|
||||
# upstream sets bf16; flip only if memory-blocked (see notes/trackwan_14b_720p_teacher_merge_plan.md).
|
||||
|
||||
distributed:
|
||||
num_gpus: 32
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8 # across the 8 nodes
|
||||
hsdp_shard_dim: 4 # within each node (over NVLink)
|
||||
|
||||
data:
|
||||
data_path: /home/hal-shared/motionstream/data/openvid-wantrack-parquets # full 259k, 720p
|
||||
dataloader_num_workers: 2
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 720
|
||||
num_width: 1280
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5 # MotionStream teacher stage-1
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4800 # upstream stage-1: 2.4 epochs over 259k at bs128; ~12 days on 32 GPUs; resumable
|
||||
gradient_accumulation_steps: 4 # 1 x 32 GPU x 4 = global batch 128
|
||||
|
||||
checkpoint:
|
||||
output_dir: /home/hal-kevin/data/motion-stream-test/openvid_stage1_14b_720p_out
|
||||
training_state_checkpointing_steps: 200
|
||||
checkpoints_total_limit: 5
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-stage1-14b-720p-d64-bias
|
||||
entity: s4duan-uc-san-diego
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
track_validation:
|
||||
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
|
||||
every_steps: 2000 # 14B/720p validation is a multi-hour job; keep it rare
|
||||
num_val_samples: 1
|
||||
val_sample_indices: [28]
|
||||
num_inference_steps: 20
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
include_heldout: false
|
||||
validate_at_start: false
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,85 @@
|
||||
# Phase 0 overfit — 14B/720p, d64 + bias, sparse conditioning, HEAD TRAINABLE.
|
||||
# Goal: co-adapt the track pathway (track_encoder + patch-embed track slot) on a small subset so
|
||||
# the model follows tracks. The base is DISCARDED by the Phase-1 merge (only the track pathway is
|
||||
# lifted), so dit_precision: bf16 here is fine + faster. Sparse/bias/freeze knobs come from the
|
||||
# launch script (02_run_overfit.sh), not this file.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /home/hal-kevin/models/trackwan_14b_i2v_d64_bias_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /home/hal-kevin/models/trackwan_14b_i2v_d64_bias_init
|
||||
dit_precision: bf16 # overfit base is thrown away by the merge; bf16 for speed
|
||||
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 2 # across the 2 racks
|
||||
hsdp_shard_dim: 4 # within each rack
|
||||
|
||||
data:
|
||||
data_path: /home/hal-kevin/data/motion-stream-synth/preprocessed_i2v_track/combined_parquet_dataset
|
||||
dataloader_num_workers: 2
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 720
|
||||
num_width: 1280
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4 # upstream overfit LR (synth_sparse_*_14b), constant
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 2000 # upstream overfit steps (per stage); stop earlier if tracks lock in
|
||||
gradient_accumulation_steps: 1 # 1 x 8 GPU x 1 = global batch 8 (fine on a tiny subset)
|
||||
|
||||
checkpoint:
|
||||
output_dir: /home/hal-kevin/data/motion-stream-test/overfit_14b_720p_d64_bias_out
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 30 # keep effectively all (20 ckpts @ every-100 for 2000 steps) so you can export the BEST, not just the latest
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: overfit-14b-720p-d64-bias
|
||||
entity: s4duan-uc-san-diego
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
track_validation:
|
||||
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
|
||||
# Validate on the overfit clips themselves — you want to SEE track-following emerge. The
|
||||
# with-track vs no-track/adversarial deltas are the signal to stop the overfit.
|
||||
every_steps: 200
|
||||
num_val_samples: 2
|
||||
val_sample_indices: [0, 10]
|
||||
num_inference_steps: 20
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
include_heldout: false
|
||||
validate_at_start: false
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,84 @@
|
||||
# Step 07 config — Stage 2 (robustness). 14B/720p, inits from the exported Stage-1 teacher, HEAD
|
||||
# still FROZEN, sparse + heavy masking/dropout. Short finetune (lr 1e-6) -> the FINAL teacher.
|
||||
#
|
||||
# DATA: upstream uses a synthetic set here for controlled motion. You don't have synth generated,
|
||||
# so this defaults to the SAME 720p openvid parquets + the stage-2 masking env (the robustness comes
|
||||
# from the masking, applied on whatever data). Swap data_path to a synth set if you generate one
|
||||
# (data_pipeline/gen_synth_worker.py). Masking knobs are set in 07_run_synth_stage2.sh, not here.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /home/hal-kevin/models/openvid_stage1_14b_export # from 06_export_stage1.sh
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /home/hal-kevin/models/openvid_stage1_14b_export
|
||||
|
||||
distributed:
|
||||
num_gpus: 32
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8 # across the 8 nodes
|
||||
hsdp_shard_dim: 4 # within each node
|
||||
|
||||
data:
|
||||
data_path: /home/hal-shared/motionstream/data/openvid-wantrack-parquets # swap to synth if available
|
||||
dataloader_num_workers: 2
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 720
|
||||
num_width: 1280
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-6 # low LR, short robustness finetune
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 600 # upstream stage-2; ~1.5 days on 32 GPUs at global bs 128
|
||||
gradient_accumulation_steps: 4 # 1 x 32 GPU x 4 = global batch 128
|
||||
|
||||
checkpoint:
|
||||
output_dir: /home/hal-kevin/data/motion-stream-test/synth_stage2_14b_720p_out
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: synth-stage2-14b-720p-d64-bias
|
||||
entity: s4duan-uc-san-diego
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
track_validation:
|
||||
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
|
||||
every_steps: 300
|
||||
num_val_samples: 1
|
||||
val_sample_indices: [28]
|
||||
num_inference_steps: 20
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
include_heldout: false
|
||||
validate_at_start: false
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,111 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Fill in real captions on a videos2caption.json built from OpenVid clips.
|
||||
|
||||
The OpenVid-WanTrack shards ship mp4s only -- no metadata -- so a manifest built by
|
||||
scanning the clip directory has empty ``cap`` fields. Stage 5 would then encode empty
|
||||
strings through T5, giving every clip an identical null text embedding: text conditioning
|
||||
(and the joint text+motion CFG that depends on a meaningful conditional/null contrast)
|
||||
would be silently dead.
|
||||
|
||||
OpenVid-1M's caption CSV keys on exactly the same filenames (``---_iRTHryQ_13_0to241.mp4``),
|
||||
so the join is 1:1 on basename -- no id parsing required.
|
||||
|
||||
python data_pipeline/add_captions.py --manifest <root>/videos2caption.json
|
||||
|
||||
The CSV is downloaded once from the Hub (~300 MB) and cached; pass --captions-csv to use
|
||||
a local copy. Clips with no caption keep "" and are counted, never silently dropped.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
CAPTION_REPO = "nkp37/OpenVid-1M"
|
||||
CAPTION_FILES = ["data/train/OpenVid-1M.csv", "data/train/OpenVidHD.csv"]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--manifest", type=Path, required=True, help="videos2caption.json to patch in place.")
|
||||
p.add_argument("--captions-csv", type=Path, nargs="*", default=None,
|
||||
help="Local caption CSV(s). Default: download+cache from the Hub.")
|
||||
p.add_argument("--cache-dir", type=Path, default=Path.home() / ".cache/openvid_captions",
|
||||
help="Where downloaded caption CSVs are cached.")
|
||||
p.add_argument("--min-coverage", type=float, default=0.9,
|
||||
help="Fail if fewer than this fraction of clips get a caption (0 = never fail).")
|
||||
p.add_argument("--dry-run", action="store_true", help="Report coverage without writing.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def caption_paths(args: argparse.Namespace) -> list[Path]:
|
||||
if args.captions_csv:
|
||||
return list(args.captions_csv)
|
||||
from huggingface_hub import hf_hub_download
|
||||
args.cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
out = []
|
||||
for f in CAPTION_FILES:
|
||||
try:
|
||||
out.append(Path(hf_hub_download(CAPTION_REPO, f, repo_type="dataset",
|
||||
cache_dir=str(args.cache_dir))))
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[caption] warn: could not fetch {f} ({e})", flush=True)
|
||||
if not out:
|
||||
sys.exit("[caption] ERROR: no caption CSV available")
|
||||
return out
|
||||
|
||||
|
||||
def load_captions(paths: list[Path]) -> dict[str, str]:
|
||||
"""basename -> caption. Later files do not overwrite earlier hits."""
|
||||
caps: dict[str, str] = {}
|
||||
csv.field_size_limit(10 * 1024 * 1024) # captions can be long
|
||||
for p in paths:
|
||||
n0 = len(caps)
|
||||
with p.open(newline="", encoding="utf-8", errors="replace") as fh:
|
||||
for row in csv.DictReader(fh):
|
||||
key, cap = row.get("video"), row.get("caption")
|
||||
if key and cap and key not in caps:
|
||||
caps[key] = cap.strip()
|
||||
print(f"[caption] {p.name}: +{len(caps) - n0} captions (total {len(caps)})", flush=True)
|
||||
return caps
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
items = json.loads(args.manifest.read_text())
|
||||
caps = load_captions(caption_paths(args))
|
||||
|
||||
hit = 0
|
||||
missing: list[str] = []
|
||||
for it in items:
|
||||
name = Path(it.get("path", "")).name
|
||||
cap = caps.get(name)
|
||||
if cap:
|
||||
it["cap"] = [cap]
|
||||
hit += 1
|
||||
else:
|
||||
it.setdefault("cap", [""])
|
||||
missing.append(name)
|
||||
|
||||
cov = hit / max(len(items), 1)
|
||||
print(f"[caption] matched {hit}/{len(items)} clips ({cov*100:.1f}%)", flush=True)
|
||||
if missing:
|
||||
print(f"[caption] first few unmatched: {missing[:3]}", flush=True)
|
||||
|
||||
if cov < args.min_coverage:
|
||||
sys.exit(f"[caption] ERROR: coverage {cov*100:.1f}% < required {args.min_coverage*100:.0f}%. "
|
||||
"Refusing to write -- training on empty captions silently breaks text conditioning.")
|
||||
|
||||
if args.dry_run:
|
||||
print("[caption] dry run, manifest not written", flush=True)
|
||||
return
|
||||
tmp = args.manifest.with_suffix(".json.tmp")
|
||||
tmp.write_text(json.dumps(items, indent=2))
|
||||
tmp.replace(args.manifest)
|
||||
print(f"[caption] wrote {args.manifest}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,148 @@
|
||||
#!/bin/bash
|
||||
# Benchmark Stage 3 (extract_tracks) + Stage 4 (segment_tracks) end-to-end on the
|
||||
# motion-stream-test source videos, writing ALL outputs to a separate bench dir so
|
||||
# the real dataset is never touched.
|
||||
#
|
||||
# Usage:
|
||||
# bash data_pipeline/benchmark_tracks.sh # 50 videos, 4 GPUs
|
||||
# GPUS=0 bash data_pipeline/benchmark_tracks.sh # single GPU
|
||||
# LIMIT=5 bash data_pipeline/benchmark_tracks.sh # smoke run
|
||||
# VIZ=1 bash data_pipeline/benchmark_tracks.sh # include viz mp4 rendering
|
||||
# FUSED=1 bash data_pipeline/benchmark_tracks.sh # fused stage 3+4 (extract --segment)
|
||||
#
|
||||
# Results append to $OUT_DIR/benchmark_results.txt (tagged with git commit); each
|
||||
# entry records gpus/videos/viz/fused so runs stay comparable.
|
||||
set -euo pipefail
|
||||
|
||||
SRC_VIDEOS=${SRC_VIDEOS:-/home/hal-kevin/data/motion-stream-test/videos}
|
||||
SRC_MANIFEST=${SRC_MANIFEST:-/home/hal-kevin/data/motion-stream-test/videos2caption.json}
|
||||
OUT_DIR=${OUT_DIR:-/home/hal-kevin/data/motion-stream-qtest}
|
||||
GPUS=${GPUS:-0,1,2,3}
|
||||
LIMIT=${LIMIT:-}
|
||||
VIZ=${VIZ:-0} # 0 = lean run (production-like for large-scale), 1 = render overlay mp4s
|
||||
FUSED=${FUSED:-0} # 1 = single fused pass (extract_tracks --segment); stage 4 not run
|
||||
AMP=${AMP:-0} # 1 = bf16 autocast for CoTracker (--amp)
|
||||
COMPILE=${COMPILE:-0} # 1 = torch.compile the main CoTracker pass (--compile)
|
||||
OVERRIDE_EVERY=${OVERRIDE_EVERY:-3} # --vis-override-every (2 rides the entry-sweep mask cache in fused mode)
|
||||
|
||||
cd "$(dirname "$0")/.."
|
||||
IFS=',' read -ra GPU_ARR <<< "$GPUS"
|
||||
WORLD_SIZE=${#GPU_ARR[@]}
|
||||
MANIFEST=bench_manifest.json
|
||||
LOG="$OUT_DIR/benchmark.log"
|
||||
RESULTS="$OUT_DIR/benchmark_results.txt"
|
||||
LIMIT_ARGS=()
|
||||
[[ -n "$LIMIT" ]] && LIMIT_ARGS=(--limit "$LIMIT")
|
||||
VIZ_ARGS=()
|
||||
[[ "$VIZ" == "1" ]] && VIZ_ARGS=(--viz --viz-dir "$OUT_DIR/bench_viz")
|
||||
FUSED_ARGS=()
|
||||
if [[ "$FUSED" == "1" ]]; then
|
||||
FUSED_ARGS=(--segment --vis-override-every "$OVERRIDE_EVERY")
|
||||
[[ "$VIZ" == "1" ]] && FUSED_ARGS+=("${VIZ_ARGS[@]}")
|
||||
fi
|
||||
SPEED_ARGS=()
|
||||
[[ "$AMP" == "1" ]] && SPEED_ARGS+=(--amp)
|
||||
if [[ "$COMPILE" == "1" ]]; then
|
||||
SPEED_ARGS+=(--compile)
|
||||
# Persist compile artifacts in $HOME (default /tmp/torchinductor_* is wiped by
|
||||
# reboots/tmp-cleaners), so repeat runs start warm instead of recompiling.
|
||||
export TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR:-$HOME/.cache/torchinductor}
|
||||
export TRITON_CACHE_DIR=${TRITON_CACHE_DIR:-$HOME/.cache/triton}
|
||||
fi
|
||||
|
||||
# --- setup: symlink source videos, copy manifest without points_path -------------
|
||||
mkdir -p "$OUT_DIR"
|
||||
ln -sfn "$SRC_VIDEOS" "$OUT_DIR/bench_videos"
|
||||
python - "$SRC_MANIFEST" "$OUT_DIR/$MANIFEST" <<'PY'
|
||||
import json, sys
|
||||
items = json.load(open(sys.argv[1]))
|
||||
for it in items:
|
||||
it.pop("points_path", None) # stage 3 re-patches these to the bench tracks dir
|
||||
json.dump(items, open(sys.argv[2], "w"), indent=2)
|
||||
print(f"[bench] manifest: {len(items)} items -> {sys.argv[2]}")
|
||||
PY
|
||||
rm -rf "$OUT_DIR/bench_tracks"
|
||||
[[ "$VIZ" == "1" ]] && rm -rf "$OUT_DIR/bench_viz"
|
||||
: > "$LOG"
|
||||
|
||||
N_VIDEOS=$(ls "$OUT_DIR"/bench_videos/*.mp4 | wc -l)
|
||||
[[ -n "$LIMIT" ]] && N_VIDEOS=$LIMIT
|
||||
COMMIT=$(git rev-parse --short HEAD 2>/dev/null || echo unknown)$(git diff --quiet 2>/dev/null || echo -dirty)
|
||||
echo "[bench] commit=$COMMIT gpus=$GPUS videos=$N_VIDEOS log=$LOG"
|
||||
|
||||
wait_workers() { # wait_workers <stage-name> <pid...>
|
||||
local stage=$1 fail=0 pid
|
||||
shift
|
||||
for pid in "$@"; do
|
||||
wait "$pid" || fail=$((fail + 1))
|
||||
done
|
||||
[[ $fail -gt 0 ]] && echo "[bench] WARNING: $fail $stage worker(s) failed — check $LOG"
|
||||
return 0
|
||||
}
|
||||
|
||||
# --- stage 3: extract tracks ------------------------------------------------------
|
||||
t0=$(date +%s)
|
||||
PIDS=()
|
||||
for i in "${!GPU_ARR[@]}"; do
|
||||
CUDA_VISIBLE_DEVICES=${GPU_ARR[$i]} python -u data_pipeline/extract_tracks.py \
|
||||
--data-dir "$OUT_DIR" \
|
||||
--videos-subdir bench_videos \
|
||||
--out-subdir bench_tracks \
|
||||
--manifest "$MANIFEST" \
|
||||
--grid-size 50 \
|
||||
--device cuda \
|
||||
--detect-entries \
|
||||
--sam-conf 0.75 --sam-iou 0.9 --sam-imgsz 1024 \
|
||||
--entry-sample-every 2 --entry-min-area 0.001 --entry-new-area 0.5 \
|
||||
"${FUSED_ARGS[@]}" \
|
||||
"${SPEED_ARGS[@]}" \
|
||||
--force \
|
||||
--rank "$i" --world-size "$WORLD_SIZE" \
|
||||
"${LIMIT_ARGS[@]}" \
|
||||
>> "$LOG" 2>&1 &
|
||||
PIDS+=($!)
|
||||
done
|
||||
wait_workers "stage-3" "${PIDS[@]}"
|
||||
t1=$(date +%s)
|
||||
S3=$((t1 - t0)); [[ $S3 -eq 0 ]] && S3=1
|
||||
N_NPZ=$(ls "$OUT_DIR"/bench_tracks/*.npz 2>/dev/null | wc -l || true)
|
||||
echo "[bench] stage 3: ${S3}s for $N_NPZ npz"
|
||||
[[ "$N_NPZ" -eq "$N_VIDEOS" ]] || echo "[bench] WARNING: expected $N_VIDEOS npz — check $LOG"
|
||||
|
||||
# --- stage 4: segment tracks (skipped when FUSED=1: stage 3 already segmented) -----
|
||||
S4=0
|
||||
if [[ "$FUSED" != "1" ]]; then
|
||||
t2=$(date +%s)
|
||||
PIDS=()
|
||||
for i in "${!GPU_ARR[@]}"; do
|
||||
CUDA_VISIBLE_DEVICES=${GPU_ARR[$i]} python -u data_pipeline/segment_tracks.py \
|
||||
--data-dir "$OUT_DIR" \
|
||||
--videos-subdir bench_videos \
|
||||
--manifest "$MANIFEST" \
|
||||
--conf 0.75 --iou 0.9 --imgsz 1024 \
|
||||
--vis-override-every "$OVERRIDE_EVERY" \
|
||||
"${VIZ_ARGS[@]}" \
|
||||
--force \
|
||||
--rank "$i" --world-size "$WORLD_SIZE" \
|
||||
"${LIMIT_ARGS[@]}" \
|
||||
>> "$LOG" 2>&1 &
|
||||
PIDS+=($!)
|
||||
done
|
||||
wait_workers "stage-4" "${PIDS[@]}"
|
||||
t3=$(date +%s)
|
||||
S4=$((t3 - t2)); [[ $S4 -eq 0 ]] && S4=1
|
||||
echo "[bench] stage 4: ${S4}s"
|
||||
fi
|
||||
|
||||
# --- summary ----------------------------------------------------------------------
|
||||
{
|
||||
echo "=== $(date -u '+%Y-%m-%d %H:%M:%S') UTC commit=$COMMIT gpus=$GPUS videos=$N_VIDEOS viz=$VIZ fused=$FUSED amp=$AMP compile=$COMPILE override=$OVERRIDE_EVERY ==="
|
||||
awk -v s3="$S3" -v s4="$S4" -v w="$WORLD_SIZE" -v n="$N_VIDEOS" -v fused="$FUSED" 'BEGIN {
|
||||
label = (fused == "1") ? "stage 3+4 (fused):" : "stage 3 (extract):"
|
||||
printf "%s %5ds total %6.1fs/video/worker %5.1f videos/min\n", label, s3, s3*w/n, 60*n/s3
|
||||
if (fused != "1")
|
||||
printf "stage 4 (segment): %5ds total %6.1fs/video/worker %5.1f videos/min\n", s4, s4*w/n, 60*n/s4
|
||||
printf "end-to-end: %5ds total\n", s3+s4
|
||||
}'
|
||||
} | tee -a "$RESULTS"
|
||||
echo "[bench] results appended to $RESULTS"
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Build videos2caption.json + merge.txt for the final track dataset.
|
||||
|
||||
Pairs each 121f/720p clip (--clips-dir) with its CoTracker npz (--tracks-dir) and
|
||||
caption (--captions JSON {video.mp4: caption}). Output is exactly the manifest
|
||||
`fastvideo/pipelines/preprocess/v1_preprocess.py --preprocess_task i2v_track` consumes
|
||||
(same format used for the overfit run), so the dataset is training-ready.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--clips-dir", required=True)
|
||||
ap.add_argument("--tracks-dir", required=True)
|
||||
ap.add_argument("--captions", required=True, help="JSON {video.mp4: caption}")
|
||||
ap.add_argument("--out-dir", required=True)
|
||||
ap.add_argument("--fps", type=int, default=24)
|
||||
ap.add_argument("--num-frames", type=int, default=121)
|
||||
ap.add_argument("--height", type=int, default=720)
|
||||
ap.add_argument("--width", type=int, default=1280)
|
||||
a = ap.parse_args()
|
||||
|
||||
caps = json.load(open(a.captions))
|
||||
tracks = {p.stem: p for p in Path(a.tracks_dir).glob("*.npz")}
|
||||
items, missing_track, missing_cap = [], 0, 0
|
||||
for clip in sorted(Path(a.clips_dir).glob("*.mp4")):
|
||||
tr = tracks.get(clip.stem)
|
||||
if tr is None:
|
||||
missing_track += 1; continue
|
||||
cap = caps.get(clip.name, "")
|
||||
if not cap:
|
||||
missing_cap += 1
|
||||
items.append({
|
||||
"path": clip.name,
|
||||
"cap": [cap],
|
||||
"points_path": str(tr.resolve()),
|
||||
"fps": float(a.fps),
|
||||
"num_frames": a.num_frames,
|
||||
"resolution": {"width": a.width, "height": a.height},
|
||||
})
|
||||
outd = Path(a.out_dir); outd.mkdir(parents=True, exist_ok=True)
|
||||
(outd / "videos2caption.json").write_text(json.dumps(items, indent=2))
|
||||
(outd / "merge.txt").write_text(f"{a.clips_dir},{outd / 'videos2caption.json'}")
|
||||
print(f"[manifest] {len(items)} clips paired (clip+track+caption); "
|
||||
f"skipped {missing_track} clips w/o track, {missing_cap} w/o caption")
|
||||
print(f"[manifest] wrote {outd}/videos2caption.json + merge.txt")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,241 @@
|
||||
#!/usr/bin/env python3
|
||||
"""CFG-ablation eval on a trackwan checkpoint.
|
||||
|
||||
Reproduces ``fastvideo/train/callbacks/track_validation.py::_sample`` denoise
|
||||
formula EXACTLY, but with configurable ``w_text`` and ``w_motion`` so we can
|
||||
isolate which CFG branch is causing artifacts.
|
||||
|
||||
Runs each val sample under N configs; uploads all videos to one wandb run for
|
||||
side-by-side visual comparison.
|
||||
|
||||
Usage (single-GPU, on any held alloc):
|
||||
python data_pipeline/cfg_ablation.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/synth_stage2_paperLR_ckpt400 \
|
||||
--yaml examples/train/scenario/worldmodel/finetune_wantrack_synth_stage2_paperLR.yaml \
|
||||
--val-parquet /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset \
|
||||
--wandb-run-name cfg_abl_paperLR_ckpt400
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__)))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
from fastvideo.forward_context import set_forward_context # noqa: E402
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import ( # noqa: E402
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
|
||||
def load_val_samples(parquet_dir: str, text_len: int) -> list:
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
fs = sorted(glob.glob(os.path.join(parquet_dir, "**", "*.parquet"), recursive=True))
|
||||
if not fs:
|
||||
raise FileNotFoundError(parquet_dir)
|
||||
rows = []
|
||||
for f in fs:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
batch = collate_rows_from_parquet_schema(rows, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in rows]
|
||||
return [{
|
||||
"id": str(infos[i].get("id", f"clip{i:03d}")) if i < len(infos) else f"clip{i:03d}",
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i:i + 1].clone(),
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(),
|
||||
"object_ids": batch["object_ids"][i:i + 1].clone() if "object_ids" in batch else None,
|
||||
"track_weights": batch["track_weights"][i:i + 1].clone() if "track_weights" in batch else None,
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
} for i in range(len(rows))]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def generate_with_cfg(model, sample: dict, *, w_text: float, w_motion: float,
|
||||
num_steps: int, seed: int,
|
||||
null_text: torch.Tensor | None = None,
|
||||
null_text_mask: torch.Tensor | None = None,
|
||||
no_track: bool = False) -> np.ndarray:
|
||||
"""Reproduce track_validation._sample denoise EXACTLY, w/ configurable CFG.
|
||||
|
||||
If ``null_text`` is provided, use it as the ∅ text (canonical Wan/T5 convention
|
||||
= UMT5-encoded empty string). Otherwise fall back to zeros_like(txt).
|
||||
If ``no_track`` is True, pass tp=None/tv=None to every forward (matches the
|
||||
val callback's `track_val/no_track` panel).
|
||||
"""
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
|
||||
ff = sample["first_frame_latent"].to(device, dtype)
|
||||
txt = sample["text_embedding"].to(device, dtype)
|
||||
mask = sample["text_attention_mask"].to(device, dtype)
|
||||
clip = sample["clip_feature"].to(device, dtype)
|
||||
if no_track:
|
||||
tp = None
|
||||
tv = None
|
||||
else:
|
||||
# CRITICAL: sparse-sample tracks the same way training + track_validation do.
|
||||
# Feeding the raw 2500 tracks is out-of-distribution and produces melted output.
|
||||
tp_raw = sample["track_points"].float()
|
||||
tv_raw = sample["track_visibility"].float()
|
||||
oid = sample.get("object_ids")
|
||||
tw = sample.get("track_weights")
|
||||
aug_gen = torch.Generator(device=device).manual_seed(int(seed))
|
||||
tp_s, tv_s = model._augment_tracks(tp_raw, tv_raw, aug_gen, object_ids=oid, track_weights=tw)
|
||||
if tp_s is None or tv_s is None: # motion_drop fired -> fall back to raw
|
||||
tp_s, tv_s = tp_raw, tv_raw
|
||||
tp = tp_s.to(device, dtype)
|
||||
tv = tv_s.to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(int(num_steps), device=device)
|
||||
|
||||
do_text_cfg = (w_text != 1.0)
|
||||
do_motion_cfg = (w_motion != 1.0) and not no_track # no tracks -> motion CFG is a no-op
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, ts_):
|
||||
with torch.autocast(device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=ts_, attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=ts_,
|
||||
encoder_hidden_states_image=clip,
|
||||
track_points=tp_e, track_visibility=tv_e,
|
||||
return_dict=False)
|
||||
|
||||
if do_text_cfg:
|
||||
if null_text is not None:
|
||||
# Broadcast the [seq_len, dim] null embedding to match txt shape [B, seq_len, dim]
|
||||
nl = null_text.to(device, dtype)
|
||||
if nl.dim() == 2:
|
||||
nl = nl.unsqueeze(0).expand_as(txt).contiguous()
|
||||
txt_null = nl
|
||||
else:
|
||||
txt_null = torch.zeros_like(txt)
|
||||
else:
|
||||
txt_null = None
|
||||
|
||||
for tt in sched.timesteps:
|
||||
mi = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
ts_ = tt.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp, tv, mi, ts_)
|
||||
if not do_text_cfg and not do_motion_cfg:
|
||||
v = v_full
|
||||
elif do_text_cfg and not do_motion_cfg:
|
||||
v_no_text = _fwd(txt_null, tp, tv, mi, ts_)
|
||||
v = v_no_text + w_text * (v_full - v_no_text)
|
||||
elif do_motion_cfg and not do_text_cfg:
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
v = v_no_motion + w_motion * (v_full - v_no_motion)
|
||||
else:
|
||||
# joint text+motion CFG (matches track_validation._sample)
|
||||
v_no_text = _fwd(txt_null, tp, tv, mi, ts_)
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
alpha = w_text / max(w_text + w_motion, 1e-6)
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + w_text * (v_full - v_no_text) + w_motion * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), tt, latents.float(), return_dict=False)[0]
|
||||
|
||||
return twi.decode_to_pixels(model, latents)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True)
|
||||
p.add_argument("--yaml", required=True)
|
||||
p.add_argument("--val-parquet", required=True)
|
||||
p.add_argument("--wandb-project", default="wantrack-bidir")
|
||||
p.add_argument("--wandb-run-name", required=True)
|
||||
p.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/cfg_ablation")
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--configs", default="3.0,1.5;1.0,1.5;1.0,1.0",
|
||||
help="';'-sep list of 'w_text,w_motion' pairs")
|
||||
p.add_argument("--null-text-pt",
|
||||
help="Path to precomputed UMT5(``'') embedding .pt file. If given, this is "
|
||||
"used as the ∅ text branch instead of zeros_like(txt).")
|
||||
p.add_argument("--no-track", action="store_true",
|
||||
help="Force track_points=None throughout denoise. Reproduces the val "
|
||||
"callback's `track_val/no_track` panel.")
|
||||
args = p.parse_args()
|
||||
|
||||
null_text = None
|
||||
null_text_mask = None
|
||||
if args.null_text_pt:
|
||||
blob = torch.load(args.null_text_pt, map_location="cpu")
|
||||
null_text = blob["embedding"]
|
||||
null_text_mask = blob.get("attention_mask")
|
||||
print(f"[abl] loaded null-text embedding: shape={list(null_text.shape)} "
|
||||
f"norm={null_text.norm().item():.3f}", flush=True)
|
||||
|
||||
import wandb
|
||||
Path(args.out_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
print(f"[abl] loading model {args.model_dir}", flush=True)
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(getattr(tc.data, "text_padding_length", 256))
|
||||
|
||||
print(f"[abl] loading val samples from {args.val_parquet}", flush=True)
|
||||
samples = load_val_samples(args.val_parquet, text_len)
|
||||
print(f"[abl] {len(samples)} val samples", flush=True)
|
||||
|
||||
configs = [tuple(float(x) for x in pair.split(","))
|
||||
for pair in args.configs.split(";")]
|
||||
print(f"[abl] configs: {configs}", flush=True)
|
||||
|
||||
run = wandb.init(project=args.wandb_project, name=args.wandb_run_name,
|
||||
config={"model_dir": args.model_dir, "steps": args.steps,
|
||||
"configs": [list(c) for c in configs]})
|
||||
|
||||
for i, s in enumerate(samples):
|
||||
sid = s["id"] or f"clip{i:03d}"
|
||||
print(f"[abl] === sample {i}: {sid} ===", flush=True)
|
||||
wandb_log = {"sample": i, "caption": s["caption"][:120]}
|
||||
for wt, wm in configs:
|
||||
t0 = time.time()
|
||||
frames = generate_with_cfg(model, s, w_text=wt, w_motion=wm,
|
||||
num_steps=args.steps, seed=args.seed + i,
|
||||
null_text=null_text, null_text_mask=null_text_mask,
|
||||
no_track=args.no_track)
|
||||
dt = time.time() - t0
|
||||
tag = f"wt{wt}_wm{wm}".replace(".", "p")
|
||||
fn = Path(args.out_dir) / f"{sid}_{tag}.mp4"
|
||||
imageio.mimsave(str(fn), frames, fps=args.fps, macro_block_size=1)
|
||||
wandb_log[f"gen_{tag}"] = wandb.Video(str(fn), fps=args.fps, format="mp4")
|
||||
print(f"[abl] {tag}: {dt:.1f}s -> {fn.name}", flush=True)
|
||||
wandb.log(wandb_log)
|
||||
|
||||
# also render GT reference once
|
||||
for i, s in enumerate(samples):
|
||||
try:
|
||||
ref = twi.decode_reference(model, s["vae_latent"].to(model.device, torch.bfloat16))
|
||||
fn = Path(args.out_dir) / f"{s['id']}_gt.mp4"
|
||||
imageio.mimsave(str(fn), ref, fps=args.fps, macro_block_size=1)
|
||||
wandb.log({"sample": i, "reference_gt": wandb.Video(str(fn), fps=args.fps, format="mp4")})
|
||||
except Exception as e:
|
||||
print(f"[abl] gt render failed for sample {i}: {e}", flush=True)
|
||||
|
||||
wandb.finish()
|
||||
print("[abl] DONE", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,94 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Report the WanTrack "gate" — patch_embedding.weight[:, 36:] — from a DCP checkpoint.
|
||||
|
||||
The overfit stages (A/B) exist to grow this slot from EXACTLY ZERO up to roughly the scale the
|
||||
successful 1.3B merged init reached (std ~0.011), co-adapting it with track_encoder. If it is
|
||||
still ~0 there is no bootstrap happening and step C would merge nothing; if it explodes past the
|
||||
pretrained channel scale (~0.0124) the track pathway is overwhelming the image prior.
|
||||
|
||||
Loads ONLY the tensors of interest (DCP resharding is per-tensor), so it costs ~MBs, not the
|
||||
~200GB a full checkpoint load would.
|
||||
|
||||
Usage: python data_pipeline/check_gate_growth.py <checkpoint-dir> [more dirs...]
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.distributed.checkpoint as dcp
|
||||
from torch.distributed.checkpoint import FileSystemReader
|
||||
|
||||
BASE_IN = 36
|
||||
# NOTE: the live module nests the conv, so a TRAINED checkpoint stores
|
||||
# 'patch_embedding.proj.weight', while the converter writes 'patch_embedding.weight' into the
|
||||
# init safetensors. Try both — silently matching neither is how the earlier LWLR experiment
|
||||
# ended up measuring nothing.
|
||||
PE_KEYS = ["patch_embedding.proj.weight", "patch_embedding.weight"]
|
||||
KEYS = PE_KEYS + [
|
||||
"track_encoder.proj.weight",
|
||||
"track_encoder.temporal_conv.weight",
|
||||
]
|
||||
|
||||
|
||||
def _find(md_keys, want: str) -> str | None:
|
||||
"""DCP nests model tensors under a role prefix (e.g. 'student.'); match on the suffix."""
|
||||
exact = [k for k in md_keys if k == want]
|
||||
if exact:
|
||||
return exact[0]
|
||||
cands = [k for k in md_keys if k.endswith(want)]
|
||||
return cands[0] if cands else None
|
||||
|
||||
|
||||
def report(ckpt: Path) -> None:
|
||||
dcp_dir = ckpt / "dcp" if (ckpt / "dcp").is_dir() else ckpt
|
||||
reader = FileSystemReader(str(dcp_dir))
|
||||
md = reader.read_metadata()
|
||||
planned = md.state_dict_metadata
|
||||
sd = {}
|
||||
resolved = {}
|
||||
for want in KEYS:
|
||||
k = _find(list(planned.keys()), want)
|
||||
if k is None:
|
||||
continue
|
||||
meta = planned[k]
|
||||
sd[k] = torch.empty(tuple(meta.size), dtype=meta.properties.dtype)
|
||||
resolved[want] = k
|
||||
if not sd:
|
||||
print(f"{ckpt.name}: none of {KEYS} found in checkpoint metadata")
|
||||
return
|
||||
dcp.load(sd, checkpoint_id=str(dcp_dir))
|
||||
|
||||
out = [f"{ckpt.name:>18s}"]
|
||||
pe_k = next((resolved[k] for k in PE_KEYS if k in resolved), None)
|
||||
if pe_k is None:
|
||||
out.append("!! patch_embedding NOT FOUND — cannot report gate")
|
||||
else:
|
||||
pe = sd[pe_k].float()
|
||||
out.append(f"gate[:, {BASE_IN}:] std={pe[:, BASE_IN:].std():.6f}")
|
||||
out.append(f"pretrained std={pe[:, :BASE_IN].std():.6f}")
|
||||
for w in ("track_encoder.proj.weight", "track_encoder.temporal_conv.weight"):
|
||||
if w in resolved:
|
||||
out.append(f"{w.split('.')[1]}_norm={sd[resolved[w]].float().norm():.4f}")
|
||||
print(" ".join(out))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if len(sys.argv) < 2:
|
||||
print(__doc__)
|
||||
sys.exit(1)
|
||||
for key, default in [("RANK", "0"), ("LOCAL_RANK", "0"), ("WORLD_SIZE", "1"),
|
||||
("MASTER_ADDR", "127.0.0.1"), ("MASTER_PORT", "29555")]:
|
||||
os.environ.setdefault(key, default)
|
||||
print(f"reference: 1.3B merged-from-overfit gate std = 0.011418 (grown from 0.000000)")
|
||||
for p in sys.argv[1:]:
|
||||
try:
|
||||
report(Path(p))
|
||||
except Exception as e: # keep going across checkpoints
|
||||
print(f"{Path(p).name}: ERROR {type(e).__name__}: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 0b (DROID): build a track-conditioning training set from DROID.
|
||||
|
||||
Goal (per the project decision): a set where the **first frame is similar** across
|
||||
clips but the **motion is diverse**, with a single generic prompt, so the model is
|
||||
forced to read the point tracks to reduce loss (content cannot determine the output).
|
||||
|
||||
DROID has no scene grouping in its LeRobot metadata, so "similar first frame" is
|
||||
obtained by *clustering frame-0*: we download a candidate pool, embed each first
|
||||
frame (downscaled RGB), and pick the ``--num-clips`` episodes whose first frames are
|
||||
mutually most similar (the tightest dense cluster -- the centroid that minimises the
|
||||
radius to its k-th nearest neighbour, then its k nearest).
|
||||
|
||||
Output mirrors ``generate_videos.py`` exactly so the existing downstream
|
||||
(``extract_tracks.py`` -> ``v1_preprocess.py --preprocess_task i2v_track``) ingests it
|
||||
unchanged:
|
||||
|
||||
<out>/videos/vid_000000.mp4 ... libx264, ``--num-frames`` frames @ ``--fps``
|
||||
<out>/videos2caption.json FastVideo manifest (one generic caption)
|
||||
<out>/merge.txt "<videos_dir>,<json_path>"
|
||||
|
||||
Source: HuggingFace LeRobot mirror ``IPEC-COMMUNITY/droid_lerobot`` (v2.0): per-episode
|
||||
av1 mp4 at ``videos/chunk-{c:03d}/{video_key}/episode_{i:06d}.mp4`` (av1 -> decoded
|
||||
with PyAV; re-encoded to libx264 so downstream decord reads it). 180x320, 15 fps.
|
||||
|
||||
Run on a compute node (network + a little CPU; no GPU)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 .venv/bin/python data_pipeline/convert_droid.py \
|
||||
--out-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200 \
|
||||
--num-clips 200 --candidate-scan 1500
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
REPO = "IPEC-COMMUNITY/droid_lerobot"
|
||||
CHUNK_SIZE = 1000 # meta/info.json chunks_size
|
||||
DEFAULT_VIDEO_KEY = "observation.images.exterior_image_1_left"
|
||||
DEFAULT_PROMPT = "a robot arm manipulating objects on a tabletop"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--out-dir", type=Path, required=True)
|
||||
p.add_argument("--num-clips", type=int, default=200, help="How many episodes to keep (the cluster size).")
|
||||
p.add_argument("--candidate-scan",
|
||||
type=int,
|
||||
default=1500,
|
||||
help="How many episodes to download+embed before clustering. Larger => tighter cluster.")
|
||||
p.add_argument("--video-key", type=str, default=DEFAULT_VIDEO_KEY)
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--fps", type=int, default=24, help="Written fps == train_fps so preprocess does NOT resample.")
|
||||
p.add_argument("--frame-start", type=int, default=0, help="First source-frame index of the kept window.")
|
||||
p.add_argument("--prompt", type=str, default=DEFAULT_PROMPT)
|
||||
p.add_argument("--embed-size", type=int, default=64, help="Frame-0 is resized to NxN for the similarity embedding.")
|
||||
p.add_argument("--no-cluster", action="store_true", help="Skip clustering; keep the first --num-clips candidates.")
|
||||
p.add_argument("--download-workers", type=int, default=16)
|
||||
p.add_argument("--seed", type=int, default=0)
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def list_candidate_episodes(num_frames: int, scan: int) -> list[int]:
|
||||
"""Episode indices (length >= num_frames), first ``scan`` in index order."""
|
||||
from huggingface_hub import hf_hub_download
|
||||
ep_path = hf_hub_download(REPO, "meta/episodes.jsonl", repo_type="dataset")
|
||||
out: list[int] = []
|
||||
with open(ep_path) as f:
|
||||
for line in f:
|
||||
d = json.loads(line)
|
||||
if int(d.get("length", 0)) >= num_frames:
|
||||
out.append(int(d["episode_index"]))
|
||||
if len(out) >= scan:
|
||||
break
|
||||
return out
|
||||
|
||||
|
||||
def episode_video_path(ep_idx: int, video_key: str) -> str:
|
||||
chunk = ep_idx // CHUNK_SIZE
|
||||
return f"videos/chunk-{chunk:03d}/{video_key}/episode_{ep_idx:06d}.mp4"
|
||||
|
||||
|
||||
def download_one(ep_idx: int, video_key: str) -> tuple[int, str | None]:
|
||||
from huggingface_hub import hf_hub_download
|
||||
try:
|
||||
p = hf_hub_download(REPO, episode_video_path(ep_idx, video_key), repo_type="dataset")
|
||||
return ep_idx, p
|
||||
except Exception: # noqa: BLE001
|
||||
return ep_idx, None
|
||||
|
||||
|
||||
def decode_frames(path: str, start: int, count: int) -> np.ndarray | None:
|
||||
"""Decode ``count`` frames starting at ``start`` (RGB uint8 [T,H,W,3]) via PyAV (handles av1)."""
|
||||
import av
|
||||
try:
|
||||
container = av.open(path)
|
||||
frames = []
|
||||
for idx, frame in enumerate(container.decode(video=0)):
|
||||
if idx >= start:
|
||||
frames.append(frame.to_ndarray(format="rgb24"))
|
||||
if len(frames) >= count:
|
||||
break
|
||||
container.close()
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if len(frames) < count:
|
||||
return None
|
||||
return np.stack(frames, axis=0)
|
||||
|
||||
|
||||
def embed_frame0(path: str, size: int) -> np.ndarray | None:
|
||||
from PIL import Image
|
||||
fr = decode_frames(path, 0, 1)
|
||||
if fr is None:
|
||||
return None
|
||||
img = Image.fromarray(fr[0]).resize((size, size), Image.BILINEAR)
|
||||
v = np.asarray(img, dtype=np.float32).reshape(-1)
|
||||
n = np.linalg.norm(v) + 1e-8
|
||||
return v / n
|
||||
|
||||
|
||||
def pick_tightest_cluster(embs: np.ndarray, k: int) -> np.ndarray:
|
||||
"""Return indices of the k mutually-most-similar embeddings.
|
||||
|
||||
Centroid = the point minimising the distance to its k-th nearest neighbour
|
||||
(densest region); then take that point's k nearest neighbours.
|
||||
"""
|
||||
# cosine distance (embs are L2-normalised) -> 1 - sim
|
||||
sim = embs @ embs.T
|
||||
dist = 1.0 - sim
|
||||
np.fill_diagonal(dist, 0.0)
|
||||
part = np.partition(dist, kth=min(k, dist.shape[0] - 1), axis=1)
|
||||
kth = part[:, min(k, dist.shape[0] - 1)] # radius to k-th NN per row
|
||||
centroid = int(np.argmin(kth))
|
||||
order = np.argsort(dist[centroid])[:k]
|
||||
return order
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
out_dir = args.out_dir
|
||||
videos_dir = out_dir / "videos"
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
json_path = out_dir / "videos2caption.json"
|
||||
merge_path = out_dir / "merge.txt"
|
||||
|
||||
print(f"[droid] listing candidates (length>={args.num_frames}, scan={args.candidate_scan})...", flush=True)
|
||||
candidates = list_candidate_episodes(args.num_frames, args.candidate_scan)
|
||||
print(f"[droid] {len(candidates)} candidate episodes", flush=True)
|
||||
|
||||
# Download (cached) + embed frame-0 in parallel.
|
||||
print(f"[droid] downloading + embedding frame-0 ({args.download_workers} workers)...", flush=True)
|
||||
ep_to_path: dict[int, str] = {}
|
||||
with ThreadPoolExecutor(max_workers=args.download_workers) as ex:
|
||||
futs = {ex.submit(download_one, ep, args.video_key): ep for ep in candidates}
|
||||
for i, fut in enumerate(as_completed(futs), 1):
|
||||
ep, path = fut.result()
|
||||
if path is not None:
|
||||
ep_to_path[ep] = path
|
||||
if i % 100 == 0:
|
||||
print(f"[droid] downloaded {i}/{len(candidates)}", flush=True)
|
||||
print(f"[droid] downloaded {len(ep_to_path)} mp4s; embedding...", flush=True)
|
||||
|
||||
eps_ok: list[int] = []
|
||||
embs: list[np.ndarray] = []
|
||||
for ep in candidates:
|
||||
if ep not in ep_to_path:
|
||||
continue
|
||||
e = embed_frame0(ep_to_path[ep], args.embed_size)
|
||||
if e is not None:
|
||||
eps_ok.append(ep)
|
||||
embs.append(e)
|
||||
embs_arr = np.stack(embs, axis=0)
|
||||
print(f"[droid] embedded {len(eps_ok)} frame-0s", flush=True)
|
||||
|
||||
if args.no_cluster or len(eps_ok) <= args.num_clips:
|
||||
sel_local = list(range(min(args.num_clips, len(eps_ok))))
|
||||
print(f"[droid] no clustering -> first {len(sel_local)} candidates", flush=True)
|
||||
else:
|
||||
sel_local = pick_tightest_cluster(embs_arr, args.num_clips).tolist()
|
||||
# report cluster tightness
|
||||
sub = embs_arr[sel_local]
|
||||
c = sub.mean(0, keepdims=True)
|
||||
c /= np.linalg.norm(c) + 1e-8
|
||||
mean_sim = float((sub @ c.T).mean())
|
||||
print(f"[droid] tightest-cluster of {len(sel_local)} clips; mean cos-sim to centroid={mean_sim:.3f}",
|
||||
flush=True)
|
||||
|
||||
selected_eps = [eps_ok[i] for i in sel_local]
|
||||
|
||||
# Decode + write libx264 mp4s + manifest.
|
||||
import imageio.v2 as imageio
|
||||
records = []
|
||||
seq = 0
|
||||
for ep in selected_eps:
|
||||
frames = decode_frames(ep_to_path[ep], args.frame_start, args.num_frames)
|
||||
if frames is None:
|
||||
print(f"[droid] skip ep{ep}: decode<{args.num_frames} frames", flush=True)
|
||||
continue
|
||||
H, W = int(frames.shape[1]), int(frames.shape[2])
|
||||
name = f"vid_{seq:06d}.mp4"
|
||||
imageio.mimsave(str(videos_dir / name),
|
||||
list(frames),
|
||||
fps=args.fps,
|
||||
codec="libx264",
|
||||
macro_block_size=1,
|
||||
output_params=["-pix_fmt", "yuv420p"])
|
||||
records.append({
|
||||
"idx": seq,
|
||||
"path": name,
|
||||
"cap": [args.prompt],
|
||||
"fps": float(args.fps),
|
||||
"duration": float(args.num_frames) / float(args.fps),
|
||||
"num_frames": int(args.num_frames),
|
||||
"resolution": {
|
||||
"width": W,
|
||||
"height": H
|
||||
},
|
||||
"source_episode": int(ep),
|
||||
})
|
||||
seq += 1
|
||||
if seq % 25 == 0:
|
||||
print(f"[droid] wrote {seq}/{len(selected_eps)}", flush=True)
|
||||
|
||||
json_path.write_text(json.dumps(records, indent=2))
|
||||
merge_path.write_text(f"{videos_dir.resolve()},{json_path.resolve()}\n")
|
||||
print(f"[droid] DONE: wrote {len(records)} clips -> {videos_dir}", flush=True)
|
||||
print(f"[droid] manifest -> {json_path}", flush=True)
|
||||
print(f"[droid] next: extract_tracks.py --data-dir {out_dir} ; then v1_preprocess i2v_track", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Convert openvid-wantrack parquets to bf16 (a COPY; never mutates the source).
|
||||
|
||||
Casts the big float tensor fields (vae_latent, first_frame_latent, text_embedding,
|
||||
clip_feature, track_points, track_visibility) from float32 to bfloat16 and tags their
|
||||
``*_dtype`` column ``"bfloat16"``. Integer / tiny fields (object_ids, track_weights) and all
|
||||
metadata are copied unchanged. Training already downcasts these fields to bf16
|
||||
(``wan.py`` / ``wantrack.py``), so the quality loss is negligible while the files ~halve.
|
||||
|
||||
Requires the loader change in ``fastvideo/dataset/utils.py`` that honors the ``*_dtype``
|
||||
column (both decoders). Without it, the bf16 bytes would be misread as float32.
|
||||
|
||||
The output mirrors the source tree (shardNNN/combined_parquet_dataset/worker_N/*.parquet), so
|
||||
point training ``data_path`` at ``--dst`` once you've converted what you want. Resumable:
|
||||
already-written files are skipped. CPU only; run on a compute node (I/O + RAM heavy).
|
||||
|
||||
Examples:
|
||||
# first half of all parquets -> a bf16 sibling dir
|
||||
python data_pipeline/convert_parquets_to_bf16.py --fraction 0.5
|
||||
|
||||
# a fixed number of files, dry-run first
|
||||
python data_pipeline/convert_parquets_to_bf16.py --limit 100 --dry-run
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
DEFAULT_SRC = "/home/hal-shared/motionstream/data/openvid-wantrack-parquets"
|
||||
DEFAULT_DST = "/home/hal-shared/motionstream/data/openvid-wantrack-parquets-bf16"
|
||||
|
||||
# Big float fields consumed at bf16 by training -> safe to store bf16.
|
||||
DEFAULT_BF16_FIELDS = [
|
||||
"vae_latent", "first_frame_latent", "text_embedding", "clip_feature",
|
||||
"track_points", "track_visibility",
|
||||
]
|
||||
# Left untouched (integer labels / tiny): object_ids, track_weights, all metadata.
|
||||
|
||||
_SRC_STR_TO_TORCH = {
|
||||
"float32": torch.float32, "float16": torch.float16, "float64": torch.float64,
|
||||
}
|
||||
|
||||
|
||||
def _to_bf16_bytes(b: bytes, dtype_str: str) -> tuple[bytes, str]:
|
||||
"""Re-encode a raw float tensor blob as bf16. Returns (bytes, dtype_label)."""
|
||||
if not b: # empty optional field -> leave as-is
|
||||
return b, (dtype_str or "")
|
||||
if dtype_str == "bfloat16": # already converted
|
||||
return b, dtype_str
|
||||
src = _SRC_STR_TO_TORCH.get(dtype_str or "float32")
|
||||
if src is None:
|
||||
raise ValueError(f"cannot convert stored dtype {dtype_str!r} to bf16")
|
||||
t = torch.frombuffer(bytearray(b), dtype=src).to(torch.bfloat16)
|
||||
return t.view(torch.uint8).numpy().tobytes(), "bfloat16"
|
||||
|
||||
|
||||
def convert_file(src_path: str, dst_path: str, fields: list[str]) -> tuple[int, int]:
|
||||
"""Convert one parquet file. Returns (src_bytes, dst_bytes)."""
|
||||
tbl = pq.read_table(src_path)
|
||||
names = list(tbl.schema.names)
|
||||
cols: dict[str, object] = {n: tbl.column(n) for n in names}
|
||||
|
||||
for fld in fields:
|
||||
bkey, dkey = f"{fld}_bytes", f"{fld}_dtype"
|
||||
if bkey not in cols or dkey not in cols:
|
||||
continue
|
||||
b_list = tbl.column(bkey).to_pylist()
|
||||
d_list = tbl.column(dkey).to_pylist()
|
||||
new_b, new_d = [], []
|
||||
for b, d in zip(b_list, d_list, strict=True):
|
||||
nb, nd = _to_bf16_bytes(b, d)
|
||||
new_b.append(nb)
|
||||
new_d.append(nd)
|
||||
cols[bkey] = pa.array(new_b, type=pa.binary())
|
||||
cols[dkey] = pa.array(new_d, type=pa.string())
|
||||
|
||||
out = pa.table([cols[n] for n in names], schema=tbl.schema)
|
||||
os.makedirs(os.path.dirname(dst_path), exist_ok=True)
|
||||
tmp = dst_path + ".tmp"
|
||||
pq.write_table(out, tmp)
|
||||
os.replace(tmp, dst_path)
|
||||
return os.path.getsize(src_path), os.path.getsize(dst_path)
|
||||
|
||||
|
||||
def verify_file(src_path: str, dst_path: str, fields: list[str]) -> None:
|
||||
"""Round-trip check: one bf16 field on row 0 must equal src fp32 -> bf16."""
|
||||
s = pq.ParquetFile(src_path).read_row_group(0).slice(0, 1).to_pylist()[0]
|
||||
d = pq.ParquetFile(dst_path).read_row_group(0).slice(0, 1).to_pylist()[0]
|
||||
for fld in fields:
|
||||
sb, db = s.get(f"{fld}_bytes"), d.get(f"{fld}_bytes")
|
||||
if not sb:
|
||||
continue
|
||||
assert d.get(f"{fld}_dtype") == "bfloat16", f"{fld}: dtype not tagged bfloat16"
|
||||
src_t = torch.frombuffer(bytearray(sb), dtype=_SRC_STR_TO_TORCH[s[f"{fld}_dtype"]]).to(torch.bfloat16)
|
||||
dst_t = torch.frombuffer(bytearray(db), dtype=torch.bfloat16)
|
||||
assert torch.equal(dst_t, src_t), f"{fld}: bf16 round-trip mismatch"
|
||||
return # one field is enough
|
||||
return
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--src", default=DEFAULT_SRC)
|
||||
p.add_argument("--dst", default=DEFAULT_DST)
|
||||
p.add_argument("--fraction", type=float, default=0.5, help="fraction of the sorted parquet list to convert (default first 0.5)")
|
||||
p.add_argument("--limit", type=int, default=None, help="convert at most N files (overrides --fraction)")
|
||||
p.add_argument("--offset", type=int, default=0, help="skip the first N files of the sorted list (parallelize by running disjoint --offset/--limit ranges)")
|
||||
p.add_argument("--fields", default=",".join(DEFAULT_BF16_FIELDS), help="comma-separated fields to cast to bf16")
|
||||
p.add_argument("--overwrite", action="store_true", help="re-convert files already present in --dst")
|
||||
p.add_argument("--no-verify", action="store_true", help="skip the per-file round-trip check")
|
||||
p.add_argument("--dry-run", action="store_true", help="list what would be converted; write nothing")
|
||||
args = p.parse_args()
|
||||
|
||||
src_root = os.path.realpath(args.src)
|
||||
dst_root = os.path.realpath(args.dst)
|
||||
fields = [f.strip() for f in args.fields.split(",") if f.strip()]
|
||||
if os.path.commonpath([src_root, dst_root]) == src_root and dst_root != src_root:
|
||||
raise SystemExit(f"--dst {dst_root} is inside --src; choose a separate directory")
|
||||
if src_root == dst_root:
|
||||
raise SystemExit("refusing to convert in place; --dst must differ from --src")
|
||||
|
||||
all_files = []
|
||||
for root, _, files in os.walk(src_root):
|
||||
for f in files:
|
||||
if f.endswith(".parquet"):
|
||||
all_files.append(os.path.join(root, f))
|
||||
all_files.sort()
|
||||
n_total = len(all_files)
|
||||
n_take = args.limit if args.limit is not None else int(args.fraction * n_total)
|
||||
selected = all_files[args.offset:args.offset + n_take]
|
||||
print(f"[bf16] {n_total} parquet(s) found; converting {len(selected)} "
|
||||
f"[offset {args.offset}, {'limit ' + str(args.limit) if args.limit is not None else f'fraction {args.fraction}'}]")
|
||||
print(f"[bf16] fields -> bf16: {fields}")
|
||||
print(f"[bf16] src={src_root}\n[bf16] dst={dst_root}")
|
||||
|
||||
if args.dry_run:
|
||||
for f in selected[:5]:
|
||||
print(" would convert:", os.path.relpath(f, src_root))
|
||||
if len(selected) > 5:
|
||||
print(f" ... and {len(selected) - 5} more")
|
||||
return
|
||||
|
||||
src_tot = dst_tot = done = skipped = 0
|
||||
for i, sp in enumerate(selected, 1):
|
||||
rel = os.path.relpath(sp, src_root)
|
||||
dp = os.path.join(dst_root, rel)
|
||||
if os.path.exists(dp) and not args.overwrite:
|
||||
skipped += 1
|
||||
continue
|
||||
sb, db = convert_file(sp, dp, fields)
|
||||
if not args.no_verify:
|
||||
verify_file(sp, dp, fields)
|
||||
src_tot += sb
|
||||
dst_tot += db
|
||||
done += 1
|
||||
if done % 20 == 0 or i == len(selected):
|
||||
gb = 1024 ** 3
|
||||
print(f"[bf16] {i}/{len(selected)} | converted {done}, skipped {skipped} | "
|
||||
f"{src_tot/gb:.1f}GB -> {dst_tot/gb:.1f}GB"
|
||||
f"{f' ({dst_tot/src_tot*100:.0f}%)' if src_tot else ''}")
|
||||
|
||||
print(f"[bf16] done: converted {done}, skipped {skipped}. Output tree at {dst_root}")
|
||||
print(f"[bf16] point training data_path at {dst_root} once you've converted enough shards.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,159 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Build a WanTrack init checkpoint from a base Wan diffusers model.
|
||||
|
||||
The WanTrack DiT widens the patch-embed input to 52 channels and adds a
|
||||
``track_encoder``. Works from either base:
|
||||
- a Wan **I2V** model (e.g. Wan2.1-Fun-1.3B-InP: in_channels=36 = 16 noisy +
|
||||
4 mask + 16 first-frame, with CLIP image cross-attention): the 36 pretrained
|
||||
image-conditioning channels are kept and only the 16 track channels are
|
||||
zero-init, so first-frame conditioning works from step 0. **Recommended.**
|
||||
- a Wan **T2V** model (in_channels=16): the 20 I2V + 16 track channels are all
|
||||
zero-init, so I2V must be learned from scratch (slower, weaker).
|
||||
|
||||
It produces a diffusers transformer dir whose weights load *strictly* into
|
||||
``TrackWanTransformer3DModel``:
|
||||
- ``patch_embedding.weight`` zero-padded base_in -> 52 input channels (pretrained
|
||||
weights occupy the first base_in channels; the new track channels start at
|
||||
zero, so a freshly converted model reproduces the base at step 0),
|
||||
- ``track_encoder.*`` added (proj normal-init; the patch-embed track channels are the
|
||||
single zero-conv -> zero track contribution at step 0, but gradient still flows so the
|
||||
track pathway can learn -- see build_track_encoder_state for the deadlock rationale),
|
||||
- ``config.json`` gets ``in_channels=52`` + ``track_config`` (CLIP ``image_dim``
|
||||
is inherited from the base config when present).
|
||||
|
||||
Other pipeline components (vae / text_encoder / tokenizer / scheduler /
|
||||
model_index.json) are symlinked so the output is a complete, loadable model dir.
|
||||
|
||||
Usage (no GPU / no fastvideo import needed):
|
||||
python data_pipeline/convert_trackwan_init.py \
|
||||
--base <base diffusers model dir> --out <output dir>
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from torch import nn
|
||||
|
||||
NEW_IN_CHANNELS = 52
|
||||
TRACK_CHANNELS = 16
|
||||
ID_DIM = 128
|
||||
VAE_T_COMP = 4
|
||||
TRACK_CONFIG = {
|
||||
"id_dim": ID_DIM,
|
||||
"track_channels": TRACK_CHANNELS,
|
||||
"vae_spatial_compression": 8,
|
||||
"vae_temporal_compression": VAE_T_COMP,
|
||||
"max_track_id": 100_000,
|
||||
# The single zero-conv is the patch-embed track channels (zero-padded below), NOT the
|
||||
# track head -> proj is normal-init so the track pathway actually receives gradient.
|
||||
"zero_init_head": False,
|
||||
}
|
||||
|
||||
|
||||
def build_track_encoder_state() -> dict[str, torch.Tensor]:
|
||||
"""Match TrackEncoder's params: temporal_conv + proj, BOTH normal-initialized.
|
||||
|
||||
IMPORTANT (deadlock fix): the track signal passes through two layers in series --
|
||||
``track_encoder.proj`` then the patch-embed track channels [36:52]. The single
|
||||
ControlNet-style zero-conv is the *patch-embed track channels* (kept at 0 by the
|
||||
zero-pad in main()), which already guarantees zero track contribution at step 0
|
||||
(teacher behavior). ``proj`` must therefore be NON-zero so gradient can reach the
|
||||
patch-embed track channels (grad ∝ proj output): if BOTH were zero, each layer's
|
||||
gradient is gated by the other being nonzero -> both stay exactly 0 forever and the
|
||||
track pathway never learns (observed: proj & patch-embed[36:52] frozen at 0.0 after
|
||||
4000 steps). So leave proj at its default Conv init here."""
|
||||
# bias=False MUST match the model (fastvideo/models/dits/trackwan/track_encoder.py): TrackEncoder
|
||||
# omits the conv bias (a bias broadcasts to every latent cell and densifies the sparse track
|
||||
# signal -- load-bearing). Emitting bias tensors breaks strict loading with
|
||||
# "track_encoder.proj.bias not found in custom model state dict", so build bias-free and write
|
||||
# only the two weight tensors the model actually has.
|
||||
temporal_conv = nn.Conv3d(ID_DIM, TRACK_CHANNELS, kernel_size=(VAE_T_COMP, 1, 1),
|
||||
stride=(VAE_T_COMP, 1, 1), bias=False)
|
||||
proj = nn.Conv3d(TRACK_CHANNELS, TRACK_CHANNELS, kernel_size=1, bias=False) # default init (NOT zero) -> breaks the deadlock
|
||||
return {
|
||||
"track_encoder.temporal_conv.weight": temporal_conv.weight.detach().clone(),
|
||||
"track_encoder.proj.weight": proj.weight.detach().clone(),
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
global ID_DIM
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--base",
|
||||
required=True,
|
||||
help="Base Wan I2V (e.g. Wan2.1-Fun-1.3B-InP) or T2V diffusers model dir (with transformer/).")
|
||||
p.add_argument("--out", required=True, help="Output model dir for the WanTrack init.")
|
||||
p.add_argument("--id-dim", type=int, default=ID_DIM,
|
||||
help="sinusoidal track-id posemb dim (MotionStream d; 64 for d64 init, default 128)")
|
||||
args = p.parse_args()
|
||||
|
||||
ID_DIM = args.id_dim
|
||||
TRACK_CONFIG["id_dim"] = args.id_dim
|
||||
|
||||
base = Path(args.base)
|
||||
out = Path(args.out)
|
||||
if not (base / "transformer" / "config.json").exists():
|
||||
raise FileNotFoundError(f"{base}/transformer/config.json not found")
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 1) Symlink every top-level entry except transformer/ (vae, text_encoder, ...).
|
||||
for entry in os.listdir(base):
|
||||
if entry == "transformer":
|
||||
continue
|
||||
src = (base / entry).resolve()
|
||||
dst = out / entry
|
||||
if dst.is_symlink() or dst.exists():
|
||||
if dst.is_dir() and not dst.is_symlink():
|
||||
shutil.rmtree(dst)
|
||||
else:
|
||||
dst.unlink()
|
||||
os.symlink(src, dst)
|
||||
|
||||
# 2) Convert transformer/.
|
||||
tdir = out / "transformer"
|
||||
tdir.mkdir(exist_ok=True)
|
||||
cfg = json.loads((base / "transformer" / "config.json").read_text())
|
||||
base_in = int(cfg["in_channels"])
|
||||
cfg["in_channels"] = NEW_IN_CHANNELS
|
||||
cfg["track_config"] = TRACK_CONFIG
|
||||
(tdir / "config.json").write_text(json.dumps(cfg, indent=2))
|
||||
|
||||
sf_files = sorted((base / "transformer").glob("*.safetensors"))
|
||||
if not sf_files:
|
||||
raise FileNotFoundError(f"No transformer safetensors under {base}/transformer")
|
||||
state: dict[str, torch.Tensor] = {}
|
||||
for sf in sf_files:
|
||||
state.update(load_file(str(sf))) # merge shards if the base is sharded
|
||||
|
||||
pe_key = "patch_embedding.weight"
|
||||
if pe_key not in state:
|
||||
cands = [k for k in state if "patch_embed" in k and k.endswith(".weight")]
|
||||
if len(cands) != 1:
|
||||
raise KeyError(f"Could not find patch_embedding weight; candidates={cands}")
|
||||
pe_key = cands[0]
|
||||
w = state[pe_key] # [out, base_in, 1, 2, 2]
|
||||
if w.shape[1] != base_in:
|
||||
raise ValueError(f"{pe_key} in_ch {w.shape[1]} != config in_channels {base_in}")
|
||||
new_w = torch.zeros((w.shape[0], NEW_IN_CHANNELS, *w.shape[2:]), dtype=w.dtype)
|
||||
new_w[:, :base_in] = w # pretrained channels first; new channels zero
|
||||
state[pe_key] = new_w
|
||||
print(f"[convert] padded {pe_key}: {tuple(w.shape)} -> {tuple(new_w.shape)}")
|
||||
|
||||
te_state = build_track_encoder_state()
|
||||
for k, v in te_state.items():
|
||||
state[k] = v.to(w.dtype)
|
||||
print(f"[convert] added {len(te_state)} track_encoder tensors")
|
||||
|
||||
save_file(state, str(tdir / "diffusion_pytorch_model.safetensors"), metadata={"format": "pt"})
|
||||
print(f"[convert] wrote {tdir/'diffusion_pytorch_model.safetensors'} ({len(state)} keys)")
|
||||
print(f"[convert] done -> {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,245 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Build a WanTrack init checkpoint from a base Wan diffusers model, with several
|
||||
strategies for the patch_embedding track-channel slot and the track_encoder.
|
||||
|
||||
Modes for patch_embedding[:, base_in:] (the new 16 in-channels):
|
||||
--pe-init zero : zero-pad (original stage-1 recipe)
|
||||
--pe-init random : normal-init with std matched to base first-N channels
|
||||
|
||||
Modes for track_encoder.{proj, temporal_conv}:
|
||||
--track-src default : default Conv init (used by original convert)
|
||||
--track-src <ckpt.safetensors> : lift both convs from an existing WanTrack safetensors
|
||||
(e.g. a trained 1.3B ckpt). Bias is copied when present
|
||||
and its shape matches.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base /mnt/lustre/vlm-s4duan/models/Wan2.1-I2V-14B-720P-Diffusers \
|
||||
--out /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_random_init \
|
||||
--id-dim 64 --pe-init random
|
||||
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base /mnt/lustre/vlm-s4duan/models/Wan2.1-I2V-14B-720P-Diffusers \
|
||||
--out /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_partial_merged \
|
||||
--id-dim 64 --pe-init random \
|
||||
--track-src /mnt/lustre/vlm-s4duan/exports/synth_stage2_paperLR_ckpt600/transformer/model.safetensors
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from torch import nn
|
||||
|
||||
NEW_IN_CHANNELS = 52
|
||||
TRACK_CHANNELS = 16
|
||||
VAE_T_COMP = 4
|
||||
DEFAULT_ID_DIM = 128
|
||||
|
||||
|
||||
def build_track_encoder_default(id_dim: int, dtype: torch.dtype, use_bias: bool = False) -> dict[str, torch.Tensor]:
|
||||
"""Default (non-zero) Conv init for track_encoder.proj + temporal_conv."""
|
||||
temporal_conv = nn.Conv3d(id_dim, TRACK_CHANNELS,
|
||||
kernel_size=(VAE_T_COMP, 1, 1), stride=(VAE_T_COMP, 1, 1),
|
||||
bias=use_bias)
|
||||
proj = nn.Conv3d(TRACK_CHANNELS, TRACK_CHANNELS, kernel_size=1, bias=use_bias)
|
||||
d = {
|
||||
"track_encoder.temporal_conv.weight": temporal_conv.weight.detach().to(dtype).clone(),
|
||||
"track_encoder.proj.weight": proj.weight.detach().to(dtype).clone(),
|
||||
}
|
||||
if use_bias:
|
||||
d["track_encoder.temporal_conv.bias"] = temporal_conv.bias.detach().to(dtype).clone()
|
||||
d["track_encoder.proj.bias"] = proj.bias.detach().to(dtype).clone()
|
||||
return d
|
||||
|
||||
|
||||
def load_src_state(src_path: str) -> dict[str, torch.Tensor]:
|
||||
"""Load a source state dict from a .safetensors FILE or a diffusers model DIR.
|
||||
|
||||
A 14B export is sharded across several *.safetensors, so accept a directory (either the
|
||||
model root or its transformer/ subdir) and merge every shard.
|
||||
"""
|
||||
p = Path(src_path)
|
||||
if p.is_file():
|
||||
return load_file(str(p))
|
||||
tdir = p / "transformer" if (p / "transformer").is_dir() else p
|
||||
shards = sorted(tdir.glob("*.safetensors"))
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"no *.safetensors under {tdir}")
|
||||
state: dict[str, torch.Tensor] = {}
|
||||
for s in shards:
|
||||
state.update(load_file(str(s)))
|
||||
print(f"[src] loaded {len(state)} tensors from {len(shards)} shard(s) in {tdir}")
|
||||
return state
|
||||
|
||||
|
||||
def lift_track_encoder_from_src(src_path: str, id_dim: int, dtype: torch.dtype) -> dict[str, torch.Tensor]:
|
||||
"""Load track_encoder.{proj, temporal_conv}.{weight, bias} from an existing WanTrack ckpt.
|
||||
|
||||
Verifies shape compatibility. If bias is missing in the source we skip it; if the
|
||||
source shapes don't match the expected shapes we fall back to default init for that
|
||||
key and print a warning.
|
||||
"""
|
||||
src_state = load_src_state(src_path)
|
||||
expected_shapes = {
|
||||
"track_encoder.temporal_conv.weight": (TRACK_CHANNELS, id_dim, VAE_T_COMP, 1, 1),
|
||||
"track_encoder.proj.weight": (TRACK_CHANNELS, TRACK_CHANNELS, 1, 1, 1),
|
||||
"track_encoder.temporal_conv.bias": (TRACK_CHANNELS,),
|
||||
"track_encoder.proj.bias": (TRACK_CHANNELS,),
|
||||
}
|
||||
picked: dict[str, torch.Tensor] = {}
|
||||
for k, want in expected_shapes.items():
|
||||
if k not in src_state:
|
||||
continue
|
||||
got = tuple(src_state[k].shape)
|
||||
if got != want:
|
||||
print(f"[track-src] SHAPE MISMATCH on {k}: src={got} vs expected={want}; SKIPPING (will need fallback)")
|
||||
continue
|
||||
picked[k] = src_state[k].detach().to(dtype).clone()
|
||||
print(f"[track-src] lifted {k} shape={got}")
|
||||
# If either weight is missing, we still need to keep the model loadable.
|
||||
fallback = build_track_encoder_default(id_dim, dtype, use_bias=False)
|
||||
for k, v in fallback.items():
|
||||
if k not in picked:
|
||||
print(f"[track-src] {k} not in source -> default init")
|
||||
picked[k] = v
|
||||
return picked
|
||||
|
||||
|
||||
def std_of_first_n(w: torch.Tensor, n: int) -> float:
|
||||
"""Std of the pretrained first-N in-channels slice of patch_embedding.weight."""
|
||||
with torch.no_grad():
|
||||
s = w[:, :n].float().std().item()
|
||||
return float(s)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--base", required=True, help="Base Wan diffusers model dir.")
|
||||
p.add_argument("--out", required=True, help="Output diffusers model dir.")
|
||||
p.add_argument("--id-dim", type=int, default=DEFAULT_ID_DIM)
|
||||
p.add_argument("--pe-init", choices=("zero", "random"), default="zero",
|
||||
help="patch_embedding init strategy for the ADDED track channels [:, base_in:]. "
|
||||
"Base pretrained channels [:, :base_in] are always preserved.")
|
||||
p.add_argument("--pe-random-seed", type=int, default=1234)
|
||||
p.add_argument("--track-src", default=None,
|
||||
help="Path to a .safetensors (or diffusers model dir) with track_encoder.* to lift.")
|
||||
p.add_argument("--pe-src", default=None,
|
||||
help="Path to a .safetensors (or diffusers model dir) to lift patch_embedding.weight[:, base_in:] "
|
||||
"from — i.e. the TRACK SLOT only; the pretrained [:, :base_in] channels always come from "
|
||||
"--base. Use together with --track-src pointing at the SAME source so the encoder and the "
|
||||
"track slot stay CO-ADAPTED (lifting the encoder alone is worse than random init).")
|
||||
p.add_argument("--use-bias-defaults", action="store_true",
|
||||
help="When --track-src is not set, build track_encoder convs WITH bias (matches TRACKWAN_TRACK_BIAS=1 training).")
|
||||
args = p.parse_args()
|
||||
|
||||
id_dim = int(args.id_dim)
|
||||
TRACK_CONFIG = {
|
||||
"id_dim": id_dim,
|
||||
"track_channels": TRACK_CHANNELS,
|
||||
"vae_spatial_compression": 8,
|
||||
"vae_temporal_compression": VAE_T_COMP,
|
||||
"max_track_id": 100_000,
|
||||
"zero_init_head": False,
|
||||
}
|
||||
|
||||
base = Path(args.base)
|
||||
out = Path(args.out)
|
||||
if not (base / "transformer" / "config.json").exists():
|
||||
raise FileNotFoundError(f"{base}/transformer/config.json not found")
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 1) Symlink every top-level entry except transformer/
|
||||
for entry in os.listdir(base):
|
||||
if entry == "transformer":
|
||||
continue
|
||||
src_p = (base / entry).resolve()
|
||||
dst = out / entry
|
||||
if dst.is_symlink() or dst.exists():
|
||||
if dst.is_dir() and not dst.is_symlink():
|
||||
shutil.rmtree(dst)
|
||||
else:
|
||||
dst.unlink()
|
||||
os.symlink(src_p, dst)
|
||||
|
||||
# 2) Rewrite transformer/config.json
|
||||
tdir = out / "transformer"
|
||||
tdir.mkdir(exist_ok=True)
|
||||
cfg = json.loads((base / "transformer" / "config.json").read_text())
|
||||
base_in = int(cfg["in_channels"])
|
||||
cfg["in_channels"] = NEW_IN_CHANNELS
|
||||
cfg["track_config"] = TRACK_CONFIG
|
||||
(tdir / "config.json").write_text(json.dumps(cfg, indent=2))
|
||||
|
||||
# 3) Load full base transformer state
|
||||
sf_files = sorted((base / "transformer").glob("*.safetensors"))
|
||||
if not sf_files:
|
||||
raise FileNotFoundError(f"No safetensors under {base}/transformer")
|
||||
state: dict[str, torch.Tensor] = {}
|
||||
for sf in sf_files:
|
||||
state.update(load_file(str(sf)))
|
||||
dtype = next(iter(state.values())).dtype
|
||||
|
||||
# 4) Grow patch_embedding.weight from base_in -> 52 input channels
|
||||
pe_key = "patch_embedding.weight"
|
||||
if pe_key not in state:
|
||||
raise KeyError(f"{pe_key} not found in base transformer")
|
||||
w = state[pe_key] # [hidden, base_in, 1, 2, 2]
|
||||
if w.shape[1] != base_in:
|
||||
raise ValueError(f"{pe_key} in_ch {w.shape[1]} != config in_channels {base_in}")
|
||||
new_w = torch.empty((w.shape[0], NEW_IN_CHANNELS, *w.shape[2:]), dtype=w.dtype)
|
||||
new_w[:, :base_in] = w # pretrained channels first
|
||||
if args.pe_init == "zero":
|
||||
new_w[:, base_in:] = 0
|
||||
print(f"[convert] pe-init=zero -> new channels are zero")
|
||||
else:
|
||||
std = std_of_first_n(w, base_in)
|
||||
g = torch.Generator().manual_seed(int(args.pe_random_seed))
|
||||
noise = torch.empty_like(new_w[:, base_in:], dtype=torch.float32)
|
||||
noise.normal_(mean=0.0, std=std, generator=g)
|
||||
new_w[:, base_in:] = noise.to(w.dtype)
|
||||
print(f"[convert] pe-init=random N(0, {std:.5f}) matching first-{base_in} std")
|
||||
# 4b) Optionally overwrite ONLY the track slot [:, base_in:] from a trained source.
|
||||
# This is the other half of the "merged" recipe: the track slot and the track_encoder were
|
||||
# co-adapted during the overfit, so they must be lifted TOGETHER (--pe-src + --track-src from
|
||||
# the same ckpt). Lifting the encoder alone leaves it talking to a random projection.
|
||||
if args.pe_src:
|
||||
src_pe_state = load_src_state(args.pe_src)
|
||||
if pe_key not in src_pe_state:
|
||||
raise KeyError(f"{pe_key} not found in --pe-src {args.pe_src}")
|
||||
src_pe = src_pe_state[pe_key]
|
||||
want = (w.shape[0], NEW_IN_CHANNELS, *w.shape[2:])
|
||||
if tuple(src_pe.shape) != want:
|
||||
raise ValueError(f"--pe-src {pe_key} shape {tuple(src_pe.shape)} != expected {want}")
|
||||
slot = src_pe[:, base_in:].to(w.dtype)
|
||||
new_w[:, base_in:] = slot
|
||||
print(f"[pe-src] lifted {pe_key}[:, {base_in}:] from source "
|
||||
f"(std={slot.float().std().item():.5f}, absmax={slot.float().abs().max().item():.5f})")
|
||||
|
||||
state[pe_key] = new_w
|
||||
print(f"[convert] patch_embedding.weight: {tuple(w.shape)} -> {tuple(new_w.shape)}")
|
||||
with torch.no_grad():
|
||||
print(f"[convert] pe[:, :{base_in}] std={new_w[:, :base_in].float().std().item():.5f} (pretrained) | "
|
||||
f"pe[:, {base_in}:] std={new_w[:, base_in:].float().std().item():.5f} (track slot)")
|
||||
|
||||
# 5) Add track_encoder
|
||||
if args.track_src:
|
||||
te = lift_track_encoder_from_src(args.track_src, id_dim, dtype)
|
||||
else:
|
||||
te = build_track_encoder_default(id_dim, dtype, use_bias=args.use_bias_defaults)
|
||||
for k, v in te.items():
|
||||
state[k] = v
|
||||
print(f"[convert] added {len(te)} track_encoder tensors")
|
||||
|
||||
save_file(state, str(tdir / "diffusion_pytorch_model.safetensors"), metadata={"format": "pt"})
|
||||
print(f"[convert] wrote {tdir/'diffusion_pytorch_model.safetensors'} ({len(state)} keys)")
|
||||
print(f"[convert] done -> {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,132 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 2: produce VAE round-trip videos for CoTracker.
|
||||
|
||||
Encode each source video through the FastVideo WanVAE (use_feature_cache=False,
|
||||
with causal-boundary fix) then immediately decode it back. The resulting videos
|
||||
differ slightly from the originals (compression artifacts, mild color shift) but
|
||||
are exactly what Stage 5 (preprocess_to_parquet) will store as latents and what
|
||||
the validation callback will decode for reference. CoTracker tracks extracted from
|
||||
these round-trip videos therefore align with training latents and the validation
|
||||
reference display.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/decode_roundtrip_videos.py \\
|
||||
--data-dir /home/hal-kevin/data/motion-physics \\
|
||||
--vae-path /home/hal-kevin/models/trackwan_1.3b_i2v_control_init/vae
|
||||
|
||||
# Re-run specific indices only
|
||||
python data_pipeline/decode_roundtrip_videos.py \\
|
||||
--data-dir /home/hal-kevin/data/motion-physics \\
|
||||
--vae-path /home/hal-kevin/models/trackwan_1.3b_i2v_control_init/vae \\
|
||||
--index 4 7 12
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
from fastvideo.dataset.transform import center_crop_th_tw, resize
|
||||
from fastvideo.models.vaes.wanvae import AutoencoderKLWan
|
||||
|
||||
TARGET_H, TARGET_W = 480, 832
|
||||
NUM_FRAMES = 121
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--data-dir", type=Path, required=True, help="Dataset root (contains videos/, etc.).")
|
||||
p.add_argument("--vae-path", type=Path, required=True, help="Path to the VAE directory (contains diffusion_pytorch_model.safetensors).")
|
||||
p.add_argument("--video-subdir", type=str, default="videos", help="Input video subdirectory.")
|
||||
p.add_argument("--out-subdir", type=str, default="roundtrip_videos", help="Output subdirectory.")
|
||||
p.add_argument("--num-frames", type=int, default=NUM_FRAMES, help="Number of frames per video.")
|
||||
p.add_argument("--height", type=int, default=TARGET_H)
|
||||
p.add_argument("--width", type=int, default=TARGET_W)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--device", type=str, default="cuda")
|
||||
p.add_argument("--index", type=int, nargs="+", default=None, metavar="IDX",
|
||||
help="Process only these video indices (e.g. --index 4 7 12).")
|
||||
p.add_argument("--limit", type=int, default=None, help="Process only first N videos (smoke test).")
|
||||
p.add_argument("--force", action="store_true", help="Re-encode even if output already exists.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def load_vae(vae_path: Path, device: str) -> AutoencoderKLWan:
|
||||
config = WanVAEConfig(use_feature_cache=False)
|
||||
vae = AutoencoderKLWan(config).to(device).eval()
|
||||
weights = safetensors_load_file(str(vae_path / "diffusion_pytorch_model.safetensors"))
|
||||
vae.load_state_dict(weights, strict=True)
|
||||
return vae
|
||||
|
||||
|
||||
def load_video(path: Path, num_frames: int, height: int, width: int) -> torch.Tensor:
|
||||
"""Return pixel tensor [1, C, T, H, W] in [-1, 1]."""
|
||||
reader = imageio.get_reader(str(path))
|
||||
frames = [reader.get_data(i) for i in range(num_frames)]
|
||||
reader.close()
|
||||
clip = torch.from_numpy(np.stack(frames)).permute(0, 3, 1, 2).float() / 255.0
|
||||
clip = center_crop_th_tw(clip, height, width, top_crop=False)
|
||||
clip = resize(clip, (height, width), interpolation_mode="bilinear")
|
||||
pixel = (clip * 2.0 - 1.0).permute(1, 0, 2, 3).unsqueeze(0) # [1,C,T,H,W]
|
||||
return pixel
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def roundtrip(vae: AutoencoderKLWan, pixel: torch.Tensor, device: str) -> np.ndarray:
|
||||
"""Return decoded frames as uint8 numpy [T, H, W, C]."""
|
||||
pixel = pixel.to(device)
|
||||
with torch.autocast(device, dtype=torch.float32):
|
||||
latent = vae.encode(pixel).mean
|
||||
decoded = vae.decode(latent)
|
||||
frames = decoded[0].permute(1, 2, 3, 0).float().cpu()
|
||||
frames = ((frames / 2 + 0.5).clamp(0, 1) * 255).byte().numpy()
|
||||
return frames
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
device = args.device
|
||||
|
||||
videos_dir = args.data_dir / args.video_subdir
|
||||
out_dir = args.data_dir / args.out_subdir
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
videos = sorted(videos_dir.glob("*.mp4"))
|
||||
if args.index is not None:
|
||||
wanted = {f"vid_{i:06d}.mp4" for i in args.index}
|
||||
videos = [v for v in videos if v.name in wanted]
|
||||
if args.limit is not None:
|
||||
videos = videos[:args.limit]
|
||||
if not videos:
|
||||
print(f"[roundtrip] no videos found in {videos_dir}", flush=True)
|
||||
return
|
||||
|
||||
print(f"[roundtrip] loading VAE from {args.vae_path} ...", flush=True)
|
||||
vae = load_vae(args.vae_path, device)
|
||||
|
||||
print(f"[roundtrip] {len(videos)} videos → {out_dir}", flush=True)
|
||||
for k, vpath in enumerate(videos, 1):
|
||||
out_path = out_dir / vpath.name
|
||||
if out_path.exists() and not args.force:
|
||||
print(f"[roundtrip] [{k}/{len(videos)}] {vpath.name} already exists, skipping", flush=True)
|
||||
continue
|
||||
|
||||
pixel = load_video(vpath, args.num_frames, args.height, args.width)
|
||||
frames = roundtrip(vae, pixel, device)
|
||||
imageio.mimsave(str(out_path), frames, fps=args.fps, macro_block_size=1)
|
||||
print(f"[roundtrip] [{k}/{len(videos)}] {vpath.name} → {out_path.name} "
|
||||
f"shape={frames.shape}", flush=True)
|
||||
|
||||
print("[roundtrip] done.", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,162 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Decode FastVideo Wan2.2-Syn-121x704x1280_32k parquet latents to 480x832 mp4s.
|
||||
|
||||
The HF dataset stores Wan 2.2 5B VAE latents (48-channel, 16x spatial, 4x
|
||||
temporal). We can't feed these to our Wan 2.1 WanTrack model directly, so we
|
||||
decode → pixels → downsize to our training resolution 480x832 → save mp4.
|
||||
Downstream (SAM + CoTracker + i2v_track preprocess) then treats these as if they
|
||||
were raw synth mp4s.
|
||||
|
||||
Idempotent + resumable: skips samples whose mp4 already exists.
|
||||
|
||||
Slurm sharded via ``--num-shards``/``--shard``. Each worker processes
|
||||
``parquets[shard::num_shards]``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _resize_frames(frames: np.ndarray, out_h: int, out_w: int) -> np.ndarray:
|
||||
"""Center-crop-then-resize a (T,H,W,3) uint8 stack to (T,out_h,out_w,3)."""
|
||||
T, H, W, _ = frames.shape
|
||||
src_ar = W / H
|
||||
dst_ar = out_w / out_h
|
||||
if src_ar > dst_ar:
|
||||
new_w = int(round(H * dst_ar))
|
||||
x0 = (W - new_w) // 2
|
||||
cropped = frames[:, :, x0:x0 + new_w, :]
|
||||
else:
|
||||
new_h = int(round(W / dst_ar))
|
||||
y0 = (H - new_h) // 2
|
||||
cropped = frames[:, y0:y0 + new_h, :, :]
|
||||
out = np.empty((T, out_h, out_w, 3), dtype=np.uint8)
|
||||
for t in range(T):
|
||||
out[t] = np.asarray(Image.fromarray(cropped[t]).resize((out_w, out_h), Image.BILINEAR))
|
||||
return out
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_row(vae: AutoencoderKLWan, latent_bytes: bytes, latent_shape: list[int],
|
||||
device: torch.device, dtype: torch.dtype) -> np.ndarray:
|
||||
"""Decode one row's Wan 2.2 5B latent to a [T,704,1280,3] uint8 pixel stack."""
|
||||
lat = np.frombuffer(latent_bytes, dtype=np.float32).reshape(latent_shape)
|
||||
lat = torch.from_numpy(lat).to(device=device, dtype=dtype).unsqueeze(0) # [1,C,T,H,W]
|
||||
# unnormalize per-channel
|
||||
m = torch.tensor(vae.config.latents_mean, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
|
||||
s = torch.tensor(vae.config.latents_std, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
|
||||
lat = lat * s + m
|
||||
pix = vae.decode(lat, return_dict=False)[0] # [1,3,T_out,H*16,W*16] in [-1,1]
|
||||
pix = pix.clamp(-1, 1).float()
|
||||
pix = ((pix + 1.0) * 127.5).round().to(torch.uint8)
|
||||
pix = pix[0].permute(1, 2, 3, 0).cpu().numpy() # [T,H,W,3]
|
||||
return pix
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--parquet-dir", required=True, help="dir containing train/Part_*/latents_chunk_*.parquet")
|
||||
p.add_argument("--vae-dir", required=True, help="Wan 2.2 5B VAE dir")
|
||||
p.add_argument("--out-dir", required=True, help="output root; writes videos/vid_*.mp4 + meta/*.json")
|
||||
p.add_argument("--out-h", type=int, default=480)
|
||||
p.add_argument("--out-w", type=int, default=832)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--shard", type=int, default=0)
|
||||
p.add_argument("--num-shards", type=int, default=1)
|
||||
p.add_argument("--dtype", default="bfloat16", choices=["float32", "bfloat16", "float16"])
|
||||
p.add_argument("--limit-rows", type=int, default=None, help="cap rows for smoke test")
|
||||
args = p.parse_args()
|
||||
|
||||
device = torch.device("cuda")
|
||||
dtype = getattr(torch, args.dtype)
|
||||
|
||||
vids_dir = Path(args.out_dir) / "videos"
|
||||
meta_dir = Path(args.out_dir) / "meta"
|
||||
vids_dir.mkdir(parents=True, exist_ok=True)
|
||||
meta_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
files = sorted(glob.glob(os.path.join(args.parquet_dir, "**", "*.parquet"), recursive=True))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"no *.parquet under {args.parquet_dir}")
|
||||
my_files = files[args.shard::args.num_shards]
|
||||
print(f"[dec {args.shard}/{args.num_shards}] {len(my_files)} parquets", flush=True)
|
||||
|
||||
vae = AutoencoderKLWan.from_pretrained(args.vae_dir, torch_dtype=dtype).to(device).eval()
|
||||
|
||||
n_ok = 0
|
||||
n_skip = 0
|
||||
n_err = 0
|
||||
t0 = time.time()
|
||||
row_budget = args.limit_rows or 10**12
|
||||
processed_rows = 0
|
||||
for fk, f in enumerate(my_files):
|
||||
try:
|
||||
table = pq.read_table(f)
|
||||
except Exception as e: # skip unreadable parquets
|
||||
print(f"[dec {args.shard}] parquet read fail {f}: {e}", flush=True)
|
||||
continue
|
||||
# Dataset's `id` field only encodes chunk index (not Part), so the same id
|
||||
# appears in many Parts. Prefix Part directory so vid_ids are unique.
|
||||
part_name = Path(f).parent.name # e.g. "Part_57"
|
||||
rows = table.to_pylist()
|
||||
for r in rows:
|
||||
if processed_rows >= row_budget:
|
||||
break
|
||||
processed_rows += 1
|
||||
row_key = str(r.get("id") or r.get("file_name") or f"{fk:05d}_{n_ok:04d}")
|
||||
vid_id = f"{part_name}_{row_key}"
|
||||
vid_path = vids_dir / f"{vid_id}.mp4"
|
||||
meta_path = meta_dir / f"{vid_id}.json"
|
||||
if vid_path.exists() and meta_path.exists():
|
||||
n_skip += 1
|
||||
continue
|
||||
try:
|
||||
pix = decode_row(vae, r["vae_latent_bytes"], list(r["vae_latent_shape"]),
|
||||
device=device, dtype=dtype)
|
||||
resized = _resize_frames(pix, args.out_h, args.out_w) # [T, out_h, out_w, 3]
|
||||
# atomic mp4 write (imageio infers format from extension → keep .mp4)
|
||||
tmp = vid_path.parent / f".{vid_path.name}.tmp.mp4"
|
||||
imageio.mimsave(str(tmp), resized, fps=args.fps, codec="libx264",
|
||||
quality=7, macro_block_size=1)
|
||||
tmp.rename(vid_path)
|
||||
meta = {
|
||||
"path": vid_path.name,
|
||||
"id": vid_id,
|
||||
"cap": [str(r.get("caption") or "")],
|
||||
"fps": float(args.fps),
|
||||
"num_frames": int(resized.shape[0]),
|
||||
"resolution": [int(args.out_h), int(args.out_w)],
|
||||
"src_resolution": [int(r.get("height", 704)), int(r.get("width", 1280))],
|
||||
"src_dataset": "FastVideo/Wan2.2-Syn-121x704x1280_32k",
|
||||
}
|
||||
meta_path.write_text(json.dumps(meta))
|
||||
n_ok += 1
|
||||
if n_ok % 20 == 0:
|
||||
dt = time.time() - t0
|
||||
rate = n_ok / max(dt, 1e-9)
|
||||
print(f"[dec {args.shard}] ok={n_ok} skip={n_skip} err={n_err} "
|
||||
f"rate={rate:.2f}/s", flush=True)
|
||||
except Exception as e: # skip a bad row, keep going
|
||||
n_err += 1
|
||||
print(f"[dec {args.shard}] row {vid_id} FAILED: {e}", flush=True)
|
||||
if processed_rows >= row_budget:
|
||||
break
|
||||
|
||||
print(f"[dec {args.shard}] DONE ok={n_ok} skip={n_skip} err={n_err} "
|
||||
f"in {(time.time()-t0)/60:.1f}min", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Diagnose the first-frame I2V conditioning + visualize synthetic tracks.
|
||||
|
||||
(1) First-frame conditioning bug: the preprocessing stored ``first_frame_latent``
|
||||
using vae.scaling_factor/shift_factor (absent for Wan -> effectively RAW),
|
||||
while the model normalizes latents with per-channel latents_mean/std. This
|
||||
script decodes the stored conditioning under both interpretations and compares
|
||||
to the GT first frame to localize the space mismatch, and shows the current
|
||||
model's generated first frame.
|
||||
|
||||
(2) Synthetic-track previews: overlays authored controls (pan/zoom/drag/...) on the
|
||||
real first frame so you can see what the control signals look like.
|
||||
|
||||
Outputs PNGs to --out (default research_log/figures/).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def _mse(a, b):
|
||||
a = a[:min(len(a), len(b))].astype(np.float32)
|
||||
b = b[:len(a)].astype(np.float32)
|
||||
return float(((a - b) ** 2).mean())
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _decode_frame0(model, latent_bcthw, *, already_normalized: bool) -> np.ndarray:
|
||||
"""latent [1,16,T,H,W] -> frame0 uint8 (H,W,3). If raw, normalize first."""
|
||||
from fastvideo.training.training_utils import normalize_dit_input
|
||||
x = latent_bcthw.to(model.device, torch.bfloat16)
|
||||
if not already_normalized:
|
||||
x = normalize_dit_input("wan", x, model.vae)
|
||||
px = model.decode_latents(x.permute(0, 2, 1, 3, 4))[0] # [3,T,H,W] in [0,1]
|
||||
f0 = (px[:, 0].clamp(0, 1).float().cpu().numpy() * 255).astype(np.uint8)
|
||||
return np.transpose(f0, (1, 2, 0)) # H,W,3
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track/combined_parquet_dataset")
|
||||
p.add_argument("--out", default="research_log/figures")
|
||||
p.add_argument("--clip", type=int, default=0)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
s = twi.load_conditioning_from_parquet(args.data, [args.clip], text_len)[0]
|
||||
|
||||
vae_latent = s["vae_latent"] # raw GT latents
|
||||
ff = s["first_frame_latent"] # stored conditioning latent
|
||||
|
||||
# GT first frame: vae_latent is raw -> normalize -> decode.
|
||||
gt0 = _decode_frame0(model, vae_latent, already_normalized=False)
|
||||
|
||||
# Stored conditioning decoded TWO ways to find its true space:
|
||||
# (a) "model's view": treat stored as already-normalized (decode_latents denorms it).
|
||||
# This is effectively what the model was conditioned on.
|
||||
cond_as_norm0 = _decode_frame0(model, ff, already_normalized=True)
|
||||
# (b) "as raw": treat stored as raw -> normalize -> decode.
|
||||
cond_as_raw0 = _decode_frame0(model, ff, already_normalized=False)
|
||||
|
||||
# Current model's generated first frame (GT tracks).
|
||||
Tpx = (vae_latent.shape[2] - 1) * int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio) + 1
|
||||
lat = twi.generate(model, first_frame_latent=ff, text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=s["track_points"][:, :Tpx], track_visibility=s["track_visibility"][:, :Tpx],
|
||||
num_steps=args.steps, seed=1000)
|
||||
gen0 = twi.decode_to_pixels(model, lat)[0]
|
||||
|
||||
# Save individual + side-by-side.
|
||||
panels = {"1_gt_first_frame": gt0,
|
||||
"2_cond_as_normalized_MODELVIEW": cond_as_norm0,
|
||||
"3_cond_as_raw_then_normalized": cond_as_raw0,
|
||||
"4_generated_first_frame": gen0}
|
||||
for name, img in panels.items():
|
||||
imageio.imwrite(os.path.join(args.out, f"firstframe_clip{args.clip}_{name}.png"), img)
|
||||
strip = np.concatenate([gt0, cond_as_norm0, cond_as_raw0, gen0], axis=1)
|
||||
imageio.imwrite(os.path.join(args.out, f"firstframe_clip{args.clip}_compare.png"), strip)
|
||||
|
||||
print("=== first-frame diagnosis (clip %d) ===" % args.clip)
|
||||
print(f" MSE(GT, cond-as-normalized / MODEL'S VIEW) = {_mse(gt0, cond_as_norm0):8.1f}")
|
||||
print(f" MSE(GT, cond-as-raw-then-normalized) = {_mse(gt0, cond_as_raw0):8.1f}")
|
||||
print(f" MSE(GT, generated first frame) = {_mse(gt0, gen0):8.1f}")
|
||||
print(" -> if 'as-raw' MSE << 'model's view' MSE, the stored conditioning is RAW")
|
||||
print(" and the model was fed a mis-normalized first frame (the bug).")
|
||||
|
||||
# ---- synthetic track previews on the real first frame ----
|
||||
for name in ["pan_right", "zoom_in", "rotate_cw", "drag_center_right", "swirl"]:
|
||||
tr, vis = st.preset(name, Tpx, 50, strength=0.25)
|
||||
H, W, _ = gt0.shape
|
||||
ov = st._overlay_preview(gt0, st.to_pixel(tr, H, W), vis, stride=3)
|
||||
imageio.imwrite(os.path.join(args.out, f"trackpreview_{name}.png"), ov)
|
||||
print(f"\n[viz] wrote first-frame panels + synthetic track previews to {args.out}/")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,47 @@
|
||||
#!/usr/bin/env bash
|
||||
# Self-healing OpenVidHD download: HF unauthenticated throttles ~per-200GB then
|
||||
# resets after a cooldown. Run N idempotent download shards; watchdog detects a
|
||||
# stall (no zip growth) and kills+cooldowns+relaunches (resume) until all 98 parts
|
||||
# done. Runs as ONE srun on the node so local kill works. Honors $HF_TOKEN if set.
|
||||
set +e
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
export HF_HOME=$WORK/.hf
|
||||
cd "$WORK/FastVideo"; source .venv/bin/activate
|
||||
NSHARD=${NSHARD:-4}
|
||||
DC(){ ls "$WORK"/openvid_1m/_extracted/*.done 2>/dev/null | wc -l; }
|
||||
ZB(){ find "$WORK"/openvid_1m/_zips -type f -printf '%s\n' 2>/dev/null | awk '{s+=$1}END{print s+0}'; }
|
||||
|
||||
pkill -f openvid_download_hd.py 2>/dev/null; sleep 3
|
||||
round=0
|
||||
while [ "$(DC)" -lt 98 ]; do
|
||||
round=$((round+1))
|
||||
echo "[$(date +%H:%M)] ROUND $round start parts=$(DC)/98 ${HF_TOKEN:+(token set)}"
|
||||
pids=()
|
||||
for s in $(seq 0 $((NSHARD-1))); do
|
||||
python data_pipeline/openvid_download_hd.py --shard "$s" --num-shards "$NSHARD" \
|
||||
--videos-dir "$WORK"/openvid_1m/videos --zip-dir "$WORK"/openvid_1m/_zips/"$s" \
|
||||
--only-list "$WORK"/openvid/OpenVidHD_filtered.txt &
|
||||
pids+=($!)
|
||||
done
|
||||
last=$(ZB); stall=0
|
||||
while :; do
|
||||
sleep 120
|
||||
[ "$(DC)" -ge 98 ] && break
|
||||
alive=0; for p in "${pids[@]}"; do kill -0 "$p" 2>/dev/null && alive=1; done
|
||||
[ "$alive" -eq 0 ] && { echo "[$(date +%H:%M)] round done naturally parts=$(DC)/98"; break; }
|
||||
now=$(ZB)
|
||||
if [ "$now" -le "$last" ]; then stall=$((stall+1)); else stall=0; fi
|
||||
last=$now
|
||||
echo "[$(date +%H:%M)] parts=$(DC)/98 zip=$((now/1000000000))GB stall=$stall"
|
||||
if [ "$stall" -ge 2 ]; then
|
||||
echo "[$(date +%H:%M)] STALL -> kill + cooldown"
|
||||
kill "${pids[@]}" 2>/dev/null; sleep 5; kill -9 "${pids[@]}" 2>/dev/null
|
||||
pkill -9 -f openvid_download_hd.py 2>/dev/null
|
||||
break
|
||||
fi
|
||||
done
|
||||
wait 2>/dev/null
|
||||
[ "$(DC)" -ge 98 ] && break
|
||||
echo "[$(date +%H:%M)] cooldown 180s (parts $(DC)/98)"; sleep 180
|
||||
done
|
||||
echo "ALL_PARTS_DONE $(DC)/98 videos=$(ls "$WORK"/openvid_1m/videos/*.mp4 2>/dev/null | wc -l)"
|
||||
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# DP concurrency sweep on ONE GPU: K workers sharing GPU 3, aggregate throughput.
|
||||
set +e
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
source "$WORK/FastVideo/.venv/bin/activate"
|
||||
export CUDA_HOME=/usr/local/cuda
|
||||
export PATH=$CUDA_HOME/bin:$PATH
|
||||
export TORCH_HOME=$WORK/.torch
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
V=$WORK/dp_test/videos; T=$WORK/dp_test/tracks; L=$WORK/dp_test/videos.txt
|
||||
mkdir -p "$V" "$T"
|
||||
if [ ! -f "$V/clip_24.mp4" ]; then
|
||||
for i in $(seq -w 1 24); do cp -f "$WORK/wan_overfit/data/videos/vid_000000.mp4" "$V/clip_$i.mp4"; done
|
||||
fi
|
||||
ls "$V"/*.mp4 > "$L"
|
||||
echo "test set: $(wc -l < "$L") videos on GPU 3 (121f 720p, grid 50)"
|
||||
for K in 1 2 3 4; do
|
||||
rm -f "$T"/*.npz
|
||||
echo "===================== K=$K procs/GPU ====================="
|
||||
pids=()
|
||||
for s in $(seq 0 $((K-1))); do
|
||||
CUDA_VISIBLE_DEVICES=3 python data_pipeline/extract_tracks_mp.py \
|
||||
--video-list "$L" --out-dir "$T" --gpus-per-node 1 \
|
||||
--shard "$s" --num-shards "$K" --fps 24 --num-frames 121 --grid-size 50 &
|
||||
pids+=($!)
|
||||
done
|
||||
wait "${pids[@]}"
|
||||
done
|
||||
echo "SWEEP_DONE"
|
||||
@@ -0,0 +1,224 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Inspect a track-conditioning dataset: are the clips diverse, and do the
|
||||
CoTracker tracks actually capture the motion?
|
||||
|
||||
Two concerns this surfaces:
|
||||
1. Near-duplicate clips (the frame-0 cluster can over-select the same scene).
|
||||
2. Weak control signal: CoTracker uses a frame-0 grid, so anything that ENTERS
|
||||
the frame later (or that the tracker loses) is never queried -> its motion is
|
||||
invisible to the tracks. A clip can look dynamic yet have near-static tracks.
|
||||
|
||||
``render`` (CPU; run on a node) writes, per clip:
|
||||
- <id>__overlay.mp4 the video with the CoTracker tracks drawn (dots + tails),
|
||||
- <id>__thumb.png first frame,
|
||||
and a ``stats.json`` with per-clip motion-coverage + a near-duplicate grouping.
|
||||
|
||||
``serve`` (no GPU) is a gradio gallery: clips sorted by motion (least first, so
|
||||
dead/duplicate clips float to the top), each with its overlay + stats.
|
||||
|
||||
# 1) render artifacts (node)
|
||||
srun --jobid=<job> --overlap --ntasks=1 .venv/bin/python data_pipeline/droid_dataset_dashboard.py render \
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200 --out <dash_dir>
|
||||
# 2) serve (login node ok)
|
||||
.venv/bin/python data_pipeline/droid_dataset_dashboard.py serve --dash-dir <dash_dir> --share
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
MOVE_THRESH_PX = 8.0 # a point "moves" if it travels > this from frame 0 (input px)
|
||||
VIS_THRESH = 0.5
|
||||
|
||||
|
||||
def _load_frames(path: str) -> np.ndarray:
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001
|
||||
import imageio.v2 as imageio
|
||||
rd = imageio.get_reader(path, format="ffmpeg")
|
||||
fr = np.stack([np.asarray(f) for f in rd], axis=0)
|
||||
rd.close()
|
||||
return fr
|
||||
|
||||
|
||||
def _clip_stats(tracks: np.ndarray, vis: np.ndarray) -> dict[str, Any]:
|
||||
"""Motion coverage from tracks [T,N,2] (px) + vis [T,N]."""
|
||||
T, N, _ = tracks.shape
|
||||
v = vis > VIS_THRESH
|
||||
# displacement of each point from its frame-0 position, only where visible
|
||||
disp = np.sqrt(((tracks - tracks[0:1])**2).sum(-1)) # [T,N]
|
||||
max_disp = np.where(v, disp, 0.0).max(0) # [N] per-point peak travel
|
||||
moving = max_disp > MOVE_THRESH_PX
|
||||
return {
|
||||
"n_points": int(N),
|
||||
"frac_visible": float(v.mean()),
|
||||
"frac_moving": float(moving.mean()), # fraction of grid points that ever move
|
||||
"n_moving": int(moving.sum()),
|
||||
"mean_motion_px": float(max_disp.mean()),
|
||||
"p95_motion_px": float(np.percentile(max_disp, 95)),
|
||||
"max_motion_px": float(max_disp.max()),
|
||||
}
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
from fastvideo.train.callbacks.track_validation import (_draw_overlay, _grid_colors, _subsample)
|
||||
data_dir = Path(args.data_dir)
|
||||
manifest = json.loads((data_dir / "videos2caption.json").read_text())
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
import imageio.v2 as imageio
|
||||
from PIL import Image
|
||||
|
||||
embs: list[np.ndarray] = []
|
||||
entries: list[dict[str, Any]] = []
|
||||
for rec in manifest:
|
||||
cid = Path(rec["path"]).stem
|
||||
vpath = str(data_dir / "videos" / rec["path"])
|
||||
ppath = rec.get("points_path")
|
||||
frames = _load_frames(vpath)
|
||||
d = np.load(ppath)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]] # [T,N,2] px
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]]
|
||||
H, W = int(frames.shape[1]), int(frames.shape[2])
|
||||
|
||||
stats = _clip_stats(tracks, vis)
|
||||
|
||||
# overlay (subsampled grid for clarity)
|
||||
grid = int(round(tracks.shape[1]**0.5))
|
||||
tn = tracks / np.array([W, H], np.float32) # normalize for _subsample
|
||||
tr_s, vs_s = _subsample(tn, vis, grid, args.stride)
|
||||
colors = _grid_colors(grid, args.stride)[:tr_s.shape[1]]
|
||||
trpx = tr_s.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
ov = _draw_overlay(frames, trpx, vs_s, colors, args.tail, 2, VIS_THRESH)
|
||||
imageio.mimsave(str(out / f"{cid}__overlay.mp4"), ov, fps=args.fps, macro_block_size=1)
|
||||
Image.fromarray(frames[0]).save(str(out / f"{cid}__thumb.png"))
|
||||
|
||||
# frame-0 embedding for near-duplicate grouping (32x32 gray, L2-norm)
|
||||
g = np.asarray(Image.fromarray(frames[0]).convert("L").resize((32, 32)), np.float32).reshape(-1)
|
||||
embs.append(g / (np.linalg.norm(g) + 1e-8))
|
||||
|
||||
entries.append({"id": cid, "episode": rec.get("source_episode"), "caption": rec["cap"][0], **stats})
|
||||
print(
|
||||
f"[dash] {cid} move%={stats['frac_moving']*100:4.1f} mean={stats['mean_motion_px']:5.1f}px "
|
||||
f"max={stats['max_motion_px']:5.1f}px",
|
||||
flush=True)
|
||||
|
||||
# near-duplicate grouping: union-find on cosine-sim >= threshold
|
||||
E = np.stack(embs, 0)
|
||||
sim = E @ E.T
|
||||
n = len(entries)
|
||||
parent = list(range(n))
|
||||
|
||||
def find(i: int) -> int:
|
||||
while parent[i] != i:
|
||||
parent[i] = parent[parent[i]]
|
||||
i = parent[i]
|
||||
return i
|
||||
|
||||
for i in range(n):
|
||||
for j in range(i + 1, n):
|
||||
if sim[i, j] >= args.dup_thresh:
|
||||
parent[find(i)] = find(j)
|
||||
groups: dict[int, list[int]] = {}
|
||||
for i in range(n):
|
||||
groups.setdefault(find(i), []).append(i)
|
||||
group_id = {}
|
||||
for gi, (_, members) in enumerate(sorted(groups.items(), key=lambda kv: -len(kv[1]))):
|
||||
for m in members:
|
||||
group_id[m] = gi
|
||||
for i, e in enumerate(entries):
|
||||
e["dup_group"] = int(group_id[i])
|
||||
|
||||
n_groups = len(groups)
|
||||
n_dead = sum(1 for e in entries if e["frac_moving"] < 0.02)
|
||||
summary = {
|
||||
"n_clips": n,
|
||||
"n_dup_groups": n_groups,
|
||||
"largest_dup_group": max(len(v) for v in groups.values()),
|
||||
"n_low_motion_clips(<2%)": n_dead,
|
||||
"median_frac_moving": float(np.median([e["frac_moving"] for e in entries])),
|
||||
"median_mean_motion_px": float(np.median([e["mean_motion_px"] for e in entries])),
|
||||
"dup_thresh": args.dup_thresh,
|
||||
"move_thresh_px": MOVE_THRESH_PX,
|
||||
}
|
||||
(out / "stats.json").write_text(json.dumps({"summary": summary, "clips": entries}, indent=2))
|
||||
print(f"[dash] SUMMARY: {json.dumps(summary, indent=2)}", flush=True)
|
||||
print(f"[dash] wrote dashboard artifacts -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
dash = Path(args.dash_dir)
|
||||
data = json.loads((dash / "stats.json").read_text())
|
||||
summary, clips = data["summary"], data["clips"]
|
||||
# sort: least motion first (dead/duplicate clips surface), then by dup group
|
||||
order = sorted(range(len(clips)), key=lambda i: (clips[i]["frac_moving"], clips[i]["dup_group"]))
|
||||
|
||||
def label(i: int) -> str:
|
||||
c = clips[i]
|
||||
return (f"#{i} ep{c['episode']} | move {c['frac_moving']*100:.0f}% | "
|
||||
f"mean {c['mean_motion_px']:.0f}px | dupgrp {c['dup_group']}")
|
||||
|
||||
gallery_items = [(str(dash / f"{clips[i]['id']}__thumb.png"), label(i)) for i in order]
|
||||
|
||||
def show(evt): # gr.SelectData; unannotated so gradio's get_type_hints doesn't eval it
|
||||
i = order[evt.index]
|
||||
c = clips[i]
|
||||
return (str(dash / f"{c['id']}__overlay.mp4"), json.dumps(c, indent=2))
|
||||
|
||||
head = (f"### DROID dataset inspection — {summary['n_clips']} clips\n"
|
||||
f"- near-duplicate groups: **{summary['n_dup_groups']}** "
|
||||
f"(largest group **{summary['largest_dup_group']}** clips)\n"
|
||||
f"- low-motion clips (<2% points move): **{summary['n_low_motion_clips(<2%)']}**\n"
|
||||
f"- median fraction of points moving: **{summary['median_frac_moving']*100:.1f}%**, "
|
||||
f"median mean-motion **{summary['median_mean_motion_px']:.1f}px**\n\n"
|
||||
f"Sorted least-motion first. Overlay = CoTracker tracks (dots + tails); "
|
||||
f"watch for moving objects with NO dots on them (off-frame entries / lost tracks).")
|
||||
|
||||
with gr.Blocks(title="DROID dataset dashboard") as demo:
|
||||
gr.Markdown(head)
|
||||
with gr.Row():
|
||||
gal = gr.Gallery(value=gallery_items, columns=6, height=560, label="clips (least motion first)")
|
||||
with gr.Column():
|
||||
vid = gr.Video(label="track overlay")
|
||||
meta = gr.Code(label="clip stats", language="json")
|
||||
gal.select(show, None, [vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True, help="dataset root (videos/ + videos2caption.json + tracks/)")
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--stride", type=int, default=3, help="subsample the NxN grid for overlay clarity")
|
||||
r.add_argument("--tail", type=int, default=12)
|
||||
r.add_argument("--fps", type=int, default=12)
|
||||
r.add_argument("--dup-thresh", type=float, default=0.985, help="cosine-sim on frame-0 to call clips duplicates")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--dash-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7872)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
args = p.parse_args()
|
||||
args.func(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,416 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Evaluate super-sparse track adherence via CoTracker EPE.
|
||||
|
||||
For each val clip, we build the exact adversarial case the user described:
|
||||
* Pick ONE active foreground trace from the GT parquet — the highest-motion
|
||||
foreground track (object_ids >= 0, visible in frame 0).
|
||||
* Pick N static background anchors — random background tracks (object_ids == -1)
|
||||
with position frozen at frame-0 (visibility=1 across all frames).
|
||||
* Generate video conditioned on (active + anchors) — this is what training
|
||||
was supposed to teach the model to handle.
|
||||
* Also generate an ablation with track_points=None (unconditional motion).
|
||||
* Run CoTracker on both generated videos, initialised at frame 0 with the
|
||||
conditioning points. Compare CoTracker-extracted tracks against the GT
|
||||
trace to compute End-Point-Error in pixel space.
|
||||
|
||||
Metrics logged per clip and aggregated to wandb:
|
||||
epe_active_full EPE(GT_active, CoTracker(gen_full))
|
||||
epe_active_notrack EPE(GT_active, CoTracker(gen_no_track))
|
||||
epe_bg EPE(static_anchor, CoTracker(gen_full)) — should be small
|
||||
DELTA = epe_active_notrack - epe_active_full — positive means training helped.
|
||||
|
||||
Usage:
|
||||
srun --overlap --jobid=529 --ntasks=1 -w hpc-rack-2-13 --chdir=$PWD bash -lc '
|
||||
source .venv/bin/activate
|
||||
export HOME=/mnt/lustre/vlm-s4duan HF_HOME=/mnt/lustre/vlm-s4duan/.hf \
|
||||
TORCH_HOME=/mnt/lustre/vlm-s4duan/.torch MPLCONFIGDIR=/mnt/lustre/vlm-s4duan/.mpl \
|
||||
TRITON_CACHE_DIR=/tmp/triton_eval TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0 \
|
||||
PYTHONPATH=$PWD TRACKWAN_TRACK_BIAS=1 CUDA_VISIBLE_DEVICES=2 \
|
||||
WANDB_API_KEY=<key> WANDB_MODE=online
|
||||
python data_pipeline/eval_sparse_track_epe.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800 \
|
||||
--yaml examples/train/scenario/worldmodel/finetune_wantrack_openvid_sparse_1p3b_merged_bias.yaml \
|
||||
--data-path /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset \
|
||||
--num-clips 5 --steps 30 --w-motion 1.5 \
|
||||
--wandb-run-name eval_mb4800_sparse
|
||||
'
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
# make trackwan_infer importable
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__)))
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# CoTracker
|
||||
# ==========================================================================
|
||||
def load_cotracker(device: str) -> torch.nn.Module:
|
||||
"""Load CoTracker v3 offline. Uses the warm hub cache under $TORCH_HOME/hub."""
|
||||
hub_dir = Path(torch.hub.get_dir())
|
||||
local = hub_dir / "facebookresearch_co-tracker_main"
|
||||
if local.exists():
|
||||
model = torch.hub.load(str(local), "cotracker3_offline", source="local", trust_repo=True)
|
||||
else:
|
||||
model = torch.hub.load("facebookresearch/co-tracker", "cotracker3_offline", trust_repo=True)
|
||||
return model.to(device).eval()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def cotracker_query(cotracker: torch.nn.Module, video: np.ndarray, queries_xy: np.ndarray,
|
||||
device: str) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Track `queries_xy` (N, 2 pixel coords) forward through `video` [T,H,W,3] uint8.
|
||||
|
||||
Returns tracks [T,N,2] (pixel coords) and visibility [T,N] (bool).
|
||||
"""
|
||||
T, H, W, C = video.shape
|
||||
v = torch.from_numpy(video).permute(0, 3, 1, 2).float()[None].to(device) # [1,T,3,H,W]
|
||||
N = queries_xy.shape[0]
|
||||
q = np.zeros((1, N, 3), dtype=np.float32)
|
||||
q[0, :, 0] = 0 # start-frame index
|
||||
q[0, :, 1] = queries_xy[:, 0] # x
|
||||
q[0, :, 2] = queries_xy[:, 1] # y
|
||||
tracks, visibility = cotracker(v, queries=torch.from_numpy(q).to(device))
|
||||
return tracks[0].cpu().numpy(), visibility[0].cpu().numpy()
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Sparse-conditioning sampler
|
||||
# ==========================================================================
|
||||
def pick_sparse_conditioning(row: dict, *, num_active: int, num_bg: int,
|
||||
seed: int) -> dict:
|
||||
"""From a preprocessed parquet row, pick the sparse case: `num_active` moving fg
|
||||
traces (highest-displacement fg tracks visible in frame 0) + `num_bg` static bg
|
||||
anchors (random bg tracks, position frozen at frame-0).
|
||||
|
||||
Everything is normalised to [0,1] to match training. Pixel-space is only used at
|
||||
the metric-computation step.
|
||||
"""
|
||||
tp = row["track_points"] # [T,N,2] normalised
|
||||
tv = row["track_visibility"] # [T,N]
|
||||
oid = row["object_ids"] # [N] int, -1 for bg
|
||||
T, N, _ = tp.shape
|
||||
rng = np.random.RandomState(seed)
|
||||
|
||||
fg = np.where((oid >= 0) & (tv[0] > 0.5))[0]
|
||||
bg = np.where((oid == -1) & (tv[0] > 0.5))[0]
|
||||
if fg.size == 0:
|
||||
raise RuntimeError("no visible foreground tracks in frame 0")
|
||||
# displacement of each fg over the full clip (only over visible frames)
|
||||
disp = np.linalg.norm(tp[-1, fg] - tp[0, fg], axis=-1)
|
||||
active_idx = fg[np.argsort(-disp)[:num_active]] # top-`num_active` moving fg
|
||||
n_bg = min(num_bg, bg.size)
|
||||
bg_idx = rng.choice(bg, size=n_bg, replace=False) if n_bg > 0 else np.zeros((0, ), dtype=int)
|
||||
|
||||
active_gt = tp[:, active_idx] # [T, num_active, 2] — GT motion
|
||||
active_vis = tv[:, active_idx]
|
||||
bg_static_pos = np.broadcast_to(tp[0:1, bg_idx], (T, n_bg, 2)).copy() # position frozen at f0
|
||||
bg_vis = np.ones((T, n_bg), dtype=np.float32)
|
||||
|
||||
cond_tracks = np.concatenate([active_gt, bg_static_pos], axis=1) # [T, num_active+n_bg, 2]
|
||||
cond_vis = np.concatenate([active_vis, bg_vis], axis=1)
|
||||
return {
|
||||
"active_idx": active_idx,
|
||||
"active_gt": active_gt, # [T,num_active,2] normalised
|
||||
"active_vis": active_vis,
|
||||
"bg_static_pos": bg_static_pos, # [T,n_bg,2] normalised (const in t)
|
||||
"cond_tracks": cond_tracks, # [T, num_active+n_bg, 2] normalised
|
||||
"cond_vis": cond_vis,
|
||||
"n_active": len(active_idx),
|
||||
"n_bg": n_bg,
|
||||
}
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Metrics
|
||||
# ==========================================================================
|
||||
def epe_norm_to_px(pred_norm: np.ndarray, gt_norm: np.ndarray, vis: np.ndarray,
|
||||
res_h: int, res_w: int) -> float:
|
||||
"""L2 distance between predicted/GT normalised tracks, evaluated in pixel space,
|
||||
averaged only over visible frames × visible points."""
|
||||
diff = (pred_norm - gt_norm) * np.array([res_w, res_h])[None, None, :]
|
||||
l2 = np.linalg.norm(diff, axis=-1) # [T, N]
|
||||
mask = vis > 0.5
|
||||
if not mask.any():
|
||||
return float("nan")
|
||||
return float(l2[mask].mean())
|
||||
|
||||
|
||||
def epe_px(pred_px: np.ndarray, gt_px: np.ndarray, vis: np.ndarray) -> float:
|
||||
"""L2 in pixel space — both inputs already in pixels."""
|
||||
l2 = np.linalg.norm(pred_px - gt_px, axis=-1)
|
||||
mask = vis > 0.5
|
||||
if not mask.any():
|
||||
return float("nan")
|
||||
return float(l2[mask].mean())
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Trace overlay for wandb visualisation
|
||||
# ==========================================================================
|
||||
def overlay_tracks(video: np.ndarray, tracks_norm: np.ndarray, vis: np.ndarray,
|
||||
colors: list[tuple[int, int, int]], radius: int = 3,
|
||||
tail: int = 12) -> np.ndarray:
|
||||
"""Draw normalised tracks on a video [T,H,W,3] uint8. Small tail + current dot."""
|
||||
from PIL import Image, ImageDraw
|
||||
T, H, W, C = video.shape
|
||||
N = tracks_norm.shape[1]
|
||||
out = []
|
||||
for t in range(T):
|
||||
img = Image.fromarray(np.ascontiguousarray(video[t]))
|
||||
dr = ImageDraw.Draw(img)
|
||||
t0 = max(0, t - tail)
|
||||
for i in range(N):
|
||||
col = colors[i % len(colors)]
|
||||
# tail
|
||||
if t - t0 >= 1:
|
||||
pts = [(float(tracks_norm[k, i, 0] * W), float(tracks_norm[k, i, 1] * H))
|
||||
for k in range(t0, t + 1)]
|
||||
dr.line(pts, fill=col, width=1)
|
||||
if vis[t, i] > 0.5:
|
||||
x = float(tracks_norm[t, i, 0] * W)
|
||||
y = float(tracks_norm[t, i, 1] * H)
|
||||
dr.ellipse([x - radius, y - radius, x + radius, y + radius], fill=col)
|
||||
out.append(np.asarray(img))
|
||||
return np.stack(out)
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Parquet loader (small subset — avoid loading the whole 259k row set)
|
||||
# ==========================================================================
|
||||
def load_first_n_from_parquet(data_path: str, n: int, text_len: int) -> list:
|
||||
import pyarrow.parquet as pq
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
files = sorted(glob.glob(os.path.join(data_path, "**", "*.parquet"), recursive=True))
|
||||
if not files:
|
||||
raise FileNotFoundError(data_path)
|
||||
rows: list = []
|
||||
for f in files:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
if len(rows) >= n:
|
||||
break
|
||||
sel = rows[:n]
|
||||
batch = collate_rows_from_parquet_schema(sel, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in sel]
|
||||
out = []
|
||||
for i in range(len(sel)):
|
||||
out.append({
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i].numpy(), # [T,N,2] normalised
|
||||
"track_visibility": batch["track_visibility"][i].numpy(), # [T,N]
|
||||
"object_ids": batch["object_ids"][i].numpy(), # [N]
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Main
|
||||
# ==========================================================================
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True)
|
||||
p.add_argument("--yaml", required=True)
|
||||
p.add_argument("--data-path", required=True)
|
||||
p.add_argument("--num-clips", type=int, default=5)
|
||||
p.add_argument("--num-active", type=int, default=1)
|
||||
p.add_argument("--num-bg", type=int, default=20)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--w-motion", type=float, default=1.5)
|
||||
p.add_argument("--seed", type=int, default=1234)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--wandb-project", default="wantrack-bidir")
|
||||
p.add_argument("--wandb-run-name", required=True)
|
||||
p.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/eval_sparse_epe")
|
||||
args = p.parse_args()
|
||||
|
||||
Path(args.out_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
import wandb
|
||||
import imageio.v2 as imageio
|
||||
import trackwan_infer as twi
|
||||
|
||||
print(f"[eval] loading model {args.model_dir}", flush=True)
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(tc.data.text_padding_length) if hasattr(tc.data, "text_padding_length") else 256
|
||||
|
||||
print(f"[eval] loading {args.num_clips} parquet clips", flush=True)
|
||||
samples = load_first_n_from_parquet(args.data_path, args.num_clips, text_len)
|
||||
|
||||
print("[eval] loading CoTracker", flush=True)
|
||||
device = model.device
|
||||
cotracker = load_cotracker(device.type + ":" + str(device.index) if device.index is not None else device.type)
|
||||
|
||||
res_h = int(tc.data.num_height)
|
||||
res_w = int(tc.data.num_width)
|
||||
palette = [(255, 60, 60), (60, 255, 60), (60, 90, 255), (255, 200, 60), (200, 60, 255), (60, 255, 220)]
|
||||
|
||||
run = wandb.init(project=args.wandb_project, name=args.wandb_run_name,
|
||||
config={"model_dir": args.model_dir, "num_clips": args.num_clips,
|
||||
"num_active": args.num_active, "num_bg": args.num_bg,
|
||||
"steps": args.steps, "w_motion": args.w_motion, "res": [res_h, res_w]})
|
||||
|
||||
per_clip: list = []
|
||||
|
||||
for idx, s in enumerate(samples):
|
||||
try:
|
||||
sc = pick_sparse_conditioning(s, num_active=args.num_active, num_bg=args.num_bg,
|
||||
seed=args.seed + idx)
|
||||
except RuntimeError as exc:
|
||||
print(f"[eval] clip {idx}: {exc}; skipping", flush=True)
|
||||
continue
|
||||
|
||||
# ── Generate two videos: with sparse tracks (motion CFG) + without any tracks ──
|
||||
tp_t = torch.from_numpy(sc["cond_tracks"])[None].float() # [1,T,N,2]
|
||||
tv_t = torch.from_numpy(sc["cond_vis"])[None].float() # [1,T,N]
|
||||
|
||||
seed = args.seed + idx
|
||||
common = dict(first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=args.steps, seed=seed)
|
||||
|
||||
print(f"[eval] clip {idx}: generating full ({sc['n_active']} active + {sc['n_bg']} anchors)", flush=True)
|
||||
t0 = time.time()
|
||||
lat_full = twi.generate(model, track_points=tp_t, track_visibility=tv_t,
|
||||
guidance_scale=args.w_motion, **common)
|
||||
gen_full = twi.decode_to_pixels(model, lat_full) # [T,H,W,3] uint8
|
||||
t_full = time.time() - t0
|
||||
print(f"[eval] clip {idx}: full took {t_full:.1f}s", flush=True)
|
||||
|
||||
print(f"[eval] clip {idx}: generating no-track", flush=True)
|
||||
t0 = time.time()
|
||||
lat_no = twi.generate(model, track_points=None, track_visibility=None,
|
||||
guidance_scale=1.0, **common)
|
||||
gen_no = twi.decode_to_pixels(model, lat_no)
|
||||
t_no = time.time() - t0
|
||||
print(f"[eval] clip {idx}: no-track took {t_no:.1f}s", flush=True)
|
||||
|
||||
T, H, W, _ = gen_full.shape
|
||||
|
||||
# ── CoTracker on both generations, queried at frame-0 with same points ──
|
||||
query_norm = sc["cond_tracks"][0] # [N_total, 2] in [0,1]
|
||||
queries_px = query_norm * np.array([W, H])[None] # pixel space
|
||||
print(f"[eval] clip {idx}: CoTracker on full", flush=True)
|
||||
tk_full_px, tk_full_vis = cotracker_query(cotracker, gen_full, queries_px, device=device.type)
|
||||
print(f"[eval] clip {idx}: CoTracker on no-track", flush=True)
|
||||
tk_no_px, tk_no_vis = cotracker_query(cotracker, gen_no, queries_px, device=device.type)
|
||||
|
||||
# GT in pixel space, at the SAME resolution as the generation
|
||||
gt_full_px = sc["cond_tracks"] * np.array([W, H])[None, None, :] # [T,N_total,2]
|
||||
active_slice = slice(0, sc["n_active"])
|
||||
bg_slice = slice(sc["n_active"], sc["n_active"] + sc["n_bg"])
|
||||
|
||||
# Active EPE: CoTracker(gen) vs the GT trace we specified.
|
||||
# Filter by GT visibility (active_vis) AND CoTracker's own visibility.
|
||||
active_vis_gt = sc["active_vis"] # [T,n_active]
|
||||
both_vis_full = (active_vis_gt > 0.5) & (tk_full_vis[:, active_slice] > 0.5)
|
||||
both_vis_no = (active_vis_gt > 0.5) & (tk_no_vis[:, active_slice] > 0.5)
|
||||
|
||||
epe_active_full = epe_px(tk_full_px[:, active_slice],
|
||||
gt_full_px[:, active_slice],
|
||||
both_vis_full.astype(np.float32))
|
||||
epe_active_no = epe_px(tk_no_px[:, active_slice],
|
||||
gt_full_px[:, active_slice],
|
||||
both_vis_no.astype(np.float32))
|
||||
|
||||
# BG EPE (only for full): CoTracker(gen_full) vs static anchor position
|
||||
bg_static_px = sc["bg_static_pos"] * np.array([W, H])[None, None, :]
|
||||
bg_vis_mask = (tk_full_vis[:, bg_slice] > 0.5).astype(np.float32)
|
||||
epe_bg = epe_px(tk_full_px[:, bg_slice], bg_static_px, bg_vis_mask) if sc["n_bg"] > 0 else float("nan")
|
||||
|
||||
delta = epe_active_no - epe_active_full # positive => tracks helped
|
||||
|
||||
# ── Save + log videos ──
|
||||
overlay_full = overlay_tracks(gen_full, sc["cond_tracks"], sc["cond_vis"], palette)
|
||||
overlay_no = overlay_tracks(gen_no, sc["cond_tracks"], sc["cond_vis"], palette)
|
||||
|
||||
# Also show the CoTracker-extracted tracks over the full generation — this is
|
||||
# what the model *actually* produced motion-wise.
|
||||
cotr_tracks_norm = tk_full_px / np.array([W, H])[None, None, :]
|
||||
overlay_cotr = overlay_tracks(gen_full, cotr_tracks_norm, tk_full_vis.astype(np.float32), palette)
|
||||
|
||||
clip_dir = Path(args.out_dir) / f"clip_{idx:03d}"
|
||||
clip_dir.mkdir(parents=True, exist_ok=True)
|
||||
p_full = clip_dir / "gen_full_overlay.mp4"
|
||||
p_no = clip_dir / "gen_notrack_overlay.mp4"
|
||||
p_cotr = clip_dir / "gen_full_cotracker.mp4"
|
||||
imageio.mimsave(p_full, overlay_full, fps=args.fps, macro_block_size=1)
|
||||
imageio.mimsave(p_no, overlay_no, fps=args.fps, macro_block_size=1)
|
||||
imageio.mimsave(p_cotr, overlay_cotr, fps=args.fps, macro_block_size=1)
|
||||
|
||||
print(f"[eval] clip {idx}: EPE_active_full={epe_active_full:.2f} "
|
||||
f"EPE_active_notrack={epe_active_no:.2f} DELTA={delta:.2f} "
|
||||
f"EPE_bg={epe_bg:.2f}", flush=True)
|
||||
|
||||
# per-clip logs
|
||||
entry = {
|
||||
"clip": idx,
|
||||
"epe_active_full_px": epe_active_full,
|
||||
"epe_active_notrack_px": epe_active_no,
|
||||
"delta_px": delta,
|
||||
"epe_bg_px": epe_bg,
|
||||
"n_active": sc["n_active"],
|
||||
"n_bg": sc["n_bg"],
|
||||
"sec_gen_full": t_full,
|
||||
"sec_gen_notrack": t_no,
|
||||
"caption": s["caption"][:120],
|
||||
}
|
||||
per_clip.append(entry)
|
||||
wandb.log({
|
||||
f"clip_{idx:03d}/gen_full": wandb.Video(str(p_full), fps=args.fps, format="mp4"),
|
||||
f"clip_{idx:03d}/gen_notrack": wandb.Video(str(p_no), fps=args.fps, format="mp4"),
|
||||
f"clip_{idx:03d}/gen_full_cotracker": wandb.Video(str(p_cotr), fps=args.fps, format="mp4"),
|
||||
"clip": idx,
|
||||
"epe_active_full_px": epe_active_full,
|
||||
"epe_active_notrack_px": epe_active_no,
|
||||
"delta_px": delta,
|
||||
"epe_bg_px": epe_bg,
|
||||
})
|
||||
|
||||
# ── Aggregate + summary ──
|
||||
if per_clip:
|
||||
arr = lambda k: np.array([r[k] for r in per_clip if not np.isnan(r[k])])
|
||||
summary = {
|
||||
"n_clips": len(per_clip),
|
||||
"mean_epe_active_full_px": float(np.mean(arr("epe_active_full_px"))),
|
||||
"median_epe_active_full_px": float(np.median(arr("epe_active_full_px"))),
|
||||
"mean_epe_active_notrack_px": float(np.mean(arr("epe_active_notrack_px"))),
|
||||
"median_epe_active_notrack_px": float(np.median(arr("epe_active_notrack_px"))),
|
||||
"mean_delta_px": float(np.mean(arr("delta_px"))),
|
||||
"median_delta_px": float(np.median(arr("delta_px"))),
|
||||
"mean_epe_bg_px": float(np.mean(arr("epe_bg_px"))),
|
||||
}
|
||||
(Path(args.out_dir) / f"{args.wandb_run_name}_summary.json").write_text(json.dumps({
|
||||
"summary": summary, "per_clip": per_clip}, indent=2))
|
||||
wandb.log({f"summary/{k}": v for k, v in summary.items()})
|
||||
wandb.summary.update(summary)
|
||||
print("=" * 70)
|
||||
print(f"[eval] summary ({len(per_clip)} clips):")
|
||||
for k, v in summary.items():
|
||||
print(f" {k}: {v:.3f}" if isinstance(v, float) else f" {k}: {v}")
|
||||
print("=" * 70)
|
||||
|
||||
wandb.finish()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,536 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 0c: extract dense point tracks from generated videos with CoTracker v3.
|
||||
|
||||
For each .mp4 produced by ``generate_videos.py`` we run CoTracker v3 (``cotracker3_offline``)
|
||||
with a ``grid_size``x``grid_size`` regular query grid (default 50x50 = 2500 points) and save
|
||||
the per-frame tracks + visibility. We then patch ``points_path`` (absolute) into the
|
||||
manifest so the future points-aware preprocess task can find them (mirrors how MatrixGame2
|
||||
references ``action_path``).
|
||||
|
||||
Tracks are stored in ORIGINAL video pixel coordinates. The full 2500-point grid + visibility
|
||||
are kept; the trainer samples 1-200 points per step.
|
||||
|
||||
If ``--detect-entries`` is set, FastSAM detects objects entering after frame 0 (frames are
|
||||
segmented in batched forwards). Grid points landing on each new object are tracked from its
|
||||
entry frame T_entry — all entry events share a single extra CoTracker pass (chunked if the
|
||||
combined query count exceeds grid_size^2) — and replace dead background slots (those
|
||||
permanently occluded from T_entry onwards), keeping N = grid_size^2 throughout.
|
||||
|
||||
Run on a GPU node (never the login node), e.g.:
|
||||
|
||||
srun --jobid=<shao_wm jobid> --overlap --ntasks=1 \\
|
||||
.venv/bin/python data_pipeline/extract_tracks.py \\
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p \\
|
||||
--grid-size 50 --detect-entries
|
||||
|
||||
torch.hub note: prefetch once on the login node (internet) so the shared cache is warm:
|
||||
.venv/bin/python -c "import torch; torch.hub.load('facebookresearch/co-tracker','cotracker3_offline')"
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
HUB_REPO = "facebookresearch/co-tracker"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--data-dir", type=Path, required=True, help="Dataset root from generate_videos.py.")
|
||||
p.add_argument("--videos-subdir", type=str, default="videos")
|
||||
p.add_argument("--out-subdir", type=str, default="tracks")
|
||||
p.add_argument("--manifest", type=str, default="videos2caption.json")
|
||||
p.add_argument("--grid-size", type=int, default=50, help="NxN query grid (N*N points).")
|
||||
p.add_argument("--model", type=str, default="cotracker3_offline")
|
||||
p.add_argument("--device", type=str, default="cuda")
|
||||
p.add_argument("--downscale", type=float, default=1.0,
|
||||
help="Run tracking at this spatial scale (coords rescaled back to original px). "
|
||||
"Use <1.0 (e.g. 0.5) if full-res OOMs.")
|
||||
p.add_argument("--limit", type=int, default=None, help="Only process the first N videos (for smoke tests).")
|
||||
p.add_argument("--index", type=int, nargs="+", default=None, metavar="IDX",
|
||||
help="Only process videos with these indices (e.g. --index 4 7 12).")
|
||||
p.add_argument("--rank", type=int, default=0, help="GPU rank for sharding (0-indexed).")
|
||||
p.add_argument("--world-size", type=int, default=1, help="Total number of parallel processes.")
|
||||
p.add_argument("--force", action="store_true", help="Re-extract even if .npz already exists.")
|
||||
p.add_argument("--verbose", action="store_true", help="Print per-mask debug info for entry detection.")
|
||||
# Entry-frame detection
|
||||
p.add_argument("--detect-entries", action="store_true",
|
||||
help="Detect objects entering after frame 0 and replace dead background slots.")
|
||||
p.add_argument("--sam-model", type=str, default="FastSAM-s.pt")
|
||||
p.add_argument("--sam-conf", type=float, default=0.75)
|
||||
p.add_argument("--sam-iou", type=float, default=0.9)
|
||||
p.add_argument("--sam-imgsz", type=int, default=1024)
|
||||
p.add_argument("--sam-batch", type=int, default=16,
|
||||
help="Frames per batched FastSAM forward during entry detection.")
|
||||
p.add_argument("--amp", action="store_true",
|
||||
help="Run CoTracker under bf16 autocast (~1.5-2x faster, slightly different coords).")
|
||||
p.add_argument("--compile", action="store_true",
|
||||
help="torch.compile the main CoTracker pass (fixed shape; compiled once per worker). "
|
||||
"The variable-size entry pass stays eager to avoid recompiles.")
|
||||
p.add_argument("--prefetch", type=int, default=2,
|
||||
help="Decode up to N videos ahead on a background thread so the GPU never waits "
|
||||
"on video decode (0 = disable).")
|
||||
# Fused Stage 4
|
||||
p.add_argument("--segment", action="store_true",
|
||||
help="Fused Stage 4: also compute object_ids/n_objects/track_weights and the "
|
||||
"vis-override sweep in this pass, reusing the decoded video and the "
|
||||
"entry-detection FastSAM masks. Uses --sam-conf/--sam-iou/--sam-imgsz; "
|
||||
"masks are unfiltered (no min-area-frac/max-masks).")
|
||||
p.add_argument("--vis-override-every", type=int, default=3,
|
||||
help="(with --segment) run FastSAM every N frames and set vis=True for object "
|
||||
"points inside masks; 0 disables.")
|
||||
p.add_argument("--viz", action="store_true",
|
||||
help="(with --segment) render a track-overlay mp4 after each video.")
|
||||
p.add_argument("--viz-dir", type=str, default=None,
|
||||
help="Output directory for viz mp4s (default: <data-dir>/viz).")
|
||||
p.add_argument("--entry-sample-every", type=int, default=5,
|
||||
help="Check for new objects every N frames.")
|
||||
p.add_argument("--entry-new-area", type=float, default=0.3,
|
||||
help="A mask triggers entry detection if at least this fraction of its area "
|
||||
"is not covered by any frame-0 mask (catches partially-entering objects).")
|
||||
p.add_argument("--entry-min-area", type=float, default=0.005,
|
||||
help="Min area fraction for a new-object mask to be considered.")
|
||||
args = p.parse_args()
|
||||
if args.viz and not args.segment:
|
||||
p.error("--viz requires --segment (object IDs are needed for the overlay)")
|
||||
return args
|
||||
|
||||
|
||||
def load_cotracker(model_name: str, device: str):
|
||||
"""Load CoTracker via torch.hub, preferring the warm local cache (offline-safe)."""
|
||||
hub_dir = Path(torch.hub.get_dir())
|
||||
local = hub_dir / (HUB_REPO.replace("/", "_") + "_main")
|
||||
try:
|
||||
if local.exists():
|
||||
model = torch.hub.load(str(local), model_name, source="local", trust_repo=True)
|
||||
else:
|
||||
model = torch.hub.load(HUB_REPO, model_name, trust_repo=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
raise RuntimeError(
|
||||
f"Failed to load CoTracker ({e!r}). On a node without internet, prefetch on the "
|
||||
"login node first: .venv/bin/python -c \"import torch; "
|
||||
"torch.hub.load('facebookresearch/co-tracker','cotracker3_offline')\""
|
||||
) from e
|
||||
return model.to(device).eval()
|
||||
|
||||
|
||||
def read_video(path: Path) -> tuple[torch.Tensor, int, int]:
|
||||
"""Return (video[1,T,C,H,W] float 0-255, H, W)."""
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(str(path), ctx=cpu(0))
|
||||
frames = vr.get_batch(list(range(len(vr)))).asnumpy() # (T, H, W, C) uint8
|
||||
except Exception: # noqa: BLE001
|
||||
import imageio
|
||||
reader = imageio.get_reader(str(path), format="ffmpeg")
|
||||
frames = np.stack([np.asarray(f) for f in reader], axis=0)
|
||||
reader.close()
|
||||
vid = torch.from_numpy(np.ascontiguousarray(frames))
|
||||
h, w = int(vid.shape[1]), int(vid.shape[2])
|
||||
video = vid.permute(0, 3, 1, 2).unsqueeze(0).float() # (1,T,C,H,W)
|
||||
return video, h, w
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def track_one(model, video: torch.Tensor, grid_size: int, downscale: float, device: str,
|
||||
amp: bool = False) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Return tracks (T,N,2) in original pixel coords and visibility (T,N)."""
|
||||
_, t, c, h, w = video.shape
|
||||
track_video = video
|
||||
if downscale != 1.0:
|
||||
sh, sw = max(1, int(round(h * downscale))), max(1, int(round(w * downscale)))
|
||||
track_video = F.interpolate(video[0], size=(sh, sw), mode="bilinear", align_corners=False).unsqueeze(0)
|
||||
else:
|
||||
sh, sw = h, w
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=amp and device.startswith("cuda")):
|
||||
pred_tracks, pred_vis = model(track_video.to(device), grid_size=grid_size)
|
||||
tracks = pred_tracks[0].float().cpu().numpy()
|
||||
vis = pred_vis[0].cpu().numpy()
|
||||
|
||||
if downscale != 1.0:
|
||||
tracks[..., 0] *= w / float(sw)
|
||||
tracks[..., 1] *= h / float(sh)
|
||||
return tracks.astype(np.float32), vis
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def track_with_queries(model, video: torch.Tensor, queries_txy: np.ndarray, downscale: float,
|
||||
device: str, H: int, W: int, amp: bool = False) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Track explicit query points. queries_txy: [K,3] as (t, x, y) in original pixel coords."""
|
||||
track_video = video
|
||||
sh, sw = H, W
|
||||
q = queries_txy.astype(np.float32).copy()
|
||||
if downscale != 1.0:
|
||||
sh = max(1, int(round(H * downscale)))
|
||||
sw = max(1, int(round(W * downscale)))
|
||||
track_video = F.interpolate(video[0], size=(sh, sw), mode="bilinear", align_corners=False).unsqueeze(0)
|
||||
q[:, 1] *= sw / float(W)
|
||||
q[:, 2] *= sh / float(H)
|
||||
|
||||
q_tensor = torch.from_numpy(q).unsqueeze(0).to(device) # [1, K, 3]
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=amp and device.startswith("cuda")):
|
||||
pred_tracks, pred_vis = model(track_video.to(device), queries=q_tensor)
|
||||
tracks = pred_tracks[0].float().cpu().numpy()
|
||||
vis = pred_vis[0].cpu().numpy()
|
||||
|
||||
if downscale != 1.0:
|
||||
tracks[..., 0] *= W / float(sw)
|
||||
tracks[..., 1] *= H / float(sh)
|
||||
return tracks.astype(np.float32), vis
|
||||
|
||||
|
||||
def make_grid_queries(grid_size: int, H: int, W: int, frame_t: int) -> np.ndarray:
|
||||
"""Generate a grid_size x grid_size uniform query grid at frame_t. Returns [N,3] (t,x,y)."""
|
||||
ys = np.linspace(0, H - 1, grid_size)
|
||||
xs = np.linspace(0, W - 1, grid_size)
|
||||
xx, yy = np.meshgrid(xs, ys)
|
||||
pts_xy = np.column_stack([xx.ravel(), yy.ravel()]) # [N, 2]
|
||||
return np.column_stack([np.full(len(pts_xy), float(frame_t)), pts_xy]).astype(np.float32)
|
||||
|
||||
|
||||
|
||||
def fastsam_masks_batch(sam_model, frames_rgb: list[np.ndarray], conf: float, iou: float,
|
||||
imgsz: int, H: int, W: int, batch: int = 16) -> list[np.ndarray]:
|
||||
"""Run FastSAM on uint8 HxWx3 RGB frames in batched forwards, return bool masks [M,H,W] per frame."""
|
||||
out: list[np.ndarray] = []
|
||||
for s in range(0, len(frames_rgb), max(1, batch)):
|
||||
res = sam_model(frames_rgb[s:s + max(1, batch)], device="cuda", retina_masks=True,
|
||||
imgsz=imgsz, conf=conf, iou=iou, verbose=False)
|
||||
for r in res:
|
||||
if r.masks is None:
|
||||
out.append(np.zeros((0, H, W), bool))
|
||||
continue
|
||||
masks = r.masks.data.cpu().numpy().astype(bool)
|
||||
if masks.shape[0] and masks.shape[1:] != (H, W):
|
||||
import cv2
|
||||
masks = np.stack([
|
||||
cv2.resize(m.astype(np.uint8), (W, H), interpolation=cv2.INTER_NEAREST).astype(bool)
|
||||
for m in masks
|
||||
])
|
||||
out.append(masks)
|
||||
return out
|
||||
|
||||
|
||||
def max_iou_with_set(mask: np.ndarray, others: np.ndarray) -> float:
|
||||
"""Max IOU of a single [H,W] bool mask against a set [M,H,W]."""
|
||||
if others.shape[0] == 0:
|
||||
return 0.0
|
||||
inter = (mask & others).reshape(others.shape[0], -1).sum(1)
|
||||
union = (mask | others).reshape(others.shape[0], -1).sum(1)
|
||||
return float(np.where(union > 0, inter / union, 0.0).max())
|
||||
|
||||
|
||||
def detect_entry_events(video: torch.Tensor, sam_model, H: int, W: int, conf: float, iou: float,
|
||||
imgsz: int, sample_every: int, new_area_thresh: float,
|
||||
min_area_frac: float, sam_batch: int = 16
|
||||
) -> tuple[dict[int, np.ndarray], np.ndarray, dict[int, np.ndarray]]:
|
||||
"""Return ({frame_t: new_masks [M,H,W]}, masks0_union [H,W], {frame_t: all masks}) for
|
||||
frames where new regions appear that weren't covered by any frame-0 mask. The third
|
||||
element caches every segmented frame's raw masks for reuse by the fused --segment pass.
|
||||
|
||||
Uses new-area fraction rather than IOU so partially-entering objects (partially visible
|
||||
at frame 0, more visible later) are correctly detected: only the newly visible region
|
||||
triggers detection and receives replacement tracks.
|
||||
"""
|
||||
T = video.shape[1]
|
||||
min_area_px = int(min_area_frac * H * W)
|
||||
|
||||
# One batched FastSAM sweep over frame 0 + all sampled frames (the per-frame
|
||||
# mask filtering below is sequential, but the model forwards are independent).
|
||||
sample_ts = list(range(sample_every, T, sample_every))
|
||||
frames = [video[0, t].permute(1, 2, 0).numpy().astype(np.uint8) for t in [0, *sample_ts]]
|
||||
all_masks = fastsam_masks_batch(sam_model, frames, conf, iou, imgsz, H, W, batch=sam_batch)
|
||||
masks_by_frame = dict(zip([0, *sample_ts], all_masks))
|
||||
masks0 = all_masks[0]
|
||||
# Exclude large background masks (table surface, floor, walls) from masks0_union.
|
||||
# Only object-sized masks define "known territory" — background covers the whole frame
|
||||
# and would suppress detection of the ball moving to a new position on that surface.
|
||||
if masks0.shape[0] > 0:
|
||||
areas = masks0.reshape(masks0.shape[0], -1).sum(1)
|
||||
object_masks0 = masks0[areas < 0.4 * H * W]
|
||||
masks0_union = object_masks0.any(axis=0) if object_masks0.shape[0] > 0 else np.zeros((H, W), bool)
|
||||
else:
|
||||
masks0_union = np.zeros((H, W), bool)
|
||||
|
||||
# Tracks newly-covered regions across entry events to avoid re-detecting the same
|
||||
# entering object at multiple sample frames.
|
||||
claimed_new = np.zeros((H, W), bool)
|
||||
|
||||
entry_events: dict[int, np.ndarray] = {}
|
||||
for frame_t, masks_t in zip(sample_ts, all_masks[1:]):
|
||||
if masks_t.shape[0] == 0:
|
||||
continue
|
||||
|
||||
new_masks = []
|
||||
for m in masks_t:
|
||||
area = int(m.sum())
|
||||
if area < min_area_px:
|
||||
continue
|
||||
# New area: part of this mask not present in ANY frame-0 mask.
|
||||
new_region = m & ~masks0_union
|
||||
new_area_frac = float(new_region.sum()) / max(area, 1)
|
||||
if new_area_frac < new_area_thresh:
|
||||
continue
|
||||
# Skip if this new region was already claimed by a prior entry event.
|
||||
if float((new_region & claimed_new).sum()) / max(int(new_region.sum()), 1) > 0.5:
|
||||
continue
|
||||
new_masks.append(m)
|
||||
# Dilate claimed region by ~2x the object radius so a fast-moving
|
||||
# object doesn't re-trigger entry detection on subsequent sample frames.
|
||||
radius = int(np.sqrt(int(new_region.sum()) / np.pi))
|
||||
dil = max(1, radius * 2)
|
||||
ys, xs = np.where(new_region)
|
||||
y0, y1 = max(0, ys.min() - dil), min(H, ys.max() + dil + 1)
|
||||
x0, x1 = max(0, xs.min() - dil), min(W, xs.max() + dil + 1)
|
||||
claimed_new[y0:y1, x0:x1] = True
|
||||
|
||||
if new_masks:
|
||||
entry_events[frame_t] = np.stack(new_masks)
|
||||
|
||||
return entry_events, masks0_union, masks_by_frame
|
||||
|
||||
|
||||
def decoded_videos(todo: list[tuple[int, Path, Path]], prefetch: int):
|
||||
"""Yield (k, vpath, out_path, video, h, w), decoding up to `prefetch` videos ahead
|
||||
on a background thread (decord/av release the GIL, so decode overlaps GPU compute)."""
|
||||
if prefetch <= 0:
|
||||
for k, vpath, out_path in todo:
|
||||
yield (k, vpath, out_path, *read_video(vpath))
|
||||
return
|
||||
|
||||
q: queue.Queue = queue.Queue(maxsize=prefetch)
|
||||
|
||||
def producer() -> None:
|
||||
for k, vpath, out_path in todo:
|
||||
try:
|
||||
payload = read_video(vpath)
|
||||
except Exception as e: # noqa: BLE001
|
||||
payload = e
|
||||
q.put((k, vpath, out_path, payload))
|
||||
q.put(None)
|
||||
|
||||
threading.Thread(target=producer, daemon=True).start()
|
||||
while (item := q.get()) is not None:
|
||||
k, vpath, out_path, payload = item
|
||||
if isinstance(payload, Exception):
|
||||
# One unreadable clip must not take down the worker and everything it has
|
||||
# left to process -- real-world shards contain corrupt/truncated files.
|
||||
print(f"[track] [{k}] {vpath.name}: DECODE FAILED ({payload}), skipping", flush=True)
|
||||
continue
|
||||
yield (k, vpath, out_path, *payload)
|
||||
|
||||
|
||||
def patch_manifest(manifest_path: Path, stem_to_points: dict[str, Path]) -> int:
|
||||
if not manifest_path.exists():
|
||||
return 0
|
||||
items = json.loads(manifest_path.read_text())
|
||||
patched = 0
|
||||
for item in items:
|
||||
stem = Path(item.get("path", "")).stem
|
||||
pts = stem_to_points.get(stem)
|
||||
if pts is not None:
|
||||
item["points_path"] = str(pts.resolve())
|
||||
patched += 1
|
||||
tmp = manifest_path.with_suffix(f".json.tmp{os.getpid()}") # pid-unique: never shared
|
||||
tmp.write_text(json.dumps(items, indent=2))
|
||||
tmp.replace(manifest_path)
|
||||
return patched
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
videos_dir = args.data_dir / args.videos_subdir
|
||||
out_dir = args.data_dir / args.out_subdir
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
videos = sorted(videos_dir.glob("*.mp4"))
|
||||
if args.index is not None:
|
||||
wanted = {f"vid_{i:06d}.mp4" for i in args.index}
|
||||
videos = [v for v in videos if v.name in wanted]
|
||||
if args.limit is not None:
|
||||
videos = videos[:args.limit]
|
||||
all_videos = videos # pre-shard list: only rank 0 patches the manifest,
|
||||
if args.world_size > 1: # and it must cover every rank's videos
|
||||
videos = videos[args.rank::args.world_size]
|
||||
if not videos:
|
||||
print(f"[track] no videos found in {videos_dir}", flush=True)
|
||||
return
|
||||
|
||||
cotracker = load_cotracker(args.model, args.device)
|
||||
# Compile only the fixed-shape main pass; entry-pass query counts vary per video and
|
||||
# would recompile every time, so track_with_queries keeps the eager model.
|
||||
cotracker_main = torch.compile(cotracker) if args.compile else cotracker
|
||||
print(f"[track] loaded {args.model}; {len(videos)} videos, grid={args.grid_size}x{args.grid_size}"
|
||||
f"{' (compiled main pass)' if args.compile else ''}", flush=True)
|
||||
|
||||
sam_model = None
|
||||
if args.detect_entries or args.segment:
|
||||
from ultralytics import FastSAM
|
||||
sam_model = FastSAM(args.sam_model)
|
||||
print(f"[track] entries={args.detect_entries} segment={args.segment}: sam={args.sam_model}, "
|
||||
f"conf={args.sam_conf}, every={args.entry_sample_every}f", flush=True)
|
||||
|
||||
stem_to_fps: dict[str, float] = {}
|
||||
if args.viz:
|
||||
mpath = args.data_dir / args.manifest
|
||||
items = json.loads(mpath.read_text()) if mpath.exists() else []
|
||||
stem_to_fps = {Path(it.get("path", "")).stem: it.get("fps", 24) for it in items}
|
||||
|
||||
# points_path is deterministic, so rank 0 can record it for every video without
|
||||
# doing their work -- avoids all ranks rewriting the manifest at once.
|
||||
stem_to_points = {v.stem: out_dir / f"{v.stem}.npz" for v in all_videos}
|
||||
todo: list[tuple[int, Path, Path]] = []
|
||||
for k, vpath in enumerate(videos, 1):
|
||||
out_path = out_dir / f"{vpath.stem}.npz"
|
||||
if out_path.exists() and not args.force:
|
||||
continue
|
||||
todo.append((k, vpath, out_path))
|
||||
|
||||
for k, vpath, out_path, video, h, w in decoded_videos(todo, args.prefetch):
|
||||
tracks, vis = track_one(cotracker_main, video, args.grid_size, args.downscale, args.device, amp=args.amp)
|
||||
|
||||
n_replaced = 0
|
||||
masks_cache: dict[int, np.ndarray] = {}
|
||||
if args.detect_entries:
|
||||
entry_events, masks0_union, masks_cache = detect_entry_events(
|
||||
video, sam_model, h, w,
|
||||
conf=args.sam_conf, iou=args.sam_iou, imgsz=args.sam_imgsz,
|
||||
sample_every=args.entry_sample_every,
|
||||
new_area_thresh=args.entry_new_area,
|
||||
min_area_frac=args.entry_min_area,
|
||||
sam_batch=args.sam_batch,
|
||||
)
|
||||
|
||||
# Build queries for all (entry frame, mask) pairs at once: only grid points
|
||||
# landing inside each new region are tracked, instead of a full grid_size^2
|
||||
# pass per mask. Groups keep disjoint column ranges [qa, qb) in the batch.
|
||||
groups: list[tuple[int, int, np.ndarray, np.ndarray, int, int]] = []
|
||||
q_parts: list[np.ndarray] = []
|
||||
n_q = 0
|
||||
for frame_t, new_masks in sorted(entry_events.items()):
|
||||
grid_q = make_grid_queries(args.grid_size, h, w, frame_t)
|
||||
gxi = np.clip(grid_q[:, 1].round().astype(int), 0, w - 1)
|
||||
gyi = np.clip(grid_q[:, 2].round().astype(int), 0, h - 1)
|
||||
for mi, new_mask in enumerate(new_masks):
|
||||
new_region = new_mask & ~masks0_union
|
||||
q = grid_q[new_region[gyi, gxi]]
|
||||
groups.append((frame_t, mi, new_mask, new_region, n_q, n_q + len(q)))
|
||||
q_parts.append(q)
|
||||
n_q += len(q)
|
||||
|
||||
e_tracks = e_vis = None
|
||||
if n_q > 0:
|
||||
all_q = np.concatenate(q_parts, axis=0)
|
||||
max_q = args.grid_size * args.grid_size # bound memory to the main pass
|
||||
et_parts, ev_parts = [], []
|
||||
for s in range(0, n_q, max_q):
|
||||
ct, cv = track_with_queries(
|
||||
cotracker, video, all_q[s:s + max_q], args.downscale,
|
||||
args.device, h, w, amp=args.amp)
|
||||
et_parts.append(ct)
|
||||
ev_parts.append(cv)
|
||||
e_tracks = np.concatenate(et_parts, axis=1)
|
||||
e_vis = np.concatenate(ev_parts, axis=1)
|
||||
|
||||
for frame_t, mi, new_mask, new_region, qa, qb in groups:
|
||||
# Covered slots: original frame-0 grid points whose position at
|
||||
# T_entry falls inside the new object's region.
|
||||
orig_xi = np.clip(tracks[frame_t, :, 0].round().astype(int), 0, w - 1)
|
||||
orig_yi = np.clip(tracks[frame_t, :, 1].round().astype(int), 0, h - 1)
|
||||
covered_dst = np.where(new_region[orig_yi, orig_xi])[0]
|
||||
dead_dst = covered_dst if covered_dst.size > 0 else \
|
||||
np.where((~vis[frame_t:].astype(bool)).all(axis=0))[0]
|
||||
if dead_dst.size == 0:
|
||||
continue
|
||||
|
||||
n_src = qb - qa
|
||||
n = min(n_src, len(dead_dst))
|
||||
if n > 0:
|
||||
e_vis[:frame_t, qa:qb] = False # object didn't exist before entry frame
|
||||
dst = dead_dst[:n]
|
||||
tracks[:, dst] = e_tracks[:, qa:qa + n]
|
||||
vis[:, dst] = e_vis[:, qa:qa + n]
|
||||
# CoTracker can mark points invisible even at their query frame
|
||||
# when the object is entering from the edge. Force vis=True at
|
||||
# T_entry so segment_tracks sees these as first-visible there.
|
||||
vis[frame_t, dst] = True
|
||||
n_replaced += n
|
||||
if args.verbose:
|
||||
mask_area = int(new_mask.sum())
|
||||
new_region_area = int(new_region.sum())
|
||||
vis_at_entry = int(vis[frame_t, dst].sum()) if n > 0 else 0
|
||||
print(f" [entry] t={frame_t} mask#{mi}: area={mask_area}px "
|
||||
f"new_region={new_region_area}px "
|
||||
f"object_src={n_src} covered_dst={len(covered_dst)} "
|
||||
f"dead_dst={len(dead_dst)} -> replacing {n}, "
|
||||
f"vis[{frame_t}, dst].sum()={vis_at_entry}", flush=True)
|
||||
|
||||
frames_str = ",".join(str(t) for t in sorted(entry_events)) if entry_events else "none"
|
||||
print(f"[track] [{k}/{len(videos)}] {vpath.name}: "
|
||||
f"{len(entry_events)} entry event(s) at frames [{frames_str}], "
|
||||
f"{n_replaced} slots replaced", flush=True)
|
||||
|
||||
# Fused Stage 4: object IDs + vis override + track weights, reusing the decoded
|
||||
# video and any masks already computed by entry detection.
|
||||
seg_extra: dict[str, np.ndarray] = {}
|
||||
if args.segment:
|
||||
import segment_tracks as seg
|
||||
|
||||
def get_masks(frame_ts: list[int], _video=video, _cache=masks_cache,
|
||||
_h=h, _w=w) -> dict[int, np.ndarray]:
|
||||
missing = [t for t in frame_ts if t not in _cache]
|
||||
if missing:
|
||||
frames_np = [_video[0, t].permute(1, 2, 0).numpy().astype(np.uint8) for t in missing]
|
||||
new_masks = fastsam_masks_batch(sam_model, frames_np, args.sam_conf, args.sam_iou,
|
||||
args.sam_imgsz, _h, _w, batch=args.sam_batch)
|
||||
_cache.update(zip(missing, new_masks))
|
||||
return {t: _cache[t] for t in frame_ts}
|
||||
|
||||
vis = vis.astype(bool)
|
||||
oid, n_objects, vis, weights, n_overrides, uframes = seg.segment_tracks_arrays(
|
||||
tracks, vis, h, w, get_masks, args.vis_override_every, verbose=args.verbose)
|
||||
seg_extra = dict(object_ids=oid, n_objects=np.int64(n_objects), track_weights=weights)
|
||||
print(f"[track] [{k}/{len(videos)}] {vpath.name} seg: {n_objects} objs across "
|
||||
f"{len(uframes)} frames, {int((oid >= 0).sum())}/{oid.shape[0]} pts labeled, "
|
||||
f"{n_overrides} vis overrides", flush=True)
|
||||
|
||||
if args.viz:
|
||||
viz_dir = Path(args.viz_dir) if args.viz_dir else (args.data_dir / "viz")
|
||||
stem_dir = viz_dir / vpath.stem
|
||||
stem_dir.mkdir(parents=True, exist_ok=True)
|
||||
frames_np = video[0].permute(0, 2, 3, 1).numpy().astype(np.uint8) # [T,H,W,3]
|
||||
seg.render_viz(frames_np[:tracks.shape[0]], tracks, vis, oid, stem_dir / "tracks.mp4",
|
||||
fps=int(stem_to_fps.get(vpath.stem, 24)))
|
||||
print(f" viz -> {stem_dir}/", flush=True)
|
||||
|
||||
tmp = out_path.with_name(out_path.stem + ".tmp.npz")
|
||||
np.savez(
|
||||
tmp, tracks=tracks, visibility=vis,
|
||||
grid_size=args.grid_size, height=h, width=w,
|
||||
num_frames=tracks.shape[0],
|
||||
**seg_extra,
|
||||
)
|
||||
tmp.replace(out_path)
|
||||
print(f"[track] [{k}/{len(videos)}] {vpath.name} -> {out_path.name} "
|
||||
f"tracks={tracks.shape} vis={vis.shape}", flush=True)
|
||||
|
||||
# Only rank 0 writes: concurrent read-modify-write from every rank corrupted the
|
||||
# manifest (interleaved writes to a shared temp file -> invalid JSON).
|
||||
if args.rank == 0:
|
||||
n = patch_manifest(args.data_dir / args.manifest, stem_to_points)
|
||||
print(f"[track] done; patched points_path into {n} manifest entries", flush=True)
|
||||
else:
|
||||
print(f"[track] done (rank {args.rank}; manifest patched by rank 0)", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,145 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Data-parallel CoTracker v3 track extraction for large video sets (OpenVid-1M).
|
||||
|
||||
CoTracker batching does NOT help (s/video is flat, memory linear), so we scale by
|
||||
DATA PARALLELISM: launch many single-video workers, several per GPU, across nodes.
|
||||
Each worker processes video_list[shard::num_shards] on one pinned GPU and is idempotent.
|
||||
|
||||
Sharding + GPU pinning come from Slurm env by default:
|
||||
shard = SLURM_PROCID (0..num_shards-1)
|
||||
num_shards = SLURM_NTASKS
|
||||
gpu = SLURM_LOCALID % gpus_per_node
|
||||
so `srun --ntasks-per-node=(gpus*procs_per_gpu)` gives procs_per_gpu workers per GPU.
|
||||
|
||||
Resamples each clip to --fps / --num-frames and resizes to --height x --width (720p)
|
||||
before tracking, matching the training spec.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, os, time
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
HUB = "facebookresearch/co-tracker"
|
||||
|
||||
|
||||
def load_cotracker(device: str):
|
||||
local = Path(torch.hub.get_dir()) / (HUB.replace("/", "_") + "_main")
|
||||
if local.exists():
|
||||
m = torch.hub.load(str(local), "cotracker3_offline", source="local", trust_repo=True)
|
||||
else:
|
||||
m = torch.hub.load(HUB, "cotracker3_offline", trust_repo=True)
|
||||
return m.to(device).eval()
|
||||
|
||||
|
||||
def read_clip(path: str, fps: int, num_frames: int, H: int, W: int,
|
||||
min_height: int = 0, aspect: float | None = None, aspect_tol: float = 0.1):
|
||||
"""Decode (PyAV — decord has no aarch64 wheels, opencv needs GUI libs), resample
|
||||
to `fps`, take `num_frames`, resize to HxW. Sequential decode so it is codec-robust
|
||||
for OpenVid's varied encodings. Enforces resolution/aspect by probing stream metadata
|
||||
(fast, no decode). Returns (1,T,C,H,W) on success, or a reason string
|
||||
('lowres' | 'aspect' | 'short') on reject."""
|
||||
import av
|
||||
container = av.open(str(path))
|
||||
try:
|
||||
stream = container.streams.video[0]
|
||||
nh, nw = int(stream.height or 0), int(stream.width or 0)
|
||||
if min_height and nh and nh < min_height:
|
||||
return "lowres"
|
||||
if aspect is not None and nw and nh and abs((nw / nh) - aspect) > aspect_tol:
|
||||
return "aspect"
|
||||
native = float(stream.average_rate) if stream.average_rate else float(fps)
|
||||
step = native / float(fps)
|
||||
idxs = [int(round(i * step)) for i in range(num_frames)]
|
||||
picked, wi, cur = [], 0, 0
|
||||
for frame in container.decode(video=0):
|
||||
img = None
|
||||
while wi < len(idxs) and idxs[wi] == cur: # handles repeated idxs (fps up-sample)
|
||||
if img is None:
|
||||
img = frame.to_ndarray(format="rgb24") # (H,W,3) uint8
|
||||
picked.append(img); wi += 1
|
||||
cur += 1
|
||||
if wi >= len(idxs):
|
||||
break
|
||||
finally:
|
||||
container.close()
|
||||
if wi < len(idxs):
|
||||
return "short" # video too short at target fps
|
||||
arr = np.ascontiguousarray(np.stack(picked)) # (T,H,W,C) RGB uint8
|
||||
vid = torch.from_numpy(arr).permute(0, 3, 1, 2).unsqueeze(0).float() # 1,T,C,h,w
|
||||
vid = F.interpolate(vid[0], size=(H, W), mode="bilinear", align_corners=False).unsqueeze(0)
|
||||
return vid
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def track(model, video, grid, device):
|
||||
tr, vis = model(video.to(device), grid_size=grid)
|
||||
return tr[0].float().cpu().numpy(), vis[0].cpu().numpy()
|
||||
|
||||
|
||||
def save_clip(video, path, fps):
|
||||
"""Write the resampled 121f/720p clip so the training video IS the tracked frames."""
|
||||
import imageio
|
||||
arr = video[0].permute(0, 2, 3, 1).clamp(0, 255).byte().cpu().numpy() # (T,C,H,W)->(T,H,W,C)
|
||||
imageio.mimwrite(path, arr, fps=fps, codec="libx264", quality=7, macro_block_size=1)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--video-list", required=True, help="Text file: one video path per line.")
|
||||
ap.add_argument("--out-dir", required=True, help="Where <stem>.npz tracks are written.")
|
||||
ap.add_argument("--grid-size", type=int, default=50)
|
||||
ap.add_argument("--num-shards", type=int, default=int(os.environ.get("SLURM_NTASKS", "1")))
|
||||
ap.add_argument("--shard", type=int, default=int(os.environ.get("SLURM_PROCID", "0")))
|
||||
ap.add_argument("--gpus-per-node", type=int, default=4)
|
||||
ap.add_argument("--fps", type=int, default=24)
|
||||
ap.add_argument("--num-frames", type=int, default=121)
|
||||
ap.add_argument("--height", type=int, default=720)
|
||||
ap.add_argument("--width", type=int, default=1280)
|
||||
ap.add_argument("--min-height", type=int, default=720, help="skip clips whose native height < this (0=off)")
|
||||
ap.add_argument("--aspect", type=float, default=16.0 / 9.0, help="target native W/H aspect")
|
||||
ap.add_argument("--aspect-tol", type=float, default=0.12, help="allowed |native_aspect - target| (large=off)")
|
||||
ap.add_argument("--clips-dir", default=None, help="also save the resampled 121f/720p clip here (aligned to tracks)")
|
||||
a = ap.parse_args()
|
||||
|
||||
local = int(os.environ.get("SLURM_LOCALID", str(a.shard)))
|
||||
gpu = local % a.gpus_per_node
|
||||
torch.cuda.set_device(gpu)
|
||||
device = f"cuda:{gpu}"
|
||||
Path(a.out_dir).mkdir(parents=True, exist_ok=True)
|
||||
if a.clips_dir:
|
||||
Path(a.clips_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
vids = [l.strip() for l in open(a.video_list) if l.strip()]
|
||||
mine = vids[a.shard::a.num_shards]
|
||||
model = load_cotracker(device)
|
||||
print(f"[shard {a.shard}/{a.num_shards}] gpu={gpu} localid={local} videos={len(mine)}", flush=True)
|
||||
|
||||
t0 = time.time(); done = 0; err = 0; rej = {"short": 0, "lowres": 0, "aspect": 0}
|
||||
for vp in mine:
|
||||
out = Path(a.out_dir) / f"{Path(vp).stem}.npz"
|
||||
if out.exists():
|
||||
continue
|
||||
try:
|
||||
res = read_clip(vp, a.fps, a.num_frames, a.height, a.width,
|
||||
a.min_height, a.aspect, a.aspect_tol)
|
||||
if isinstance(res, str):
|
||||
rej[res] = rej.get(res, 0) + 1; continue
|
||||
tr, vis = track(model, res, a.grid_size, device)
|
||||
if a.clips_dir:
|
||||
save_clip(res, str(Path(a.clips_dir) / f"{Path(vp).stem}.mp4"), a.fps)
|
||||
tmp = out.with_suffix(".tmp.npz")
|
||||
np.savez(tmp, tracks=tr.astype(np.float32), visibility=vis, grid_size=a.grid_size,
|
||||
height=a.height, width=a.width, num_frames=tr.shape[0], fps=a.fps)
|
||||
tmp.replace(out); done += 1
|
||||
except Exception as e: # noqa: BLE001
|
||||
err += 1
|
||||
print(f"[shard {a.shard}] ERR {Path(vp).name}: {repr(e)[:110]}", flush=True)
|
||||
dt = time.time() - t0
|
||||
rate = done / dt if dt > 0 else 0.0
|
||||
print(f"[shard {a.shard}] DONE done={done} rej={rej} err={err} in {dt:.1f}s -> {rate:.3f} vid/s", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Select clips that already match the target geometry; symlink them for downstream stages.
|
||||
|
||||
Real shards are usually uniform but not guaranteed. Rather than re-encoding every clip to
|
||||
force conformity (a lossy no-op when the clip already matches -- measured 40 dB on the
|
||||
OpenVid shard, worse than the VAE round-trip's own distortion), this scans container
|
||||
metadata (fast: no frame decode) and links through only the clips that conform.
|
||||
|
||||
Non-conforming clips are reported and listed in ``skipped_clips.json`` so nothing is
|
||||
silently dropped -- re-run them through ``resize_videos.py`` if you want them included.
|
||||
|
||||
python data_pipeline/filter_clips.py \\
|
||||
--src-dir <root>/raw_videos --out-dir <root>/videos \\
|
||||
--height 720 --width 1280 --num-frames 121
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import cv2
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--src-dir", type=Path, required=True)
|
||||
p.add_argument("--out-dir", type=Path, required=True, help="Symlinks to conforming clips land here.")
|
||||
p.add_argument("--height", type=int, required=True)
|
||||
p.add_argument("--width", type=int, required=True)
|
||||
p.add_argument("--num-frames", type=int, default=121,
|
||||
help="Required exact frame count (0 = don't check).")
|
||||
p.add_argument("--report", type=str, default="skipped_clips.json",
|
||||
help="Written next to --out-dir; lists every skipped clip and why.")
|
||||
p.add_argument("--needs-resize-list", type=str, default="needs_resize.txt",
|
||||
help="Written next to --out-dir; names of clips that are readable but at the "
|
||||
"wrong geometry, i.e. rescuable by resize_videos.py --include-list. "
|
||||
"Unreadable clips are excluded (nothing can rescue those).")
|
||||
p.add_argument("--clean", action="store_true",
|
||||
help="Empty --out-dir of *.mp4 first (links AND regular files -- a stale real "
|
||||
"file would otherwise shadow the link and be used silently). --out-dir is "
|
||||
"a derived directory; never point it at original footage.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def probe(path: Path) -> tuple[int, int, int]:
|
||||
"""Return (width, height, n_frames) from container metadata; (0,0,0) if unreadable."""
|
||||
cap = cv2.VideoCapture(str(path))
|
||||
if not cap.isOpened():
|
||||
return (0, 0, 0)
|
||||
wh = (int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)), int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)),
|
||||
int(cap.get(cv2.CAP_PROP_FRAME_COUNT)))
|
||||
cap.release()
|
||||
return wh
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
args.out_dir.mkdir(parents=True, exist_ok=True)
|
||||
if args.clean:
|
||||
for old in args.out_dir.glob("*.mp4"):
|
||||
old.unlink()
|
||||
|
||||
clips = sorted(args.src_dir.glob("*.mp4"))
|
||||
ok = 0
|
||||
skipped: list[dict] = []
|
||||
for c in clips:
|
||||
w, h, n = probe(c)
|
||||
if (w, h) == (0, 0):
|
||||
skipped.append({"clip": c.name, "reason": "unreadable"})
|
||||
continue
|
||||
if (w, h) != (args.width, args.height):
|
||||
skipped.append({"clip": c.name, "reason": "resolution", "got": f"{w}x{h}"})
|
||||
continue
|
||||
if args.num_frames and n != args.num_frames:
|
||||
skipped.append({"clip": c.name, "reason": "frames", "got": n})
|
||||
continue
|
||||
link = args.out_dir / c.name
|
||||
target = c.resolve()
|
||||
if link.is_symlink() and link.readlink() == target:
|
||||
pass # already correct
|
||||
else:
|
||||
if link.exists() or link.is_symlink():
|
||||
link.unlink() # replace a stale file/link rather than trusting it
|
||||
link.symlink_to(target)
|
||||
ok += 1
|
||||
|
||||
report_path = args.out_dir.parent / args.report
|
||||
report_path.write_text(json.dumps(
|
||||
{"src": str(args.src_dir), "required": f"{args.width}x{args.height}@{args.num_frames}f",
|
||||
"total": len(clips), "kept": ok, "skipped": skipped}, indent=2))
|
||||
|
||||
rescuable = [s["clip"] for s in skipped if s["reason"] in ("resolution", "frames")]
|
||||
list_path = args.out_dir.parent / args.needs_resize_list
|
||||
list_path.write_text("\n".join(rescuable) + ("\n" if rescuable else ""))
|
||||
|
||||
by_reason: dict[str, int] = {}
|
||||
for s in skipped:
|
||||
by_reason[s["reason"]] = by_reason.get(s["reason"], 0) + 1
|
||||
detail = ", ".join(f"{k}={v}" for k, v in sorted(by_reason.items())) or "none"
|
||||
print(f"[filter] {ok}/{len(clips)} clips match {args.width}x{args.height}"
|
||||
f"@{args.num_frames}f -> {args.out_dir}", flush=True)
|
||||
print(f"[filter] skipped: {detail} (details in {report_path})", flush=True)
|
||||
if rescuable:
|
||||
print(f"[filter] {len(rescuable)} rescuable by resize -> {list_path}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,353 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Find relatively-static WINDOWS inside (long) videos, before tracking.
|
||||
|
||||
The camera in egocentric video turns and the whole view sweeps off-frame; a filter that
|
||||
seeds points once and tracks the whole clip would die at the first turn and mis-score
|
||||
everything after. Instead we compute a PER-FRAME global-motion signal that RE-SEEDS every
|
||||
frame (so it recovers after a turn: motion spikes during the turn, drops when static), then
|
||||
slide a window and locate the calm stretches.
|
||||
|
||||
Per-frame motion m[t] = median Lucas-Kanade optical-flow magnitude between frames t-1 and t
|
||||
(fresh goodFeaturesToTrack each step), normalized by the frame diagonal. Camera pan/turn ->
|
||||
large; static -> small. A window is static iff m[t] stays low across ALL its frames (we score
|
||||
by a high percentile so one calm-but-not-perfect frame is fine but a turn is not). We also
|
||||
report per-window SURVIVAL: seed a grid at the window START, LK-track to the window END, and
|
||||
measure the fraction still tracked & in-frame -> "would CoTracker keep traces here?".
|
||||
|
||||
.venv/bin/python data_pipeline/find_static_windows.py scan --videos-dir <dir> --out <dir> --window-sec 5 --limit 12
|
||||
.venv/bin/python data_pipeline/find_static_windows.py serve --viz-dir <dir> --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
|
||||
def read_gray_lowres(path: str, max_frames: int = 900, target_w: int = 320) -> tuple[list, float, int]:
|
||||
"""Decode DIRECTLY at low res (fast + tiny memory) -> gray frames + fps + total_frames.
|
||||
Only used for the motion analysis; full-res frames are read on demand for extraction."""
|
||||
import cv2
|
||||
from decord import VideoReader, cpu
|
||||
cap = cv2.VideoCapture(path) # fast metadata read (no full decode)
|
||||
W0 = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) or 640
|
||||
H0 = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) or 360
|
||||
fps = float(cap.get(cv2.CAP_PROP_FPS)) or 30.0
|
||||
cap.release()
|
||||
tw = min(target_w, W0)
|
||||
th = int(round(H0 * tw / max(W0, 1)))
|
||||
try:
|
||||
vr = VideoReader(path, ctx=cpu(0), width=tw, height=th)
|
||||
except Exception: # noqa: BLE001 (older decord without width/height)
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
n_total = len(vr)
|
||||
n = min(n_total, max_frames)
|
||||
batch = vr.get_batch(list(range(n))).asnumpy()
|
||||
gray = [cv2.cvtColor(f, cv2.COLOR_RGB2GRAY) for f in batch]
|
||||
return gray, fps, n_total
|
||||
|
||||
|
||||
def read_rgb_frames(path: str, indices: Any) -> np.ndarray:
|
||||
"""Full-res RGB for a specific set of frame indices (decord random access)."""
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(indices)).asnumpy()
|
||||
|
||||
|
||||
def motion_series(gray: list) -> tuple[np.ndarray, np.ndarray, float]:
|
||||
"""Per-frame normalized global camera motion via phase correlation (FFT-based global
|
||||
shift between consecutive frames, ~1ms/frame). Recovers after turns: spikes during a
|
||||
turn, drops when static. Returns (m, shifts, diag): m[t] = |shift|/diag, shifts[t] =
|
||||
signed (dx,dy)/diag (used to accumulate net view drift over a window)."""
|
||||
import cv2
|
||||
H, W = gray[0].shape
|
||||
diag = float(np.hypot(H, W))
|
||||
hann = cv2.createHanningWindow((W, H), cv2.CV_32F)
|
||||
m = np.zeros(len(gray), np.float32)
|
||||
shifts = np.zeros((len(gray), 2), np.float32)
|
||||
prev = gray[0].astype(np.float32)
|
||||
for t in range(1, len(gray)):
|
||||
cur = gray[t].astype(np.float32)
|
||||
(dx, dy), _ = cv2.phaseCorrelate(prev, cur, hann)
|
||||
shifts[t] = (dx / diag, dy / diag)
|
||||
m[t] = float(np.hypot(dx, dy)) / diag
|
||||
prev = cur
|
||||
m[0] = m[1] if len(m) > 1 else 0.0
|
||||
return m, shifts, diag
|
||||
|
||||
|
||||
def window_drift(shifts: np.ndarray, s: int, L: int) -> float:
|
||||
"""Net view drift over a window: max cumulative excursion of the accumulated per-frame
|
||||
shifts, normalized by diagonal. Large => the camera slowly panned the view off-screen
|
||||
even if no single frame moved much. Free (from the phase-correlation shifts)."""
|
||||
c = np.cumsum(shifts[s:s + L], axis=0)
|
||||
return float(np.max(np.hypot(c[:, 0], c[:, 1]))) if len(c) else 0.0
|
||||
|
||||
|
||||
def window_survival(gray: list, s: int, L: int, grid: int = 24) -> float:
|
||||
"""Seed a grid at frame s, LK-track to s+L-1; fraction still tracked & in-frame."""
|
||||
import cv2
|
||||
H, W = gray[0].shape
|
||||
xs = np.linspace(W * 0.05, W * 0.95, grid)
|
||||
ys = np.linspace(H * 0.05, H * 0.95, grid)
|
||||
gx, gy = np.meshgrid(xs, ys)
|
||||
pts = np.stack([gx.ravel(), gy.ravel()], 1).astype(np.float32).reshape(-1, 1, 2)
|
||||
alive = np.ones(pts.shape[0], bool)
|
||||
cur = pts.copy()
|
||||
lk = dict(winSize=(21, 21), maxLevel=3)
|
||||
for t in range(s + 1, min(s + L, len(gray))):
|
||||
nxt, stt, _ = cv2.calcOpticalFlowPyrLK(gray[t - 1], gray[t], cur, None, **lk)
|
||||
stt = stt.ravel().astype(bool)
|
||||
x, y = nxt[:, 0, 0], nxt[:, 0, 1]
|
||||
inb = (x >= 0) & (x < W) & (y >= 0) & (y < H)
|
||||
alive &= stt & inb
|
||||
cur = nxt
|
||||
return float(alive.mean())
|
||||
|
||||
|
||||
def best_windows(m: np.ndarray, L: int, pct: float = 95, top: int = 1, min_gap: int | None = None) -> list:
|
||||
"""Return up to `top` low-motion windows (start index) minimizing the pct-percentile of
|
||||
per-frame motion, non-overlapping by min_gap (default L)."""
|
||||
if len(m) < L:
|
||||
return []
|
||||
min_gap = min_gap or L
|
||||
scores = np.array([np.percentile(m[s:s + L], pct) for s in range(len(m) - L + 1)])
|
||||
order = np.argsort(scores)
|
||||
picks: list[int] = []
|
||||
for s in order:
|
||||
if all(abs(s - p) >= min_gap for p in picks):
|
||||
picks.append(int(s))
|
||||
if len(picks) >= top:
|
||||
break
|
||||
return picks
|
||||
|
||||
|
||||
def greedy_valid_windows(m: np.ndarray,
|
||||
shifts: np.ndarray,
|
||||
L: int,
|
||||
pct: float,
|
||||
p95_thresh: float,
|
||||
drift_thresh: float,
|
||||
cand_step: int | None = None) -> list:
|
||||
"""Greedily take ALL non-overlapping windows passing p95_motion<=p95_thresh AND
|
||||
net-drift<=drift_thresh, calmest-first. Returns list of (start, p95, drift). All metrics
|
||||
come from the phase-correlation signal (no LK), so this is cheap."""
|
||||
if len(m) < L:
|
||||
return []
|
||||
cand_step = cand_step or max(1, L // 6)
|
||||
cand = [(s, float(np.percentile(m[s:s + L], pct))) for s in range(0, len(m) - L + 1, cand_step)]
|
||||
cand.sort(key=lambda x: x[1]) # calmest first
|
||||
accepted: list[tuple[int, float, float]] = []
|
||||
for s, p in cand:
|
||||
if p > p95_thresh:
|
||||
break # sorted -> everything after is worse
|
||||
if any(abs(s - a[0]) < L for a in accepted):
|
||||
continue # overlaps an accepted window
|
||||
drift = window_drift(shifts, s, L)
|
||||
if drift > drift_thresh:
|
||||
continue
|
||||
accepted.append((int(s), round(p, 4), round(drift, 3)))
|
||||
return sorted(accepted)
|
||||
|
||||
|
||||
def cmd_scan(args: argparse.Namespace) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import imageio.v2 as imageio
|
||||
|
||||
vids = sorted(Path(args.videos_dir).glob("**/*.mp4"))[:args.limit]
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
entries = []
|
||||
for i, vp in enumerate(vids, 1):
|
||||
gray, fps, _ = read_gray_lowres(str(vp), args.max_frames)
|
||||
L = int(round(args.window_sec * fps))
|
||||
if len(gray) < L:
|
||||
print(f"[win] [{i}/{len(vids)}] {vp.name}: only {len(gray)}f < {L}, skip", flush=True)
|
||||
continue
|
||||
m, shifts, diag = motion_series(gray)
|
||||
picks = best_windows(m, L, args.pct, top=args.top)
|
||||
stem = vp.stem
|
||||
# motion curve with chosen windows shaded
|
||||
fig, ax = plt.subplots(figsize=(12, 3.2))
|
||||
ax.plot(np.arange(len(m)) / fps, m, lw=0.8, color="steelblue")
|
||||
ax.axhline(args.calm_thresh, color="red", ls="--", lw=0.8, label=f"calm thresh {args.calm_thresh}")
|
||||
wstats = []
|
||||
for j, s in enumerate(picks):
|
||||
drift = window_drift(shifts, s, L)
|
||||
p95 = float(np.percentile(m[s:s + L], args.pct))
|
||||
ax.axvspan(s / fps, (s + L) / fps, color="limegreen", alpha=0.3)
|
||||
ax.text((s + L / 2) / fps, m.max() * 0.9, f"drift {drift:.3f}", ha="center", fontsize=8)
|
||||
wstats.append({
|
||||
"start_frame": s,
|
||||
"start_sec": round(s / fps, 2),
|
||||
"len_frames": L,
|
||||
"p95_motion": round(p95, 4),
|
||||
"drift": round(drift, 3)
|
||||
})
|
||||
ax.set_xlabel("time (s)")
|
||||
ax.set_ylabel("per-frame motion (norm)")
|
||||
ax.set_title(f"{stem} ({len(m)}f @ {fps:.0f}fps) green = static {args.window_sec}s window")
|
||||
ax.legend(fontsize=8)
|
||||
fig.tight_layout()
|
||||
curve = f"{stem}_motion.png"
|
||||
fig.savefig(str(out / curve), dpi=80)
|
||||
plt.close(fig)
|
||||
# preview: the best window as a short clip (full-res, read on demand)
|
||||
prev = ""
|
||||
if picks:
|
||||
s = picks[0]
|
||||
rgbw = read_rgb_frames(str(vp), range(s, s + L))
|
||||
imageio.mimsave(str(out / f"{stem}_win.mp4"), rgbw, fps=int(round(fps)), macro_block_size=1)
|
||||
prev = f"{stem}_win.mp4"
|
||||
entries.append({
|
||||
"id": stem,
|
||||
"src": str(vp),
|
||||
"fps": round(fps, 2),
|
||||
"n_frames": len(m),
|
||||
"curve": curve,
|
||||
"preview": prev,
|
||||
"windows": wstats,
|
||||
"median_motion": round(float(np.median(m)), 4)
|
||||
})
|
||||
best = wstats[0] if wstats else {}
|
||||
print(
|
||||
f"[win] [{i}/{len(vids)}] {stem}: {len(m)}f, best window @{best.get('start_sec')}s "
|
||||
f"p95={best.get('p95_motion')} surv={best.get('survival')}",
|
||||
flush=True)
|
||||
(out / "manifest.json").write_text(json.dumps(entries, indent=2))
|
||||
print(f"[win] scanned {len(entries)} videos -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_build(args: argparse.Namespace) -> None:
|
||||
"""Scan untrimmed source clips, greedily extract all non-overlapping static windows, and
|
||||
write them as target_frames-frame clips (subsampled to span the whole window) + manifest."""
|
||||
import imageio.v2 as imageio
|
||||
vids = sorted(Path(args.videos_dir).glob("**/*.mp4"))
|
||||
out = Path(args.out)
|
||||
(out / "videos").mkdir(parents=True, exist_ok=True)
|
||||
GENERIC = "a first-person egocentric view of a person performing a kitchen task with their hands"
|
||||
man = []
|
||||
idx = 0
|
||||
for vi, vp in enumerate(vids, 1):
|
||||
if idx >= args.target_count:
|
||||
break
|
||||
try:
|
||||
gray, fps, _ = read_gray_lowres(str(vp), args.max_frames)
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[build] {vp.name}: read fail {e}", flush=True)
|
||||
continue
|
||||
L = int(round(args.window_sec * fps))
|
||||
if len(gray) < L:
|
||||
continue
|
||||
m, shifts, _ = motion_series(gray)
|
||||
wins = greedy_valid_windows(m, shifts, L, args.pct, args.p95_thresh, args.drift_thresh)
|
||||
stride = max(1, round(L / args.target_frames))
|
||||
out_fps = round(fps / stride, 3)
|
||||
part = vp.stem.split("_")[0]
|
||||
n_from_clip = 0
|
||||
for (s, p95, drift) in wins:
|
||||
if idx >= args.target_count:
|
||||
break
|
||||
sel = np.arange(s, s + L, stride)[:args.target_frames]
|
||||
if sel.size < args.target_frames:
|
||||
continue
|
||||
f = read_rgb_frames(str(vp), sel) # full-res, only the window frames
|
||||
f = f[:, :f.shape[1] // 2 * 2, :f.shape[2] // 2 * 2] # even dims for libx264
|
||||
name = f"vid_{idx:06d}.mp4"
|
||||
imageio.mimsave(str(out / "videos" / name), f, fps=int(round(out_fps)), macro_block_size=1)
|
||||
man.append({
|
||||
"idx": idx,
|
||||
"path": name,
|
||||
"cap": [GENERIC],
|
||||
"fps": out_fps,
|
||||
"num_frames": int(args.target_frames),
|
||||
"duration": round(args.target_frames / out_fps, 3),
|
||||
"resolution": [int(f.shape[1]), int(f.shape[2])],
|
||||
"participant": part,
|
||||
"orig": vp.stem,
|
||||
"win_start_sec": round(s / fps, 2),
|
||||
"p95_motion": p95,
|
||||
"drift": drift
|
||||
})
|
||||
idx += 1
|
||||
n_from_clip += 1
|
||||
if n_from_clip:
|
||||
print(f"[build] [{vi}/{len(vids)}] {vp.stem}: +{n_from_clip} windows (total {idx})", flush=True)
|
||||
(out / "videos2caption.json").write_text(json.dumps(man, indent=2))
|
||||
(out / "merge.txt").write_text(f"{out}/videos,{out}/videos2caption.json\n")
|
||||
from collections import Counter
|
||||
print(
|
||||
f"[build] wrote {len(man)} static clips -> {out} | per-participant {dict(Counter(x['participant'] for x in man))}",
|
||||
flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
entries = json.loads((viz / "manifest.json").read_text())
|
||||
gal = [(str(viz / e["curve"]), e["id"]) for e in entries]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = entries[evt.index]
|
||||
vp = str(viz / e["preview"]) if e.get("preview") else None
|
||||
return str(viz / e["curve"]), vp, json.dumps(e["windows"], indent=2)
|
||||
|
||||
with gr.Blocks(title="Static-window finder") as demo:
|
||||
gr.Markdown("### Per-frame camera motion + static-window finder\n"
|
||||
"The curve is per-frame global motion (re-seeded each frame, so it **recovers after a turn** "
|
||||
"— spikes = turns, valleys = static). Green span = the chosen static window; `surv` = grid "
|
||||
"survival seeded at the window start. Click a clip -> motion curve + a preview of the best "
|
||||
"window + per-window stats. We keep windows with low `p95_motion` and high `survival`.")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gal, columns=2, height=560, label="videos (click)")
|
||||
with gr.Column():
|
||||
curve = gr.Image(label="per-frame motion (green = static window)")
|
||||
vid = gr.Video(label="preview of best static window")
|
||||
meta = gr.Code(label="windows", language="json")
|
||||
g.select(show, None, [curve, vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("scan")
|
||||
r.add_argument("--videos-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--window-sec", type=float, default=5.0)
|
||||
r.add_argument("--pct", type=float, default=95.0, help="percentile of per-frame motion used to score a window")
|
||||
r.add_argument("--calm-thresh", type=float, default=0.01, help="reference 'calm' motion line for the plot")
|
||||
r.add_argument("--top", type=int, default=2, help="static windows to find per video")
|
||||
r.add_argument("--max-frames", type=int, default=900, help="cap frames scanned per video (bounds cost)")
|
||||
r.add_argument("--limit", type=int, default=12)
|
||||
r.set_defaults(func=cmd_scan)
|
||||
b = sub.add_parser("build")
|
||||
b.add_argument("--videos-dir", required=True)
|
||||
b.add_argument("--out", required=True)
|
||||
b.add_argument("--window-sec", type=float, default=5.0)
|
||||
b.add_argument("--target-frames", type=int, default=121, help="frames per output clip (window subsampled to span)")
|
||||
b.add_argument("--target-count", type=int, default=200)
|
||||
b.add_argument("--pct", type=float, default=95.0)
|
||||
b.add_argument("--p95-thresh", type=float, default=0.004, help="max p95 per-frame motion for a valid window")
|
||||
b.add_argument("--drift-thresh", type=float, default=0.35, help="max net view drift (frac of diagonal)")
|
||||
b.add_argument("--max-frames", type=int, default=1500, help="cap frames scanned per source clip")
|
||||
b.set_defaults(func=cmd_build)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7890)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,239 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pre-generate controllability-eval artifacts for the viewer app.
|
||||
|
||||
For one checkpoint, over clips x counterfactuals, generate the video and save:
|
||||
- <clip>_<ctrl>__gen.mp4 raw generation
|
||||
- <clip>_<ctrl>__input.mp4 gen + the INPUT control tracks (what we asked for)
|
||||
- <clip>_<ctrl>__tracked.mp4 gen + tracks RE-EXTRACTED from the gen (CoTracker3)
|
||||
- <clip>_<ctrl>__heat.mp4 gen + points colored by per-point EPE (green=followed,
|
||||
red=ignored) -- the EPE heatmap (when tracks apply)
|
||||
plus per clip: <clip>__original.mp4 (real clip) and <clip>__original_tracks.mp4.
|
||||
|
||||
A manifest.json records EPE / n_points / coverage and the file names so the viewer
|
||||
(examples/inference/gradio/trackwan/app_viewer.py) can browse without a GPU.
|
||||
|
||||
Run one process per checkpoint (parallel across GPUs)::
|
||||
|
||||
srun ... env CUDA_VISIBLE_DEVICES=5 .venv/bin/python data_pipeline/gen_eval_artifacts.py \
|
||||
--export <export_3000> --name "step 3000" --out <artifacts_root>/step3000 \
|
||||
--data <funinp_now parquet> --clips 0 1 2 3
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
from fastvideo.eval.metrics.motion.control_sensitivity.metric import ( # noqa: E402
|
||||
compute_control_sensitivity, render_diff_heatmap)
|
||||
|
||||
HEAT_SCALE = 40.0 # px error mapped green(0)->red(>=40)
|
||||
|
||||
|
||||
def _draw(frames: Any, trpx: Any, vis: Any, colors: Any) -> Any:
|
||||
from fastvideo.train.callbacks.track_validation import _draw_overlay
|
||||
return _draw_overlay(frames, trpx, vis, colors, 9, 2, 0.65)
|
||||
|
||||
|
||||
def _solid(m: Any, rgb: Any) -> np.ndarray:
|
||||
return np.tile(np.array([rgb], np.uint8), (m, 1))
|
||||
|
||||
|
||||
def _heat_colors(perpt: Any) -> np.ndarray:
|
||||
e = np.clip(np.nan_to_num(perpt, nan=HEAT_SCALE) / HEAT_SCALE, 0.0, 1.0)
|
||||
r = (255 * e).astype(np.uint8)
|
||||
g = (255 * (1 - e)).astype(np.uint8)
|
||||
b = np.full_like(r, 40)
|
||||
return np.stack([r, g, b], 1)
|
||||
|
||||
|
||||
def _retrack_core(frames: Any,
|
||||
tr_px: Any,
|
||||
vs: Any,
|
||||
ct: Any,
|
||||
device: Any,
|
||||
max_points: int = 600,
|
||||
seed: int = 0,
|
||||
qf: int = 0,
|
||||
vt: float = 0.5) -> dict[str, Any] | None:
|
||||
"""Mirror compute_epe but also return the extracted tracks + per-point error."""
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import _retrack
|
||||
T = min(frames.shape[0], tr_px.shape[0])
|
||||
frames, tr_px, vs = frames[:T], tr_px[:T], vs[:T]
|
||||
idx = np.nonzero(vs[qf] > vt)[0]
|
||||
if idx.size == 0:
|
||||
return None
|
||||
if max_points and idx.size > max_points:
|
||||
rng = np.random.default_rng(seed)
|
||||
idx = np.sort(rng.choice(idx, size=max_points, replace=False))
|
||||
q_xy = tr_px[qf, idx]
|
||||
queries = np.concatenate([np.full((idx.size, 1), qf, np.float32), q_xy], axis=1)
|
||||
rt, _ = _retrack(ct, frames, queries, device) # [T,M,2] px
|
||||
Tr = min(T, rt.shape[0])
|
||||
tgt = tr_px[:Tr][:, idx]
|
||||
pred = rt[:Tr]
|
||||
mask = vs[:Tr][:, idx] > vt
|
||||
d = np.sqrt(((pred - tgt)**2).sum(-1)) # [Tr,M]
|
||||
valid = d[mask]
|
||||
epe = float(valid.mean()) if valid.size else None
|
||||
perpt = np.array([d[:, m][mask[:, m]].mean() if mask[:, m].any() else np.nan for m in range(idx.size)])
|
||||
disp = np.sqrt(((tgt - tgt[0:1])**2).sum(-1)) # input travel from frame 0
|
||||
moving = disp.max(0) > 8.0
|
||||
mm = mask & moving[None, :]
|
||||
epe_moving = float(d[mm].mean()) if d[mm].size else None
|
||||
return dict(tgt=tgt,
|
||||
pred=pred,
|
||||
mask=mask,
|
||||
epe=epe,
|
||||
epe_moving=epe_moving,
|
||||
perpt=perpt,
|
||||
n=int(idx.size),
|
||||
n_moving=int(moving.sum()),
|
||||
cov=float(mask.mean()))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument("--name", required=True, help="checkpoint label e.g. 'step 3000'")
|
||||
p.add_argument("--out", required=True)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", required=True)
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1, 2, 3])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import load_cotracker
|
||||
ct = load_cotracker(model.device)
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
Tpx = (num_lat_t - 1) * ratio + 1
|
||||
dev = model.device
|
||||
|
||||
def save(frames: Any, name: str) -> str:
|
||||
path = os.path.join(args.out, name)
|
||||
imageio.mimsave(path, frames, fps=args.fps, macro_block_size=1)
|
||||
return name
|
||||
|
||||
manifest = {"name": args.name, "clips": {}}
|
||||
for ci, s in zip(args.clips, samples, strict=False):
|
||||
gt_t = s["track_points"][0].numpy()[:Tpx]
|
||||
gt_v = s["track_visibility"][0].numpy()[:Tpx]
|
||||
swap = samples[(args.clips.index(ci) + 1) % len(samples)]
|
||||
g = st.make_grid(50)
|
||||
pan_t, pan_v = st.pan(g, Tpx, 0.25, 0.0)
|
||||
zoom_t, zoom_v = st.zoom(g, Tpx, 1.25)
|
||||
drag_t, drag_v = st.drag(g, Tpx, center=(0.5, 0.5), dx=0.3, dy=0.0, radius=0.2)
|
||||
drag_vs = st.select_radius(drag_v, g, center=(0.5, 0.5), radius=0.2)
|
||||
controls = {
|
||||
"gt": (gt_t, gt_v),
|
||||
"none": (None, None),
|
||||
"pan_right": (pan_t, pan_v),
|
||||
"zoom_in": (zoom_t, zoom_v),
|
||||
"drag_dense": (drag_t, drag_v),
|
||||
"drag_sparse": (drag_t, drag_vs),
|
||||
"swap": (swap["track_points"][0].numpy()[:Tpx], swap["track_visibility"][0].numpy()[:Tpx]),
|
||||
}
|
||||
|
||||
ref = twi.decode_reference(model, s["vae_latent"])
|
||||
H, W = ref.shape[1], ref.shape[2]
|
||||
gtpx = gt_t.copy()
|
||||
gtpx[..., 0] *= W
|
||||
gtpx[..., 1] *= H
|
||||
entry = {
|
||||
"caption":
|
||||
s["caption"][:240],
|
||||
"H":
|
||||
int(H),
|
||||
"W":
|
||||
int(W),
|
||||
"original":
|
||||
save(ref, f"clip{ci}__original.mp4"),
|
||||
"original_tracks":
|
||||
save(_draw(ref, gtpx, gt_v, _solid(gt_t.shape[1], [0, 220, 255])), f"clip{ci}__original_tracks.mp4"),
|
||||
"controls": {}
|
||||
}
|
||||
|
||||
frames_gt = None
|
||||
for cname, (tr, vs) in controls.items():
|
||||
tp = torch.from_numpy(tr)[None].float() if tr is not None else None
|
||||
tv = torch.from_numpy(vs)[None].float() if vs is not None else None
|
||||
lat = twi.generate(model,
|
||||
first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp,
|
||||
track_visibility=tv,
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=args.steps,
|
||||
seed=args.seed)
|
||||
frames = twi.decode_to_pixels(model, lat)
|
||||
if cname == "gt":
|
||||
frames_gt = frames
|
||||
trpx = None
|
||||
if tr is not None:
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
rec = {
|
||||
"gen": save(frames, f"clip{ci}_{cname}__gen.mp4"),
|
||||
"epe": None,
|
||||
"epe_moving": None,
|
||||
"n_points": 0,
|
||||
"coverage": 0.0
|
||||
}
|
||||
if trpx is not None:
|
||||
core = _retrack_core(frames, trpx, vs, ct, dev)
|
||||
if core is not None:
|
||||
m = core["tgt"].shape[1]
|
||||
rec["input"] = save(_draw(frames, core["tgt"], core["mask"], _solid(m, [0, 220, 255])),
|
||||
f"clip{ci}_{cname}__input.mp4")
|
||||
rec["tracked"] = save(_draw(frames, core["pred"], core["mask"], _solid(m, [255, 220, 0])),
|
||||
f"clip{ci}_{cname}__tracked.mp4")
|
||||
rec["heat"] = save(_draw(frames, core["pred"], core["mask"], _heat_colors(core["perpt"])),
|
||||
f"clip{ci}_{cname}__heat.mp4")
|
||||
rec.update(epe=core["epe"],
|
||||
epe_moving=core["epe_moving"],
|
||||
n_points=core["n"],
|
||||
n_moving=core["n_moving"],
|
||||
coverage=core["cov"])
|
||||
# intervention diff vs the GT-track generation (counterfactual control sensitivity)
|
||||
if cname != "gt" and frames_gt is not None:
|
||||
sens = compute_control_sensitivity(frames_gt, frames, trpx, vs)
|
||||
rec["diff"] = save(render_diff_heatmap(frames_gt, frames), f"clip{ci}_{cname}__diff.mp4")
|
||||
sr = sens["sensitivity_roi"]
|
||||
rec.update(sensitivity_roi=(round(sr, 4) if sr == sr else None),
|
||||
bg_leakage=round(sens["bg_leakage"], 4),
|
||||
localization_iou=round(sens["localization_iou"], 4),
|
||||
mean_diff=round(sens["mean_diff"], 4))
|
||||
entry["controls"][cname] = rec
|
||||
print(
|
||||
f"[{args.name}] clip{ci} {cname:12s} "
|
||||
f"EPE={rec['epe'] if rec['epe'] is None else round(float(rec['epe']),2)} "
|
||||
f"EPEmv={rec['epe_moving'] if rec['epe_moving'] is None else round(float(rec['epe_moving']),2)} "
|
||||
f"sens={rec.get('sensitivity_roi')} iou={rec.get('localization_iou')}",
|
||||
flush=True)
|
||||
manifest["clips"][str(ci)] = entry
|
||||
|
||||
with open(os.path.join(args.out, "manifest.json"), "w") as f:
|
||||
json.dump(manifest, f, indent=2)
|
||||
print("ARTIFACTS_DONE", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,123 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-GPU I2V synthetic gen: one SHARED seed image + a fixed caption list -> videos.
|
||||
|
||||
Reproduces the wantrack_synth_toy setup (all clips start from the same seed frame; each caption
|
||||
varies only the motion) but at 720p/24fps. Uses Wan2.1-I2V-14B-720P: for every caption it
|
||||
generates a clip conditioned on --seed-image at --height x --width.
|
||||
|
||||
Unlike the T2V worker there is NO first-frame drop: I2V's frame 0 IS the (clean) seed, so we
|
||||
keep all num_frames. Output layout matches gen_synth_worker.py (videos/, meta/, manifest_shards/)
|
||||
so merge_synth_manifests.py + the tracks/preprocess stages work unchanged.
|
||||
|
||||
Idempotent/resumable (skips finished mp4s). Parallelize across GPUs with --worker-id/--num-workers.
|
||||
|
||||
CUDA_VISIBLE_DEVICES=0 python data_pipeline/gen_synth_i2v_worker.py \
|
||||
--seed-image data_pipeline/synth_toy_720p/synthetic_seed.png \
|
||||
--captions data_pipeline/synth_toy_720p/captions.txt \
|
||||
--output-dir /home/hal-kevin/data/motion-stream-synth
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
MODEL_DEFAULT = "/home/hal-kevin/models/Wan2.1-I2V-14B-720P-Diffusers"
|
||||
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--seed-image", type=Path, required=True, help="shared first-frame conditioning image")
|
||||
p.add_argument("--captions", type=Path, required=True, help="one caption per line")
|
||||
p.add_argument("--output-dir", type=Path, required=True)
|
||||
p.add_argument("--model", default=MODEL_DEFAULT)
|
||||
p.add_argument("--worker-id", type=int, default=0)
|
||||
p.add_argument("--num-workers", type=int, default=1)
|
||||
p.add_argument("--height", type=int, default=720)
|
||||
p.add_argument("--width", type=int, default=1280)
|
||||
p.add_argument("--num-frames", type=int, default=121, help="kept 1:1 (I2V, no drop); must be 4k+1")
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--steps", type=int, default=40)
|
||||
p.add_argument("--guidance-scale", type=float, default=5.0)
|
||||
p.add_argument("--seed-base", type=int, default=1024, help="per-clip seed = seed_base + caption_idx")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
a = parse_args()
|
||||
assert (a.num_frames - 1) % 4 == 0, f"num_frames {a.num_frames} must be 4k+1"
|
||||
assert a.seed_image.exists(), f"seed image not found: {a.seed_image}"
|
||||
|
||||
videos_dir = a.output_dir / "videos"
|
||||
meta_dir = a.output_dir / "meta"
|
||||
shard_dir = a.output_dir / "manifest_shards"
|
||||
for d in (videos_dir, meta_dir, shard_dir):
|
||||
d.mkdir(parents=True, exist_ok=True)
|
||||
shard_manifest = shard_dir / f"worker_{a.worker_id:04d}.jsonl"
|
||||
fail_log = a.output_dir / f"failures_worker_{a.worker_id:04d}.log"
|
||||
|
||||
captions = [ln.strip() for ln in a.captions.read_text().splitlines() if ln.strip()]
|
||||
# Fixed order (caption line -> vid index); this worker owns a stride slice.
|
||||
my_idx = list(range(a.worker_id, len(captions), a.num_workers))
|
||||
print(f"[w{a.worker_id}/{a.num_workers}] {len(my_idx)} caption(s) "
|
||||
f"seed={a.seed_image.name} {a.width}x{a.height}@{a.fps}fps x{a.num_frames}f", flush=True)
|
||||
|
||||
import imageio.v2 as imageio
|
||||
from fastvideo import VideoGenerator
|
||||
t0 = time.time()
|
||||
g = VideoGenerator.from_pretrained(
|
||||
a.model, num_gpus=1, use_fsdp_inference=False,
|
||||
dit_cpu_offload=False, vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True, pin_cpu_memory=True,
|
||||
)
|
||||
print(f"[w{a.worker_id}] model ready in {time.time()-t0:.1f}s", flush=True)
|
||||
|
||||
seed_path = str(a.seed_image.resolve())
|
||||
n_ok = 0
|
||||
for gi in my_idx:
|
||||
fp = videos_dir / f"vid_{gi:06d}.mp4"
|
||||
if fp.exists():
|
||||
continue
|
||||
prompt = captions[gi]
|
||||
tmp = videos_dir / f".tmp_w{a.worker_id}_{gi:06d}.mp4"
|
||||
t = time.time()
|
||||
try:
|
||||
res = g.generate_video(
|
||||
prompt, image_path=seed_path, save_video=False, return_frames=True,
|
||||
height=a.height, width=a.width, num_frames=a.num_frames, fps=a.fps,
|
||||
seed=a.seed_base + gi, num_inference_steps=a.steps,
|
||||
guidance_scale=a.guidance_scale,
|
||||
)
|
||||
if isinstance(res, list):
|
||||
res = res[0]
|
||||
frames = np.asarray(res["frames"]) # I2V: keep all frames, no drop
|
||||
if frames.shape[0] != a.num_frames:
|
||||
raise RuntimeError(f"got {frames.shape[0]} frames, want {a.num_frames}")
|
||||
imageio.mimsave(tmp, list(frames), fps=a.fps, format="mp4")
|
||||
os.replace(tmp, fp)
|
||||
except Exception as e: # keep the worker alive
|
||||
with fail_log.open("a") as f:
|
||||
f.write(json.dumps({"idx": gi, "err": repr(e)[:500]}) + "\n")
|
||||
print(f"[w{a.worker_id}] FAIL idx={gi}: {e!r}", flush=True)
|
||||
if tmp.exists():
|
||||
tmp.unlink()
|
||||
continue
|
||||
dt = time.time() - t
|
||||
rec = {"idx": gi, "path": fp.name, "cap": [prompt], "fps": float(a.fps),
|
||||
"num_frames": a.num_frames, "duration": a.num_frames / float(a.fps),
|
||||
"resolution": {"width": a.width, "height": a.height},
|
||||
"gen_seconds": round(dt, 1)}
|
||||
with shard_manifest.open("a") as f:
|
||||
f.write(json.dumps(rec) + "\n")
|
||||
(meta_dir / f"vid_{gi:06d}.json").write_text(json.dumps(rec))
|
||||
n_ok += 1
|
||||
print(f"[w{a.worker_id}] {n_ok}/{len(my_idx)} done, {dt:.0f}s idx={gi}", flush=True)
|
||||
print(f"[w{a.worker_id}] DONE_WORKER made {n_ok} new videos", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-GPU data-parallel worker for large-scale Wan2.2-T2V-A14B synthetic gen.
|
||||
|
||||
One VideoGenerator per GPU (num_gpus=1, no FSDP/SP/TP). Each worker owns a stride
|
||||
slice of the prompt list: prompts[worker_id :: num_workers]. Idempotent/resumable:
|
||||
a video whose final mp4 exists is skipped, so a requeue after a cordon just continues.
|
||||
|
||||
First-frame saturation fix baked in: generate (num_frames + drop) frames, drop the
|
||||
first `drop` decoded frames, keep `num_frames`. (See research_log/08-first-frame-saturation.md.)
|
||||
|
||||
Per-worker manifest shard (manifest_shards/worker_<id>.jsonl) avoids 48 procs racing on
|
||||
one file; merge_manifests.py compiles them into the FastVideo videos2caption.json.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, json, os, time, traceback
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
|
||||
MODEL = "Wan-AI/Wan2.2-T2V-A14B-Diffusers"
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--prompts", type=Path, required=True)
|
||||
p.add_argument("--output-dir", type=Path, required=True)
|
||||
p.add_argument("--worker-id", type=int, required=True)
|
||||
p.add_argument("--num-workers", type=int, required=True)
|
||||
p.add_argument("--max-videos", type=int, default=None, help="global cap on #prompts consumed (across all workers)")
|
||||
p.add_argument("--height", type=int, default=720)
|
||||
p.add_argument("--width", type=int, default=1280)
|
||||
p.add_argument("--num-frames", type=int, default=121, help="frames to KEEP")
|
||||
p.add_argument("--drop", type=int, default=8, help="leading frames dropped (gen = keep+drop, must be 4k+1)")
|
||||
p.add_argument("--fps", type=int, default=16)
|
||||
p.add_argument("--steps", type=int, default=40)
|
||||
p.add_argument("--guidance-scale", type=float, default=4.0)
|
||||
p.add_argument("--guidance-scale-2", type=float, default=3.0)
|
||||
p.add_argument("--seed-base", type=int, default=1024, help="per-video seed = seed_base + global_prompt_idx")
|
||||
p.add_argument("--shuffle-seed", type=int, default=1234, help="deterministic prompt shuffle (same across workers)")
|
||||
return p.parse_args()
|
||||
|
||||
def load_prompts(path: Path):
|
||||
lines = [ln.strip() for ln in path.read_text(encoding="utf-8").splitlines() if ln.strip()]
|
||||
return lines
|
||||
|
||||
def main():
|
||||
a = parse_args()
|
||||
gen_frames = a.num_frames + a.drop
|
||||
assert (gen_frames - 1) % 4 == 0, f"gen frames {gen_frames} must be 4k+1"
|
||||
|
||||
videos_dir = a.output_dir / "videos"
|
||||
meta_dir = a.output_dir / "meta"
|
||||
shard_dir = a.output_dir / "manifest_shards"
|
||||
for d in (videos_dir, meta_dir, shard_dir):
|
||||
d.mkdir(parents=True, exist_ok=True)
|
||||
shard_manifest = shard_dir / f"worker_{a.worker_id:04d}.jsonl"
|
||||
fail_log = a.output_dir / f"failures_worker_{a.worker_id:04d}.log"
|
||||
|
||||
prompts = load_prompts(a.prompts)
|
||||
# Deterministic global shuffle so prompt order is stable across restarts and workers.
|
||||
order = np.random.RandomState(a.shuffle_seed).permutation(len(prompts))
|
||||
if a.max_videos is not None:
|
||||
order = order[:a.max_videos]
|
||||
# This worker's stride slice.
|
||||
my_positions = list(range(a.worker_id, len(order), a.num_workers))
|
||||
my_global_idx = [int(order[pos]) for pos in my_positions]
|
||||
print(f"[w{a.worker_id}/{a.num_workers}] assigned {len(my_global_idx)} prompts "
|
||||
f"(gen {gen_frames}->keep {a.num_frames}, {a.width}x{a.height}@{a.fps}fps)", flush=True)
|
||||
|
||||
# already-done set from existing mp4s
|
||||
def final_path(idx): return videos_dir / f"vid_{idx:06d}.mp4"
|
||||
|
||||
import imageio.v2 as imageio
|
||||
from fastvideo import VideoGenerator
|
||||
t0 = time.time()
|
||||
g = VideoGenerator.from_pretrained(
|
||||
MODEL, num_gpus=1, use_fsdp_inference=False,
|
||||
dit_cpu_offload=False, vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True, pin_cpu_memory=True,
|
||||
)
|
||||
print(f"[w{a.worker_id}] model ready in {time.time()-t0:.1f}s", flush=True)
|
||||
|
||||
n_ok = 0
|
||||
for gi in my_global_idx:
|
||||
fp = final_path(gi)
|
||||
if fp.exists():
|
||||
continue
|
||||
prompt = prompts[gi]
|
||||
tmp = videos_dir / f".tmp_w{a.worker_id}_{gi:06d}.mp4"
|
||||
t = time.time()
|
||||
try:
|
||||
res = g.generate_video(
|
||||
prompt, save_video=False, return_frames=True,
|
||||
height=a.height, width=a.width, num_frames=gen_frames, fps=a.fps,
|
||||
seed=a.seed_base + gi, num_inference_steps=a.steps,
|
||||
guidance_scale=a.guidance_scale, guidance_scale_2=a.guidance_scale_2,
|
||||
)
|
||||
if isinstance(res, list):
|
||||
res = res[0]
|
||||
frames = np.asarray(res["frames"])[a.drop:a.drop + a.num_frames]
|
||||
if frames.shape[0] != a.num_frames:
|
||||
raise RuntimeError(f"got {frames.shape[0]} frames after drop, want {a.num_frames}")
|
||||
imageio.mimsave(tmp, list(frames), fps=a.fps, format="mp4")
|
||||
os.replace(tmp, fp) # atomic publish
|
||||
except Exception as e: # keep the worker alive
|
||||
with fail_log.open("a") as f:
|
||||
f.write(json.dumps({"idx": gi, "err": repr(e)[:500]}) + "\n")
|
||||
print(f"[w{a.worker_id}] FAIL idx={gi}: {e!r}", flush=True)
|
||||
if tmp.exists(): tmp.unlink()
|
||||
continue
|
||||
dt = time.time() - t
|
||||
rec = {"idx": gi, "path": fp.name, "cap": [prompt], "fps": float(a.fps),
|
||||
"num_frames": a.num_frames, "resolution": {"width": a.width, "height": a.height},
|
||||
"gen_seconds": round(dt, 1)}
|
||||
with shard_manifest.open("a") as f:
|
||||
f.write(json.dumps(rec) + "\n")
|
||||
(meta_dir / f"vid_{gi:06d}.json").write_text(json.dumps(rec))
|
||||
n_ok += 1
|
||||
if n_ok % 10 == 0:
|
||||
print(f"[w{a.worker_id}] {n_ok} done, last {dt:.0f}s idx={gi}", flush=True)
|
||||
print(f"[w{a.worker_id}] DONE_WORKER made {n_ok} new videos", flush=True)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,235 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 0b: generate a synthetic (video, prompt) dataset with Wan2.2-T2V-A14B.
|
||||
|
||||
Text prompts -> .mp4 videos + a FastVideo-compatible manifest, so the existing
|
||||
preprocess pipeline can ingest the result directly. CoTracker point extraction is a
|
||||
separate stage (``extract_tracks.py``).
|
||||
|
||||
Design (see notes/DECISIONS.md):
|
||||
- Generate at the *training* fps/length (default 16 fps, 81 frames ~= 5 s) so per-frame
|
||||
point tracks align 1:1 with frames and no resampling is needed downstream.
|
||||
- T2V only. The eventual I2V+points model uses frame 0 as the conditioning image at
|
||||
training time, so no input image is needed here.
|
||||
- Idempotent/resumable: each finished video is appended to ``manifest.jsonl`` and skipped
|
||||
on re-run.
|
||||
|
||||
Run on a GPU node (never the login node), e.g.:
|
||||
|
||||
srun --jobid=<shao_wm jobid> --overlap --ntasks=1 \
|
||||
.venv/bin/python data_pipeline/generate_videos.py \
|
||||
--prompts examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt \
|
||||
--output-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p \
|
||||
--num-videos 50 --num-gpus 8
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import random
|
||||
import shutil
|
||||
import time
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
|
||||
DEFAULT_MODEL = "Wan-AI/Wan2.2-T2V-A14B-Diffusers"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--prompts", type=Path, required=True, help="Text file, one prompt per line.")
|
||||
p.add_argument("--output-dir", type=Path, required=True, help="Dataset root (videos/, manifest, ...).")
|
||||
p.add_argument("--model", type=str, default=DEFAULT_MODEL)
|
||||
p.add_argument("--num-videos", type=int, default=50, help="How many prompts to generate (after --start).")
|
||||
p.add_argument("--start", type=int, default=0, help="Offset into the prompt list (for sharding).")
|
||||
p.add_argument("--shuffle", action="store_true", help="Deterministically shuffle prompts before slicing.")
|
||||
p.add_argument("--num-gpus", type=int, default=8)
|
||||
p.add_argument("--height", type=int, default=720)
|
||||
p.add_argument("--width", type=int, default=1280)
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--trim-start-frames", type=int, default=0,
|
||||
help="Drop this many frames from the start of each generated video (VAE warm-up artifact). "
|
||||
"Generation runs for num_frames+trim_start_frames and the head is discarded.")
|
||||
p.add_argument("--seed", type=int, default=1024, help="Base seed; per-video seed = seed + global index.")
|
||||
p.add_argument("--num-inference-steps", type=int, default=None, help="Override model default if set.")
|
||||
p.add_argument("--negative-prompt", type=str, default=None)
|
||||
p.add_argument("--image",
|
||||
type=str,
|
||||
default=None,
|
||||
help="If set, do I2V from this image (same for every prompt); else T2V.")
|
||||
# Offload controls. Default OFF: Wan2.2-A14B is ~56GB bf16 and fits on a single H200 (143GB),
|
||||
# so offloading (esp. layerwise) only makes generation ~10x slower. Enable on small GPUs.
|
||||
p.add_argument("--dit-cpu-offload", action="store_true", help="Offload DiT to CPU (slow).")
|
||||
p.add_argument("--dit-layerwise-offload",
|
||||
action="store_true",
|
||||
help="Stream DiT layers from CPU per step (very slow; only for tiny GPUs).")
|
||||
p.add_argument("--text-encoder-cpu-offload", action="store_true", help="Offload text encoder to CPU.")
|
||||
p.add_argument("--vae-cpu-offload", action="store_true", help="Offload VAE to CPU.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def load_prompts(path: Path, start: int, num: int, shuffle: bool, seed: int) -> list[tuple[int, str]]:
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Prompts file not found: {path}\n"
|
||||
"Download it first (login node): cd examples/dataset/vidprom && ./download_dataset.sh")
|
||||
lines = [ln.strip() for ln in path.read_text(encoding="utf-8").splitlines() if ln.strip()]
|
||||
indexed = list(enumerate(lines)) # global index is stable w.r.t. the raw file order
|
||||
if shuffle:
|
||||
random.Random(seed).shuffle(indexed)
|
||||
return indexed[start:start + num]
|
||||
|
||||
|
||||
def read_done_indices(manifest_jsonl: Path) -> set[int]:
|
||||
done: set[int] = set()
|
||||
if manifest_jsonl.exists():
|
||||
for ln in manifest_jsonl.read_text().splitlines():
|
||||
ln = ln.strip()
|
||||
if not ln:
|
||||
continue
|
||||
try:
|
||||
done.add(int(json.loads(ln)["idx"]))
|
||||
except (json.JSONDecodeError, KeyError, ValueError):
|
||||
continue
|
||||
return done
|
||||
|
||||
|
||||
def rebuild_manifest(manifest_jsonl: Path, videos_dir: Path, json_path: Path, merge_path: Path) -> int:
|
||||
"""Compile manifest.jsonl -> videos2caption.json + merge.txt (FastVideo format)."""
|
||||
records: dict[int, dict] = {}
|
||||
if manifest_jsonl.exists():
|
||||
for ln in manifest_jsonl.read_text().splitlines():
|
||||
ln = ln.strip()
|
||||
if not ln:
|
||||
continue
|
||||
try:
|
||||
rec = json.loads(ln)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
records[int(rec["idx"])] = rec
|
||||
ordered = [records[k] for k in sorted(records)]
|
||||
|
||||
tmp = json_path.with_suffix(".json.tmp")
|
||||
tmp.write_text(json.dumps(ordered, indent=2))
|
||||
tmp.replace(json_path)
|
||||
merge_path.write_text(f"{videos_dir.resolve()},{json_path.resolve()}\n")
|
||||
return len(ordered)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning, module="fastvideo.*")
|
||||
|
||||
output_dir = args.output_dir
|
||||
videos_dir = output_dir / "videos"
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
manifest_jsonl = output_dir / "manifest.jsonl"
|
||||
json_path = output_dir / "videos2caption.json"
|
||||
merge_path = output_dir / "merge.txt"
|
||||
fail_log = output_dir / "failures.log"
|
||||
|
||||
selected = load_prompts(args.prompts, args.start, args.num_videos, args.shuffle, args.seed)
|
||||
done = read_done_indices(manifest_jsonl)
|
||||
todo = [(i, pr) for (i, pr) in selected if i not in done]
|
||||
print(f"[gen] {len(selected)} selected, {len(done)} already done, {len(todo)} to generate", flush=True)
|
||||
|
||||
if not todo:
|
||||
n = rebuild_manifest(manifest_jsonl, videos_dir, json_path, merge_path)
|
||||
print(f"[gen] nothing to do; manifest has {n} entries -> {json_path}", flush=True)
|
||||
return
|
||||
|
||||
# Import here so --help works without loading torch/fastvideo.
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model,
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=args.dit_cpu_offload,
|
||||
dit_layerwise_offload=args.dit_layerwise_offload,
|
||||
vae_cpu_offload=args.vae_cpu_offload,
|
||||
text_encoder_cpu_offload=args.text_encoder_cpu_offload,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
extra: dict = {}
|
||||
if args.num_inference_steps is not None:
|
||||
extra["num_inference_steps"] = args.num_inference_steps
|
||||
if args.negative_prompt is not None:
|
||||
extra["negative_prompt"] = args.negative_prompt
|
||||
|
||||
# Wan VAE requires num_frames = 4k+1. Round up gen_frames to satisfy this,
|
||||
# then trim the actual excess (may be more than trim_start_frames).
|
||||
_raw = args.num_frames + args.trim_start_frames
|
||||
gen_frames = _raw if (_raw - 1) % 4 == 0 else _raw + (4 - (_raw - 1) % 4)
|
||||
actual_trim = gen_frames - args.num_frames
|
||||
duration = float(args.num_frames) / float(args.fps)
|
||||
for n_done, (idx, prompt) in enumerate(todo, 1):
|
||||
final_path = videos_dir / f"vid_{idx:06d}.mp4"
|
||||
if final_path.exists(): # belt-and-suspenders vs manifest
|
||||
continue
|
||||
tmp_dir = videos_dir / f".tmp_{idx:06d}"
|
||||
if tmp_dir.exists():
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
tmp_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
t0 = time.time()
|
||||
try:
|
||||
i2v_kwargs = {"image_path": args.image} if args.image else {}
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
output_path=str(tmp_dir),
|
||||
save_video=True,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=gen_frames,
|
||||
fps=args.fps,
|
||||
seed=args.seed + idx,
|
||||
**i2v_kwargs,
|
||||
**extra,
|
||||
)
|
||||
produced = sorted(tmp_dir.glob("*.mp4"))
|
||||
if not produced:
|
||||
raise RuntimeError("no .mp4 produced by generate_video")
|
||||
src = str(produced[0])
|
||||
if actual_trim > 0:
|
||||
import subprocess as _sp
|
||||
tmp_trim = str(final_path) + ".trim.mp4"
|
||||
_sp.run(
|
||||
["ffmpeg", "-y", "-i", src,
|
||||
"-vf", f"trim=start_frame={actual_trim},setpts=PTS-STARTPTS",
|
||||
"-c:v", "libx264", "-pix_fmt", "yuv420p", "-crf", "18", "-an", tmp_trim],
|
||||
check=True, capture_output=True,
|
||||
)
|
||||
Path(tmp_trim).replace(final_path)
|
||||
else:
|
||||
shutil.move(src, str(final_path))
|
||||
except Exception as e: # noqa: BLE001 - keep the batch alive, log and move on
|
||||
with fail_log.open("a") as f:
|
||||
f.write(json.dumps({"idx": idx, "error": repr(e), "prompt": prompt}) + "\n")
|
||||
print(f"[gen] FAILED idx={idx}: {e!r}", flush=True)
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
continue
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
|
||||
record = {
|
||||
"idx": idx,
|
||||
"path": final_path.name, # basename; folder in merge.txt is videos_dir
|
||||
"cap": [prompt],
|
||||
"fps": float(args.fps),
|
||||
"duration": duration,
|
||||
"num_frames": int(args.num_frames),
|
||||
"resolution": {
|
||||
"width": args.width,
|
||||
"height": args.height
|
||||
},
|
||||
}
|
||||
with manifest_jsonl.open("a") as f:
|
||||
f.write(json.dumps(record) + "\n")
|
||||
print(f"[gen] [{n_done}/{len(todo)}] idx={idx} {time.time()-t0:.1f}s -> {final_path.name}", flush=True)
|
||||
|
||||
n = rebuild_manifest(manifest_jsonl, videos_dir, json_path, merge_path)
|
||||
print(f"[gen] done; manifest has {n} entries -> {json_path}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,246 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Layer-wise learning-rate (LWLR) bootstrap experiment.
|
||||
|
||||
Question: on a 14B random-init WanTrack model, does giving the track pathway
|
||||
(track_encoder.* + patch_embedding.weight[:, 36:]) a 100x higher LR than the
|
||||
rest actually let it grow into functional gradient magnitudes within a small
|
||||
number of steps? Or is the encoder stuck regardless?
|
||||
|
||||
Runs a short training loop on real OpenVid samples, tracking weight+grad norms
|
||||
of the track pathway at each step. Compares against the "known good" magnitude
|
||||
observed in 1.3B_merged (per-sample track_encoder.proj grad ~0.5).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from safetensors.torch import load_file as st_load
|
||||
|
||||
# Distributed init BEFORE importing fastvideo (same pattern as gradient_flow_analysis)
|
||||
def _init_distributed_1gpu() -> None:
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
)
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
|
||||
|
||||
_init_distributed_1gpu()
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
from gradient_flow_analysis import ( # noqa: E402
|
||||
build_model_from_diffusers,
|
||||
load_samples,
|
||||
build_i2v_cond_concat,
|
||||
flow_matching_noisy,
|
||||
Sample,
|
||||
)
|
||||
|
||||
|
||||
TRACK_PATH_PARAMS = (
|
||||
"track_encoder.proj.weight",
|
||||
"track_encoder.temporal_conv.weight",
|
||||
"track_encoder.proj.bias", # may not exist depending on TRACKWAN_TRACK_BIAS
|
||||
"track_encoder.temporal_conv.bias",
|
||||
)
|
||||
|
||||
|
||||
def split_param_groups(model, high_lr: float, base_lr: float) -> list[dict]:
|
||||
"""Give track pathway (track_encoder.* + patch_embedding[:, 36:]) a high LR.
|
||||
|
||||
patch_embedding.weight is a SINGLE tensor of shape [hidden, 52, 1, 2, 2] — we can't
|
||||
give sub-slice a different LR through standard param groups. Workaround: put the
|
||||
whole patch_embedding in a "hybrid" group with high LR (the base 36 channels will
|
||||
also move faster, but by only 100x for a short warmup that's small — 100 steps of
|
||||
high LR is still << stage-1 4800 steps of normal LR).
|
||||
"""
|
||||
high_group, base_group = [], []
|
||||
named = dict(model.named_parameters())
|
||||
for n, p in named.items():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
# High-LR pathway: encoder convs (+ bias if present) + the whole patch_embedding.weight
|
||||
is_encoder = n.startswith("track_encoder.")
|
||||
is_patch_embed = (n == "patch_embedding.weight")
|
||||
if is_encoder or is_patch_embed:
|
||||
high_group.append(p)
|
||||
else:
|
||||
base_group.append(p)
|
||||
return [
|
||||
{"params": high_group, "lr": high_lr, "name": "track_pathway"},
|
||||
{"params": base_group, "lr": base_lr, "name": "dit_body"},
|
||||
]
|
||||
|
||||
|
||||
def snapshot_norms(model, sample_grad_i: int) -> dict:
|
||||
"""Return weight/grad norms for the diagnostic params."""
|
||||
out = {}
|
||||
for n, p in model.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
# Track params of interest
|
||||
keep = (
|
||||
n.startswith("track_encoder.")
|
||||
or n == "patch_embedding.weight"
|
||||
)
|
||||
if not keep:
|
||||
continue
|
||||
w_norm = p.detach().float().norm().item()
|
||||
g_norm = p.grad.detach().float().norm().item() if p.grad is not None else 0.0
|
||||
# For patch_embedding, split into base/track slices
|
||||
if n == "patch_embedding.weight":
|
||||
w_base = p.detach()[:, :36].float()
|
||||
w_track = p.detach()[:, 36:].float()
|
||||
out["patch_embedding.weight[:, :36]"] = {
|
||||
"w_norm": w_base.norm().item(), "g_norm": 0.0,
|
||||
}
|
||||
out["patch_embedding.weight[:, 36:]"] = {
|
||||
"w_norm": w_track.norm().item(), "g_norm": 0.0,
|
||||
}
|
||||
if p.grad is not None:
|
||||
g_base = p.grad.detach()[:, :36].float()
|
||||
g_track = p.grad.detach()[:, 36:].float()
|
||||
out["patch_embedding.weight[:, :36]"]["g_norm"] = g_base.norm().item()
|
||||
out["patch_embedding.weight[:, 36:]"]["g_norm"] = g_track.norm().item()
|
||||
else:
|
||||
out[n] = {"w_norm": w_norm, "g_norm": g_norm}
|
||||
return out
|
||||
|
||||
|
||||
def one_step(model, s: Sample, device: torch.device, flow_shift: float,
|
||||
from_ctx_ns) -> float:
|
||||
"""Forward+backward on one sample; grads accumulate."""
|
||||
vae = s.vae_latent.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
ff = s.first_frame_latent.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
clip = s.clip_feature.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
text = s.text_embedding.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
mask = torch.ones(1, s.text_embedding.shape[0], device=device, dtype=torch.bfloat16)
|
||||
tp = s.track_points.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
tv = s.track_visibility.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
|
||||
num_latent_t = vae.shape[2]
|
||||
expected_T = (num_latent_t - 1) * 4 + 1
|
||||
tp = tp[:, :expected_T]
|
||||
tv = tv[:, :expected_T]
|
||||
|
||||
sigma_u = float(torch.rand(1).item())
|
||||
noise = torch.randn_like(vae, dtype=torch.float32)
|
||||
clean_f = vae.float()
|
||||
noisy_f, _, ts_val = flow_matching_noisy(clean_f, noise, sigma_u,
|
||||
flow_shift=flow_shift)
|
||||
noisy = noisy_f.to(torch.bfloat16)
|
||||
cond20 = build_i2v_cond_concat(ff, vae_temporal_compression=4)
|
||||
hs = torch.cat([noisy, cond20], dim=1)
|
||||
ts = torch.tensor([ts_val], device=device, dtype=torch.bfloat16)
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16), from_ctx_ns(current_timestep=ts, attn_metadata=None):
|
||||
pred = model(
|
||||
hidden_states=hs, encoder_hidden_states=text, encoder_attention_mask=mask,
|
||||
timestep=ts, encoder_hidden_states_image=clip,
|
||||
track_points=tp, track_visibility=tv, return_dict=False,
|
||||
)
|
||||
target = noise - clean_f
|
||||
loss = F.mse_loss(pred.float(), target.float())
|
||||
loss.backward()
|
||||
return float(loss.item())
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--model", default="/mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_random_init")
|
||||
ap.add_argument("--n-samples", type=int, default=16)
|
||||
ap.add_argument("--n-steps", type=int, default=60)
|
||||
ap.add_argument("--batch-accum", type=int, default=4, help="gradient accumulation steps per optim step")
|
||||
ap.add_argument("--high-lr", type=float, default=1e-3, help="LR for track pathway")
|
||||
ap.add_argument("--base-lr", type=float, default=1e-5, help="LR for the rest of DiT")
|
||||
ap.add_argument("--flow-shift", type=float, default=6.0)
|
||||
ap.add_argument("--parquet-glob", type=str,
|
||||
default="/mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset/shard_*/**/*.parquet")
|
||||
ap.add_argument("--out", type=str,
|
||||
default="/mnt/lustre/vlm-s4duan/gradient_analysis_14b/lwlr_trace.json")
|
||||
ap.add_argument("--seed", type=int, default=42)
|
||||
args = ap.parse_args()
|
||||
|
||||
random.seed(args.seed)
|
||||
np.random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
device = torch.device("cuda")
|
||||
|
||||
samples = load_samples(args.parquet_glob, args.n_samples, seed=args.seed)
|
||||
print(f"loaded {len(samples)} samples")
|
||||
assert len(samples) >= args.batch_accum, "need enough samples for accumulation"
|
||||
|
||||
print(f"loading {args.model} ...")
|
||||
model = build_model_from_diffusers(args.model, device)
|
||||
|
||||
# Set up param groups
|
||||
param_groups = split_param_groups(model, high_lr=args.high_lr, base_lr=args.base_lr)
|
||||
n_high = sum(p.numel() for p in param_groups[0]["params"])
|
||||
n_base = sum(p.numel() for p in param_groups[1]["params"])
|
||||
print(f"track_pathway params: {n_high/1e6:.2f}M (LR={args.high_lr:g})")
|
||||
print(f"dit_body params: {n_base/1e6:.2f}M (LR={args.base_lr:g})")
|
||||
optim = torch.optim.AdamW(param_groups, weight_decay=0.0) # no wd so we see raw drift
|
||||
|
||||
trace: list[dict] = []
|
||||
# step 0 (before any update)
|
||||
print("\n=== step 0 (init) ===")
|
||||
norms0 = snapshot_norms(model, -1)
|
||||
for k, v in norms0.items():
|
||||
print(f" {k:<45s} w={v['w_norm']:.4f} g={v['g_norm']:.4f}")
|
||||
trace.append({"step": 0, "loss": None, "norms": norms0})
|
||||
|
||||
step = 0
|
||||
accum = 0
|
||||
losses_accum = []
|
||||
for outer in range(args.n_steps * args.batch_accum):
|
||||
s = samples[outer % len(samples)]
|
||||
loss = one_step(model, s, device, args.flow_shift, set_forward_context)
|
||||
losses_accum.append(loss)
|
||||
accum += 1
|
||||
if accum < args.batch_accum:
|
||||
continue
|
||||
step += 1
|
||||
# snapshot BEFORE step (grads are accumulated over the batch)
|
||||
gs = snapshot_norms(model, -1)
|
||||
optim.step()
|
||||
optim.zero_grad(set_to_none=True)
|
||||
avg_loss = float(np.mean(losses_accum))
|
||||
losses_accum = []
|
||||
accum = 0
|
||||
if step % 5 == 0 or step <= 5:
|
||||
print(f"\n=== step {step} (loss={avg_loss:.4f}) ===")
|
||||
for k, v in gs.items():
|
||||
marker = ""
|
||||
if k == "track_encoder.proj.weight":
|
||||
marker = " <-- 1.3B_merged saw per_sample_norm ~0.50"
|
||||
elif k == "patch_embedding.weight[:, 36:]":
|
||||
marker = " <-- 1.3B_merged saw per_sample_norm ~0.73"
|
||||
print(f" {k:<45s} w={v['w_norm']:.4f} g={v['g_norm']:.4f}{marker}")
|
||||
trace.append({"step": step, "loss": avg_loss, "norms": gs})
|
||||
|
||||
Path(args.out).parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(args.out, "w") as f:
|
||||
json.dump({"config": vars(args), "trace": trace}, f, indent=2)
|
||||
print(f"\nwrote trace -> {args.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Compile per-worker manifest shards -> FastVideo videos2caption.json + merge.txt.
|
||||
Idempotent; run anytime (progress check) or at the end. Dedups by idx, verifies mp4 exists."""
|
||||
from __future__ import annotations
|
||||
import argparse, json
|
||||
from pathlib import Path
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--output-dir", type=Path, required=True)
|
||||
a = p.parse_args()
|
||||
videos_dir = a.output_dir / "videos"
|
||||
shard_dir = a.output_dir / "manifest_shards"
|
||||
recs: dict[int, dict] = {}
|
||||
for sh in sorted(shard_dir.glob("worker_*.jsonl")):
|
||||
for ln in sh.read_text().splitlines():
|
||||
ln = ln.strip()
|
||||
if not ln:
|
||||
continue
|
||||
try:
|
||||
r = json.loads(ln)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if (videos_dir / r["path"]).exists():
|
||||
recs[int(r["idx"])] = r
|
||||
ordered = [recs[k] for k in sorted(recs)]
|
||||
j = a.output_dir / "videos2caption.json"
|
||||
j.write_text(json.dumps(ordered, indent=2))
|
||||
(a.output_dir / "merge.txt").write_text(f"{videos_dir.resolve()},{j.resolve()}\n")
|
||||
# also count mp4s on disk (may exceed manifest if a worker died mid-write)
|
||||
on_disk = len(list(videos_dir.glob("vid_*.mp4")))
|
||||
print(f"manifest: {len(ordered)} entries; {on_disk} mp4s on disk -> {j}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,65 @@
|
||||
# Data Pipeline — Decisions & Rationale
|
||||
|
||||
## Why generation + tracking live OUTSIDE `fastvideo/pipelines/preprocess/`
|
||||
FastVideo's preprocess pipeline is structurally a *consumer of existing .mp4 files*
|
||||
(`VideoCaptionMergedDataset` reads a `merge.txt`→`annotation.json` over a folder and
|
||||
asserts each video path exists). There is no seam to *produce* a video inside it. So
|
||||
generation (Wan2.2) and tracking (CoTracker) are standalone Stage-0 scripts that
|
||||
produce the files preprocess later ingests — matching the repo convention (everything
|
||||
upstream of preprocess is plain scripts, not pipeline stages).
|
||||
|
||||
## Two scripts, not one fused process
|
||||
- The `.mp4` on disk is the checkpoint; generation is the expensive part. Tracking
|
||||
config will be tweaked many times — must not re-run Wan each time.
|
||||
- Wan2.2-14B is huge (multi-GPU + offload); CoTracker is tiny. "Generate all → free
|
||||
model → track all" is the natural ordering and is exactly two phases.
|
||||
- Independent restart/parallelism. Both scripts are idempotent (skip done work).
|
||||
|
||||
## Generate at training fps/length (no temporal alignment headache)
|
||||
CoTracker tracks per *video* frame, but the VAE compresses time (~4× for Wan). If we
|
||||
generated at arbitrary fps we'd have to re-index tracks by the preprocess frame-sampler
|
||||
(`sample_frame_index`) and then fold onto latent time. Since we *control* generation,
|
||||
we emit videos already at the train fps/length (default 16 fps, 81 frames ≈ 5 s) so
|
||||
source frames == sampled frames, 1:1. Tracks then align to frames trivially; the only
|
||||
remaining fold (frames→latent-frames) happens in the model/trainer.
|
||||
|
||||
## Generation is T2V even though the model is I2V+points
|
||||
Synthetic videos come from text→video (Wan2.2-T2V-A14B). At *training* time the model
|
||||
uses frame 0 as the I2V conditioning image + the point tracks as motion control. So we
|
||||
never need an input image during generation.
|
||||
|
||||
## Manifest format (FastVideo-compatible)
|
||||
`generate_videos.py` emits the same shape the existing loader expects:
|
||||
- `videos/vid_NNNNNN.mp4`
|
||||
- `videos2caption.json`: list of `{path(basename), cap:[...], fps, duration, num_frames, resolution}`
|
||||
- `merge.txt`: one line `<videos_dir>,<videos2caption.json>`
|
||||
Loader actually consumes only `path, cap, fps, duration` (+ optional conditioning path);
|
||||
`resolution`/`num_frames` are informational. We know all of these at generation time, so
|
||||
we write the manifest directly and SKIP `scripts/dataset_preparation/prepare_json_file.py`
|
||||
(it exists only to *recover* fps/duration by re-probing videos of unknown provenance).
|
||||
|
||||
## points_path stored as ABSOLUTE path
|
||||
The future preprocess loader joins `folder + relative_path`. Tracks live in a sibling
|
||||
`tracks/` dir, not under `videos/`. Storing `points_path` absolute makes `os.path.join`
|
||||
return it unchanged, so it resolves regardless of the manifest folder. Mirrors how
|
||||
MatrixGame2 references `action_path`, which is the pattern the points preprocess task
|
||||
will copy.
|
||||
|
||||
## CoTracker v3 via torch.hub (`cotracker3_offline`)
|
||||
The repo only vendors CoTracker **v2** under `third_party/eval/vbench` (eval-only). For
|
||||
v3 we pull `facebookresearch/co-tracker` `cotracker3_offline` via torch.hub. Prefetch on
|
||||
the login node (has internet); the hub cache lives on shared weka so compute nodes reuse
|
||||
it offline (`extract_tracks.py` falls back to `source="local"` from the cache dir).
|
||||
|
||||
## Output location OUTSIDE the git repo
|
||||
Generated videos are large (≈2–5 MB each → ~25 GB at 5k). Default output dir is
|
||||
`/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/...` (outside `FastVideo/`) to avoid
|
||||
bloating the working tree / accidental commits.
|
||||
|
||||
## Cluster
|
||||
- `shao_wm` = SLURM job **1788946** on node **fs-mbz-gpu-538**: 8×GPU (~140 GB each),
|
||||
128 CPU, 768 GB RAM, 5-day limit. As of setup all 8 GPUs were idle.
|
||||
- Attach with `srun --jobid=1788946 --overlap ...`. Never run heavy work on the login
|
||||
node (`fs-mbz-login-big-001`) — it lags the shared connection.
|
||||
- `sqs` is a `.bashrc` alias (`squeue -u hao.zhang`); not available in non-interactive
|
||||
shells — use `squeue -u $USER` directly.
|
||||
@@ -0,0 +1,40 @@
|
||||
# Data Pipeline — Progress Log
|
||||
|
||||
Project: build a `(video, prompt, point-tracks)` dataset to train a motion-controlled
|
||||
real-time interactive video model (bidirectional finetune stage first).
|
||||
|
||||
Branch: `shao/realtime-bidir`. This `data_pipeline/` dir holds the **Stage-0** data
|
||||
generation scripts (everything *upstream* of FastVideo's `preprocess` pipeline).
|
||||
|
||||
## Pipeline shape (where we are)
|
||||
|
||||
```
|
||||
vidprom prompts ──> generate_videos.py ──> videos/*.mp4 + videos2caption.json + merge.txt
|
||||
│
|
||||
extract_tracks.py ──> tracks/*.npz (+ patch points_path into the json)
|
||||
│
|
||||
[TODO] v1_preprocess.py --preprocess_task i2v_points ──> parquet
|
||||
│
|
||||
[TODO] bidir I2V+points trainer
|
||||
```
|
||||
|
||||
## Status
|
||||
|
||||
- [x] Located existing FastVideo preprocess stack + MatrixGame2 control-signal pattern (the template).
|
||||
- [x] Confirmed env = repo `.venv` (`/mnt/weka/home/hao.zhang/shao/FastVideo/.venv/bin/python`).
|
||||
- [x] Confirmed Wan2.2-T2V-A14B weights cached (`~/.cache/huggingface/hub/models--Wan-AI--Wan2.2-T2V-A14B-Diffusers`).
|
||||
- [x] Wrote `generate_videos.py` (Wan2.2-14B T2V → mp4 + manifest).
|
||||
- [x] Wrote `extract_tracks.py` (CoTracker v3 50×50 grid → npz, patch manifest).
|
||||
- [~] Smoke-test launched on fs-mbz-gpu-538 (`run.sh --smoke`, background). Node validated:
|
||||
internet=yes (downloads OK on the node), 8 GPUs idle, `.venv` imports fastvideo + torch
|
||||
2.11.0/CUDA, cotracker3_offline loads (25.4M params). Prompts downloaded: 248,221 (140 MB).
|
||||
- [ ] Scale to 50, then 5k.
|
||||
- [ ] Extend preprocess with an `i2v_points` task (mirror MatrixGame2).
|
||||
|
||||
## Open questions / decisions pending
|
||||
|
||||
- Training model for the bidir stage: 1.3B/480p for plumbing bring-up vs straight to 14B (data gen is fixed at 14B/720p regardless — user confirmed).
|
||||
- Confirm conditioning shape: I2V + points (input image + per-frame points → video). Assumed yes.
|
||||
- Point sampling at train time: store full 2500-pt grid + visibility; sample 1–200 in the trainer.
|
||||
|
||||
See `DECISIONS.md` for rationale, `RUNBOOK.md` for exact commands.
|
||||
@@ -0,0 +1,70 @@
|
||||
# Data Pipeline — Runbook
|
||||
|
||||
All heavy work runs on the `shao_wm` allocation via `srun --overlap`. Do NOT run model
|
||||
code on the login node.
|
||||
|
||||
## 0. Env + node handles
|
||||
```bash
|
||||
PY=/mnt/weka/home/hao.zhang/shao/FastVideo/.venv/bin/python
|
||||
JOBID=$(squeue -u "$USER" -n shao_wm -h -o %i | head -1) # shao_wm allocation id
|
||||
echo "jobid=$JOBID"
|
||||
# live GPU usage on the node (lightweight):
|
||||
srun --jobid="$JOBID" --overlap nvidia-smi \
|
||||
--query-gpu=index,memory.used,memory.total,utilization.gpu --format=csv,noheader
|
||||
```
|
||||
|
||||
## 1. One-time setup (login node — has internet)
|
||||
```bash
|
||||
cd /mnt/weka/home/hao.zhang/shao/FastVideo
|
||||
# prompts (self-forcing / vidprom):
|
||||
[ -f examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt ] || \
|
||||
( cd examples/dataset/vidprom && ./download_dataset.sh )
|
||||
# prefetch CoTracker v3 into the shared torch.hub cache so compute nodes work offline:
|
||||
$PY -c "import torch; torch.hub.load('facebookresearch/co-tracker','cotracker3_offline'); print('cotracker ok')"
|
||||
```
|
||||
|
||||
**Pick free GPUs yourself** from step 0 and pass them via `CUDA_VISIBLE_DEVICES` — the node
|
||||
is shared and the scripts do NOT auto-select. Pinning to busy GPUs OOMs.
|
||||
|
||||
## 2. Generate videos (Wan2.2-T2V-A14B, 720p, 16fps, 81f)
|
||||
```bash
|
||||
GPUS=2,3 # <- two currently-free GPUs from step 0
|
||||
srun --jobid="$JOBID" --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=$GPUS \
|
||||
$PY data_pipeline/generate_videos.py \
|
||||
--prompts examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt \
|
||||
--output-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p \
|
||||
--num-videos 50 --num-gpus 2 # num-gpus must match the device count
|
||||
```
|
||||
- Smoke first: `--num-videos 2`.
|
||||
- Idempotent: re-run resumes from `manifest.jsonl`.
|
||||
|
||||
## 3. Extract point tracks (CoTracker v3, 50×50 grid)
|
||||
```bash
|
||||
srun --jobid="$JOBID" --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=$GPUS \
|
||||
$PY data_pipeline/extract_tracks.py \
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p \
|
||||
--grid-size 50
|
||||
```
|
||||
- Idempotent: skips videos whose `tracks/<stem>.npz` exists.
|
||||
- If 720p OOMs CoTracker, add `--downscale 0.5` (coords are rescaled back to original px).
|
||||
|
||||
## 4. One-shot wrapper (derives num_gpus from CUDA_VISIBLE_DEVICES)
|
||||
```bash
|
||||
CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh # full run (defaults)
|
||||
CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh --smoke # 2 videos
|
||||
```
|
||||
|
||||
## Outputs
|
||||
```
|
||||
<output-dir>/
|
||||
videos/vid_000000.mp4 ...
|
||||
tracks/vid_000000.npz ... # keys: tracks (T,N,2 px), visibility (T,N), grid_size, height, width, fps, num_frames
|
||||
videos2caption.json # FastVideo manifest (+ points_path after step 3)
|
||||
merge.txt
|
||||
manifest.jsonl # incremental resume log (generation)
|
||||
```
|
||||
|
||||
## Gotchas
|
||||
- Check GPU availability before launching (step 0) — the node may not be fully free.
|
||||
- `generate_video(...)` logs a Deprecation warning (use of legacy API) — harmless.
|
||||
- First Wan2.2 run loads ~14B MoE weights; expect a few min of startup before frame 1.
|
||||
@@ -0,0 +1,148 @@
|
||||
# Data Pipeline
|
||||
|
||||
All commands run from `/home/hal-kevin/FastVideo`.
|
||||
|
||||
---
|
||||
|
||||
## Stage 1 — Generate Videos
|
||||
|
||||
```bash
|
||||
python data_pipeline/generate_videos.py \
|
||||
--prompts examples/dataset/motion-test/prompts.txt \
|
||||
--output-dir /home/hal-kevin/data/motion-stream-test \
|
||||
--num-videos 100 \
|
||||
--num-gpus 4 \
|
||||
--trim-start-frames 10 \
|
||||
--num-inference-steps 40
|
||||
```
|
||||
|
||||
Output: `motion-physics/videos/vid_000000.mp4 ...`
|
||||
|
||||
---
|
||||
|
||||
## Stage 2 — VAE Round-Trip Videos
|
||||
|
||||
Encode then decode each video through the FastVideo WanVAE so CoTracker runs on
|
||||
the same frames that training will see (Stage 5 re-encodes these, so tracks must
|
||||
be extracted from the decoded version, not the raw source).
|
||||
|
||||
```bash
|
||||
python data_pipeline/decode_roundtrip_videos.py \
|
||||
--data-dir /home/hal-kevin/data/motion-stream-test \
|
||||
--vae-path /home/hal-kevin/models/trackwan_1.3b_i2v_control_init/vae
|
||||
```
|
||||
|
||||
Output: `motion-physics/roundtrip_videos/vid_000000.mp4 ...`
|
||||
|
||||
---
|
||||
|
||||
## Stage 3 — Extract Tracks
|
||||
|
||||
Run CoTracker on the round-trip videos, parallelized across 4 GPUs:
|
||||
|
||||
```bash
|
||||
bash data_pipeline/run_extract_tracks.sh
|
||||
```
|
||||
|
||||
Pass extra args (e.g. `--force`, `--limit 5`) directly — they are forwarded to each worker.
|
||||
The script no longer hardcodes `--force`; existing `.npz` are skipped unless you pass it:
|
||||
|
||||
```bash
|
||||
bash data_pipeline/run_extract_tracks.sh --force
|
||||
```
|
||||
|
||||
Speed knobs: `--sam-batch 16` (frames per batched FastSAM forward, default 16) and `--amp`
|
||||
(bf16 autocast for CoTracker, ~1.5-2x faster but slightly different coords — validate before
|
||||
adopting). Entry events now share one extra CoTracker pass (queries filtered to the new-object
|
||||
regions) instead of a full 2500-point pass per mask.
|
||||
|
||||
**Fused mode:** `--segment` (with `--vis-override-every 2`) runs Stage 4 inside this pass —
|
||||
object IDs, vis override, and track weights — reusing the decoded video and the entry-detection
|
||||
FastSAM masks, so Stage 4 does not need to run at all. Same results as the standalone stage
|
||||
(shared implementation). `--viz`/`--viz-dir` render the same overlay mp4s as standalone Stage 4
|
||||
(slow — skip for large-scale runs); `--min-area-frac`/`--max-masks` remain standalone-only.
|
||||
Benchmark it with `FUSED=1 bash data_pipeline/benchmark_tracks.sh` (add `VIZ=1` for renders).
|
||||
|
||||
Single-GPU alternative:
|
||||
|
||||
```bash
|
||||
python data_pipeline/extract_tracks.py \
|
||||
--data-dir /home/hal-kevin/data/motion-stream-test \
|
||||
--videos-subdir roundtrip_videos \
|
||||
--grid-size 50 \
|
||||
--device cuda \
|
||||
--detect-entries \
|
||||
--sam-conf 0.75 \
|
||||
--sam-iou 0.9 \
|
||||
--sam-imgsz 1024
|
||||
```
|
||||
|
||||
Output: `motion-stream-test/tracks/vid_000000.npz ...`
|
||||
|
||||
---
|
||||
|
||||
## Stage 4 — Segment Tracks
|
||||
|
||||
Assign object IDs and compute motion weights, parallelized across 4 GPUs:
|
||||
|
||||
```bash
|
||||
bash data_pipeline/run_segment_tracks.sh
|
||||
```
|
||||
|
||||
The script no longer hardcodes `--force` (pass it to re-process npz that already have
|
||||
`object_ids`). Each video is now decoded once and FastSAM runs in batched forwards
|
||||
(`--sam-batch 16`), shared between object-ID assignment and the vis override sweep.
|
||||
|
||||
Single-GPU alternative:
|
||||
|
||||
```bash
|
||||
python data_pipeline/segment_tracks.py \
|
||||
--data-dir /home/hal-kevin/data/motion-stream-test \
|
||||
--videos-subdir roundtrip_videos \
|
||||
--conf 0.75 --iou 0.9 --imgsz 1024 \
|
||||
--vis-override-every 2 \
|
||||
--force \
|
||||
--viz
|
||||
```
|
||||
|
||||
Adds `object_ids`, `n_objects`, `track_weights` to each `.npz`.
|
||||
|
||||
---
|
||||
|
||||
## Stage 5 — Preprocess to Parquet
|
||||
|
||||
Reads raw `videos/` for VAE encoding and `tracks/` for track data.
|
||||
|
||||
```bash
|
||||
torchrun --nproc_per_node=1 -m fastvideo.pipelines.preprocess.v1_preprocess \
|
||||
--model_path /home/hal-kevin/models/trackwan_1.3b_i2v_control_init \
|
||||
--data_merge_path /home/hal-kevin/data/motion-stream-test/data_merge.txt \
|
||||
--output_dir /home/hal-kevin/data/motion-stream-test/preprocessed_i2v_track \
|
||||
--preprocess_task i2v_track \
|
||||
--num_frames 121 \
|
||||
--num_latent_t 31 \
|
||||
--train_fps 24 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--preprocess_video_batch_size 1 \
|
||||
--samples_per_file 64
|
||||
```
|
||||
|
||||
`--train_fps` must match the source video fps (24 here). If omitted it defaults to 30,
|
||||
and `FrameSamplingStage` resamples with interval `fps/train_fps = 0.8` — duplicating
|
||||
every 5th frame and covering only the first ~97 of 121 frames. The stored latents then
|
||||
encode slowed, stuttering motion that no longer aligns with the tracks (extracted at
|
||||
native fps), which shows up as drifting motion in validation reference videos.
|
||||
|
||||
Output: `motion-physics/preprocessed_i2v_track/combined_parquet_dataset/`
|
||||
|
||||
---
|
||||
|
||||
## Stage 6 — Training
|
||||
|
||||
```bash
|
||||
python -m fastvideo.train.train \
|
||||
--config examples/train/scenario/worldmodel/finetune_wantrack_golf_overfit.yaml
|
||||
```
|
||||
|
||||
Update the yaml to point at your data and checkpoint directories.
|
||||
@@ -0,0 +1,97 @@
|
||||
---
|
||||
license: apache-2.0
|
||||
task_categories:
|
||||
- text-to-video
|
||||
- image-to-video
|
||||
tags:
|
||||
- video-generation
|
||||
- point-tracking
|
||||
- motionstream
|
||||
- wantrack
|
||||
- fastvideo
|
||||
---
|
||||
|
||||
# OpenVid-WanTrack Processed (v2, 720p, **bf16**)
|
||||
|
||||
FastVideo preprocessing parquets for training the TrackWan point-track-conditioned I2V model on
|
||||
the OpenVid-derived WanTrack set. Each row is one 121-frame clip with its VAE latents, text and
|
||||
image conditioning, and dense CoTracker3 tracks — everything the trainer memory-maps, so no video
|
||||
decoding happens at train time.
|
||||
|
||||
**This is the `bfloat16` variant** of `…/openvid-wantrack-processed` (v2, 720p): the large float
|
||||
tensor fields are stored in **bf16** instead of float32, so the dataset is roughly **half the size**.
|
||||
Everything else (clips, ids, shapes, layout) is identical.
|
||||
|
||||
## ⚠️ Precision — read before loading
|
||||
|
||||
- The big tensors — `vae_latent`, `first_frame_latent`, `clip_feature`, `text_embedding`,
|
||||
`track_points`, `track_visibility` — are **`bfloat16`**. Each field's `_dtype` column says so.
|
||||
- `object_ids` and `track_weights` are kept **`float32`** (small integer/label fields).
|
||||
- **You must honor the per-field `_dtype` when decoding.** numpy has **no** bfloat16, so
|
||||
`np.frombuffer(bytes, "bfloat16")` fails — decode via `torch.frombuffer` (see Loading below).
|
||||
- **Quality is unaffected for training:** the TrackWan trainer already downcasts these fields to
|
||||
bf16 before use, so storing bf16 just pre-applies the exact rounding the model does anyway.
|
||||
- The FastVideo trainer's loader honors `_dtype`, so pointing `data_path` at this set "just works".
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
shard000/combined_parquet_dataset/worker_*/data_chunk_*.parquet
|
||||
shard001/combined_parquet_dataset/worker_*/data_chunk_*.parquet
|
||||
...
|
||||
shard259/... # shard259 is a 110-clip remainder; all others are 1000
|
||||
```
|
||||
|
||||
~259,110 clips across 260 shards. Clip ids join 1:1 with `openvid-wantrack-clips` (videos),
|
||||
`openvid-wantrack-tracks-v2` (raw npz tracks), and OpenVid-1M captions.
|
||||
|
||||
## Row schema (`pyarrow_schema_i2v_track`, 33 columns)
|
||||
|
||||
Scalars: `id, file_name, caption, media_type, width, height, num_frames, duration_sec, fps`.
|
||||
|
||||
Tensors — each stored as a triplet `<name>_bytes` (raw buffer), `_shape` (list<int64>), `_dtype`:
|
||||
|
||||
| tensor | shape (720p) | dtype | description |
|
||||
|--------|--------------|-------|-------------|
|
||||
| `vae_latent` | `[16, 31, 90, 160]` | **bfloat16** | WanVAE latent of the clip (training target) |
|
||||
| `first_frame_latent` | `[16, 31, 90, 160]` | **bfloat16** | I2V conditioning: VAE-encode of `[frame0, zeros...]` |
|
||||
| `clip_feature` | `[257, 1280]` | **bfloat16** | CLIP image embedding of frame 0 |
|
||||
| `text_embedding` | `[L, 4096]` | **bfloat16** | T5 caption embedding (variable length `L`, padding stripped) |
|
||||
| `track_points` | `[121, 2500, 2]` | **bfloat16** | CoTracker tracks, **normalized [0,1]** |
|
||||
| `track_visibility` | `[121, 2500]` | **bfloat16** | per-frame visibility |
|
||||
| `object_ids` | `[2500]` | float32 | FastSAM object id per track (-1 = background) |
|
||||
| `track_weights` | `[2500]` | float32 | low-rank motion weight in [0,1] |
|
||||
|
||||
`num_frames=31` for the latents (VAE 4x temporal compression: `(121-1)/4+1`); `track_points`
|
||||
stay at native `121`. Text embedding length varies per row (padding removed), so read the
|
||||
per-row `_shape`.
|
||||
|
||||
## Config
|
||||
|
||||
- Video: 1280x720, 121 frames, 24 fps
|
||||
- VAE: FastVideo WanVAE (latents encoded in fp32, **stored as bf16**), `use_feature_cache=True`
|
||||
- CLIP: frame-0 image embedding; T5: caption text embedding
|
||||
- Tracks: CoTracker3, 50x50 grid (2500 points), FastSAM segmentation
|
||||
|
||||
## Loading
|
||||
|
||||
numpy cannot represent bfloat16, so decode through torch, honoring each field's `_dtype`:
|
||||
|
||||
```python
|
||||
import glob, torch, pyarrow.parquet as pq
|
||||
|
||||
_STR2T = {"float32": torch.float32, "bfloat16": torch.bfloat16, "float16": torch.float16}
|
||||
|
||||
def decode(row, name):
|
||||
dt = _STR2T[row[f"{name}_dtype"]]
|
||||
# bytearray() -> writable buffer that doesn't alias the parquet row
|
||||
return torch.frombuffer(bytearray(row[f"{name}_bytes"]), dtype=dt).reshape(row[f"{name}_shape"])
|
||||
|
||||
files = glob.glob("**/*.parquet", recursive=True) # all shards
|
||||
row = pq.read_table(files[0]).slice(0, 1).to_pylist()[0]
|
||||
lat = decode(row, "vae_latent") # torch.bfloat16, shape [16, 31, 90, 160]
|
||||
tracks = decode(row, "track_points") # torch.bfloat16, normalized [0,1]
|
||||
```
|
||||
|
||||
The FastVideo trainer discovers all parquets under the dataset root via `os.walk` and its loader
|
||||
honors the `_dtype` column, so point `data_path` at the directory containing the `shard*/` folders.
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Download + extract OpenVid / OpenVidHD video shards from HuggingFace.
|
||||
|
||||
Videos are only distributed as ~30-50 GB zip shards (no per-file fetch), so this
|
||||
pulls whole shards, extracts the .mp4s into --videos-dir, and (unless --keep-zip)
|
||||
deletes the zip. Pair with openvid_filter.py (which clips to keep) + extract_tracks_mp.py
|
||||
(resolution/aspect probe + tracking). Idempotent: skips shards already extracted.
|
||||
|
||||
Some parts are HF-split (e.g. OpenVid_part102_partaa/ab) — those must be
|
||||
concatenated before unzip; this handles single-file parts (all OpenVidHD parts).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, zipfile
|
||||
from pathlib import Path
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
REPO = "nkp37/OpenVid-1M"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--shards", nargs="+", required=True,
|
||||
help="repo paths, e.g. OpenVidHD/OpenVidHD_part_1.zip OpenVid_part0.zip")
|
||||
ap.add_argument("--videos-dir", required=True, help="where .mp4s are extracted")
|
||||
ap.add_argument("--zip-dir", default=None, help="scratch for downloaded zips (default: <videos-dir>/../_zips)")
|
||||
ap.add_argument("--keep-zip", action="store_true", help="don't delete the zip after extract")
|
||||
ap.add_argument("--limit", type=int, default=0, help="extract at most N mp4s per shard (0=all; for smoke tests)")
|
||||
ap.add_argument("--only-list", default=None, help="file of basenames; extract only these (e.g. the ≥5s filtered list)")
|
||||
a = ap.parse_args()
|
||||
only = None
|
||||
if a.only_list:
|
||||
only = {l.strip() for l in open(a.only_list) if l.strip()}
|
||||
|
||||
vdir = Path(a.videos_dir); vdir.mkdir(parents=True, exist_ok=True)
|
||||
zdir = Path(a.zip_dir) if a.zip_dir else vdir.parent / "_zips"
|
||||
zdir.mkdir(parents=True, exist_ok=True)
|
||||
marker_dir = vdir.parent / "_extracted"; marker_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for shard in a.shards:
|
||||
done_marker = marker_dir / (Path(shard).name + ".done")
|
||||
if done_marker.exists():
|
||||
print(f"[dl] skip {shard} (already extracted)", flush=True); continue
|
||||
print(f"[dl] downloading {shard} ...", flush=True)
|
||||
zp = hf_hub_download(REPO, shard, repo_type="dataset", local_dir=str(zdir))
|
||||
print(f"[dl] extracting {zp} -> {vdir}", flush=True)
|
||||
n = 0
|
||||
with zipfile.ZipFile(zp) as z:
|
||||
for m in z.namelist():
|
||||
if m.lower().endswith(".mp4"):
|
||||
if only is not None and Path(m).name not in only:
|
||||
continue
|
||||
# flatten: write basename directly into videos-dir
|
||||
(vdir / Path(m).name).write_bytes(z.read(m))
|
||||
n += 1
|
||||
if a.limit and n >= a.limit:
|
||||
break
|
||||
done_marker.write_text(str(n))
|
||||
if not a.keep_zip:
|
||||
Path(zp).unlink(missing_ok=True)
|
||||
print(f"[dl] {shard}: extracted {n} mp4s", flush=True)
|
||||
print("[dl] done", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,100 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Robust OpenVidHD (1080p, 16:9) full downloader for the trackwan bidir dataset.
|
||||
|
||||
OpenVidHD parts are a mix of single-file `.zip` (parts 1-14) and SPLIT parts
|
||||
(`OpenVidHD_part_<i>_part_aa` + `_part_ab`) that must be concatenated before unzip
|
||||
(mirrors the official download_scripts/download_OpenVid.py). This enumerates the HF
|
||||
repo, plans each part, downloads (single or split+cat), extracts only the mp4s in
|
||||
--only-list (the >=5s filtered set), writes them into --videos-dir, and deletes the
|
||||
zip. Idempotent via per-part .done markers, so it resumes after interruption.
|
||||
|
||||
Shard by --shard/--num-shards (e.g. SLURM env) to spread the download across nodes.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, os, re, subprocess, zipfile, collections
|
||||
from pathlib import Path
|
||||
from huggingface_hub import HfApi, hf_hub_download
|
||||
|
||||
REPO = "nkp37/OpenVid-1M"
|
||||
|
||||
|
||||
def plan_parts():
|
||||
"""Return {part_num: [repo_files]} for OpenVidHD (single .zip or split _part_aa/ab)."""
|
||||
files = HfApi().list_repo_files(REPO, repo_type="dataset")
|
||||
parts = collections.defaultdict(list)
|
||||
for f in files:
|
||||
if not f.startswith("OpenVidHD/"):
|
||||
continue
|
||||
m = re.search(r"OpenVidHD_part_(\d+)", f)
|
||||
if m:
|
||||
parts[int(m.group(1))].append(f)
|
||||
return {k: sorted(v) for k, v in parts.items()}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--videos-dir", required=True)
|
||||
ap.add_argument("--zip-dir", required=True, help="scratch for zips (deleted after extract unless --keep-zip)")
|
||||
ap.add_argument("--only-list", required=True, help="basenames to keep (the >=5s filtered list)")
|
||||
ap.add_argument("--keep-zip", action="store_true")
|
||||
ap.add_argument("--num-shards", type=int, default=int(os.environ.get("SLURM_NTASKS", "1")))
|
||||
ap.add_argument("--shard", type=int, default=int(os.environ.get("SLURM_PROCID", "0")))
|
||||
ap.add_argument("--parts", default="all", help="'all' or comma range like 1-14 / 1,2,3")
|
||||
a = ap.parse_args()
|
||||
|
||||
vdir = Path(a.videos_dir); vdir.mkdir(parents=True, exist_ok=True)
|
||||
zdir = Path(a.zip_dir); zdir.mkdir(parents=True, exist_ok=True)
|
||||
mdir = vdir.parent / "_extracted"; mdir.mkdir(parents=True, exist_ok=True)
|
||||
only = {l.strip() for l in open(a.only_list) if l.strip()}
|
||||
|
||||
parts = plan_parts()
|
||||
keys = sorted(parts)
|
||||
if a.parts != "all":
|
||||
want = set()
|
||||
for tok in a.parts.split(","):
|
||||
if "-" in tok:
|
||||
lo, hi = tok.split("-"); want |= set(range(int(lo), int(hi) + 1))
|
||||
else:
|
||||
want.add(int(tok))
|
||||
keys = [k for k in keys if k in want]
|
||||
keys = keys[a.shard::a.num_shards]
|
||||
print(f"[hd-dl shard {a.shard}/{a.num_shards}] {len(keys)} parts: {keys[:6]}{'...' if len(keys)>6 else ''}", flush=True)
|
||||
|
||||
for i in keys:
|
||||
marker = mdir / f"OpenVidHD_part_{i}.done"
|
||||
if marker.exists():
|
||||
continue
|
||||
fs = parts[i]
|
||||
zip_path = zdir / f"OpenVidHD_part_{i}.zip"
|
||||
try:
|
||||
if len(fs) == 1 and fs[0].endswith(".zip"):
|
||||
p = hf_hub_download(REPO, fs[0], repo_type="dataset", local_dir=str(zdir))
|
||||
zip_path = Path(p)
|
||||
else: # split: download _part_aa/_part_ab, concat
|
||||
locals_ = []
|
||||
for f in fs:
|
||||
locals_.append(hf_hub_download(REPO, f, repo_type="dataset", local_dir=str(zdir)))
|
||||
with open(zip_path, "wb") as out:
|
||||
for lp in sorted(locals_):
|
||||
with open(lp, "rb") as src:
|
||||
while (chunk := src.read(1 << 24)):
|
||||
out.write(chunk)
|
||||
for lp in locals_:
|
||||
Path(lp).unlink(missing_ok=True)
|
||||
n = 0
|
||||
with zipfile.ZipFile(zip_path) as z:
|
||||
for m in z.namelist():
|
||||
if m.lower().endswith(".mp4") and Path(m).name in only:
|
||||
(vdir / Path(m).name).write_bytes(z.read(m))
|
||||
n += 1
|
||||
marker.write_text(str(n))
|
||||
if not a.keep_zip:
|
||||
zip_path.unlink(missing_ok=True)
|
||||
print(f"[hd-dl] part_{i}: extracted {n} clips", flush=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[hd-dl] part_{i} FAILED: {repr(e)[:150]}", flush=True)
|
||||
print(f"[hd-dl shard {a.shard}] done", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage-1 (metadata) filter for OpenVid — from the dataset CSV.
|
||||
|
||||
The OpenVid CSV columns are: video, caption, aesthetic score, motion score,
|
||||
temporal consistency score, camera motion, frame, fps, seconds.
|
||||
There is NO resolution / aspect-ratio column, so 720p + 16:9 CANNOT be filtered
|
||||
here — those are enforced downstream by probing (extract_tracks_mp.py
|
||||
--min-height / --aspect-tol). This stage keeps clips that are long enough to yield
|
||||
`num_frames` at the target fps, with optional motion/quality gates (tracking wants
|
||||
motion; static/low-motion clips give degenerate tracks).
|
||||
|
||||
Output: newline-delimited video basenames to keep (feeds the shard extraction +
|
||||
tracking stages). Also writes a captions sidecar (video -> caption) for later use.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, csv, json, sys
|
||||
|
||||
csv.field_size_limit(sys.maxsize) # captions are long
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--csv", required=True, help="OpenVid-1M.csv or OpenVidHD.csv")
|
||||
ap.add_argument("--out", required=True, help="output: one video basename per line")
|
||||
ap.add_argument("--captions-out", default=None, help="optional JSON {video: caption}")
|
||||
ap.add_argument("--num-frames", type=int, default=121)
|
||||
ap.add_argument("--target-fps", type=int, default=24)
|
||||
ap.add_argument("--min-seconds", type=float, default=None,
|
||||
help="override; default = num_frames/target_fps (+2%% margin)")
|
||||
ap.add_argument("--min-motion", type=float, default=0.0, help="drop clips below this motion score")
|
||||
ap.add_argument("--min-aesthetic", type=float, default=0.0)
|
||||
ap.add_argument("--exclude-static", action="store_true", help="drop camera motion == 'static'")
|
||||
a = ap.parse_args()
|
||||
|
||||
min_secs = a.min_seconds if a.min_seconds is not None else (a.num_frames / a.target_fps) * 1.02
|
||||
kept = 0; total = 0; caps = {}
|
||||
with open(a.csv, newline="") as f, open(a.out, "w") as o:
|
||||
r = csv.DictReader(f)
|
||||
for row in r:
|
||||
total += 1
|
||||
try:
|
||||
secs = float(row["seconds"]); motion = float(row["motion score"]); aes = float(row["aesthetic score"])
|
||||
except (KeyError, ValueError):
|
||||
continue
|
||||
if secs < min_secs: continue
|
||||
if motion < a.min_motion: continue
|
||||
if aes < a.min_aesthetic: continue
|
||||
if a.exclude_static and row.get("camera motion", "").strip().lower() == "static":
|
||||
continue
|
||||
o.write(row["video"] + "\n"); kept += 1
|
||||
if a.captions_out is not None:
|
||||
caps[row["video"]] = row.get("caption", "")
|
||||
if a.captions_out is not None:
|
||||
json.dump(caps, open(a.captions_out, "w"))
|
||||
print(f"[filter] min_seconds={min_secs:.3f} min_motion={a.min_motion} min_aes={a.min_aesthetic} "
|
||||
f"exclude_static={a.exclude_static}")
|
||||
print(f"[filter] kept {kept}/{total} ({100*kept/max(total,1):.1f}%) -> {a.out}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,81 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Precompute the UMT5 embedding of an empty string.
|
||||
|
||||
This is the canonical Wan/T5 null-text convention: feed "" to the text encoder
|
||||
and use the resulting non-zero embedding as the ∅ in classifier-free guidance.
|
||||
Contrast with `torch.zeros_like(text_embedding)`, which is a different tensor
|
||||
distribution that the base model was never asked to handle.
|
||||
|
||||
Saves a dict:
|
||||
{"embedding": [seq_len, dim], "attention_mask": [seq_len]}
|
||||
so the val callback can just index into these when it needs a null text.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/precompute_null_text.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/synth_stage2_paperLR_ckpt400 \
|
||||
--output /mnt/lustre/vlm-s4duan/exports/null_text_umt5.pt
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True,
|
||||
help="Diffusers export dir containing text_encoder/ and tokenizer/")
|
||||
p.add_argument("--output", required=True)
|
||||
p.add_argument("--text-len", type=int, default=256,
|
||||
help="Pad/truncate to this length (matches training text_padding_length)")
|
||||
p.add_argument("--device", default="cuda")
|
||||
args = p.parse_args()
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
try:
|
||||
from transformers import UMT5EncoderModel # type: ignore
|
||||
except ImportError:
|
||||
UMT5EncoderModel = None # type: ignore
|
||||
from transformers import AutoModel
|
||||
|
||||
model_dir = Path(args.model_dir)
|
||||
tok = AutoTokenizer.from_pretrained(model_dir / "tokenizer")
|
||||
if UMT5EncoderModel is not None:
|
||||
try:
|
||||
enc = UMT5EncoderModel.from_pretrained(model_dir / "text_encoder", torch_dtype=torch.bfloat16)
|
||||
except Exception:
|
||||
enc = AutoModel.from_pretrained(model_dir / "text_encoder", torch_dtype=torch.bfloat16)
|
||||
else:
|
||||
enc = AutoModel.from_pretrained(model_dir / "text_encoder", torch_dtype=torch.bfloat16)
|
||||
enc = enc.to(args.device).eval()
|
||||
|
||||
tokens = tok(
|
||||
[""],
|
||||
max_length=args.text_len,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
return_attention_mask=True,
|
||||
)
|
||||
input_ids = tokens["input_ids"].to(args.device)
|
||||
attn_mask = tokens["attention_mask"].to(args.device)
|
||||
print(f"tokenized '' -> {input_ids.shape}, non-pad tokens = {int(attn_mask.sum().item())}")
|
||||
print(f" first 5 token ids: {input_ids[0, :5].tolist()}")
|
||||
|
||||
with torch.no_grad():
|
||||
out = enc(input_ids=input_ids, attention_mask=attn_mask)
|
||||
# take last_hidden_state; strip the batch dim
|
||||
hs = out.last_hidden_state[0].float().cpu() # [seq_len, dim]
|
||||
am = attn_mask[0].cpu()
|
||||
print(f"embedding shape={list(hs.shape)} norm={hs.norm().item():.3f} mean={hs.mean().item():.5f}")
|
||||
print(f"attn_mask shape={list(am.shape)} sum={int(am.sum().item())}")
|
||||
|
||||
Path(args.output).parent.mkdir(parents=True, exist_ok=True)
|
||||
torch.save({"embedding": hs, "attention_mask": am, "text_len": args.text_len}, args.output)
|
||||
print(f"[null-text] saved -> {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,137 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 2 (VAE-free variant): crop+resize source videos to the training geometry.
|
||||
|
||||
Applies exactly the same ``center_crop_th_tw`` + ``resize`` transform as
|
||||
``decode_roundtrip_videos.py`` (and as Stage 5's ``CenterCropResizeVideo``), but skips
|
||||
the VAE encode/decode. Tracks extracted from these videos land in the same coordinate
|
||||
frame as the training latents; they just don't carry the VAE's reconstruction artifacts.
|
||||
|
||||
Purpose: the geometry is what track alignment *requires*; the VAE round-trip is what it
|
||||
*may* require. This script exists so the two can be A/B'd -- extract tracks from
|
||||
``resized_videos/`` and from ``roundtrip_videos/``, then diff the npz. If the track delta
|
||||
is small, large-scale runs can skip the VAE pass entirely (it is pure GPU cost per clip).
|
||||
|
||||
Usage:
|
||||
python data_pipeline/resize_videos.py \\
|
||||
--data-dir /home/hal-kevin/data/motion-stream-test
|
||||
|
||||
CPU-only; parallelize with --index / --limit sharding if needed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||||
|
||||
from fastvideo.dataset.transform import center_crop_th_tw, resize
|
||||
|
||||
TARGET_H, TARGET_W = 480, 832
|
||||
NUM_FRAMES = 121
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--data-dir", type=Path, required=True, help="Dataset root (contains videos/, etc.).")
|
||||
p.add_argument("--video-subdir", type=str, default="videos", help="Input video subdirectory.")
|
||||
p.add_argument("--out-subdir", type=str, default="resized_videos", help="Output subdirectory.")
|
||||
p.add_argument("--num-frames", type=int, default=NUM_FRAMES)
|
||||
p.add_argument("--height", type=int, default=TARGET_H)
|
||||
p.add_argument("--width", type=int, default=TARGET_W)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--index", type=int, nargs="+", default=None, metavar="IDX",
|
||||
help="Process only these video indices (e.g. --index 4 7 12). Assumes vid_%06d naming.")
|
||||
p.add_argument("--include-list", type=Path, default=None,
|
||||
help="Text file of clip filenames (one per line) to process; everything else in "
|
||||
"--video-subdir is ignored. Pairs with filter_clips.py's needs_resize.txt "
|
||||
"so only off-spec clips are re-encoded.")
|
||||
p.add_argument("--limit", type=int, default=None, help="Process only first N videos (smoke test).")
|
||||
p.add_argument("--rank", type=int, default=0, help="Shard index for CPU-parallel runs (0-indexed).")
|
||||
p.add_argument("--world-size", type=int, default=1, help="Total number of parallel processes.")
|
||||
p.add_argument("--min-frames", type=int, default=None,
|
||||
help="Skip clips with fewer than this many frames (default: --num-frames). "
|
||||
"Set 0 to keep short clips (output T then varies per clip).")
|
||||
p.add_argument("--force", action="store_true", help="Re-write even if output already exists.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def resize_video(path: Path, num_frames: int, height: int, width: int) -> np.ndarray:
|
||||
"""Crop+resize to the training geometry. Returns uint8 frames [T, H, W, C].
|
||||
|
||||
Reads sequentially and stops at num_frames or end-of-file, so clips shorter than
|
||||
num_frames yield what they have rather than raising (real-world shards are ragged).
|
||||
"""
|
||||
reader = imageio.get_reader(str(path))
|
||||
frames = []
|
||||
for i, frame in enumerate(reader):
|
||||
if i >= num_frames:
|
||||
break
|
||||
frames.append(np.asarray(frame))
|
||||
reader.close()
|
||||
if not frames:
|
||||
raise ValueError(f"no frames decoded from {path}")
|
||||
clip = torch.from_numpy(np.stack(frames)).permute(0, 3, 1, 2).float() / 255.0
|
||||
clip = center_crop_th_tw(clip, height, width, top_crop=False)
|
||||
clip = resize(clip, (height, width), interpolation_mode="bilinear")
|
||||
return (clip.clamp(0, 1) * 255).byte().permute(0, 2, 3, 1).numpy()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
videos_dir = args.data_dir / args.video_subdir
|
||||
out_dir = args.data_dir / args.out_subdir
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
videos = sorted(videos_dir.glob("*.mp4"))
|
||||
if args.include_list is not None:
|
||||
wanted = {ln.strip() for ln in args.include_list.read_text().splitlines() if ln.strip()}
|
||||
videos = [v for v in videos if v.name in wanted]
|
||||
if args.index is not None:
|
||||
wanted = {f"vid_{i:06d}.mp4" for i in args.index}
|
||||
videos = [v for v in videos if v.name in wanted]
|
||||
if args.limit is not None:
|
||||
videos = videos[:args.limit]
|
||||
if args.world_size > 1:
|
||||
videos = videos[args.rank::args.world_size]
|
||||
if not videos:
|
||||
print(f"[resize] no videos found in {videos_dir}", flush=True)
|
||||
return
|
||||
|
||||
min_frames = args.num_frames if args.min_frames is None else args.min_frames
|
||||
print(f"[resize] {len(videos)} videos → {out_dir} ({args.height}x{args.width}, no VAE)"
|
||||
f"{f' [shard {args.rank}/{args.world_size}]' if args.world_size > 1 else ''}", flush=True)
|
||||
n_ok = n_short = n_err = 0
|
||||
for k, vpath in enumerate(videos, 1):
|
||||
out_path = out_dir / vpath.name
|
||||
if out_path.exists() and not args.force:
|
||||
continue
|
||||
try:
|
||||
frames = resize_video(vpath, args.num_frames, args.height, args.width)
|
||||
except Exception as e: # noqa: BLE001
|
||||
n_err += 1
|
||||
print(f"[resize] [{k}/{len(videos)}] {vpath.name}: DECODE FAILED ({e}), skipping", flush=True)
|
||||
continue
|
||||
if frames.shape[0] < min_frames:
|
||||
n_short += 1
|
||||
print(f"[resize] [{k}/{len(videos)}] {vpath.name}: only {frames.shape[0]} frames "
|
||||
f"(< {min_frames}), skipping", flush=True)
|
||||
continue
|
||||
# Dot-prefixed so a leftover temp is NOT picked up by downstream `*.mp4` globs
|
||||
# (a stale "<name>.tmp.mp4" once got fed to the tracker and killed the worker).
|
||||
tmp = out_path.with_name(f".{out_path.stem}.tmp.mp4")
|
||||
imageio.mimsave(str(tmp), frames, fps=args.fps, macro_block_size=1)
|
||||
tmp.replace(out_path)
|
||||
n_ok += 1
|
||||
if k % 50 == 0 or k == len(videos):
|
||||
print(f"[resize] [{k}/{len(videos)}] ok={n_ok} short={n_short} err={n_err}", flush=True)
|
||||
|
||||
print(f"[resize] done. ok={n_ok} short={n_short} err={n_err}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+43
@@ -0,0 +1,43 @@
|
||||
#!/bin/bash
|
||||
# Stage-0 data pipeline: generate videos (Wan2.2-14B T2V) -> extract tracks (CoTracker v3).
|
||||
# Runs both steps on the shao_wm allocation via `srun --overlap`.
|
||||
#
|
||||
# You pick the GPUs (the node is shared); pass them explicitly:
|
||||
# CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh # full run (defaults below)
|
||||
# CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh --smoke # 2 videos
|
||||
# CUDA_VISIBLE_DEVICES=2,3 NUM_VIDEOS=50 OUTPUT_DIR=... bash data_pipeline/run.sh
|
||||
#
|
||||
# num_gpus is derived from CUDA_VISIBLE_DEVICES. Never run model code on the login node.
|
||||
set -euo pipefail
|
||||
|
||||
REPO=/mnt/weka/home/hao.zhang/shao/FastVideo
|
||||
PY="$REPO/.venv/bin/python"
|
||||
PROMPTS="${PROMPTS:-$REPO/examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt}"
|
||||
OUTPUT_DIR="${OUTPUT_DIR:-/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p}"
|
||||
JOB_NAME="${JOB_NAME:-shao_wm}"
|
||||
NUM_VIDEOS="${NUM_VIDEOS:-50}"
|
||||
GRID_SIZE="${GRID_SIZE:-50}"
|
||||
|
||||
[[ "${1:-}" == "--smoke" ]] && NUM_VIDEOS=2
|
||||
|
||||
: "${CUDA_VISIBLE_DEVICES:?set CUDA_VISIBLE_DEVICES to the GPUs to use, e.g. CUDA_VISIBLE_DEVICES=2,3}"
|
||||
IFS=',' read -ra _gpus <<< "$CUDA_VISIBLE_DEVICES"
|
||||
NUM_GPUS="${NUM_GPUS:-${#_gpus[@]}}"
|
||||
|
||||
JOBID="${JOBID:-$(squeue -u "$USER" -n "$JOB_NAME" -h -o %i | head -1)}"
|
||||
[[ -n "$JOBID" ]] || { echo "[run] ERROR: no running '$JOB_NAME' allocation" >&2; exit 1; }
|
||||
echo "[run] jobid=$JOBID CUDA_VISIBLE_DEVICES=$CUDA_VISIBLE_DEVICES num_gpus=$NUM_GPUS num_videos=$NUM_VIDEOS"
|
||||
|
||||
# /usr/bin/env (absolute: bare `env` may resolve to a broken ~/.local/bin/env) pins the
|
||||
# device list inside the step. Not SLURM --export: its comma parsing splits "2,3".
|
||||
SRUN=(srun --jobid="$JOBID" --overlap --ntasks=1 /usr/bin/env "CUDA_VISIBLE_DEVICES=$CUDA_VISIBLE_DEVICES")
|
||||
|
||||
echo "[run] === generate videos ==="
|
||||
"${SRUN[@]}" "$PY" "$REPO/data_pipeline/generate_videos.py" \
|
||||
--prompts "$PROMPTS" --output-dir "$OUTPUT_DIR" --num-videos "$NUM_VIDEOS" --num-gpus "$NUM_GPUS"
|
||||
|
||||
echo "[run] === extract tracks ==="
|
||||
"${SRUN[@]}" "$PY" "$REPO/data_pipeline/extract_tracks.py" \
|
||||
--data-dir "$OUTPUT_DIR" --grid-size "$GRID_SIZE"
|
||||
|
||||
echo "[run] done -> $OUTPUT_DIR"
|
||||
Executable
+25
@@ -0,0 +1,25 @@
|
||||
#!/usr/bin/env bash
|
||||
# Wan 2.2 5B VAE decode of the 32k HF latent dataset -> 480x832 mp4s.
|
||||
# Sharded across NODES * GPUS * PROCS_PER_GPU workers via SLURM_PROCID.
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${PARQUET_DIR:?}"; : "${OUT_DIR:?}"; : "${VAE_DIR:?}"
|
||||
NODES=${NODES:-8}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
NUM_SHARDS=$(( NODES * TASKS_PER_NODE ))
|
||||
|
||||
mkdir -p "$OUT_DIR" "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $NUM_SHARDS workers"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 12:00:00 -J decode_wansyn32k \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/decode_wansyn32k_%j.out" -e "$WORK/logs/decode_wansyn32k_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && \
|
||||
export HOME=$WORK HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch \
|
||||
MPLCONFIGDIR=$WORK/.mpl TOKENIZERS_PARALLELISM=false && \
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS )) && \
|
||||
python data_pipeline/decode_wansyn32k.py \
|
||||
--parquet-dir $PARQUET_DIR --vae-dir $VAE_DIR --out-dir $OUT_DIR \
|
||||
--num-shards $NUM_SHARDS --shard \$SLURM_PROCID'"
|
||||
@@ -0,0 +1,21 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node OpenVidHD download+extract (>=5s clips only), split-shard aware.
|
||||
# Parts are sharded across tasks (SLURM_PROCID/NTASKS) so the ~8.8TB download runs
|
||||
# in parallel across nodes. CPU/network only — no GPU. Idempotent (resumes).
|
||||
#
|
||||
# Usage:
|
||||
# NODES=8 TASKS_PER_NODE=3 bash data_pipeline/run_download_hd_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
NODES=${NODES:-8}; TASKS_PER_NODE=${TASKS_PER_NODE:-3} # concurrent HF downloads per node
|
||||
OUT=${OUT:-$WORK/openvid_1m} # goal folder
|
||||
FILTER=${FILTER:-$WORK/openvid/OpenVidHD_filtered.txt}
|
||||
mkdir -p "$OUT/videos" "$OUT/_zips" "$WORK/logs"
|
||||
echo "download OpenVidHD -> $OUT/videos ($NODES nodes x $TASKS_PER_NODE = $((NODES*TASKS_PER_NODE)) parallel downloads)"
|
||||
|
||||
sbatch -N "$NODES" --ntasks-per-node=$TASKS_PER_NODE --cpus-per-task=16 --mem=0 --exclusive \
|
||||
-t 48:00:00 -J openvid_dl --chdir="$WORK/FastVideo" \
|
||||
-o "$WORK/logs/openvid_dl_%j.out" -e "$WORK/logs/openvid_dl_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo bash -lc 'source .venv/bin/activate && export HF_HOME=$WORK/.hf && \
|
||||
python data_pipeline/openvid_download_hd.py \
|
||||
--videos-dir $OUT/videos --zip-dir $OUT/_zips --only-list $FILTER'"
|
||||
Executable
+32
@@ -0,0 +1,32 @@
|
||||
#!/bin/bash
|
||||
# Run extract_tracks.py in parallel across 4 GPUs.
|
||||
# Usage: bash data_pipeline/run_extract_tracks.sh [extra args]
|
||||
|
||||
DATA_DIR=/home/hal-kevin/data/motion-stream-test
|
||||
WORLD_SIZE=4
|
||||
LOG_FILE=data_pipeline/extract_tracks.log
|
||||
|
||||
> $LOG_FILE # truncate on each run
|
||||
|
||||
echo "[track] launching $WORLD_SIZE workers... logging to $LOG_FILE"
|
||||
|
||||
for RANK in $(seq 0 $((WORLD_SIZE - 1))); do
|
||||
CUDA_VISIBLE_DEVICES=$RANK python -u data_pipeline/extract_tracks.py \
|
||||
--data-dir $DATA_DIR \
|
||||
--videos-subdir roundtrip_videos \
|
||||
--grid-size 50 \
|
||||
--device cuda \
|
||||
--detect-entries \
|
||||
--sam-conf 0.75 \
|
||||
--sam-iou 0.9 \
|
||||
--sam-imgsz 1024 \
|
||||
--entry-sample-every 2 \
|
||||
--entry-min-area 0.001 \
|
||||
--entry-new-area 0.5 \
|
||||
--rank $RANK --world-size $WORLD_SIZE \
|
||||
"$@" \
|
||||
>> $LOG_FILE 2>&1 &
|
||||
done
|
||||
|
||||
wait
|
||||
echo "[track] all done. log at $LOG_FILE"
|
||||
@@ -0,0 +1,405 @@
|
||||
#!/bin/bash
|
||||
# End-to-end preprocessing of one OpenVid-WanTrack shard, WITHOUT the VAE round-trip:
|
||||
#
|
||||
# download shard -> extract -> crop+resize to 720p (CPU, parallel) -> fused tracks (GPU)
|
||||
#
|
||||
# NOTE (deliberate, per request): this skips Stage 2's VAE encode/decode and tracks the
|
||||
# resized source frames directly. The 50-clip A/B (ab_vae_roundtrip.sh) found tracks then
|
||||
# differ from round-trip tracks by ~5.5px on shared grid points, plus different entry-object
|
||||
# sets (n_objects differed on 23/50 clips). Fine for a throughput measurement or a training
|
||||
# A/B; see notes before adopting for a production set.
|
||||
#
|
||||
# NOTE (geometry): --height/--width define the coordinate frame the tracks live in. They
|
||||
# must match what Stage 5 crops/resizes to, or tracks won't align with the latents. 720p
|
||||
# here is NOT the current training geometry (480x832) -- set HEIGHT/WIDTH accordingly if
|
||||
# these tracks are meant to feed the existing training config.
|
||||
#
|
||||
# Usage:
|
||||
# bash data_pipeline/run_openvid_shard.sh # shard 0, 720p, 4 GPUs
|
||||
# LIMIT=50 bash data_pipeline/run_openvid_shard.sh # quick smoke run
|
||||
# SHARD=3 HEIGHT=480 WIDTH=832 bash data_pipeline/run_openvid_shard.sh
|
||||
# SKIP_DOWNLOAD=1 bash data_pipeline/run_openvid_shard.sh # shard already on disk
|
||||
set -euo pipefail
|
||||
|
||||
REPO_ID=${REPO_ID:-noctuashap/openvid-wantrack-clips}
|
||||
SHARD=${SHARD:-0}
|
||||
SHARD_NAME=$(printf "clips-%05d.tar" "$SHARD")
|
||||
DATA_ROOT=${DATA_ROOT:-$(printf "/home/hal-shared/motionstream/data/openvid-wantrack/shard%03d" "$SHARD")}
|
||||
HEIGHT=${HEIGHT:-720}
|
||||
WIDTH=${WIDTH:-1280}
|
||||
NUM_FRAMES=${NUM_FRAMES:-121}
|
||||
FPS=${FPS:-24}
|
||||
GPUS=${GPUS:-0,1,2,3}
|
||||
CPU_WORKERS=${CPU_WORKERS:-$(( $(nproc) > 16 ? 16 : $(nproc) ))}
|
||||
LIMIT=${LIMIT:-}
|
||||
AMP=${AMP:-1} # bf16 CoTracker (validated: ~1.1px delta, ~1.5x faster)
|
||||
COMPILE=${COMPILE:-0} # torch.compile main pass (~10-15%; pays a per-worker warmup)
|
||||
VIZ=${VIZ:-0}
|
||||
SKIP_DOWNLOAD=${SKIP_DOWNLOAD:-0}
|
||||
TRACKS=${TRACKS:-1} # 0 = skip tracking (parquet-only pass over existing npz)
|
||||
PARQUET=${PARQUET:-0} # 1 = also run Stage 5 (v1_preprocess) -> training parquets
|
||||
# Stage 5's dataloader defaults to 1 worker, so video decode blocks the GPU between clips
|
||||
# (the same bubble --prefetch removes in tracking). More workers overlap decode with encode,
|
||||
# but each holds a decoded 720p clip (~334MB of raw frames), so this also drives host RAM.
|
||||
PARQUET_WORKERS=${PARQUET_WORKERS:-2}
|
||||
# How many samples the parquet writer buffers before flushing to disk. At 720p each sample's
|
||||
# latents are ~2.4x the 480p reference, so the default (256) can OOM a 127GB node. Flushing
|
||||
# every samples_per_file keeps the in-RAM buffer small.
|
||||
PARQUET_FLUSH=${PARQUET_FLUSH:-64}
|
||||
MODEL_PATH=${MODEL_PATH:-/home/hal-kevin/models/trackwan_1.3b_i2v_control_init}
|
||||
RESUME=${RESUME:-0} # 1 = continue an interrupted run: skip finished phases AND
|
||||
# already-tracked clips (implies FORCE_TRACKS=0)
|
||||
FORCE_TRACKS=${FORCE_TRACKS:-$([[ "$RESUME" == "1" ]] && echo 0 || echo 1)}
|
||||
DRY_RUN=${DRY_RUN:-0} # 1 = print the worker command lines and exit (no work done)
|
||||
|
||||
cd "$(dirname "$0")/.."
|
||||
IFS=',' read -ra GPU_ARR <<< "$GPUS"
|
||||
WORLD_SIZE=${#GPU_ARR[@]}
|
||||
DL_DIR="$DATA_ROOT/download"
|
||||
RAW_DIR="$DATA_ROOT/raw_videos"
|
||||
LOG_DIR="$DATA_ROOT/logs"
|
||||
mkdir -p "$DATA_ROOT" "$LOG_DIR"
|
||||
|
||||
# Two concurrent runs would interleave resize writes with the tracker's video glob, so the
|
||||
# tracker would see a partially-populated dir and silently process a subset. Refuse to overlap.
|
||||
LOCK="$DATA_ROOT/.run.lock"
|
||||
if ! mkdir "$LOCK" 2>/dev/null; then
|
||||
echo "[openvid] ERROR: another run holds $LOCK (remove it if stale)" >&2; exit 1
|
||||
fi
|
||||
trap 'rmdir "$LOCK" 2>/dev/null || true' EXIT
|
||||
|
||||
LIMIT_ARGS=(); [[ -n "$LIMIT" ]] && LIMIT_ARGS=(--limit "$LIMIT")
|
||||
secs() { date +%s; }
|
||||
hms() { awk -v s="$1" 'BEGIN{printf "%dm%02ds", s/60, s%60}'; }
|
||||
|
||||
# --- progress tracking: survives SIGKILL (cluster reapers), enables RESUME=1 --------
|
||||
PROGRESS="$DATA_ROOT/progress.json"
|
||||
prog_set() { # prog_set <phase> <json-object>
|
||||
python - "$PROGRESS" "$1" "$2" "$SHARD_NAME" "${HEIGHT}x${WIDTH}@${NUM_FRAMES}f" <<'PY'
|
||||
import datetime, json, sys
|
||||
from pathlib import Path
|
||||
p, phase, payload, shard, target = Path(sys.argv[1]), sys.argv[2], json.loads(sys.argv[3]), sys.argv[4], sys.argv[5]
|
||||
d = json.loads(p.read_text()) if p.exists() else {}
|
||||
d.update(shard=shard, target=target) # informational: the most recent run's target
|
||||
now = datetime.datetime.now().isoformat(timespec="seconds")
|
||||
# `target` is recorded PER PHASE: a later run at a different geometry must not be able to
|
||||
# reuse filter/track outputs produced for the old one.
|
||||
d.setdefault("phases", {})[phase] = {**payload, "target": target, "ts": now}
|
||||
d["updated"] = now
|
||||
tmp = p.with_suffix(".json.tmp") # atomic: a kill mid-write must not corrupt state
|
||||
tmp.write_text(json.dumps(d, indent=2))
|
||||
tmp.replace(p)
|
||||
PY
|
||||
}
|
||||
prog_done() { # prog_done <phase> -> 0 if that phase completed for this target
|
||||
python - "$PROGRESS" "$1" "${HEIGHT}x${WIDTH}@${NUM_FRAMES}f" <<'PY'
|
||||
import json, sys
|
||||
from pathlib import Path
|
||||
p = Path(sys.argv[1])
|
||||
if not p.exists():
|
||||
sys.exit(1)
|
||||
ph = json.loads(p.read_text()).get("phases", {}).get(sys.argv[2], {})
|
||||
# a different target geometry invalidates that phase's outputs
|
||||
sys.exit(0 if (ph.get("done") and ph.get("target") == sys.argv[3]) else 1)
|
||||
PY
|
||||
}
|
||||
[[ "$RESUME" == "1" ]] && echo "[openvid] RESUME=1 -- finished phases and already-tracked clips will be skipped"
|
||||
|
||||
echo "[openvid] shard=$SHARD_NAME target=${HEIGHT}x${WIDTH} gpus=$GPUS cpu_workers=$CPU_WORKERS"
|
||||
echo "[openvid] data root: $DATA_ROOT"
|
||||
|
||||
# --- 1. download -------------------------------------------------------------------
|
||||
if [[ "$SKIP_DOWNLOAD" != "1" && ! -f "$DL_DIR/$SHARD_NAME" ]]; then
|
||||
echo "[openvid] downloading $SHARD_NAME (~3.3 GB) ..."
|
||||
t=$(secs)
|
||||
if command -v hf >/dev/null 2>&1; then
|
||||
hf download "$REPO_ID" "$SHARD_NAME" --repo-type dataset --local-dir "$DL_DIR"
|
||||
elif command -v huggingface-cli >/dev/null 2>&1; then
|
||||
huggingface-cli download "$REPO_ID" "$SHARD_NAME" --repo-type dataset --local-dir "$DL_DIR"
|
||||
else
|
||||
echo "[openvid] ERROR: neither 'hf' nor 'huggingface-cli' found (pip install -U huggingface_hub)" >&2
|
||||
exit 1
|
||||
fi
|
||||
echo "[openvid] download: $(hms $(( $(secs) - t )))"
|
||||
prog_set download "{\"done\": true, \"secs\": $(( $(secs) - t ))}"
|
||||
else
|
||||
echo "[openvid] download: skipped (have $DL_DIR/$SHARD_NAME)"
|
||||
prog_set download '{"done": true, "note": "pre-existing"}'
|
||||
fi
|
||||
|
||||
# --- 2. extract --------------------------------------------------------------------
|
||||
if [[ -z "$(ls -A "$RAW_DIR" 2>/dev/null)" ]]; then
|
||||
echo "[openvid] extracting ..."
|
||||
t=$(secs)
|
||||
mkdir -p "$RAW_DIR"
|
||||
tar -xf "$DL_DIR/$SHARD_NAME" -C "$RAW_DIR"
|
||||
# flatten any nested layout so *.mp4 all sit directly in RAW_DIR
|
||||
find "$RAW_DIR" -mindepth 2 -name '*.mp4' -exec mv -t "$RAW_DIR" {} + 2>/dev/null || true
|
||||
find "$RAW_DIR" -mindepth 1 -type d -empty -delete 2>/dev/null || true
|
||||
echo "[openvid] extract: $(hms $(( $(secs) - t )))"
|
||||
else
|
||||
echo "[openvid] extract: skipped (raw_videos/ non-empty)"
|
||||
fi
|
||||
N_RAW=$(ls "$RAW_DIR"/*.mp4 2>/dev/null | wc -l || true)
|
||||
echo "[openvid] raw clips: $N_RAW"
|
||||
[[ "$N_RAW" -gt 0 ]] || { echo "[openvid] ERROR: no mp4s extracted" >&2; exit 1; }
|
||||
prog_set extract "{\"done\": true, \"raw_clips\": $N_RAW}"
|
||||
|
||||
# --- 3. crop+resize to target geometry (CPU, parallel) ------------------------------
|
||||
# If the clips are already at the target geometry, resizing is a pure lossy re-encode
|
||||
# (measured 40 dB / 1.9-per-255 on this shard -- the same order as the VAE round-trip's
|
||||
# distortion, for zero benefit) plus a duplicate copy on disk. Track the raw clips instead.
|
||||
VID_SUBDIR=videos
|
||||
if [[ "$RESUME" == "1" ]] && prog_done filter && [[ -n "$(ls -A "$DATA_ROOT/videos" 2>/dev/null)" ]]; then
|
||||
T_RESIZE=0
|
||||
N_VID=$(ls "$DATA_ROOT"/videos/*.mp4 2>/dev/null | wc -l || true)
|
||||
echo "[openvid] filter: skipped (progress.json says done; $N_VID clips staged)"
|
||||
elif [[ "${FORCE_RESIZE:-0}" != "1" ]]; then
|
||||
# Scan every clip's metadata (no decode) and symlink through the conforming ones.
|
||||
# Clips at other resolutions/lengths are skipped and listed in skipped_clips.json.
|
||||
t=$(secs)
|
||||
python -u data_pipeline/filter_clips.py \
|
||||
--src-dir "$RAW_DIR" --out-dir "$DATA_ROOT/videos" \
|
||||
--height "$HEIGHT" --width "$WIDTH" --num-frames "$NUM_FRAMES" --clean \
|
||||
2>&1 | tee -a "$LOG_DIR/filter.log"
|
||||
# Rescue pass: clips that are readable but off-spec get re-encoded to the target
|
||||
# (only these -- the conforming majority stays symlinked, never re-encoded).
|
||||
NEEDS="$DATA_ROOT/needs_resize.txt"
|
||||
N_RESCUE=$(wc -l < "$NEEDS" 2>/dev/null || echo 0)
|
||||
if [[ "$N_RESCUE" -gt 0 ]]; then
|
||||
echo "[openvid] rescuing $N_RESCUE off-spec clip(s) by resize -> ${HEIGHT}x${WIDTH} ..."
|
||||
pids=()
|
||||
for i in $(seq 0 $((CPU_WORKERS - 1))); do
|
||||
python -u data_pipeline/resize_videos.py \
|
||||
--data-dir "$DATA_ROOT" \
|
||||
--video-subdir raw_videos \
|
||||
--out-subdir videos \
|
||||
--include-list "$NEEDS" \
|
||||
--height "$HEIGHT" --width "$WIDTH" \
|
||||
--num-frames "$NUM_FRAMES" --fps "$FPS" \
|
||||
--rank "$i" --world-size "$CPU_WORKERS" \
|
||||
>> "$LOG_DIR/resize.log" 2>&1 &
|
||||
pids+=($!)
|
||||
done
|
||||
rfail=0; for p in "${pids[@]}"; do wait "$p" || rfail=$((rfail + 1)); done
|
||||
[[ $rfail -gt 0 ]] && echo "[openvid] WARNING: $rfail rescue worker(s) failed -- see $LOG_DIR/resize.log"
|
||||
fi
|
||||
T_RESIZE=$(( $(secs) - t ))
|
||||
N_VID=$(ls "$DATA_ROOT"/videos/*.mp4 2>/dev/null | wc -l || true)
|
||||
N_SKIP=$(( N_RAW - N_VID ))
|
||||
echo "[openvid] filter: $(hms $T_RESIZE) conforming: $N_VID / $N_RAW skipped: $N_SKIP"
|
||||
# only record success if something was actually staged
|
||||
[[ "$N_VID" -gt 0 ]] && prog_set filter "{\"done\": true, \"kept\": $N_VID, \"skipped\": $N_SKIP, \"secs\": $T_RESIZE}"
|
||||
if [[ "$N_VID" -eq 0 ]]; then
|
||||
echo "[openvid] ERROR: no clips match ${HEIGHT}x${WIDTH}@${NUM_FRAMES}f." >&2
|
||||
echo "[openvid] Run with FORCE_RESIZE=1 to re-encode them to the target instead." >&2
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
echo "[openvid] resizing $NATIVE -> ${HEIGHT}x${WIDTH} across $CPU_WORKERS CPU workers ..."
|
||||
t=$(secs)
|
||||
pids=()
|
||||
for i in $(seq 0 $((CPU_WORKERS - 1))); do
|
||||
python -u data_pipeline/resize_videos.py \
|
||||
--data-dir "$DATA_ROOT" \
|
||||
--video-subdir raw_videos \
|
||||
--out-subdir videos \
|
||||
--height "$HEIGHT" --width "$WIDTH" \
|
||||
--num-frames "$NUM_FRAMES" --fps "$FPS" \
|
||||
--rank "$i" --world-size "$CPU_WORKERS" \
|
||||
"${LIMIT_ARGS[@]}" \
|
||||
>> "$LOG_DIR/resize.log" 2>&1 &
|
||||
pids+=($!)
|
||||
done
|
||||
rfail=0; for p in "${pids[@]}"; do wait "$p" || rfail=$((rfail + 1)); done
|
||||
[[ $rfail -gt 0 ]] && echo "[openvid] WARNING: $rfail resize worker(s) failed -- see $LOG_DIR/resize.log"
|
||||
T_RESIZE=$(( $(secs) - t ))
|
||||
find "$DATA_ROOT/$VID_SUBDIR" -name '.*.tmp.mp4' -delete 2>/dev/null || true # drop any stale temps
|
||||
N_VID=$(ls "$DATA_ROOT"/$VID_SUBDIR/*.mp4 2>/dev/null | wc -l || true)
|
||||
echo "[openvid] resize: $(hms $T_RESIZE) usable clips: $N_VID / $N_RAW"
|
||||
fi
|
||||
[[ "$N_VID" -gt 0 ]] || { echo "[openvid] ERROR: no clips available to track" >&2; exit 1; }
|
||||
|
||||
# --- 4. manifest (so tracks get points_path patched + Stage 5 has an entry point) ---
|
||||
python - "$DATA_ROOT" "$FPS" "$NUM_FRAMES" "$HEIGHT" "$WIDTH" "$VID_SUBDIR" <<'PY'
|
||||
import json, sys
|
||||
from pathlib import Path
|
||||
root, fps, nf, h, w = Path(sys.argv[1]), float(sys.argv[2]), int(sys.argv[3]), int(sys.argv[4]), int(sys.argv[5])
|
||||
tracks_dir = root / "tracks"
|
||||
items = []
|
||||
for i, p in enumerate(sorted((root / sys.argv[6]).glob("*.mp4"))):
|
||||
it = {"idx": i, "path": p.name, "cap": [""], "fps": fps, "num_frames": nf,
|
||||
"duration": nf / fps, "resolution": {"width": w, "height": h}}
|
||||
# This manifest gets rewritten every run, so re-attach points_path here whenever the
|
||||
# npz exists. In PHASE=parquet, tracking is skipped and never patches it back, so Stage 5
|
||||
# would otherwise see no track sidecar (PreprocessPipeline_I2V_Track then errors).
|
||||
npz = tracks_dir / f"{p.stem}.npz"
|
||||
if npz.exists():
|
||||
it["points_path"] = str(npz.resolve())
|
||||
items.append(it)
|
||||
(root / "videos2caption.json").write_text(json.dumps(items, indent=2))
|
||||
n_pts = sum(1 for it in items if "points_path" in it)
|
||||
print(f"[openvid] manifest: {len(items)} entries ({n_pts} with points_path) -> {root/'videos2caption.json'}")
|
||||
PY
|
||||
|
||||
# Real captions from OpenVid-1M (joins 1:1 on clip filename). Without this every clip
|
||||
# carries an empty prompt and Stage 5 bakes identical null T5 embeddings into the parquets.
|
||||
if [[ "${CAPTIONS:-1}" == "1" ]]; then
|
||||
python -u data_pipeline/add_captions.py \
|
||||
--manifest "$DATA_ROOT/videos2caption.json" \
|
||||
--min-coverage "${MIN_CAPTION_COVERAGE:-0.9}" \
|
||||
2>&1 | tee -a "$LOG_DIR/captions.log" | tail -3
|
||||
cap_rc=${PIPESTATUS[0]}
|
||||
if [[ "$cap_rc" -ne 0 ]]; then
|
||||
echo "[openvid] ERROR: caption join failed -- see $LOG_DIR/captions.log" >&2
|
||||
echo "[openvid] set CAPTIONS=0 to proceed with empty captions (tracks are still valid;" >&2
|
||||
echo "[openvid] parquets built from them would have dead text conditioning)." >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
# --- 5. fused tracks (stages 3+4 in one pass, no VAE round-trip) --------------------
|
||||
if [[ "$TRACKS" != "1" ]]; then
|
||||
T_TRACK=0; tfail=0
|
||||
N_NPZ=0
|
||||
N_NPZ_ALL=$(ls "$DATA_ROOT"/tracks/*.npz 2>/dev/null | wc -l || true)
|
||||
echo "[openvid] tracks: SKIPPED (TRACKS=0) existing npz: $N_NPZ_ALL / $N_VID"
|
||||
else
|
||||
echo "[openvid] extracting tracks across $WORLD_SIZE GPUs ..."
|
||||
SPEED=(); [[ "$AMP" == "1" ]] && SPEED+=(--amp); [[ "$COMPILE" == "1" ]] && SPEED+=(--compile)
|
||||
VIZ_ARGS=(); [[ "$VIZ" == "1" ]] && VIZ_ARGS=(--viz --viz-dir "$DATA_ROOT/viz")
|
||||
if [[ "$COMPILE" == "1" ]]; then
|
||||
export TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR:-$HOME/.cache/torchinductor}
|
||||
export TRITON_CACHE_DIR=${TRITON_CACHE_DIR:-$HOME/.cache/triton}
|
||||
fi
|
||||
FORCE_ARGS=(); [[ "$FORCE_TRACKS" == "1" ]] && FORCE_ARGS=(--force)
|
||||
if [[ "$DRY_RUN" == "1" ]]; then
|
||||
echo "[openvid] DRY RUN -- rank 0 track worker would be:"
|
||||
echo " CUDA_VISIBLE_DEVICES=${GPU_ARR[0]} python -u data_pipeline/extract_tracks.py" \
|
||||
"--data-dir $DATA_ROOT --videos-subdir $VID_SUBDIR --out-subdir tracks" \
|
||||
"--grid-size 50 --device cuda --detect-entries --sam-conf 0.75 --sam-iou 0.9 --sam-imgsz 1024" \
|
||||
"--entry-sample-every 2 --entry-min-area 0.001 --entry-new-area 0.5" \
|
||||
"--segment --vis-override-every 3 ${SPEED[*]} ${VIZ_ARGS[*]} ${FORCE_ARGS[*]} ${LIMIT_ARGS[*]}" \
|
||||
"--rank 0 --world-size $WORLD_SIZE"
|
||||
exit 0
|
||||
fi
|
||||
t=$(secs)
|
||||
pids=()
|
||||
for i in "${!GPU_ARR[@]}"; do
|
||||
CUDA_VISIBLE_DEVICES=${GPU_ARR[$i]} python -u data_pipeline/extract_tracks.py \
|
||||
--data-dir "$DATA_ROOT" \
|
||||
--videos-subdir "$VID_SUBDIR" \
|
||||
--out-subdir tracks \
|
||||
--grid-size 50 --device cuda \
|
||||
--detect-entries --sam-conf 0.75 --sam-iou 0.9 --sam-imgsz 1024 \
|
||||
--entry-sample-every 2 --entry-min-area 0.001 --entry-new-area 0.5 \
|
||||
--segment --vis-override-every 3 \
|
||||
"${SPEED[@]}" "${VIZ_ARGS[@]}" "${FORCE_ARGS[@]}" "${LIMIT_ARGS[@]}" \
|
||||
--rank "$i" --world-size "$WORLD_SIZE" \
|
||||
>> "$LOG_DIR/tracks.log" 2>&1 &
|
||||
pids+=($!)
|
||||
done
|
||||
tfail=0; for p in "${pids[@]}"; do wait "$p" || tfail=$((tfail + 1)); done
|
||||
[[ $tfail -gt 0 ]] && echo "[openvid] WARNING: $tfail track worker(s) failed -- see $LOG_DIR/tracks.log"
|
||||
T_TRACK=$(( $(secs) - t ))
|
||||
# Count only npz written by THIS run -- counting the whole dir would fold in earlier runs
|
||||
# and (with FORCE_TRACKS=0) report a rate for work that was skipped.
|
||||
N_NPZ=$(find "$DATA_ROOT/tracks" -name '*.npz' -newermt "@$t" 2>/dev/null | wc -l || true)
|
||||
N_NPZ_ALL=$(ls "$DATA_ROOT"/tracks/*.npz 2>/dev/null | wc -l || true)
|
||||
echo "[openvid] tracks: $(hms $T_TRACK) npz this run: $N_NPZ total in dir: $N_NPZ_ALL / $N_VID"
|
||||
if [[ "$N_NPZ_ALL" -ge "$N_VID" && "$tfail" -eq 0 && -z "$LIMIT" ]]; then
|
||||
prog_set tracks "{\"done\": true, \"npz\": $N_NPZ_ALL, \"clips\": $N_VID, \"secs\": $T_TRACK}"
|
||||
echo "[openvid] shard COMPLETE"
|
||||
else
|
||||
prog_set tracks "{\"done\": false, \"npz\": $N_NPZ_ALL, \"clips\": $N_VID, \"secs\": $T_TRACK}"
|
||||
[[ "$N_NPZ_ALL" -lt "$N_VID" ]] && \
|
||||
echo "[openvid] INCOMPLETE: $(( N_VID - N_NPZ_ALL )) clips remain -- resume with: RESUME=1 SKIP_DOWNLOAD=1 bash $0"
|
||||
fi
|
||||
fi
|
||||
|
||||
# --- 6. Stage 5: parquets (opt-in; the tracks above are already usable without this) --
|
||||
T_PARQUET=0
|
||||
if [[ "$PARQUET" == "1" ]]; then
|
||||
if [[ "$N_NPZ_ALL" -lt "$N_VID" ]]; then
|
||||
echo "[openvid] SKIPPING parquets: tracks incomplete ($N_NPZ_ALL/$N_VID)" >&2
|
||||
else
|
||||
# num_latent_t = (num_frames - 1)/4 + 1 for WanVAE's 4x temporal compression.
|
||||
NLT=$(( (NUM_FRAMES - 1) / 4 + 1 ))
|
||||
echo "$DATA_ROOT/$VID_SUBDIR,$DATA_ROOT/videos2caption.json" > "$DATA_ROOT/data_merge.txt"
|
||||
# Parquet output location. Default keeps it beside the shard's other data; set
|
||||
# PARQUET_ROOT to collect all shards' parquets under one tree (their own dir),
|
||||
# each in a per-shard subdir so shards stay independent (wipe/verify are per-shard).
|
||||
if [[ -n "${PARQUET_ROOT:-}" ]]; then
|
||||
PQ_OUT="$PARQUET_ROOT/$(printf 'shard%03d' "$SHARD")"
|
||||
else
|
||||
PQ_OUT="$DATA_ROOT/preprocessed_i2v_track"
|
||||
fi
|
||||
# Stage 5 now resumes by clip id (preprocess_pipeline_base.py): a re-run skips clips
|
||||
# already written and appends only the rest, so a reaper kill costs minutes (the
|
||||
# unflushed buffer), not the whole shard. Do NOT wipe -- that would discard progress.
|
||||
echo "[openvid] Stage 5: parquets (${HEIGHT}x${WIDTH}, ${NUM_FRAMES}f, num_latent_t=$NLT, train_fps=$FPS) ..."
|
||||
t=$(secs)
|
||||
# --train_fps MUST equal the source fps: a mismatch makes FrameSamplingStage resample
|
||||
# (duplicating/dropping frames) so latents no longer align with the tracks in the same row.
|
||||
# v1_preprocess asserts num_gpus == 1 (fastvideo/pipelines/preprocess/v1_preprocess.py:27),
|
||||
# so Stage 5 is single-GPU per shard. Scale it by running shards concurrently
|
||||
# (one GPU each) rather than by raising nproc_per_node.
|
||||
CUDA_VISIBLE_DEVICES="${PARQUET_GPU:-${GPU_ARR[0]}}" \
|
||||
torchrun --nproc_per_node=1 -m fastvideo.pipelines.preprocess.v1_preprocess \
|
||||
--model_path "$MODEL_PATH" \
|
||||
--data_merge_path "$DATA_ROOT/data_merge.txt" \
|
||||
--output_dir "$PQ_OUT" \
|
||||
--preprocess_task i2v_track \
|
||||
--num_frames "$NUM_FRAMES" \
|
||||
--num_latent_t "$NLT" \
|
||||
--train_fps "$FPS" \
|
||||
--max_height "$HEIGHT" \
|
||||
--max_width "$WIDTH" \
|
||||
--preprocess_video_batch_size 1 \
|
||||
--dataloader_num_workers "$PARQUET_WORKERS" \
|
||||
--samples_per_file "${PARQUET_SAMPLES:-64}" \
|
||||
--flush_frequency "$PARQUET_FLUSH" \
|
||||
>> "$LOG_DIR/parquet.log" 2>&1
|
||||
prc=$?
|
||||
T_PARQUET=$(( $(secs) - t ))
|
||||
N_PQ=$(find "$PQ_OUT" -name '*.parquet' 2>/dev/null | wc -l || true)
|
||||
# Verify integrity: total rows must equal clip count AND be free of duplicate ids.
|
||||
# A reaper kill mid-run leaves a partial (rows < N_VID) which the next attempt wipes+redoes.
|
||||
read -r N_ROWS N_UNIQ < <(python - "$PQ_OUT" <<'PY'
|
||||
import glob, sys
|
||||
import pyarrow.parquet as pq
|
||||
ids = []
|
||||
for f in glob.glob(f"{sys.argv[1]}/**/*.parquet", recursive=True):
|
||||
ids += pq.read_table(f, columns=["id"]).column("id").to_pylist()
|
||||
print(len(ids), len(set(ids)))
|
||||
PY
|
||||
)
|
||||
if [[ $prc -eq 0 && "$N_ROWS" == "$N_VID" && "$N_UNIQ" == "$N_VID" ]]; then
|
||||
echo "[openvid] parquets: $(hms $T_PARQUET) rows=$N_ROWS unique=$N_UNIQ files=$N_PQ"
|
||||
prog_set parquet "{\"done\": true, \"rows\": $N_ROWS, \"files\": $N_PQ, \"secs\": $T_PARQUET}"
|
||||
else
|
||||
echo "[openvid] Stage 5 INCOMPLETE (rc=$prc, rows=$N_ROWS unique=$N_UNIQ, want $N_VID) -- see $LOG_DIR/parquet.log" >&2
|
||||
prog_set parquet "{\"done\": false, \"rows\": $N_ROWS, \"unique\": $N_UNIQ, \"secs\": $T_PARQUET}"
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
|
||||
# --- 7. summary --------------------------------------------------------------------
|
||||
RESULTS="$DATA_ROOT/shard_results.txt"
|
||||
{
|
||||
echo "=== $(date -u '+%Y-%m-%d %H:%M:%S') UTC shard=$SHARD_NAME ${HEIGHT}x${WIDTH} gpus=$GPUS cpu=$CPU_WORKERS amp=$AMP compile=$COMPILE viz=$VIZ (no VAE round-trip) ==="
|
||||
awk -v r="$T_RESIZE" -v tk="$T_TRACK" -v n="$N_NPZ" -v w="$WORLD_SIZE" -v c="$CPU_WORKERS" 'BEGIN {
|
||||
if (n == 0) { print "no npz produced"; exit }
|
||||
printf "resize (CPU): %6ds %5.2fs/clip/worker\n", r, r*c/n
|
||||
printf "tracks (GPU): %6ds %5.2fs/clip/worker %5.1f clips/min\n", tk, tk*w/n, 60*n/tk
|
||||
printf "total: %6ds for %d clips\n", r+tk, n
|
||||
printf " -> 259k clips on %d GPUs (tracks only): %.1f h\n", w, 259000*(tk*w/n)/w/3600
|
||||
printf " -> full shard (1000 clips) at this rate: %.1f min\n", (tk/n)*1000/60
|
||||
}'
|
||||
# `|| true`: a false [[ ]] would make this block (and so the piped tee) exit non-zero
|
||||
# under `set -e -o pipefail`, failing the whole shard after the work already succeeded.
|
||||
{ [[ "$VIZ" == "1" ]] && echo " NOTE: viz=1 -- overlay rendering dominates the GPU phase (~4x); projections are pessimistic, not a throughput measurement."; } || true
|
||||
{ [[ -n "$LIMIT" ]] && echo " NOTE: limit=$LIMIT -- per-worker startup is a large share at this size; totals understate steady-state throughput."; } || true
|
||||
} | tee -a "$RESULTS"
|
||||
echo "[openvid] appended to $RESULTS"
|
||||
echo "[openvid] outputs: $DATA_ROOT/{videos,tracks}$([[ "$VIZ" == "1" ]] && echo ",viz") logs: $LOG_DIR"
|
||||
@@ -0,0 +1,215 @@
|
||||
#!/bin/bash
|
||||
# Drive run_openvid_shard.sh over a range of shards, with resume and per-shard accounting.
|
||||
#
|
||||
# Every shard runs with RESUME=1, so re-invoking after an interruption (cluster reaper,
|
||||
# node loss, Ctrl-C) picks up exactly where it stopped: finished shards are skipped via
|
||||
# their progress.json, and a half-finished shard resumes at the first untracked clip.
|
||||
#
|
||||
# Usage:
|
||||
# SHARDS=0-130 bash data_pipeline/run_openvid_shards.sh
|
||||
# SHARDS=0-9,20,30-35 bash data_pipeline/run_openvid_shards.sh
|
||||
# SHARDS=0-130 CLEANUP=1 bash data_pipeline/run_openvid_shards.sh # drop tar+raw after each
|
||||
# SHARDS=0-3 LIMIT=20 bash data_pipeline/run_openvid_shards.sh # smoke run
|
||||
#
|
||||
# Passes through the per-shard knobs (GPUS, HEIGHT/WIDTH, AMP, COMPILE, VIZ, CPU_WORKERS,
|
||||
# LIMIT, DATA_ROOT_BASE); see run_openvid_shard.sh for their meanings.
|
||||
set -uo pipefail # NOT -e: one bad shard must not kill a 130-shard run
|
||||
|
||||
SHARDS=${SHARDS:-0}
|
||||
CLEANUP=${CLEANUP:-0} # 1 = delete download/ and raw_videos/ once a shard completes
|
||||
STOP_ON_FAIL=${STOP_ON_FAIL:-0}
|
||||
# PARALLEL=1 runs one shard pipeline per GPU concurrently instead of one shard at a time
|
||||
# across all GPUs. Stage 5 (v1_preprocess) asserts a single GPU, so sequential mode leaves
|
||||
# 3 of 4 GPUs idle for the whole parquet phase; shard-level parallelism keeps them all busy.
|
||||
PARALLEL=${PARALLEL:-0}
|
||||
GPUS=${GPUS:-0,1,2,3}
|
||||
# PHASE picks what this invocation produces:
|
||||
# tracks -- tracks only, one shard at a time across all GPUs (fastest per-shard: ~13 min)
|
||||
# parquet -- Stage 5 only, over shards that already have tracks; forces PARALLEL=1 because
|
||||
# v1_preprocess is single-GPU, so concurrency has to come from running shards
|
||||
# both -- everything per shard (honours PARALLEL as set)
|
||||
PHASE=${PHASE:-both}
|
||||
case "$PHASE" in
|
||||
tracks) export TRACKS=1 PARQUET=0 ;;
|
||||
parquet) export TRACKS=0 PARQUET=1; PARALLEL=1 ;;
|
||||
both) export TRACKS=1 ;;
|
||||
*) echo "[shards] ERROR: PHASE must be tracks|parquet|both (got '$PHASE')" >&2; exit 1 ;;
|
||||
esac
|
||||
DATA_ROOT_BASE=${DATA_ROOT_BASE:-/home/hal-shared/motionstream/data/openvid-wantrack/shard}
|
||||
shard_root() { printf "%s%03d" "$DATA_ROOT_BASE" "$1"; } # zero-padded: shard000 .. shard259
|
||||
|
||||
cd "$(dirname "$0")/.."
|
||||
SHARD_SCRIPT=data_pipeline/run_openvid_shard.sh
|
||||
|
||||
# --- expand "0-9,20,30-35" into a list ---------------------------------------------
|
||||
expand() {
|
||||
local spec=$1 tok lo hi out=()
|
||||
IFS=',' read -ra toks <<< "$spec"
|
||||
for tok in "${toks[@]}"; do
|
||||
if [[ "$tok" =~ ^([0-9]+)-([0-9]+)$ ]]; then
|
||||
lo=${BASH_REMATCH[1]}; hi=${BASH_REMATCH[2]}
|
||||
(( lo <= hi )) || { echo "[shards] ERROR: bad range '$tok'" >&2; exit 1; }
|
||||
for ((i = lo; i <= hi; i++)); do out+=("$i"); done
|
||||
elif [[ "$tok" =~ ^[0-9]+$ ]]; then
|
||||
out+=("$tok")
|
||||
else
|
||||
echo "[shards] ERROR: cannot parse '$tok' (want N or A-B, comma-separated)" >&2; exit 1
|
||||
fi
|
||||
done
|
||||
printf '%s\n' "${out[@]}"
|
||||
}
|
||||
mapfile -t SHARD_LIST < <(expand "$SHARDS")
|
||||
N_SHARDS=${#SHARD_LIST[@]}
|
||||
|
||||
# 260 shards exist: clips-00000.tar .. clips-00259.tar
|
||||
for s in "${SHARD_LIST[@]}"; do
|
||||
(( s <= 259 )) || { echo "[shards] ERROR: shard $s out of range (max 259)" >&2; exit 1; }
|
||||
done
|
||||
|
||||
# A shard counts as complete only once every phase this run is producing has finished --
|
||||
# with PARQUET=1 that includes Stage 5, so a shard whose tracks landed but whose parquets
|
||||
# failed is retried rather than skipped.
|
||||
case "$PHASE" in
|
||||
tracks) REQUIRED_PHASES="tracks" ;;
|
||||
parquet) REQUIRED_PHASES="parquet" ;; # tracks were the previous pass's job
|
||||
both) REQUIRED_PHASES="tracks"; [[ "${PARQUET:-0}" == "1" ]] && REQUIRED_PHASES="tracks parquet" ;;
|
||||
esac
|
||||
shard_complete() { # shard_complete <n> -> 0 if all required phases are done
|
||||
python - "$(shard_root "$1")/progress.json" $REQUIRED_PHASES <<'PY'
|
||||
import json, sys
|
||||
from pathlib import Path
|
||||
p = Path(sys.argv[1])
|
||||
if not p.exists():
|
||||
sys.exit(1)
|
||||
phases = json.loads(p.read_text()).get("phases", {})
|
||||
sys.exit(0 if all(phases.get(ph, {}).get("done") for ph in sys.argv[2:]) else 1)
|
||||
PY
|
||||
}
|
||||
|
||||
# Roll-up across all shards, so overall progress is one file rather than 131.
|
||||
OVERALL=${OVERALL:-$(dirname "$DATA_ROOT_BASE")/progress.json}
|
||||
overall_set() { # overall_set <done> <skipped> <failed> <total> <clips> <elapsed> <failed-list>
|
||||
python - "$OVERALL" "$@" "$DATA_ROOT_BASE" "$SHARDS" <<'PY'
|
||||
import datetime, json, os, sys
|
||||
from pathlib import Path
|
||||
p = Path(sys.argv[1])
|
||||
done, skipped, failed, total, clips, elapsed = (int(x) for x in sys.argv[2:8])
|
||||
failed_list, base, spec = sys.argv[8], sys.argv[9], sys.argv[10]
|
||||
processed = done + skipped
|
||||
d = {
|
||||
"spec": spec, "data_root_base": base,
|
||||
"shards_total": total, "complete": processed, "processed_this_session": done,
|
||||
"skipped_already_done": skipped, "failed": failed,
|
||||
"failed_shards": [int(x) for x in failed_list.split() if x],
|
||||
"npz_this_session": clips,
|
||||
"elapsed_min": round(elapsed / 60, 1),
|
||||
"avg_min_per_shard": round(elapsed / done / 60, 1) if done else None,
|
||||
"eta_hours": round((elapsed / done) * (total - processed - failed) / 3600, 1) if done else None,
|
||||
"updated": datetime.datetime.now().isoformat(timespec="seconds"),
|
||||
}
|
||||
tmp = p.with_suffix(f".json.tmp{os.getpid()}"); tmp.write_text(json.dumps(d, indent=2)); tmp.replace(p)
|
||||
PY
|
||||
}
|
||||
|
||||
echo "[shards] $N_SHARDS shard(s): ${SHARD_LIST[0]}..${SHARD_LIST[-1]} phase=$PHASE parallel=$PARALLEL cleanup=$CLEANUP"
|
||||
echo "[shards] overall progress: $OVERALL per-shard: ${DATA_ROOT_BASE}<N>/progress.json"
|
||||
|
||||
# ---- parallel mode: one shard pipeline per GPU -------------------------------------
|
||||
if [[ "$PARALLEL" == "1" ]]; then
|
||||
IFS=',' read -ra GPU_ARR <<< "$GPUS"
|
||||
NG=${#GPU_ARR[@]}
|
||||
echo "[shards] parallel: $NG concurrent pipelines, one GPU each (${GPUS})"
|
||||
T0=$(date +%s)
|
||||
wpids=()
|
||||
for gi in "${!GPU_ARR[@]}"; do
|
||||
(
|
||||
gpu=${GPU_ARR[$gi]}
|
||||
mine=()
|
||||
for ((j = gi; j < N_SHARDS; j += NG)); do mine+=("${SHARD_LIST[$j]}"); done
|
||||
echo "[gpu$gpu] ${#mine[@]} shard(s): ${mine[*]:0:6}$([[ ${#mine[@]} -gt 6 ]] && echo ' ...')"
|
||||
for s in "${mine[@]}"; do
|
||||
root="$(shard_root "$s")"
|
||||
if shard_complete "$s"; then echo "[gpu$gpu] shard $s already complete"; continue; fi
|
||||
t=$(date +%s)
|
||||
if SHARD="$s" RESUME=1 DATA_ROOT="$root" GPUS="$gpu" PARQUET_GPU="$gpu" \
|
||||
bash "$SHARD_SCRIPT" >> "$root.log" 2>&1; then
|
||||
echo "[gpu$gpu] shard $s OK in $(( ($(date +%s) - t) / 60 ))m"
|
||||
[[ "$CLEANUP" == "1" ]] && shard_complete "$s" && rm -rf "$root/download" "$root/raw_videos"
|
||||
else
|
||||
echo "[gpu$gpu] shard $s FAILED -- see $root.log" >&2
|
||||
fi
|
||||
done
|
||||
) &
|
||||
wpids+=($!)
|
||||
done
|
||||
for p in "${wpids[@]}"; do wait "$p"; done
|
||||
|
||||
# tally from the per-shard progress files (authoritative, survives restarts)
|
||||
done_n=0; fail_n=0; failed_list=()
|
||||
for s in "${SHARD_LIST[@]}"; do
|
||||
if shard_complete "$s"; then done_n=$((done_n + 1)); else fail_n=$((fail_n + 1)); failed_list+=("$s"); fi
|
||||
done
|
||||
npz_n=$(find "$(dirname "$DATA_ROOT_BASE")" -name '*.npz' -path '*/tracks/*' 2>/dev/null | wc -l || echo 0)
|
||||
overall_set "$done_n" 0 "$fail_n" "$N_SHARDS" "$npz_n" "$(( $(date +%s) - T0 ))" "${failed_list[*]:-}"
|
||||
echo "=============================================================="
|
||||
echo "[shards] done in $(( ($(date +%s) - T0) / 60 ))m: $done_n complete, $fail_n incomplete"
|
||||
echo "[shards] per-shard console logs: ${DATA_ROOT_BASE}<N>.log"
|
||||
(( fail_n > 0 )) && { echo "[shards] incomplete: ${failed_list[*]}"; echo "[shards] re-run to retry"; exit 1; }
|
||||
exit 0
|
||||
fi
|
||||
T0=$(date +%s)
|
||||
n_done=0 n_skip=0 n_fail=0 clips_total=0
|
||||
FAILED=()
|
||||
|
||||
for s in "${SHARD_LIST[@]}"; do
|
||||
root="$(shard_root "$s")"
|
||||
if shard_complete "$s"; then
|
||||
n_skip=$((n_skip + 1))
|
||||
echo "[shards] shard $s: already complete, skipping"
|
||||
continue
|
||||
fi
|
||||
|
||||
echo "=============================================================="
|
||||
echo "[shards] shard $s ($((n_done + n_skip + n_fail + 1))/$N_SHARDS) elapsed $(( ($(date +%s) - T0) / 60 ))m"
|
||||
t=$(date +%s)
|
||||
rc=0
|
||||
SHARD="$s" RESUME=1 DATA_ROOT="$root" bash "$SHARD_SCRIPT" || rc=$?
|
||||
# count npz regardless of outcome: a shard can produce tracks and still fail a later phase
|
||||
c=$(ls "$root"/tracks/*.npz 2>/dev/null | wc -l || echo 0)
|
||||
clips_total=$((clips_total + c))
|
||||
if [[ $rc -eq 0 ]]; then
|
||||
n_done=$((n_done + 1))
|
||||
echo "[shards] shard $s OK in $(( ($(date +%s) - t) / 60 ))m ($c npz)"
|
||||
if [[ "$CLEANUP" == "1" ]] && shard_complete "$s"; then
|
||||
# only after progress.json confirms completion -- never delete inputs for a
|
||||
# shard that would need re-processing
|
||||
rm -rf "$root/download" "$root/raw_videos"
|
||||
echo "[shards] shard $s: removed download/ and raw_videos/ (tracks kept)"
|
||||
fi
|
||||
else
|
||||
n_fail=$((n_fail + 1)); FAILED+=("$s")
|
||||
echo "[shards] shard $s FAILED -- see $root/logs/" >&2
|
||||
[[ "$STOP_ON_FAIL" == "1" ]] && { echo "[shards] stopping (STOP_ON_FAIL=1)" >&2; break; }
|
||||
fi
|
||||
|
||||
overall_set "$n_done" "$n_skip" "$n_fail" "$N_SHARDS" "$clips_total" \
|
||||
"$(( $(date +%s) - T0 ))" "${FAILED[*]:-}"
|
||||
|
||||
# rolling ETA from shards actually processed this session
|
||||
if (( n_done > 0 )); then
|
||||
avg=$(( ($(date +%s) - T0) / n_done ))
|
||||
left=$(( N_SHARDS - n_done - n_skip - n_fail ))
|
||||
echo "[shards] avg $(( avg / 60 ))m/shard, $left left, ETA $(( avg * left / 3600 ))h"
|
||||
fi
|
||||
done
|
||||
|
||||
overall_set "$n_done" "$n_skip" "$n_fail" "$N_SHARDS" "$clips_total" \
|
||||
"$(( $(date +%s) - T0 ))" "${FAILED[*]:-}"
|
||||
echo "=============================================================="
|
||||
echo "[shards] done in $(( ($(date +%s) - T0) / 60 ))m: $n_done processed, $n_skip skipped, $n_fail failed"
|
||||
echo "[shards] npz produced this session: $clips_total"
|
||||
if (( n_fail > 0 )); then
|
||||
echo "[shards] failed shards: ${FAILED[*]}"
|
||||
echo "[shards] re-run the same command to retry them (completed shards are skipped)"
|
||||
exit 1
|
||||
fi
|
||||
Executable
+88
@@ -0,0 +1,88 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node data-parallel i2v_track PREPROCESS over OpenVid-1M.
|
||||
# v1_preprocess.py asserts WORLD_SIZE==1 (one GPU per process), so scale-out is
|
||||
# N nodes x 4 GPU x 1 proc/GPU = 4N independent single-GPU processes, each over a
|
||||
# SHARD of the clips, writing parquet into <combined>/shard_<idx>/. The trainer
|
||||
# reads <combined> as ONE dataset (get_parquet_files_and_length os.walks recursively).
|
||||
#
|
||||
# Usage:
|
||||
# NODES=16 bash data_pipeline/run_preprocess_track_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
DATA_DIR=${DATA_DIR:-$WORK/openvid_1m}
|
||||
MODEL=${MODEL:-$WORK/models/trackwan_1.3b_i2v_d64_nobias_init}
|
||||
CLIPS_DIR=${CLIPS_DIR:-$DATA_DIR/clips}
|
||||
MANIFEST=${MANIFEST:-$DATA_DIR/videos2caption.json}
|
||||
COMBINED=${COMBINED:-$DATA_DIR/combined_parquet_dataset} # trainer points here
|
||||
SHARDS_DIR=${SHARDS_DIR:-$DATA_DIR/preprocess_shards}
|
||||
NODES=${NODES:-14}; GPUS=${GPUS:-4}
|
||||
PARTITION=${PARTITION:-all}
|
||||
CONC=${CONC:-$(( NODES * GPUS ))} # max concurrent shards (= total GPUs)
|
||||
# MANY small shards (not one-per-GPU) so a node glitch only loses its in-flight
|
||||
# ~CONC/NODES shards, and .done markers make a resubmit skip finished work.
|
||||
NUM_SHARDS=${NUM_SHARDS:-1024} # ~253 clips/shard (fine-grained .done resume)
|
||||
# Training-target geometry (must match trainer):
|
||||
MAX_H=${MAX_H:-480}; MAX_W=${MAX_W:-832}; NUM_FRAMES=${NUM_FRAMES:-121}
|
||||
TRAIN_FPS=${TRAIN_FPS:-24}; NUM_LATENT_T=${NUM_LATENT_T:-31}
|
||||
# batch_size MUST be 1: the T5 tokenizer_kwargs pad-free config (configs/models/
|
||||
# encoders/t5.py) can't batch variable-length captions with return_tensors="pt".
|
||||
BATCH=${BATCH:-1}
|
||||
VAE_PREC=${VAE_PREC:-fp32} # bf16 gives NO speedup here (decode-bound, not VAE-compute-bound)
|
||||
NW=${NW:-8} # decode-prefetch workers. SAFE now: patched VideoCaptionMergedDataset.__iter__
|
||||
# to shard by get_worker_info() (else num_workers>0 duplicates). Validated:
|
||||
# 16 clips -> 16 rows, latents byte-identical to nw=0. Overlaps decode w/ GPU encode.
|
||||
TIME=${TIME:-24:00:00}
|
||||
CPUS_PER_TASK=$(( 128 / GPUS )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
mkdir -p "$WORK/logs" "$COMBINED"
|
||||
|
||||
# 1) Build/refresh shard manifests (idempotent; cheap).
|
||||
source "$WORK/FastVideo/.venv/bin/activate"
|
||||
python "$WORK/FastVideo/data_pipeline/split_manifest_shards.py" \
|
||||
--manifest "$MANIFEST" --clips-dir "$CLIPS_DIR" \
|
||||
--out-dir "$SHARDS_DIR" --num-shards "$NUM_SHARDS"
|
||||
|
||||
echo "array of $NUM_SHARDS shards, up to $CONC concurrent (~$((NODES)) nodes x $GPUS GPU); out=$COMBINED"
|
||||
|
||||
# 2) Launch. This cluster gives a GPU job the WHOLE node (all 4 GPU + 144 CPU), so we
|
||||
# use -N nodes x ntasks-per-node=4 (= 4N=56 workers) rather than a job array (which
|
||||
# would get 1 whole node per task = only 14 concurrent). Each worker loops over its
|
||||
# slice shards[PROCID::NW], running a FRESH python per SMALL (~260-clip) shard:
|
||||
# the fresh process bounds the pipeline's ~0.75 GB/clip RSS growth (VAE feature cache)
|
||||
# -> 4 workers x ~214 GB peak = ~856 GB < 979 GB/node. Per-shard .done makes resubmit
|
||||
# skip finished work; --no-requeue avoids a destructive auto-restart on node glitches.
|
||||
WORKERS=$(( NODES * GPUS ))
|
||||
# CORDON-PROOF: the k8s operator cordons worker pods mid-run (node-pool scale-down);
|
||||
# a gang -N$NODES job dies entirely if ANY one node is cordoned. So submit a JOB ARRAY
|
||||
# of single-node tasks (--array=0..NODES-1, each -N1): a cordon kills+requeues only ONE
|
||||
# array task; the other nodes keep running. Global worker idx = ARRAY_TASK_ID*GPUS+LOCALID
|
||||
# over WORKERS=NODES*GPUS; per-shard .done makes every (re)start skip finished work.
|
||||
WORKERS=$(( NODES * GPUS ))
|
||||
sbatch --array=0-$(( NODES - 1 )) -N1 --gres=gpu:$GPUS --ntasks-per-node=$GPUS --exclusive --mem=0 \
|
||||
-t "$TIME" --requeue -p "$PARTITION" -J preprocess_track \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/preprocess_track_%A_%a.out" -e "$WORK/logs/preprocess_track_%A_%a.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo bash -lc '
|
||||
set -uo pipefail
|
||||
source .venv/bin/activate
|
||||
export HOME=$WORK TRITON_CACHE_DIR=$WORK/.triton XDG_CACHE_HOME=$WORK/.cache \
|
||||
HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch TOKENIZERS_PARALLELISM=false
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS ))
|
||||
export WORLD_SIZE=1 RANK=0 LOCAL_RANK=0 MASTER_ADDR=127.0.0.1
|
||||
export MASTER_PORT=\$(( 29500 + SLURM_LOCALID ))
|
||||
W=\$(( SLURM_ARRAY_TASK_ID * $GPUS + SLURM_LOCALID )); NW=$WORKERS
|
||||
echo \"[worker \$W/\$NW arr=\$SLURM_ARRAY_TASK_ID] host=\$(hostname) CVD=\$CUDA_VISIBLE_DEVICES\"
|
||||
for IDX in \$(seq \$W \$NW $(( NUM_SHARDS - 1 ))); do
|
||||
SHARD=\$(printf shard_%05d \$IDX); SDIR=$SHARDS_DIR/\$SHARD; ODIR=$COMBINED/\$SHARD
|
||||
[ -f \"\$ODIR/.done\" ] && continue
|
||||
rm -rf \"\$ODIR\"
|
||||
if python fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL --preprocess_task i2v_track \
|
||||
--data_merge_path \$SDIR/merge.txt --output_dir \$ODIR \
|
||||
--max_height $MAX_H --max_width $MAX_W --num_frames $NUM_FRAMES \
|
||||
--train_fps $TRAIN_FPS --num_latent_t $NUM_LATENT_T --vae_precision $VAE_PREC \
|
||||
--preprocess_video_batch_size $BATCH --dataloader_num_workers $NW; then
|
||||
touch \"\$ODIR/.done\"; echo \"[worker \$W] \$SHARD done\"
|
||||
else
|
||||
echo \"[worker \$W] \$SHARD FAILED (rc=\$?), skipping\"
|
||||
fi
|
||||
done
|
||||
echo \"[worker \$W] all shards done\"'"
|
||||
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node data-parallel frame-0 segmentation (SAM2.1-b+) over OpenVid-1M.
|
||||
# Bare node (torch-only ultralytics needs just the driver + Lustre venv, no container).
|
||||
# Each Slurm task = one SAM worker; PROCS_PER_GPU tasks share each GPU.
|
||||
#
|
||||
# Usage:
|
||||
# DATA_DIR=/mnt/lustre/vlm-s4duan/openvid_1m NODES=6 PROCS_PER_GPU=2 \
|
||||
# bash data_pipeline/run_seg_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
DATA_DIR=${DATA_DIR:-$WORK/openvid_1m}
|
||||
NODES=${NODES:-1}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
MODEL=${MODEL:-sam2.1_b.pt}; VIDEOS_SUBDIR=${VIDEOS_SUBDIR:-clips}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
TIME=${TIME:-12:00:00}
|
||||
mkdir -p "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $((NODES*TASKS_PER_NODE)) workers, model=$MODEL"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t "$TIME" -J seg_sam_dp \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/seg_sam_dp_%j.out" -e "$WORK/logs/seg_sam_dp_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && \
|
||||
export TORCH_HOME=$WORK/.torch HF_HOME=$WORK/.hf MPLCONFIGDIR=$WORK/.mpl \
|
||||
YOLO_CONFIG_DIR=$WORK/.ultralytics TOKENIZERS_PARALLELISM=false \
|
||||
PYTHONPATH=$WORK/FastVideo:$WORK/FastVideo/data_pipeline && \
|
||||
python data_pipeline/segment_tracks_mp.py \
|
||||
--data-dir $DATA_DIR --videos-subdir $VIDEOS_SUBDIR --model $MODEL --gpus-per-node $GPUS'"
|
||||
Executable
+25
@@ -0,0 +1,25 @@
|
||||
#!/usr/bin/env bash
|
||||
# SAM 2.1-b+ segmentation of frame-0 for our synth mp4s. Adds object_ids +
|
||||
# track_weights to each tracks .npz. Idempotent (skips npz already labeled).
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${DATA_DIR:?}"
|
||||
NODES=${NODES:-4}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
NUM_SHARDS=$(( NODES * TASKS_PER_NODE ))
|
||||
|
||||
mkdir -p "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $NUM_SHARDS workers"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 4:00:00 -J segment_synth \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/segment_synth_%j.out" -e "$WORK/logs/segment_synth_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && \
|
||||
export TORCH_HOME=$WORK/.torch HF_HOME=$WORK/.hf HOME=$WORK MPLCONFIGDIR=$WORK/.mpl \
|
||||
YOLO_CONFIG_DIR=$WORK/.ultralytics TOKENIZERS_PARALLELISM=false && \
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS )) && \
|
||||
python data_pipeline/segment_tracks.py \
|
||||
--data-dir $DATA_DIR \
|
||||
--num-shards $NUM_SHARDS --shard \$SLURM_PROCID'"
|
||||
Executable
+26
@@ -0,0 +1,26 @@
|
||||
#!/bin/bash
|
||||
# Run segment_tracks.py in parallel across 4 GPUs.
|
||||
# Usage: bash data_pipeline/run_segment_tracks.sh [extra args]
|
||||
# Example: bash data_pipeline/run_segment_tracks.sh --limit 20
|
||||
|
||||
DATA_DIR=/home/hal-kevin/data/motion-stream-test
|
||||
WORLD_SIZE=4
|
||||
LOG_FILE=data_pipeline/segment_tracks.log
|
||||
|
||||
> $LOG_FILE # truncate on each run
|
||||
|
||||
echo "[seg] launching $WORLD_SIZE workers... logging to $LOG_FILE"
|
||||
|
||||
for RANK in $(seq 0 $((WORLD_SIZE - 1))); do
|
||||
CUDA_VISIBLE_DEVICES=$RANK python -u data_pipeline/segment_tracks.py \
|
||||
--data-dir $DATA_DIR \
|
||||
--videos-subdir roundtrip_videos \
|
||||
--conf 0.75 --iou 0.9 --imgsz 1024 \
|
||||
--vis-override-every 3 --viz \
|
||||
--rank $RANK --world-size $WORLD_SIZE \
|
||||
"$@" \
|
||||
>> $LOG_FILE 2>&1 &
|
||||
done
|
||||
|
||||
wait
|
||||
echo "[seg] all done. log at $LOG_FILE"
|
||||
@@ -0,0 +1,62 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node data-parallel Wan2.2-T2V-A14B synthetic video generation.
|
||||
#
|
||||
# One VideoGenerator per GPU (num_gpus=1, NO SP/TP -- data-parallel is faster at scale).
|
||||
# Node layout: NODES x 4 GPU x 1 proc/GPU = 4*NODES independent single-GPU workers.
|
||||
# Worker W in [0, 4*NODES) owns prompts[W :: 4*NODES] (stride slice of a fixed shuffle).
|
||||
#
|
||||
# CORDON-PROOF: k8s cordons worker pods mid-run; a gang -N$NODES job dies if ANY node is
|
||||
# cordoned. So we submit a JOB ARRAY of single-node tasks (--array=0..NODES-1, each -N1):
|
||||
# a cordon kills+requeues only ONE array task; the others keep running. Per-video mp4
|
||||
# existence makes every (re)start skip finished work -> fully resumable.
|
||||
#
|
||||
# Usage:
|
||||
# NODES=12 MAX_VIDEOS=100000 bash data_pipeline/run_synth_gen_slurm.sh
|
||||
# NODES=1 MAX_VIDEOS=8 SMOKE=1 bash data_pipeline/run_synth_gen_slurm.sh # smoke test
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
FV=$WORK/FastVideo
|
||||
PROMPTS=${PROMPTS:-$FV/examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt}
|
||||
OUTDIR=${OUTDIR:-$WORK/data/wan22_synth_720p}
|
||||
NODES=${NODES:-12}
|
||||
GPUS=${GPUS:-4}
|
||||
WORKERS=$(( NODES * GPUS ))
|
||||
MAX_VIDEOS=${MAX_VIDEOS:-100000}
|
||||
HEIGHT=${HEIGHT:-720}; WIDTH=${WIDTH:-1280}
|
||||
NUM_FRAMES=${NUM_FRAMES:-121}; DROP=${DROP:-8}; FPS=${FPS:-16}
|
||||
STEPS=${STEPS:-40}; GS=${GS:-4.0}; GS2=${GS2:-3.0}
|
||||
SEED_BASE=${SEED_BASE:-1024}; SHUFFLE_SEED=${SHUFFLE_SEED:-1234}
|
||||
PARTITION=${PARTITION:-all}
|
||||
TIME=${TIME:-24:00:00}
|
||||
SMOKE=${SMOKE:-0}
|
||||
CPUS_PER_TASK=$(( 128 / GPUS )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
mkdir -p "$WORK/logs" "$OUTDIR"
|
||||
|
||||
[ -f "$PROMPTS" ] || { echo "PROMPTS not found: $PROMPTS"; exit 1; }
|
||||
echo "NODES=$NODES GPUS=$GPUS -> $WORKERS workers | target=$MAX_VIDEOS videos | out=$OUTDIR"
|
||||
echo "gen $((NUM_FRAMES+DROP))f -> keep ${NUM_FRAMES} (drop ${DROP}) @ ${WIDTH}x${HEIGHT} ${FPS}fps, ${STEPS} steps, CFG ${GS}/${GS2}"
|
||||
|
||||
sbatch --array=0-$(( NODES - 1 )) -N1 --gres=gpu:$GPUS --ntasks-per-node=$GPUS --exclusive --mem=0 \
|
||||
--cpus-per-task=$CPUS_PER_TASK -t "$TIME" --requeue -p "$PARTITION" -J synth_gen \
|
||||
--chdir="$FV" -o "$WORK/logs/synth_gen_%A_%a.out" -e "$WORK/logs/synth_gen_%A_%a.out" \
|
||||
--wrap "srun --chdir=$FV bash -lc '
|
||||
set -uo pipefail
|
||||
source .venv/bin/activate
|
||||
export HOME=$WORK HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch MPLCONFIGDIR=$WORK/.mpl \
|
||||
XDG_CACHE_HOME=$WORK/.cache TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0 \
|
||||
PYTHONPATH=$FV
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS ))
|
||||
export TRITON_CACHE_DIR=/tmp/triton_synth_\${SLURM_LOCALID}
|
||||
export WORLD_SIZE=1 RANK=0 LOCAL_RANK=0 MASTER_ADDR=127.0.0.1 MASTER_PORT=\$(( 29700 + SLURM_LOCALID ))
|
||||
mkdir -p \$TRITON_CACHE_DIR
|
||||
W=\$(( SLURM_ARRAY_TASK_ID * $GPUS + SLURM_LOCALID ))
|
||||
echo \"[launch] worker \$W/$WORKERS host=\$(hostname) CVD=\$CUDA_VISIBLE_DEVICES\"
|
||||
python data_pipeline/gen_synth_worker.py \
|
||||
--prompts $PROMPTS --output-dir $OUTDIR \
|
||||
--worker-id \$W --num-workers $WORKERS --max-videos $MAX_VIDEOS \
|
||||
--height $HEIGHT --width $WIDTH --num-frames $NUM_FRAMES --drop $DROP --fps $FPS \
|
||||
--steps $STEPS --guidance-scale $GS --guidance-scale-2 $GS2 \
|
||||
--seed-base $SEED_BASE --shuffle-seed $SHUFFLE_SEED
|
||||
'"
|
||||
echo "submitted. monitor: squeue -u \$USER -n synth_gen ; tail -f $WORK/logs/synth_gen_*.out"
|
||||
echo "progress: python data_pipeline/merge_synth_manifests.py --output-dir $OUTDIR"
|
||||
Executable
+119
@@ -0,0 +1,119 @@
|
||||
#!/bin/bash
|
||||
# Synth toy post-processing: tracks (+SAM) and parquet on the SINGLE synth video dir (not sharded).
|
||||
# Mirrors run_openvid_shard.sh's PHASE interface for the 720p/24fps synth set produced by
|
||||
# gen_synth_i2v_worker.py, reusing the same extract_tracks.py --segment and v1_preprocess calls.
|
||||
#
|
||||
# PHASE=tracks COMPILE=0 bash data_pipeline/run_synth_pipeline.sh # CoTracker + SAM -> tracks/*.npz
|
||||
# PHASE=parquet bash data_pipeline/run_synth_pipeline.sh # captions + points_path + Stage 5 -> parquet
|
||||
# PHASE=both bash data_pipeline/run_synth_pipeline.sh
|
||||
#
|
||||
# Resumable: extract_tracks skips clips whose npz exists (unless FORCE_TRACKS=1); v1_preprocess
|
||||
# skips clip ids already written. Run tracks first (across all GPUs), then parquet (single-GPU).
|
||||
set -uo pipefail
|
||||
cd "$(dirname "$0")/.."
|
||||
|
||||
DATA_ROOT=${DATA_ROOT:-/home/hal-kevin/data/motion-stream-synth}
|
||||
HEIGHT=${HEIGHT:-720}; WIDTH=${WIDTH:-1280}
|
||||
NUM_FRAMES=${NUM_FRAMES:-121}; FPS=${FPS:-24} # MUST match the generated videos (24fps)
|
||||
GRID=${GRID:-50}
|
||||
GPUS=${GPUS:-0,1,2,3}
|
||||
AMP=${AMP:-1}; COMPILE=${COMPILE:-0}; FORCE_TRACKS=${FORCE_TRACKS:-0}
|
||||
# v1_preprocess only uses the VAE / T5 / CLIP from MODEL_PATH (same Wan2.1 VAE across sizes), so
|
||||
# any Wan2.1 model dir with those encoders works; override to a lighter one if you have it.
|
||||
MODEL_PATH=${MODEL_PATH:-/home/hal-kevin/models/Wan2.1-I2V-14B-720P-Diffusers}
|
||||
PARQUET_WORKERS=${PARQUET_WORKERS:-2}
|
||||
VID_SUBDIR=videos
|
||||
PHASE=${PHASE:-both}
|
||||
case "$PHASE" in
|
||||
tracks) TRACKS=1; PARQUET=0 ;;
|
||||
parquet) TRACKS=0; PARQUET=1 ;;
|
||||
both) TRACKS=1; PARQUET=1 ;;
|
||||
*) echo "[synth] ERROR: PHASE must be tracks|parquet|both (got '$PHASE')" >&2; exit 1 ;;
|
||||
esac
|
||||
|
||||
IFS=',' read -ra GPU_ARR <<< "$GPUS"; WORLD_SIZE=${#GPU_ARR[@]}
|
||||
LOG_DIR="$DATA_ROOT/logs"; mkdir -p "$LOG_DIR" "$DATA_ROOT/tracks"
|
||||
N_VID=$(ls "$DATA_ROOT/$VID_SUBDIR"/*.mp4 2>/dev/null | wc -l || true)
|
||||
[ "$N_VID" -gt 0 ] || { echo "[synth] no videos under $DATA_ROOT/$VID_SUBDIR -- run gen_synth_i2v_worker first" >&2; exit 1; }
|
||||
echo "[synth] DATA_ROOT=$DATA_ROOT videos=$N_VID PHASE=$PHASE ${WIDTH}x${HEIGHT}@${FPS}fps x${NUM_FRAMES}f GPUS=$GPUS"
|
||||
|
||||
# --- tracks (+ SAM object_ids / track_weights, fused) across all GPUs -------------------
|
||||
if [[ "$TRACKS" == "1" ]]; then
|
||||
echo "[synth] extracting tracks (+segment) across $WORLD_SIZE GPU(s) ..."
|
||||
SPEED=(); [[ "$AMP" == "1" ]] && SPEED+=(--amp); [[ "$COMPILE" == "1" ]] && SPEED+=(--compile)
|
||||
FORCE_ARGS=(); [[ "$FORCE_TRACKS" == "1" ]] && FORCE_ARGS=(--force)
|
||||
if [[ "$COMPILE" == "1" ]]; then
|
||||
export TORCHINDUCTOR_CACHE_DIR=${TORCHINDUCTOR_CACHE_DIR:-$HOME/.cache/torchinductor}
|
||||
export TRITON_CACHE_DIR=${TRITON_CACHE_DIR:-$HOME/.cache/triton}
|
||||
fi
|
||||
pids=()
|
||||
for i in "${!GPU_ARR[@]}"; do
|
||||
CUDA_VISIBLE_DEVICES=${GPU_ARR[$i]} python -u data_pipeline/extract_tracks.py \
|
||||
--data-dir "$DATA_ROOT" --videos-subdir "$VID_SUBDIR" --out-subdir tracks \
|
||||
--grid-size "$GRID" --device cuda \
|
||||
--detect-entries --sam-conf 0.75 --sam-iou 0.9 --sam-imgsz 1024 \
|
||||
--entry-sample-every 2 --entry-min-area 0.001 --entry-new-area 0.5 \
|
||||
--segment --vis-override-every 3 \
|
||||
"${SPEED[@]}" "${FORCE_ARGS[@]}" \
|
||||
--rank "$i" --world-size "$WORLD_SIZE" \
|
||||
>> "$LOG_DIR/tracks.log" 2>&1 &
|
||||
pids+=($!)
|
||||
done
|
||||
tfail=0; for p in "${pids[@]}"; do wait "$p" || tfail=$((tfail + 1)); done
|
||||
N_NPZ=$(ls "$DATA_ROOT"/tracks/*.npz 2>/dev/null | wc -l || true)
|
||||
echo "[synth] tracks: $N_NPZ/$N_VID npz (failed workers: $tfail) -- log: $LOG_DIR/tracks.log"
|
||||
[[ "$tfail" -gt 0 ]] && echo "[synth] WARNING: some track workers failed; inspect the log and re-run (resumable)"
|
||||
fi
|
||||
|
||||
# --- parquet (Stage 5): captions from gen + points_path patch + v1_preprocess -----------
|
||||
if [[ "$PARQUET" == "1" ]]; then
|
||||
N_NPZ=$(ls "$DATA_ROOT"/tracks/*.npz 2>/dev/null | wc -l || true)
|
||||
if [[ "$N_NPZ" -lt "$N_VID" ]]; then
|
||||
echo "[synth] SKIPPING parquet: tracks incomplete ($N_NPZ/$N_VID) -- run PHASE=tracks first" >&2
|
||||
exit 1
|
||||
fi
|
||||
# 1. compile the gen manifest shards -> videos2caption.json (real captions) + merge.txt
|
||||
python data_pipeline/merge_synth_manifests.py --output-dir "$DATA_ROOT"
|
||||
# 2. patch points_path (tracks/<stem>.npz) into each entry, preserving the captions
|
||||
python - "$DATA_ROOT" <<'PY'
|
||||
import json, sys
|
||||
from pathlib import Path
|
||||
root = Path(sys.argv[1]); j = root / "videos2caption.json"; td = root / "tracks"
|
||||
items = json.loads(j.read_text())
|
||||
n = 0
|
||||
for it in items:
|
||||
# preprocess validation computes num_frames = ceil(fps*duration); the gen manifest omits
|
||||
# 'duration', which makes it 0 and rejects every clip -- derive it from num_frames/fps.
|
||||
if not it.get("duration") and it.get("fps"):
|
||||
it["duration"] = it["num_frames"] / float(it["fps"])
|
||||
npz = td / (Path(it["path"]).stem + ".npz")
|
||||
if npz.exists():
|
||||
it["points_path"] = str(npz.resolve()); n += 1
|
||||
j.write_text(json.dumps(items, indent=2))
|
||||
print(f"[synth] manifest: {len(items)} entries, {n} with points_path (duration patched)")
|
||||
PY
|
||||
# 3. Stage 5. --train_fps MUST equal the generated fps (24) or FrameSamplingStage resamples
|
||||
# and the latents stop aligning with the tracks. num_latent_t = (num_frames-1)/4 + 1.
|
||||
NLT=$(( (NUM_FRAMES - 1) / 4 + 1 ))
|
||||
PQ_OUT="${PARQUET_ROOT:-$DATA_ROOT/preprocessed_i2v_track}"
|
||||
echo "[synth] Stage 5: v1_preprocess (${HEIGHT}x${WIDTH}, ${NUM_FRAMES}f, num_latent_t=$NLT, train_fps=$FPS) -> $PQ_OUT"
|
||||
CUDA_VISIBLE_DEVICES="${PARQUET_GPU:-${GPU_ARR[0]}}" \
|
||||
torchrun --nproc_per_node=1 -m fastvideo.pipelines.preprocess.v1_preprocess \
|
||||
--model_path "$MODEL_PATH" \
|
||||
--data_merge_path "$DATA_ROOT/merge.txt" \
|
||||
--output_dir "$PQ_OUT" \
|
||||
--preprocess_task i2v_track \
|
||||
--num_frames "$NUM_FRAMES" \
|
||||
--num_latent_t "$NLT" \
|
||||
--train_fps "$FPS" \
|
||||
--max_height "$HEIGHT" \
|
||||
--max_width "$WIDTH" \
|
||||
--preprocess_video_batch_size 1 \
|
||||
--dataloader_num_workers "$PARQUET_WORKERS" \
|
||||
--samples_per_file "${PARQUET_SAMPLES:-64}" \
|
||||
--flush_frequency "${PARQUET_FLUSH:-8}" \
|
||||
2>&1 | tee -a "$LOG_DIR/parquet.log"
|
||||
N_PQ=$(find "$PQ_OUT" -name '*.parquet' 2>/dev/null | wc -l || true)
|
||||
echo "[synth] parquet: $N_PQ file(s) -> $PQ_OUT/combined_parquet_dataset"
|
||||
echo "[synth] point the overfit config data_path at: $PQ_OUT/combined_parquet_dataset"
|
||||
fi
|
||||
@@ -0,0 +1,32 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node / multi-GPU data-parallel CoTracker on this Slinky GB200 cluster.
|
||||
# Each Slurm task = one CoTracker worker (B=1); PROCS_PER_GPU tasks share each GPU.
|
||||
# Runs inside the enroot `fvbuild` container (needed for the venv's CUDA runtime).
|
||||
#
|
||||
# Usage:
|
||||
# VIDEO_LIST=/mnt/lustre/vlm-s4duan/openvid/videos.txt \
|
||||
# OUT_DIR=/mnt/lustre/vlm-s4duan/openvid/tracks \
|
||||
# NODES=4 PROCS_PER_GPU=2 FPS=24 NUM_FRAMES=121 GRID=50 \
|
||||
# bash data_pipeline/run_tracks_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${VIDEO_LIST:?set VIDEO_LIST}"; : "${OUT_DIR:?set OUT_DIR}"
|
||||
NODES=${NODES:-1}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
FPS=${FPS:-24}; NUM_FRAMES=${NUM_FRAMES:-121}; GRID=${GRID:-50}
|
||||
CLIPS_ARG=""; [ -n "${CLIPS_DIR:-}" ] && CLIPS_ARG="--clips-dir $CLIPS_DIR"
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
|
||||
mkdir -p "$OUT_DIR" "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $((NODES*TASKS_PER_NODE)) workers"
|
||||
|
||||
# Bare-node: torch (self-contained cu128 wheels) + CoTracker need only the driver +
|
||||
# the Lustre venv — NO container (avoids a per-node image pull).
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 24:00:00 -J cotracker_dp \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/cotracker_dp_%j.out" -e "$WORK/logs/cotracker_dp_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && export TORCH_HOME=$WORK/.torch TOKENIZERS_PARALLELISM=false && \
|
||||
python data_pipeline/extract_tracks_mp.py \
|
||||
--video-list $VIDEO_LIST --out-dir $OUT_DIR $CLIPS_ARG \
|
||||
--gpus-per-node $GPUS --fps $FPS --num-frames $NUM_FRAMES --grid-size $GRID'"
|
||||
Executable
+25
@@ -0,0 +1,25 @@
|
||||
#!/usr/bin/env bash
|
||||
# CoTracker extraction for our synth mp4s. Same pattern as run_tracks_slurm.sh but
|
||||
# with the aspect/lowres filters disabled (synth videos are exactly 720x1280) and
|
||||
# H/W set to our training resolution 480x832.
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${VIDEO_LIST:?}"; : "${OUT_DIR:?}"
|
||||
NODES=${NODES:-4}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
FPS=${FPS:-24}; NUM_FRAMES=${NUM_FRAMES:-121}; GRID=${GRID:-50}
|
||||
HEIGHT=${HEIGHT:-480}; WIDTH=${WIDTH:-832}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
|
||||
mkdir -p "$OUT_DIR" "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $((NODES*TASKS_PER_NODE)) workers at ${HEIGHT}x${WIDTH}"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 6:00:00 -J cotracker_synth \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/cotracker_synth_%j.out" -e "$WORK/logs/cotracker_synth_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && export TORCH_HOME=$WORK/.torch TOKENIZERS_PARALLELISM=false && \
|
||||
python data_pipeline/extract_tracks_mp.py \
|
||||
--video-list $VIDEO_LIST --out-dir $OUT_DIR \
|
||||
--gpus-per-node $GPUS --fps $FPS --num-frames $NUM_FRAMES --grid-size $GRID \
|
||||
--height $HEIGHT --width $WIDTH --min-height 0 --aspect-tol 10.0'"
|
||||
@@ -0,0 +1,213 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Visualize the training-time track sampler on a preprocessed dataset.
|
||||
|
||||
For each clip it shows, straight from the stored ``tracks.npz`` (tracks + visibility +
|
||||
``object_ids`` + ``track_weights``):
|
||||
(1) the low-rank informativeness HEATMAP over the 50x50 grid on frame 0,
|
||||
(2) the points the SAMPLER KEEPS (exact port of WanTrack ``_augment_tracks``:
|
||||
>=1 pt per SAM object + low-rank-weighted draw (1-uniform_frac) + uniform draw),
|
||||
(3) an overlay video of ONLY the kept tracks over the source video, colored by low-rank
|
||||
weight (blue = generic, red = informative), so you can see if good traces are sampled.
|
||||
|
||||
CPU only (no GPU / no model needed).
|
||||
|
||||
.venv/bin/python data_pipeline/sampling_viz.py render --data-dir <dataset root> --out <dir> --limit 8
|
||||
.venv/bin/python data_pipeline/sampling_viz.py serve --viz-dir <dir> --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from track_informativeness import _draw_overlay, _read_frames, _norm # noqa: E402
|
||||
|
||||
|
||||
def sample_kept(tracks: np.ndarray,
|
||||
vis: np.ndarray,
|
||||
oid: np.ndarray | None,
|
||||
tw: np.ndarray | None,
|
||||
k: int,
|
||||
uniform_frac: float = 0.3,
|
||||
diversity: float = 1.0,
|
||||
seed: int = 0) -> np.ndarray:
|
||||
"""Exact numpy port of WanTrack._augment_tracks sampling -> kept bool [N]."""
|
||||
rng = np.random.default_rng(seed)
|
||||
N = tracks.shape[1]
|
||||
valid = (vis[0] > 0.5).astype(np.float64) # only frame-0-visible are valid queries
|
||||
if tw is not None and tw.size == N:
|
||||
w = (0.05 + diversity * tw.astype(np.float64)) * valid
|
||||
else:
|
||||
disp = np.sqrt(((tracks - tracks[0:1])**2).sum(-1))
|
||||
motion = (disp * (vis > 0.5)).max(0)
|
||||
w = (1.0 + diversity * (motion / (motion.mean() + 1e-6))) * valid
|
||||
if w.sum() <= 0:
|
||||
w = valid.copy()
|
||||
if w.sum() <= 0:
|
||||
w = np.ones(N)
|
||||
keep = np.zeros(N, bool)
|
||||
# (a) object coverage: one weighted pick per present segment
|
||||
if oid is not None:
|
||||
for o in np.unique(oid):
|
||||
if int(o) < 0:
|
||||
continue
|
||||
idx = np.where(oid == o)[0]
|
||||
ww = w[idx]
|
||||
ww = np.ones_like(ww) if ww.sum() <= 0 else ww
|
||||
keep[rng.choice(idx, p=ww / ww.sum())] = True
|
||||
# (b) low-rank-weighted draw for (1 - uniform_frac) of the remaining budget
|
||||
rem = max(0, k - int(keep.sum()))
|
||||
n_weighted = rem - int(round(rem * uniform_frac))
|
||||
pool = np.where((~keep) & (w > 0))[0]
|
||||
nw = min(n_weighted, pool.size)
|
||||
if nw > 0:
|
||||
pw = w[pool] / w[pool].sum()
|
||||
keep[rng.choice(pool, size=nw, replace=False, p=pw)] = True
|
||||
# (c) uniform draw to fill up to k (also backfills weighted shortfall)
|
||||
pool2 = np.where((~keep) & (valid > 0))[0]
|
||||
nu = min(max(0, k - int(keep.sum())), pool2.size)
|
||||
if nu > 0:
|
||||
keep[rng.choice(pool2, size=nu, replace=False)] = True
|
||||
return keep
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import imageio.v2 as imageio
|
||||
turbo = matplotlib.colormaps["turbo"]
|
||||
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
entries = []
|
||||
for i, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
vpath = str(data / "videos" / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = _read_frames(vpath)
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]]
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]]
|
||||
oid = d["object_ids"].astype(np.int64) if "object_ids" in d else None
|
||||
tw = d["track_weights"].astype(np.float32) if "track_weights" in d else None
|
||||
N = tracks.shape[1]
|
||||
k = int(args.k)
|
||||
keep = sample_kept(tracks, vis, oid, tw, k, args.uniform_frac, args.diversity, args.seed)
|
||||
twn = _norm(tw) if tw is not None else np.zeros(N)
|
||||
x, y = tracks[0, :, 0], tracks[0, :, 1]
|
||||
|
||||
# ---- figure: (1) lowrank heatmap (2) kept vs dropped
|
||||
fig, axes = plt.subplots(1, 2, figsize=(17, 5.6))
|
||||
axes[0].imshow(frames[0])
|
||||
sc = axes[0].scatter(x, y, c=twn, cmap="turbo", s=12, vmin=0, vmax=1)
|
||||
axes[0].set_title(f"low-rank informativeness heatmap (all {N} pts)", fontsize=11)
|
||||
axes[0].axis("off")
|
||||
fig.colorbar(sc, ax=axes[0], fraction=0.03)
|
||||
axes[1].imshow(frames[0])
|
||||
axes[1].scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.5) # dropped
|
||||
axes[1].scatter(x[keep],
|
||||
y[keep],
|
||||
c=twn[keep],
|
||||
cmap="turbo",
|
||||
s=16,
|
||||
vmin=0,
|
||||
vmax=1,
|
||||
edgecolors="white",
|
||||
linewidths=0.3) # kept, colored by weight
|
||||
n_obj = len(np.unique(oid[oid >= 0])) if oid is not None else 0
|
||||
n_cov = len(np.unique(oid[keep][oid[keep] >= 0])) if oid is not None else 0
|
||||
axes[1].set_title(f"kept by sampler: {int(keep.sum())}/{N} (objects covered {n_cov}/{n_obj})", fontsize=11)
|
||||
axes[1].axis("off")
|
||||
fig.suptitle(f"{stem} | {(it['cap'][0] if isinstance(it.get('cap'), list) else '')[:80]}", fontsize=10)
|
||||
fig.tight_layout()
|
||||
heat = f"{stem}_sampling.png"
|
||||
fig.savefig(str(out / heat), dpi=80, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
|
||||
# ---- overlay video: ONLY kept tracks, colored by low-rank weight
|
||||
kidx = np.where(keep)[0]
|
||||
# subsample for legibility if many kept
|
||||
if kidx.size > 700:
|
||||
kidx = kidx[np.linspace(0, kidx.size - 1, 700).round().astype(int)]
|
||||
pcols = (np.array([turbo(v)[:3] for v in twn[kidx]]) * 255).astype(np.uint8)
|
||||
ov = _draw_overlay(frames, tracks[:, kidx].copy(), vis[:, kidx], pcols, 14, 2, 0.5)
|
||||
vid = f"{stem}_kept.mp4"
|
||||
imageio.mimsave(str(out / vid), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
|
||||
stats = {
|
||||
"id": stem,
|
||||
"n_points": int(N),
|
||||
"k_sampled": k,
|
||||
"kept": int(keep.sum()),
|
||||
"objects_covered": f"{n_cov}/{n_obj}",
|
||||
"mean_weight_kept": round(float(twn[keep].mean()), 3) if keep.any() else 0.0,
|
||||
"mean_weight_dropped": round(float(twn[~keep].mean()), 3) if (~keep).any() else 0.0,
|
||||
"frac_kept_informative(>0.5)": round(float((twn[keep] > 0.5).mean()), 3) if keep.any() else 0.0,
|
||||
}
|
||||
entries.append({"id": stem, "heat": heat, "video": vid, "stats": stats})
|
||||
print(
|
||||
f"[samp] [{i}/{len(items)}] {stem}: kept {int(keep.sum())}/{N}, cov {n_cov}/{n_obj}, "
|
||||
f"w_kept={stats['mean_weight_kept']} vs w_drop={stats['mean_weight_dropped']}",
|
||||
flush=True)
|
||||
|
||||
(out / "manifest.json").write_text(json.dumps(entries, indent=2))
|
||||
print(f"[samp] rendered {len(entries)} clips -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
entries = json.loads((viz / "manifest.json").read_text())
|
||||
gal = [(str(viz / e["heat"]), e["id"]) for e in entries]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = entries[evt.index]
|
||||
return str(viz / e["heat"]), str(viz / e["video"]), json.dumps(e["stats"], indent=2)
|
||||
|
||||
with gr.Blocks(title="Track sampling viz") as demo:
|
||||
gr.Markdown("### Training-time track sampler: low-rank heatmap + kept points + kept-track overlay\n"
|
||||
"Left panel: low-rank informativeness heatmap (bright = informative). Right panel: the points "
|
||||
"the sampler **keeps** (colored by weight, gray = dropped) — should cover every object and "
|
||||
"favor bright/informative points while keeping some uniform. Click a clip -> the overlay video "
|
||||
"of ONLY the kept tracks (blue = generic, red = informative). Good sampling = kept traces follow "
|
||||
"the moving hands/objects. `stats.mean_weight_kept` should exceed `mean_weight_dropped`.")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gal, columns=2, height=620, label="clips (click)")
|
||||
with gr.Column():
|
||||
heat = gr.Image(label="heatmap (left) + kept points (right)")
|
||||
vid = gr.Video(label="kept tracks over video (colored by low-rank weight)")
|
||||
meta = gr.Code(label="stats", language="json")
|
||||
g.select(show, None, [heat, vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--k", type=int, default=1500, help="tracks to keep (training samples U[1000,2500])")
|
||||
r.add_argument("--uniform-frac", type=float, default=0.3)
|
||||
r.add_argument("--diversity", type=float, default=1.0)
|
||||
r.add_argument("--seed", type=int, default=0)
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7889)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,377 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Compare segmentation-model *variants* (not just FastSAM configs) for frame-0
|
||||
object segmentation + sparse track sampling, on real openvid clips.
|
||||
|
||||
``render`` (GPU) runs each model in no-prompt "everything" mode on frame-0 of each
|
||||
clip, labels the CoTracker grid points by object (smallest containing mask wins),
|
||||
runs the 1-per-object + extras sparse sampler, and writes per (model, clip):
|
||||
- <..>_masks.png frame-0 with colored SAM masks
|
||||
- <..>_sparse.png frame-0 with kept sparse tracks (colored by object) vs dropped (gray)
|
||||
- <..>_tracks.mp4 kept sparse tracks over the whole clip
|
||||
plus a per-entry .npz (oid/tw/vis0/xy0) so the dashboard can re-sample live.
|
||||
|
||||
``serve`` (CPU, login node) is a gradio dashboard: pick a clip, see every model's
|
||||
masks / sparse-sampling side by side, tweak num_sampled live, click to inspect.
|
||||
|
||||
aarch64-safe: PyAV decode, headless cv2, no decord.
|
||||
|
||||
# render on a GPU node (piggyback job 365)
|
||||
srun --overlap --jobid=365 --ntasks=1 --chdir=/mnt/lustre/vlm-s4duan/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && export CUDA_VISIBLE_DEVICES=1 \
|
||||
YOLO_CONFIG_DIR=/mnt/lustre/vlm-s4duan/.ultralytics \
|
||||
TORCH_HOME=/mnt/lustre/vlm-s4duan/.torch HF_HOME=/mnt/lustre/vlm-s4duan/.hf \
|
||||
PYTHONPATH=$PWD:$PWD/data_pipeline MPLCONFIGDIR=/mnt/lustre/vlm-s4duan/.mpl && \
|
||||
python -u data_pipeline/seg_compare.py render --limit 8'
|
||||
|
||||
# serve on the login node, then ssh -L 7862:localhost:7862 <login>
|
||||
.venv/bin/python data_pipeline/seg_compare.py serve
|
||||
"""
|
||||
# NOTE: no ``from __future__ import annotations`` — gradio needs the real
|
||||
# gr.SelectData annotation on the click handler.
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from segment_tracks import object_ids_for_points, lowrank_track_weights # noqa: E402
|
||||
from sparse_sampling_dashboard import sparse_sample # noqa: E402
|
||||
from track_informativeness import _draw_overlay # noqa: E402
|
||||
|
||||
DATA_DEFAULT = "/mnt/lustre/vlm-s4duan/openvid_1m"
|
||||
WEIGHTS_DIR = "/mnt/lustre/vlm-s4duan/models/seg"
|
||||
OUT_DEFAULT = "/mnt/lustre/vlm-s4duan/seg_compare_out"
|
||||
|
||||
# The model zoo. loader in {FastSAM, SAM}. All run no-prompt everything-mode.
|
||||
# min_area_frac drops mask specks (< frac of frame); nms_contain drops a mask if
|
||||
# > that fraction of it is already covered by a larger kept mask (collapses
|
||||
# part-level fragments into their object). max_masks caps to N largest (0=all).
|
||||
DEFAULT_MODELS = [
|
||||
{"name": "fastsam-s", "loader": "FastSAM", "weight": "FastSAM-s.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "fastsam-x", "loader": "FastSAM", "weight": "FastSAM-x.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "mobilesam", "loader": "SAM", "weight": "mobile_sam.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam-b", "loader": "SAM", "weight": "sam_b.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam2-b", "loader": "SAM", "weight": "sam2_b.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam2.1-b+", "loader": "SAM", "weight": "sam2.1_b.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam2.1-l", "loader": "SAM", "weight": "sam2.1_l.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
]
|
||||
|
||||
# sparse-sampler defaults baked into the rendered PNG/mp4 (dashboard can re-sample the PNG live)
|
||||
SPARSE_DEFAULT = {"num_sampled": 20, "mode": "weighted", "seed": 0}
|
||||
|
||||
|
||||
def read_frames(path: str) -> np.ndarray:
|
||||
"""Full clip -> [T,H,W,3] uint8 via PyAV (aarch64-safe)."""
|
||||
import av
|
||||
c = av.open(path)
|
||||
frames = [f.to_ndarray(format="rgb24") for f in c.decode(video=0)]
|
||||
c.close()
|
||||
return np.stack(frames)
|
||||
|
||||
|
||||
def colors(n: int) -> np.ndarray:
|
||||
import colorsys
|
||||
return np.array([[int(255 * v) for v in colorsys.hsv_to_rgb((i * 0.61803) % 1.0, 0.65, 1.0)]
|
||||
for i in range(max(1, n))], np.uint8)
|
||||
|
||||
|
||||
def extract_masks(res, H: int, W: int) -> np.ndarray:
|
||||
if not res or res[0].masks is None:
|
||||
return np.zeros((0, H, W), bool)
|
||||
m = res[0].masks.data.cpu().numpy().astype(bool) # [M,h,w]
|
||||
if m.shape[0] and m.shape[1:] != (H, W): # nearest resize w/o cv2
|
||||
import torch
|
||||
t = torch.from_numpy(m.astype(np.uint8))[None].float()
|
||||
t = torch.nn.functional.interpolate(t, size=(H, W), mode="nearest")[0]
|
||||
m = t.numpy().astype(bool)
|
||||
return m
|
||||
|
||||
|
||||
def filter_masks(masks: np.ndarray, min_area_frac: float, nms_contain: float, max_masks: int) -> np.ndarray:
|
||||
if masks.shape[0] == 0:
|
||||
return masks
|
||||
H, W = masks.shape[1], masks.shape[2]
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
if min_area_frac > 0:
|
||||
keep = (areas / float(H * W)) >= min_area_frac
|
||||
masks, areas = masks[keep], areas[keep]
|
||||
if masks.shape[0] and nms_contain > 0: # drop fragments mostly inside a larger kept mask
|
||||
order = np.argsort(-areas) # large -> small
|
||||
kept_idx, covered = [], np.zeros((H, W), bool)
|
||||
for i in order:
|
||||
m = masks[i]
|
||||
a = areas[i]
|
||||
if a > 0 and (m & covered).sum() / a > nms_contain:
|
||||
continue
|
||||
kept_idx.append(i)
|
||||
covered |= m
|
||||
masks, areas = masks[kept_idx], areas[kept_idx]
|
||||
if max_masks and masks.shape[0] > max_masks:
|
||||
masks = masks[np.argsort(-areas)[:max_masks]]
|
||||
return masks
|
||||
|
||||
|
||||
def mask_panel(frame0: np.ndarray, masks: np.ndarray, oid: np.ndarray, xy0: np.ndarray) -> "Image.Image":
|
||||
from PIL import Image, ImageDraw
|
||||
mcols = colors(len(masks) + 1)
|
||||
base = frame0.astype(np.float32)
|
||||
for mi, m in enumerate(masks):
|
||||
base[m] = 0.55 * base[m] + 0.45 * mcols[mi % len(mcols)][None].astype(np.float32)
|
||||
img = Image.fromarray(base.clip(0, 255).astype(np.uint8))
|
||||
draw = ImageDraw.Draw(img)
|
||||
for o in sorted({int(v) for v in np.unique(oid) if int(v) >= 0}): # 1 dot per object (its centroid pt)
|
||||
idx = np.where(oid == o)[0]
|
||||
cx, cy = float(xy0[idx, 0].mean()), float(xy0[idx, 1].mean())
|
||||
draw.ellipse([cx - 4, cy - 4, cx + 4, cy + 4], fill=(255, 255, 255), outline=(0, 0, 0))
|
||||
return img
|
||||
|
||||
|
||||
def sparse_panel(frame0: np.ndarray, xy0: np.ndarray, oid: np.ndarray, keep: np.ndarray, title: str,
|
||||
out_path: str) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
cmap = plt.get_cmap("tab20")
|
||||
fig, ax = plt.subplots(figsize=(9, 5.2))
|
||||
ax.imshow(frame0)
|
||||
x, y = xy0[:, 0], xy0[:, 1]
|
||||
ax.scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.30)
|
||||
for o in np.unique(oid[keep]):
|
||||
m = keep & (oid == o)
|
||||
col = "white" if int(o) < 0 else cmap(int(o) % 20)
|
||||
ax.scatter(x[m], y[m], c=[col], s=80, edgecolors="black", linewidths=1.1)
|
||||
ax.set_title(title, fontsize=10)
|
||||
ax.axis("off")
|
||||
fig.savefig(out_path, dpi=90, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
import torch
|
||||
import imageio.v2 as imageio
|
||||
from ultralytics import FastSAM, SAM
|
||||
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
clips_dir = data / ("clips" if (data / "clips").exists() else "videos")
|
||||
models_cfg = json.loads(Path(args.models_json).read_text()) if args.models_json else DEFAULT_MODELS
|
||||
if args.models: # subset by name
|
||||
want = set(args.models.split(","))
|
||||
models_cfg = [m for m in models_cfg if m["name"] in want]
|
||||
loaders = {"FastSAM": FastSAM, "SAM": SAM}
|
||||
|
||||
# load every model once (reused across all clips)
|
||||
print(f"[seg] loading {len(models_cfg)} models ...", flush=True)
|
||||
loaded = {}
|
||||
for mc in models_cfg:
|
||||
wp = Path(WEIGHTS_DIR) / mc["weight"]
|
||||
loaded[mc["name"]] = loaders[mc["loader"]](str(wp) if wp.exists() else mc["weight"])
|
||||
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
print(f"[seg] {len(items)} clips x {len(models_cfg)} models", flush=True)
|
||||
entries, clip_ids = [], []
|
||||
for k, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
clip_ids.append(stem)
|
||||
vpath = str(clips_dir / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = read_frames(vpath)
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]] # [T,N,2] px
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]] # [T,N]
|
||||
xy0, vis0 = tracks[0], vis[0]
|
||||
tw = lowrank_track_weights(tracks) # [N] informativeness in [0,1]
|
||||
# frame-0 png (shared, for dashboard live re-sampling background)
|
||||
from PIL import Image
|
||||
f0png = f"{stem}__frame0.png"
|
||||
Image.fromarray(frames[0]).save(str(out / f0png))
|
||||
|
||||
for mc in models_cfg:
|
||||
t0 = time.time()
|
||||
res = loaded[mc["name"]](frames[0], device=args.device, retina_masks=True,
|
||||
imgsz=mc["imgsz"], conf=mc["conf"], iou=mc["iou"], verbose=False)
|
||||
masks = filter_masks(extract_masks(res, H, W), mc["min_area_frac"], mc["nms_contain"], mc["max_masks"])
|
||||
oid = object_ids_for_points(masks, xy0, H, W)
|
||||
n_obj = int(len({int(v) for v in np.unique(oid) if int(v) >= 0}))
|
||||
keep = sparse_sample(oid, tw, vis0, SPARSE_DEFAULT["num_sampled"], SPARSE_DEFAULT["mode"],
|
||||
SPARSE_DEFAULT["seed"])
|
||||
dt = time.time() - t0
|
||||
pfx = f"{mc['name']}__{stem}"
|
||||
mask_panel(frames[0], masks, oid, xy0).save(str(out / f"{pfx}_masks.png"))
|
||||
sparse_panel(frames[0], xy0, oid, keep,
|
||||
f"{stem} | {mc['name']}: {int(keep.sum())} kept ({n_obj} objs +"
|
||||
f"{SPARSE_DEFAULT['num_sampled']} extra)", str(out / f"{pfx}_sparse.png"))
|
||||
trk_name = ""
|
||||
if not args.no_video:
|
||||
import matplotlib.pyplot as plt
|
||||
cmap = plt.get_cmap("tab20")
|
||||
kidx = np.where(keep)[0]
|
||||
cols = np.zeros((kidx.size, 3), np.uint8)
|
||||
for o in np.unique(oid[kidx]):
|
||||
mm = oid[kidx] == o
|
||||
rgba = (1., 1., 1., 1.) if int(o) < 0 else cmap(int(o) % 20)
|
||||
cols[mm] = (np.array(rgba[:3]) * 255).astype(np.uint8)
|
||||
ov = _draw_overlay(frames, tracks[:, kidx].copy(), vis[:, kidx], cols, 14, 3, 0.5)
|
||||
trk_name = f"{pfx}_tracks.mp4"
|
||||
imageio.mimsave(str(out / trk_name), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
# per-entry npz so the dashboard can re-sample live (no GPU)
|
||||
np.savez(str(out / f"{pfx}.npz"), oid=oid.astype(np.int64), tw=tw, vis0=vis0, xy0=xy0)
|
||||
entries.append({
|
||||
"model": mc["name"], "clip": stem,
|
||||
"caption": (it["cap"][0] if isinstance(it.get("cap"), list) else str(it.get("cap", "")))[:140],
|
||||
"n_masks": int(masks.shape[0]), "n_objects": n_obj,
|
||||
"n_labeled": int((oid >= 0).sum()), "n_points": int(oid.shape[0]),
|
||||
"n_kept": int(keep.sum()), "sec": round(dt, 2),
|
||||
"masks_png": f"{pfx}_masks.png", "sparse_png": f"{pfx}_sparse.png",
|
||||
"tracks_mp4": trk_name, "frame0_png": f0png, "npz": f"{pfx}.npz",
|
||||
})
|
||||
print(f"[seg] [{k}/{len(items)}] {stem} [{mc['name']}]: {masks.shape[0]} masks, "
|
||||
f"{n_obj} objs, {(oid >= 0).sum()}/{oid.shape[0]} pts labeled, {dt:.1f}s", flush=True)
|
||||
del frames
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
manifest = {"models": models_cfg, "clips": clip_ids, "sparse_default": SPARSE_DEFAULT, "entries": entries}
|
||||
(out / "manifest.json").write_text(json.dumps(manifest, indent=2))
|
||||
print(f"[seg] wrote {len(entries)} entries -> {out}/manifest.json", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
manifest = json.loads((viz / "manifest.json").read_text())
|
||||
entries = manifest["entries"]
|
||||
model_names = [m["name"] for m in manifest["models"]]
|
||||
clip_ids = manifest["clips"]
|
||||
by_clip = {c: [e for e in entries if e["clip"] == c] for c in clip_ids}
|
||||
caption_of = {c: (by_clip[c][0]["caption"] if by_clip.get(c) else "") for c in clip_ids}
|
||||
npz_cache = {}
|
||||
|
||||
def _npz(e):
|
||||
p = e["npz"]
|
||||
if p not in npz_cache:
|
||||
npz_cache[p] = {k: v for k, v in np.load(str(viz / p)).items()}
|
||||
return npz_cache[p]
|
||||
|
||||
def masks_gallery(clip):
|
||||
es = sorted(by_clip.get(clip, []), key=lambda e: model_names.index(e["model"]))
|
||||
return [(str(viz / e["masks_png"]), f'{e["model"]} • {e["n_masks"]}m / {e["n_objects"]}o') for e in es]
|
||||
|
||||
def resample_png(e, num_sampled, mode, seed):
|
||||
z = _npz(e)
|
||||
oid, tw, vis0, xy0 = z["oid"], z["tw"], z["vis0"], z["xy0"]
|
||||
keep = sparse_sample(oid, tw, vis0, int(num_sampled), mode, int(seed))
|
||||
frame0 = plt.imread(str(viz / e["frame0_png"]))
|
||||
cmap = plt.get_cmap("tab20")
|
||||
fig, ax = plt.subplots(figsize=(9, 5.2))
|
||||
ax.imshow(frame0)
|
||||
x, y = xy0[:, 0], xy0[:, 1]
|
||||
ax.scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.30)
|
||||
for o in np.unique(oid[keep]):
|
||||
m = keep & (oid == o)
|
||||
col = "white" if int(o) < 0 else cmap(int(o) % 20)
|
||||
ax.scatter(x[m], y[m], c=[col], s=80, edgecolors="black", linewidths=1.1)
|
||||
n_obj = int(len({int(v) for v in np.unique(oid) if int(v) >= 0}))
|
||||
ax.set_title(f'{e["clip"]} | {e["model"]}: {int(keep.sum())} kept ({n_obj} objs +{num_sampled}, {mode})',
|
||||
fontsize=10)
|
||||
ax.axis("off")
|
||||
import numpy as _np
|
||||
fig.canvas.draw()
|
||||
buf = _np.asarray(fig.canvas.buffer_rgba())[..., :3].copy()
|
||||
plt.close(fig)
|
||||
return buf
|
||||
|
||||
def sparse_gallery(clip, num_sampled, mode, seed):
|
||||
es = sorted(by_clip.get(clip, []), key=lambda e: model_names.index(e["model"]))
|
||||
return [(resample_png(e, num_sampled, mode, seed), f'{e["model"]} • {e["n_kept"]}→') for e in es]
|
||||
|
||||
with gr.Blocks(title="Segmentation model comparison") as demo:
|
||||
gr.Markdown(
|
||||
"## Segmentation model comparison — frame-0 masks + sparse track sampling\n"
|
||||
"Pick a **clip**; every model is shown side by side. **Masks** tab = raw everything-mode "
|
||||
"segmentation (tile label `<#masks>m / <#objects-with-points>o`). **Sparse sampling** tab = "
|
||||
"1 track per object + N extras (tweak live). Fewer, cleaner object masks = better for our "
|
||||
"1-per-object recipe. Click a tile to inspect it big + its track-overlay video.")
|
||||
clip_dd = gr.Dropdown(clip_ids, value=clip_ids[0], label=f"clip (1 of {len(clip_ids)})")
|
||||
clip_cap = gr.Markdown(f"*{caption_of[clip_ids[0]]}*")
|
||||
with gr.Tab("Masks (everything-mode)"):
|
||||
mg = gr.Gallery(value=masks_gallery(clip_ids[0]), columns=3, height=680,
|
||||
label="SAM masks per model (click a tile)")
|
||||
with gr.Tab("Sparse sampling"):
|
||||
with gr.Row():
|
||||
num_sl = gr.Slider(0, 150, value=manifest["sparse_default"]["num_sampled"], step=1,
|
||||
label="num_sampled (extras beyond 1-per-object)")
|
||||
mode_rd = gr.Radio(["weighted", "random"], value=manifest["sparse_default"]["mode"], label="extras mode")
|
||||
seed_sl = gr.Slider(0, 50, value=manifest["sparse_default"]["seed"], step=1, label="seed")
|
||||
sg = gr.Gallery(value=sparse_gallery(clip_ids[0], manifest["sparse_default"]["num_sampled"],
|
||||
manifest["sparse_default"]["mode"],
|
||||
manifest["sparse_default"]["seed"]),
|
||||
columns=3, height=640, label="sparse-sampled tracks per model (click a tile)")
|
||||
gr.Markdown("### Inspect")
|
||||
with gr.Row():
|
||||
big = gr.Image(label="selected panel", height=460)
|
||||
vid = gr.Video(label="kept-tracks overlay over the clip", height=460)
|
||||
meta = gr.Code(label="metrics", language="json")
|
||||
|
||||
def on_clip(clip, num_sampled, mode, seed):
|
||||
return masks_gallery(clip), sparse_gallery(clip, num_sampled, mode, seed), f"*{caption_of.get(clip, '')}*"
|
||||
|
||||
def on_sparse_ctrl(clip, num_sampled, mode, seed):
|
||||
return sparse_gallery(clip, num_sampled, mode, seed)
|
||||
|
||||
def pick(clip, evt: gr.SelectData):
|
||||
es = sorted(by_clip.get(clip, []), key=lambda e: model_names.index(e["model"]))
|
||||
e = es[evt.index]
|
||||
v = str(viz / e["tracks_mp4"]) if e.get("tracks_mp4") else None
|
||||
return str(viz / e["masks_png"]), v, json.dumps(e, indent=2)
|
||||
|
||||
clip_dd.change(on_clip, [clip_dd, num_sl, mode_rd, seed_sl], [mg, sg, clip_cap])
|
||||
for c in (num_sl, mode_rd, seed_sl):
|
||||
c.change(on_sparse_ctrl, [clip_dd, num_sl, mode_rd, seed_sl], sg)
|
||||
mg.select(pick, clip_dd, [big, vid, meta])
|
||||
sg.select(pick, clip_dd, [big, vid, meta])
|
||||
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", default=DATA_DEFAULT)
|
||||
r.add_argument("--out", default=OUT_DEFAULT)
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--device", default="cuda")
|
||||
r.add_argument("--models", default=None, help="comma-separated subset of model names")
|
||||
r.add_argument("--models-json", default=None, help="path to a JSON list of model config dicts")
|
||||
r.add_argument("--no-video", action="store_true", help="skip the track-overlay mp4s (faster)")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", default=OUT_DEFAULT)
|
||||
s.add_argument("--host", default="127.0.0.1")
|
||||
s.add_argument("--port", type=int, default=7862)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,78 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Smoke-test: run several ultralytics segmentation backends in no-prompt
|
||||
'everything' mode on one real openvid frame. Report #masks + latency so we can
|
||||
pick which variants are worth a full render. aarch64-safe (PyAV decode, no cv2)."""
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
DATA = Path("/mnt/lustre/vlm-s4duan/openvid_1m")
|
||||
WEIGHTS = Path("/mnt/lustre/vlm-s4duan/models/seg")
|
||||
WEIGHTS.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# (display name, weight file, loader class)
|
||||
MODELS = [
|
||||
("fastsam-s", "FastSAM-s.pt", "FastSAM"),
|
||||
("fastsam-x", "FastSAM-x.pt", "FastSAM"),
|
||||
("mobilesam", "mobile_sam.pt", "SAM"),
|
||||
("sam-b", "sam_b.pt", "SAM"),
|
||||
("sam2-t", "sam2_t.pt", "SAM"),
|
||||
("sam2-b", "sam2_b.pt", "SAM"),
|
||||
("sam2.1-l", "sam2.1_l.pt", "SAM"),
|
||||
]
|
||||
|
||||
|
||||
def read_frame0(path: str) -> np.ndarray:
|
||||
import av
|
||||
c = av.open(path)
|
||||
for f in c.decode(video=0):
|
||||
return f.to_ndarray(format="rgb24")
|
||||
raise RuntimeError(f"no frames in {path}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
from ultralytics import FastSAM, SAM # noqa: F401
|
||||
import torch
|
||||
|
||||
items = json.loads((DATA / "videos2caption.json").read_text())
|
||||
item = items[0]
|
||||
vpath = str(DATA / "clips" / item["path"]) if (DATA / "clips" / item["path"]).exists() \
|
||||
else str(DATA / "videos" / item["path"])
|
||||
frame0 = read_frame0(vpath)
|
||||
H, W = frame0.shape[:2]
|
||||
print(f"frame: {vpath} {W}x{H}", flush=True)
|
||||
|
||||
loaders = {"FastSAM": FastSAM, "SAM": SAM}
|
||||
rows = []
|
||||
for name, wf, cls in MODELS:
|
||||
wp = WEIGHTS / wf
|
||||
try:
|
||||
t0 = time.time()
|
||||
model = loaders[cls](str(wp) if wp.exists() else wf)
|
||||
# everything-mode: no prompts. imgsz 1024, retina full-res masks.
|
||||
res = model(frame0, device="cuda", retina_masks=True, imgsz=1024,
|
||||
conf=0.4, iou=0.9, verbose=False)
|
||||
dt = time.time() - t0
|
||||
m = res[0].masks
|
||||
nm = 0 if m is None else int(m.data.shape[0])
|
||||
mh, mw = (0, 0) if m is None else tuple(m.data.shape[1:])
|
||||
print(f"[OK] {name:10s} weights={wf:14s} masks={nm:4d} maskres={mw}x{mh} "
|
||||
f"time={dt:6.2f}s", flush=True)
|
||||
rows.append((name, nm, dt))
|
||||
# cache downloaded weight into WEIGHTS dir
|
||||
if not wp.exists() and Path(wf).exists():
|
||||
Path(wf).replace(wp)
|
||||
del model
|
||||
torch.cuda.empty_cache()
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[FAIL] {name:10s} weights={wf:14s} -> {type(e).__name__}: {e}", flush=True)
|
||||
print("\nsummary (name, n_masks, sec):", flush=True)
|
||||
for r in rows:
|
||||
print(f" {r[0]:10s} {r[1]:4d} {r[2]:6.2f}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,363 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 0d: segment frames (FastSAM) and label each CoTracker grid point by object.
|
||||
|
||||
Adds an ``object_ids`` array ([N] int, -1 = background/none) to each tracks ``.npz``,
|
||||
so the trainer's object-coverage sampling can guarantee >=1 track per object.
|
||||
|
||||
Each point is assigned based on its FIRST-VISIBLE frame (from CoTracker visibility):
|
||||
FastSAM runs on each unique first-visible frame, and the point is assigned the smallest
|
||||
mask containing its position at that frame. This correctly handles objects that enter
|
||||
the scene after frame 0.
|
||||
|
||||
Run on a GPU node (FastSAM is light). Idempotent (skips npz that already have object_ids)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=0 PYTHONPATH=$PWD \
|
||||
.venv/bin/python data_pipeline/segment_tracks.py --data-dir <dataset root>
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def read_all_frames(path: str) -> np.ndarray:
|
||||
"""Read all frames from a video, return [T,H,W,3] uint8."""
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001
|
||||
import av
|
||||
c = av.open(path)
|
||||
return np.stack([f.to_ndarray(format="rgb24") for f in c.decode(video=0)])
|
||||
|
||||
|
||||
def _colors(n: int) -> np.ndarray:
|
||||
import colorsys
|
||||
return np.array(
|
||||
[[int(255 * c) for c in colorsys.hsv_to_rgb((i * 0.61803) % 1.0, 0.65, 1.0)]
|
||||
for i in range(max(1, n))],
|
||||
np.uint8,
|
||||
)
|
||||
|
||||
|
||||
def render_viz(frames: np.ndarray, tracks: np.ndarray, vis: np.ndarray,
|
||||
object_ids: np.ndarray, out_path: Path, fps: int = 24) -> None:
|
||||
import imageio.v2 as imageio
|
||||
from fastvideo.train.callbacks.track_validation import _draw_overlay
|
||||
|
||||
objs = sorted(int(o) for o in np.unique(object_ids) if int(o) >= 0)
|
||||
ocols = _colors(len(objs) + 1)
|
||||
N = tracks.shape[1]
|
||||
pcols = np.tile(np.array([[110, 110, 110]], np.uint8), (N, 1))
|
||||
for oi, o in enumerate(objs):
|
||||
pcols[object_ids == o] = ocols[oi % len(ocols)]
|
||||
|
||||
G = int(round(N**0.5))
|
||||
if G * G == N:
|
||||
k = max(1, G // 50)
|
||||
sel = np.arange(N).reshape(G, G)[::k, ::k].reshape(-1)
|
||||
else:
|
||||
st = max(1, N // 1500)
|
||||
sel = np.arange(0, N, st)
|
||||
|
||||
ov = _draw_overlay(frames, tracks[:, sel].copy(), vis[:, sel], pcols[sel], 12, 2, 0.5)
|
||||
out_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
imageio.mimsave(str(out_path), ov, fps=fps, macro_block_size=1)
|
||||
|
||||
|
||||
def render_seg_viz(frame: np.ndarray, masks: np.ndarray, out_path: Path) -> None:
|
||||
import imageio.v2 as imageio
|
||||
img = frame.copy()
|
||||
if masks.shape[0] > 0:
|
||||
colors = _colors(masks.shape[0])
|
||||
for i, mask in enumerate(masks):
|
||||
img[mask] = (img[mask] * 0.5 + colors[i % len(colors)] * 0.5).astype(np.uint8)
|
||||
out_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
imageio.imwrite(str(out_path), img)
|
||||
|
||||
|
||||
def object_ids_for_points(masks: np.ndarray, pts_xy: np.ndarray, H: int, W: int) -> np.ndarray:
|
||||
"""masks [M,H,W] bool, pts_xy [N,2] px -> object id per point (-1 none), smallest mask wins."""
|
||||
N = pts_xy.shape[0]
|
||||
oid = np.full(N, -1, np.int64)
|
||||
if masks.size == 0:
|
||||
return oid
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1) # [M]
|
||||
order = np.argsort(areas) # smallest first -> assign, larger won't overwrite
|
||||
xi = np.clip(pts_xy[:, 0].round().astype(int), 0, W - 1)
|
||||
yi = np.clip(pts_xy[:, 1].round().astype(int), 0, H - 1)
|
||||
assigned = np.zeros(N, bool)
|
||||
for m in order:
|
||||
inside = masks[m][yi, xi] & (~assigned)
|
||||
oid[inside] = int(m)
|
||||
assigned |= inside
|
||||
return oid
|
||||
|
||||
|
||||
def extract_masks(result, H: int, W: int, min_area_frac: float, max_masks: int) -> np.ndarray:
|
||||
"""Pull masks out of a single FastSAM Result, resize if needed, and apply filtering."""
|
||||
masks = np.zeros((0, H, W), bool)
|
||||
if result is not None and result.masks is not None:
|
||||
masks = result.masks.data.cpu().numpy().astype(bool)
|
||||
if masks.shape[0] and masks.shape[1:] != (H, W):
|
||||
import cv2
|
||||
masks = np.stack([
|
||||
cv2.resize(m.astype(np.uint8), (W, H), interpolation=cv2.INTER_NEAREST).astype(bool)
|
||||
for m in masks
|
||||
])
|
||||
if masks.shape[0] and (min_area_frac > 0 or max_masks):
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
if min_area_frac > 0:
|
||||
keep = (areas / float(H * W)) >= min_area_frac
|
||||
masks, areas = masks[keep], areas[keep]
|
||||
if max_masks and masks.shape[0] > max_masks:
|
||||
masks = masks[np.argsort(-areas)[:max_masks]]
|
||||
return masks
|
||||
|
||||
|
||||
def masks_for_frames(model, frames: np.ndarray, frame_ts: list[int], cache: dict[int, np.ndarray],
|
||||
args) -> dict[int, np.ndarray]:
|
||||
"""Segment the requested frames in batched FastSAM forwards, filling/reusing `cache`."""
|
||||
todo = [t for t in frame_ts if t not in cache]
|
||||
for s in range(0, len(todo), max(1, args.sam_batch)):
|
||||
chunk = todo[s:s + max(1, args.sam_batch)]
|
||||
res = model([frames[t] for t in chunk], device=args.device, retina_masks=True,
|
||||
imgsz=args.imgsz, conf=args.conf, iou=args.iou, verbose=False)
|
||||
for t, r in zip(chunk, res):
|
||||
cache[t] = extract_masks(r, frames.shape[1], frames.shape[2],
|
||||
args.min_area_frac, args.max_masks)
|
||||
return cache
|
||||
|
||||
|
||||
def assign_object_ids_multiframe(
|
||||
frame_masks: dict[int, np.ndarray],
|
||||
first_visible: np.ndarray,
|
||||
tracks: np.ndarray,
|
||||
H: int,
|
||||
W: int,
|
||||
) -> np.ndarray:
|
||||
"""Assign globally-unique object IDs using each point's first-visible frame.
|
||||
|
||||
frame_masks: {frame_idx: masks [M,H,W] bool}
|
||||
first_visible: [N] int, -1 = never visible
|
||||
tracks: [T,N,2] px
|
||||
Returns object_ids [N] int64, -1 = background/never visible.
|
||||
"""
|
||||
N = first_visible.shape[0]
|
||||
oid = np.full(N, -1, np.int64)
|
||||
global_offset = 0
|
||||
for frame_t, masks in sorted(frame_masks.items()):
|
||||
point_sel = first_visible == frame_t
|
||||
if point_sel.any() and masks.shape[0] > 0:
|
||||
pts_xy = tracks[frame_t, point_sel]
|
||||
local_oid = object_ids_for_points(masks, pts_xy, H, W)
|
||||
oid[point_sel] = np.where(local_oid >= 0, local_oid + global_offset, -1)
|
||||
global_offset += masks.shape[0]
|
||||
return oid
|
||||
|
||||
|
||||
def segment_tracks_arrays(tracks: np.ndarray, vis: np.ndarray, H: int, W: int, get_masks,
|
||||
vis_override_every: int, verbose: bool = False
|
||||
) -> tuple[np.ndarray, int, np.ndarray, np.ndarray, int, list[int]]:
|
||||
"""Core of Stage 4: object IDs from first-visible frames, vis override, low-rank weights.
|
||||
|
||||
Shared by segment_tracks.py and extract_tracks.py --segment (fused mode) so both paths
|
||||
produce identical results. ``get_masks(frame_ts)`` must return {frame_t: masks [M,H,W] bool}.
|
||||
tracks: [T,N,2] px. vis: [T,N] bool, updated in place by the override sweep.
|
||||
Returns (object_ids, n_objects, vis, track_weights, n_overrides, unique_frames).
|
||||
"""
|
||||
T = tracks.shape[0]
|
||||
# Per-point first-visible frame; -1 for points CoTracker never marks visible
|
||||
ever_visible = vis.any(axis=0) # [N]
|
||||
first_visible = np.where(ever_visible, np.argmax(vis, axis=0), -1) # [N]
|
||||
unique_frames = sorted(set(first_visible[ever_visible].tolist()))
|
||||
|
||||
if verbose:
|
||||
fv_counts = {f: int((first_visible == f).sum()) for f in unique_frames}
|
||||
print(f" [seg] first_visible frames: {fv_counts}", flush=True)
|
||||
|
||||
# FastSAM masks for each unique first-visible frame
|
||||
frame_masks = get_masks(unique_frames)
|
||||
if verbose:
|
||||
for frame_t, masks in frame_masks.items():
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).tolist() if masks.shape[0] else []
|
||||
print(f" [seg] frame {frame_t}: {masks.shape[0]} masks, areas={[int(a) for a in areas]}", flush=True)
|
||||
|
||||
oid = assign_object_ids_multiframe(frame_masks, first_visible, tracks, H, W)
|
||||
n_objects = int(np.unique(oid[oid >= 0]).shape[0]) if (oid >= 0).any() else 0
|
||||
|
||||
if verbose:
|
||||
for frame_t in unique_frames:
|
||||
pt_sel = first_visible == frame_t
|
||||
labeled = int((oid[pt_sel] >= 0).sum())
|
||||
print(f" [seg] frame {frame_t}: {pt_sel.sum()} pts, {labeled} got oid>=0", flush=True)
|
||||
|
||||
# Vis override: run FastSAM every N frames and set vis=True for object points
|
||||
# that fall inside any mask — fixes CoTracker vis=0 on edge-of-frame objects.
|
||||
n_overrides = 0
|
||||
if vis_override_every > 0 and (oid >= 0).any():
|
||||
object_pts = np.where(oid >= 0)[0]
|
||||
override_ts = list(range(0, T, vis_override_every))
|
||||
override_masks = get_masks(override_ts)
|
||||
for frame_t in override_ts:
|
||||
masks = override_masks[frame_t]
|
||||
if masks.shape[0] == 0:
|
||||
continue
|
||||
pts = tracks[frame_t, object_pts] # [K,2]
|
||||
xi = np.clip(pts[:, 0].round().astype(int), 0, W - 1)
|
||||
yi = np.clip(pts[:, 1].round().astype(int), 0, H - 1)
|
||||
in_any_mask = masks[:, yi, xi].any(axis=0) # [K] bool
|
||||
newly_visible = in_any_mask & ~vis[frame_t, object_pts]
|
||||
vis[frame_t, object_pts] |= in_any_mask
|
||||
n_overrides += int(newly_visible.sum())
|
||||
|
||||
# Fill gaps between True frames caused by the sampling interval.
|
||||
# If vis is True at frame T and True again at T+k (k <= override_every),
|
||||
# the frames in between should also be True — the object didn't disappear.
|
||||
V = vis[:, object_pts] # [T,K]
|
||||
t_idx = np.arange(T)[:, None]
|
||||
prev = np.maximum.accumulate(np.where(V, t_idx, -1), axis=0)
|
||||
nxt = np.minimum.accumulate(np.where(V, t_idx, 2 * T)[::-1], axis=0)[::-1]
|
||||
fill = (~V) & (prev >= 0) & (nxt < T) & ((nxt - prev) <= vis_override_every)
|
||||
vis[:, object_pts] = V | fill
|
||||
|
||||
weights = lowrank_track_weights(tracks)
|
||||
return oid.astype(np.int64), n_objects, vis, weights, n_overrides, unique_frames
|
||||
|
||||
|
||||
def lowrank_track_weights(tracks: np.ndarray, rank: int = 3, pct: float = 97.0) -> np.ndarray:
|
||||
"""Per-point sampling weight in [0,1] = percentile-normalized low-rank motion residual.
|
||||
|
||||
Stack per-point displacement [N, 2T], subtract the mean trajectory + top-`rank` shared
|
||||
SVD modes (camera / dominant scene motion), and take the residual norm. A point is heavy
|
||||
iff it moves *uniquely* relative to all other points (independent object motion), not just
|
||||
a lot -- so on egocentric/moving-camera clips the head-motion background is down-weighted.
|
||||
"""
|
||||
T, N, _ = tracks.shape
|
||||
D = (tracks - tracks[0:1]).transpose(1, 0, 2).reshape(N, 2 * T).astype(np.float64)
|
||||
Dc = D - D.mean(0, keepdims=True)
|
||||
U, S, Vt = np.linalg.svd(Dc, full_matrices=False)
|
||||
r = int(min(rank, S.shape[0]))
|
||||
resid = np.sqrt(((Dc - (U[:, :r] * S[:r]) @ Vt[:r])**2).sum(1))
|
||||
lo, hi = float(resid.min()), float(np.percentile(resid, pct))
|
||||
w = np.zeros_like(resid) if hi <= lo else np.clip((resid - lo) / (hi - lo), 0.0, 1.0)
|
||||
return w.astype(np.float32)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--data-dir", type=Path, required=True)
|
||||
p.add_argument("--videos-subdir", type=str, default="videos")
|
||||
p.add_argument("--tracks-subdir", type=str, default="tracks")
|
||||
p.add_argument("--manifest", type=str, default="videos2caption.json")
|
||||
p.add_argument("--model", type=str, default="sam2.1_b.pt",
|
||||
help="ultralytics weight; FastSAM-*.pt use the FastSAM loader, everything else "
|
||||
"(sam2.1_b.pt, sam2.1_l.pt, sam_b.pt, mobile_sam.pt, ...) use the SAM loader")
|
||||
p.add_argument("--weights-dir", type=str, default="/mnt/lustre/vlm-s4duan/models/seg",
|
||||
help="dir holding cached weights; falls back to auto-download if missing")
|
||||
p.add_argument("--device", type=str, default="cuda")
|
||||
p.add_argument("--imgsz", type=int, default=1024)
|
||||
p.add_argument("--conf", type=float, default=0.4)
|
||||
p.add_argument("--iou", type=float, default=0.9)
|
||||
p.add_argument("--sam-batch", type=int, default=16,
|
||||
help="Frames per batched FastSAM forward.")
|
||||
p.add_argument("--min-area-frac", type=float, default=0.0,
|
||||
help="drop masks smaller than this fraction of the frame")
|
||||
p.add_argument("--max-masks", type=int, default=0, help="keep only the N largest masks (0 = keep all)")
|
||||
p.add_argument("--limit", type=int, default=None)
|
||||
p.add_argument("--index", type=int, nargs="+", default=None, metavar="IDX",
|
||||
help="Only process videos at these manifest indices (e.g. --index 4 7 12).")
|
||||
p.add_argument("--rank", type=int, default=0, help="GPU rank for sharding (0-indexed).")
|
||||
p.add_argument("--world-size", type=int, default=1, help="Total number of parallel processes.")
|
||||
p.add_argument("--force", action="store_true", help="re-run even if object_ids already present")
|
||||
p.add_argument("--vis-override-every", type=int, default=0,
|
||||
help="Run FastSAM every N frames and set vis=True for object points inside masks. "
|
||||
"Fixes CoTracker vis=0 on edge-of-frame objects. 0 = disabled.")
|
||||
p.add_argument("--viz", action="store_true", help="Render a track-overlay mp4 after each video.")
|
||||
p.add_argument("--viz-dir", type=str, default=None,
|
||||
help="Output directory for viz mp4s (default: <data-dir>/viz).")
|
||||
p.add_argument("--verbose", action="store_true", help="Print per-frame debug info.")
|
||||
args = p.parse_args()
|
||||
if args.viz and args.viz_dir is None:
|
||||
args.viz_dir = str(args.data_dir / "viz")
|
||||
|
||||
# FastSAM-*.pt -> FastSAM loader; everything else (SAM/SAM2/SAM2.1/MobileSAM) -> SAM loader.
|
||||
# Both expose the same no-prompt "everything" call used below.
|
||||
from ultralytics import FastSAM, SAM
|
||||
wp = Path(args.weights_dir) / args.model
|
||||
weight = str(wp) if wp.exists() else args.model
|
||||
model = (FastSAM if Path(args.model).name.lower().startswith("fastsam") else SAM)(weight)
|
||||
|
||||
manifest_path = args.data_dir / args.manifest
|
||||
items = json.loads(manifest_path.read_text()) if manifest_path.exists() else []
|
||||
if args.index is not None:
|
||||
items = [items[i] for i in args.index if i < len(items)]
|
||||
if args.limit:
|
||||
items = items[:args.limit]
|
||||
if args.world_size > 1:
|
||||
items = items[args.rank::args.world_size]
|
||||
n_ok = 0
|
||||
for k, item in enumerate(items, 1):
|
||||
vpath = args.data_dir / args.videos_subdir / item["path"]
|
||||
npz_path = Path(item.get("points_path") or (args.data_dir / args.tracks_subdir / f"{vpath.stem}.npz"))
|
||||
if not npz_path.exists():
|
||||
print(f"[seg] [{k}/{len(items)}] {vpath.name}: no npz, skip", flush=True)
|
||||
continue
|
||||
d = dict(np.load(npz_path))
|
||||
if not args.force and "object_ids" in d and "track_weights" in d:
|
||||
n_ok += 1
|
||||
continue
|
||||
|
||||
tracks = d["tracks"].astype(np.float32) # [T,N,2] px
|
||||
vis = d["visibility"].astype(bool) # [T,N]
|
||||
H, W = int(d["height"]), int(d["width"])
|
||||
|
||||
# Decode the video once; frames are reused for segmentation, vis override, and viz.
|
||||
frames = read_all_frames(str(vpath))
|
||||
mask_cache: dict[int, np.ndarray] = {}
|
||||
|
||||
def get_masks(frame_ts: list[int], _frames=frames, _cache=mask_cache) -> dict[int, np.ndarray]:
|
||||
masks_for_frames(model, _frames, frame_ts, _cache, args)
|
||||
return {t: _cache[t] for t in frame_ts}
|
||||
|
||||
oid, n_objects, vis, weights, n_overrides, unique_frames = segment_tracks_arrays(
|
||||
tracks, vis, H, W, get_masks, args.vis_override_every, verbose=args.verbose)
|
||||
|
||||
if args.vis_override_every > 0:
|
||||
d["visibility"] = vis
|
||||
d["object_ids"] = oid
|
||||
d["n_objects"] = np.int64(n_objects)
|
||||
d["track_weights"] = weights
|
||||
tmp = npz_path.with_suffix(".tmp.npz")
|
||||
np.savez(tmp, **d)
|
||||
tmp.replace(npz_path)
|
||||
n_ok += 1
|
||||
if args.viz:
|
||||
stem_dir = Path(args.viz_dir) / vpath.stem
|
||||
stem_dir.mkdir(parents=True, exist_ok=True)
|
||||
render_viz(frames[:tracks.shape[0]], tracks, vis, oid, stem_dir / "tracks.mp4",
|
||||
fps=int(item.get("fps", 24)))
|
||||
# for label, fidx in [("000", 0), ("mid", T_v // 2), ("last", T_v - 1)]:
|
||||
# frame = frames[fidx]
|
||||
# res = model(frame, device=args.device, retina_masks=True,
|
||||
# imgsz=args.imgsz, conf=args.conf, iou=args.iou, verbose=False)
|
||||
# fmasks = extract_masks(res, H, W, args.min_area_frac, args.max_masks)
|
||||
# render_seg_viz(frame, fmasks, stem_dir / f"seg_frame{label}.jpg")
|
||||
print(f" viz -> {stem_dir}/", flush=True)
|
||||
|
||||
cov = int((oid >= 0).sum())
|
||||
override_str = f", {n_overrides} vis overrides" if args.vis_override_every > 0 else ""
|
||||
print(
|
||||
f"[seg] [{k}/{len(items)}] {vpath.name}: {n_objects} objs across {len(unique_frames)} frames, "
|
||||
f"{cov}/{oid.shape[0]} grid pts labeled{override_str}, "
|
||||
f"w[mean={d['track_weights'].mean():.3f} >0.5={(d['track_weights'] > 0.5).mean():.2f}]",
|
||||
flush=True)
|
||||
print(f"[seg] done; {n_ok}/{len(items)} npz have object_ids", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Data-parallel frame-0 segmentation (SAM2.1-b+ everything-mode) for OpenVid-1M.
|
||||
|
||||
Mirrors extract_tracks_mp.py: each Slurm task segments manifest[shard::num_shards]
|
||||
on one pinned GPU and writes object_ids + track_weights into the EXISTING tracks
|
||||
.npz (adds keys, keeps tracks/visibility). Idempotent — skips npz that already have
|
||||
both keys, so it is resumable.
|
||||
|
||||
Sharding + GPU pinning from Slurm env:
|
||||
shard = SLURM_PROCID (0..num_shards-1)
|
||||
num_shards = SLURM_NTASKS
|
||||
gpu = SLURM_LOCALID % gpus_per_node
|
||||
|
||||
aarch64-safe: PyAV frame-0 decode, torch mask resize (no cv2/decord).
|
||||
|
||||
srun --ntasks-per-node=<gpus*procs_per_gpu> ... \
|
||||
python data_pipeline/segment_tracks_mp.py --data-dir <root> --videos-subdir clips
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from segment_tracks import object_ids_for_points, lowrank_track_weights # noqa: E402
|
||||
|
||||
|
||||
def read_frame0(path: str) -> np.ndarray:
|
||||
import av
|
||||
c = av.open(str(path))
|
||||
try:
|
||||
for f in c.decode(video=0):
|
||||
return f.to_ndarray(format="rgb24")
|
||||
finally:
|
||||
c.close()
|
||||
raise RuntimeError(f"no frames in {path}")
|
||||
|
||||
|
||||
def extract_masks(res, H: int, W: int) -> np.ndarray:
|
||||
if not res or res[0].masks is None:
|
||||
return np.zeros((0, H, W), bool)
|
||||
import torch
|
||||
m = res[0].masks.data
|
||||
if m.shape[1:] != (H, W):
|
||||
m = torch.nn.functional.interpolate(m[None].float(), size=(H, W), mode="nearest")[0]
|
||||
return m.cpu().numpy().astype(bool)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--data-dir", required=True)
|
||||
ap.add_argument("--videos-subdir", default="clips", help="frame-0 source; MUST match track coords (720p clips)")
|
||||
ap.add_argument("--tracks-subdir", default="tracks")
|
||||
ap.add_argument("--manifest", default="videos2caption.json")
|
||||
ap.add_argument("--model", default="sam2.1_b.pt")
|
||||
ap.add_argument("--weights-dir", default="/mnt/lustre/vlm-s4duan/models/seg")
|
||||
ap.add_argument("--imgsz", type=int, default=1024)
|
||||
ap.add_argument("--conf", type=float, default=0.4)
|
||||
ap.add_argument("--iou", type=float, default=0.9)
|
||||
ap.add_argument("--min-area-frac", type=float, default=0.0)
|
||||
ap.add_argument("--num-shards", type=int, default=int(os.environ.get("SLURM_NTASKS", "1")))
|
||||
ap.add_argument("--shard", type=int, default=int(os.environ.get("SLURM_PROCID", "0")))
|
||||
ap.add_argument("--gpus-per-node", type=int, default=4)
|
||||
ap.add_argument("--limit", type=int, default=None)
|
||||
a = ap.parse_args()
|
||||
|
||||
import torch
|
||||
local = int(os.environ.get("SLURM_LOCALID", str(a.shard)))
|
||||
gpu = local % a.gpus_per_node
|
||||
torch.cuda.set_device(gpu)
|
||||
device = f"cuda:{gpu}"
|
||||
|
||||
from ultralytics import FastSAM, SAM
|
||||
wp = Path(a.weights_dir) / a.model
|
||||
weight = str(wp) if wp.exists() else a.model
|
||||
model = (FastSAM if Path(a.model).name.lower().startswith("fastsam") else SAM)(weight)
|
||||
|
||||
data = Path(a.data_dir)
|
||||
items = json.loads((data / a.manifest).read_text())
|
||||
if a.limit:
|
||||
items = items[:a.limit]
|
||||
mine = items[a.shard::a.num_shards]
|
||||
print(f"[seg {a.shard}/{a.num_shards}] gpu={gpu} localid={local} items={len(mine)}", flush=True)
|
||||
|
||||
t0 = time.time()
|
||||
done = skip = notrk = err = 0
|
||||
for it in mine:
|
||||
stem = Path(it["path"]).stem
|
||||
npz_path = Path(it.get("points_path") or (data / a.tracks_subdir / f"{stem}.npz"))
|
||||
if not npz_path.exists():
|
||||
notrk += 1
|
||||
continue
|
||||
try:
|
||||
d = dict(np.load(npz_path))
|
||||
if "object_ids" in d and "track_weights" in d:
|
||||
skip += 1
|
||||
continue
|
||||
vpath = data / a.videos_subdir / it["path"]
|
||||
frame0 = read_frame0(str(vpath))
|
||||
H, W = frame0.shape[0], frame0.shape[1]
|
||||
res = model(frame0, device=device, retina_masks=True, imgsz=a.imgsz,
|
||||
conf=a.conf, iou=a.iou, verbose=False)
|
||||
masks = extract_masks(res, H, W)
|
||||
if masks.shape[0] and a.min_area_frac > 0:
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
masks = masks[(areas / float(H * W)) >= a.min_area_frac]
|
||||
tracks = d["tracks"].astype(np.float32)
|
||||
oid = object_ids_for_points(masks, tracks[0], H, W)
|
||||
d["object_ids"] = oid.astype(np.int64)
|
||||
d["n_objects"] = np.int64(masks.shape[0])
|
||||
d["track_weights"] = lowrank_track_weights(tracks)
|
||||
tmp = npz_path.with_suffix(".seg.tmp.npz")
|
||||
np.savez(tmp, **d)
|
||||
tmp.replace(npz_path)
|
||||
done += 1
|
||||
if done % 200 == 0:
|
||||
r = done / (time.time() - t0)
|
||||
print(f"[seg {a.shard}] {done} done ({r:.2f}/s), skip={skip} notrk={notrk} err={err}", flush=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
err += 1
|
||||
print(f"[seg {a.shard}] ERR {stem}: {repr(e)[:120]}", flush=True)
|
||||
dt = time.time() - t0
|
||||
print(f"[seg {a.shard}] DONE done={done} skip={skip} notrk={notrk} err={err} in {dt:.1f}s "
|
||||
f"-> {done / dt if dt > 0 else 0:.3f} clip/s", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,346 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Visualize / sweep the segmentation pipeline: FastSAM masks + chosen per-object
|
||||
points + CoTracker tracks overlaid on the source video.
|
||||
|
||||
``render`` (GPU) sweeps a set of FastSAM configs over the clips and writes one
|
||||
SAM-panel PNG + one track-overlay mp4 per (config, clip). ``serve`` (no GPU) is a
|
||||
gradio gallery with config/clip filters so you can scroll and compare any combo.
|
||||
|
||||
srun ... env CUDA_VISIBLE_DEVICES=0 PYTHONPATH=$PWD .venv/bin/python data_pipeline/segment_viz.py render \
|
||||
--data-dir <dataset root> --out <viz dir> --limit 8
|
||||
.venv/bin/python data_pipeline/segment_viz.py serve --viz-dir <viz dir> --share
|
||||
|
||||
Configs default to DEFAULT_SWEEP (below); override with --configs '<json>' or
|
||||
--configs-json <file>. Each config: {name, conf, iou, imgsz, min_area_frac, max_masks}.
|
||||
``min_area_frac`` drops masks smaller than that fraction of the frame; ``max_masks``
|
||||
keeps only the N largest — both fight FastSAM over-segmentation.
|
||||
"""
|
||||
# NOTE: intentionally no ``from __future__ import annotations`` — gradio needs the
|
||||
# real ``gr.SelectData`` annotation object on the click handler to inject the event.
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
# name, FastSAM conf/iou/imgsz, then two over-segmentation knobs:
|
||||
# min_area_frac: drop masks whose area < this fraction of the frame (0 = keep all)
|
||||
# max_masks: keep only the N largest masks (0 = keep all)
|
||||
DEFAULT_SWEEP = [
|
||||
{
|
||||
"name": "baseline",
|
||||
"conf": 0.4,
|
||||
"iou": 0.9,
|
||||
"imgsz": 1024,
|
||||
"min_area_frac": 0.0,
|
||||
"max_masks": 0
|
||||
},
|
||||
{
|
||||
"name": "conf0.6",
|
||||
"conf": 0.6,
|
||||
"iou": 0.9,
|
||||
"imgsz": 1024,
|
||||
"min_area_frac": 0.0,
|
||||
"max_masks": 0
|
||||
},
|
||||
{
|
||||
"name": "conf0.75",
|
||||
"conf": 0.75,
|
||||
"iou": 0.9,
|
||||
"imgsz": 1024,
|
||||
"min_area_frac": 0.0,
|
||||
"max_masks": 0
|
||||
},
|
||||
{
|
||||
"name": "iou0.6",
|
||||
"conf": 0.4,
|
||||
"iou": 0.6,
|
||||
"imgsz": 1024,
|
||||
"min_area_frac": 0.0,
|
||||
"max_masks": 0
|
||||
},
|
||||
{
|
||||
"name": "areafloor",
|
||||
"conf": 0.4,
|
||||
"iou": 0.9,
|
||||
"imgsz": 1024,
|
||||
"min_area_frac": 0.006,
|
||||
"max_masks": 0
|
||||
},
|
||||
{
|
||||
"name": "clean",
|
||||
"conf": 0.6,
|
||||
"iou": 0.7,
|
||||
"imgsz": 1024,
|
||||
"min_area_frac": 0.004,
|
||||
"max_masks": 25
|
||||
},
|
||||
{
|
||||
"name": "img1536",
|
||||
"conf": 0.5,
|
||||
"iou": 0.8,
|
||||
"imgsz": 1536,
|
||||
"min_area_frac": 0.003,
|
||||
"max_masks": 0
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _colors(n: int) -> np.ndarray:
|
||||
import colorsys
|
||||
return np.array([[int(255 * c) for c in colorsys.hsv_to_rgb((i * 0.61803) % 1.0, 0.65, 1.0)]
|
||||
for i in range(max(1, n))], np.uint8)
|
||||
|
||||
|
||||
def _read_frames(path: str) -> np.ndarray:
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001
|
||||
import av
|
||||
c = av.open(path)
|
||||
return np.stack([f.to_ndarray(format="rgb24") for f in c.decode(video=0)])
|
||||
|
||||
|
||||
def _load_configs(args: argparse.Namespace) -> list[dict]:
|
||||
if args.configs:
|
||||
raw = json.loads(args.configs)
|
||||
elif args.configs_json:
|
||||
raw = json.loads(Path(args.configs_json).read_text())
|
||||
else:
|
||||
raw = DEFAULT_SWEEP
|
||||
cfgs = []
|
||||
for i, c in enumerate(raw):
|
||||
cfgs.append({
|
||||
"name": str(c.get("name", f"cfg{i}")),
|
||||
"conf": float(c.get("conf", 0.4)),
|
||||
"iou": float(c.get("iou", 0.9)),
|
||||
"imgsz": int(c.get("imgsz", 1024)),
|
||||
"min_area_frac": float(c.get("min_area_frac", 0.0)),
|
||||
"max_masks": int(c.get("max_masks", 0)),
|
||||
})
|
||||
return cfgs
|
||||
|
||||
|
||||
def _extract_masks(res, H: int, W: int) -> np.ndarray:
|
||||
if not res or res[0].masks is None:
|
||||
return np.zeros((0, H, W), bool)
|
||||
masks = res[0].masks.data.cpu().numpy().astype(bool) # [M,h,w]
|
||||
if masks.shape[0] and masks.shape[1:] != (H, W):
|
||||
import cv2
|
||||
masks = np.stack(
|
||||
[cv2.resize(m.astype(np.uint8), (W, H), interpolation=cv2.INTER_NEAREST).astype(bool) for m in masks])
|
||||
return masks
|
||||
|
||||
|
||||
def filter_masks(masks: np.ndarray, min_area_frac: float, max_masks: int) -> np.ndarray:
|
||||
"""Drop tiny masks (< min_area_frac of frame) and keep only the N largest."""
|
||||
if masks.shape[0] == 0:
|
||||
return masks
|
||||
H, W = masks.shape[1], masks.shape[2]
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
if min_area_frac > 0:
|
||||
keep = (areas / float(H * W)) >= min_area_frac
|
||||
masks, areas = masks[keep], areas[keep]
|
||||
if max_masks and masks.shape[0] > max_masks:
|
||||
masks = masks[np.argsort(-areas)[:max_masks]]
|
||||
return masks
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
from PIL import Image, ImageDraw
|
||||
from ultralytics import FastSAM
|
||||
from fastvideo.train.callbacks.track_validation import _draw_overlay
|
||||
from segment_tracks import assign_object_ids_multiframe, extract_masks
|
||||
import imageio.v2 as imageio
|
||||
|
||||
configs = _load_configs(args)
|
||||
model = FastSAM(args.model)
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
print(f"[viz] {len(items)} clips x {len(configs)} configs = {len(items) * len(configs)} panels", flush=True)
|
||||
|
||||
entries: list[dict] = []
|
||||
clip_ids: list[str] = []
|
||||
for k, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
clip_ids.append(stem)
|
||||
vpath = str(data / "videos" / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = _read_frames(vpath)
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]] # [T,N,2] px
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]] # [T,N]
|
||||
disp = np.sqrt(((tracks - tracks[0:1])**2).sum(-1)).max(0) # [N] max displacement
|
||||
|
||||
ever_visible = vis.astype(bool).any(axis=0) # [N]
|
||||
first_visible = np.where(ever_visible, np.argmax(vis.astype(bool), axis=0), -1) # [N]
|
||||
unique_frames = sorted(set(first_visible[ever_visible].tolist()))
|
||||
|
||||
for cfg in configs:
|
||||
# Run FastSAM on each unique first-visible frame for this config.
|
||||
frame_masks: dict[int, np.ndarray] = {}
|
||||
for frame_t in unique_frames:
|
||||
res = model(frames[frame_t],
|
||||
device=args.device,
|
||||
retina_masks=True,
|
||||
imgsz=cfg["imgsz"],
|
||||
conf=cfg["conf"],
|
||||
iou=cfg["iou"],
|
||||
verbose=False)
|
||||
frame_masks[frame_t] = extract_masks(res, H, W, cfg["min_area_frac"], cfg["max_masks"])
|
||||
|
||||
oid = assign_object_ids_multiframe(frame_masks, first_visible, tracks, H, W)
|
||||
masks = frame_masks.get(0, np.zeros((0, H, W), bool)) # frame-0 masks for SAM panel
|
||||
objs = sorted(int(o) for o in np.unique(oid) if int(o) >= 0)
|
||||
|
||||
# (1) SAM panel: colored masks + chosen point (highest-motion track) per object
|
||||
mcols = _colors(len(masks) + 1)
|
||||
base = frames[0].astype(np.float32)
|
||||
for mi, m in enumerate(masks):
|
||||
base[m] = 0.55 * base[m] + 0.45 * mcols[mi % len(mcols)][None].astype(np.float32)
|
||||
img = Image.fromarray(base.clip(0, 255).astype(np.uint8))
|
||||
draw = ImageDraw.Draw(img)
|
||||
for o in objs:
|
||||
idx = np.where(oid == o)[0]
|
||||
pick = idx[np.argmax(disp[idx])]
|
||||
x, y = float(tracks[0, pick, 0]), float(tracks[0, pick, 1])
|
||||
draw.ellipse([x - 5, y - 5, x + 5, y + 5], fill=(255, 255, 255), outline=(0, 0, 0))
|
||||
sam_name = f"{cfg['name']}__{stem}_sam.png"
|
||||
img.save(str(out / sam_name))
|
||||
|
||||
# (2) CoTracker tracks over the video, colored by object (background gray)
|
||||
trk_name = ""
|
||||
if not args.no_tracks:
|
||||
ocols = _colors(len(objs) + 1)
|
||||
pcols = np.tile(np.array([[110, 110, 110]], np.uint8), (tracks.shape[1], 1))
|
||||
for oi, o in enumerate(objs):
|
||||
pcols[oid == o] = ocols[oi % len(ocols)]
|
||||
# grid-aware subsample: a flat row-major stride staggers columns and
|
||||
# looks like half the grid; stride rows AND cols equally instead so the
|
||||
# true (e.g. 50x50) grid stays visible and aligned.
|
||||
N = tracks.shape[1]
|
||||
G = int(round(N**0.5))
|
||||
if G * G == N:
|
||||
k = max(1, G // 50) # keep full grid up to 50x50
|
||||
sel = np.arange(N).reshape(G, G)[::k, ::k].reshape(-1)
|
||||
else:
|
||||
st = max(1, N // 1500)
|
||||
sel = np.arange(0, N, st)
|
||||
ov = _draw_overlay(frames, tracks[:, sel].copy(), vis[:, sel], pcols[sel], 12, 2, 0.5)
|
||||
trk_name = f"{cfg['name']}__{stem}_tracks.mp4"
|
||||
imageio.mimsave(str(out / trk_name), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
|
||||
entries.append({
|
||||
"config":
|
||||
cfg["name"],
|
||||
"clip":
|
||||
stem,
|
||||
"caption": (it["cap"][0] if isinstance(it.get("cap"), list) else str(it.get("cap", "")))[:120],
|
||||
"n_masks":
|
||||
int(len(masks)),
|
||||
"n_objects":
|
||||
len(objs),
|
||||
"n_labeled":
|
||||
int((oid >= 0).sum()),
|
||||
"n_points":
|
||||
int(oid.shape[0]),
|
||||
"sam":
|
||||
sam_name,
|
||||
"tracks":
|
||||
trk_name,
|
||||
})
|
||||
print(
|
||||
f"[viz] [{k}/{len(items)}] {stem} [{cfg['name']}]: "
|
||||
f"{len(masks)} masks, {len(objs)} objs, {(oid >= 0).sum()}/{oid.shape[0]} pts labeled",
|
||||
flush=True)
|
||||
|
||||
manifest = {"configs": configs, "clips": clip_ids, "entries": entries}
|
||||
(out / "manifest.json").write_text(json.dumps(manifest, indent=2))
|
||||
print(f"[viz] rendered {len(entries)} panels -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
manifest = json.loads((viz / "manifest.json").read_text())
|
||||
entries = manifest["entries"]
|
||||
cfg_names = [c["name"] for c in manifest["configs"]]
|
||||
clip_ids = manifest["clips"]
|
||||
cfg_params = {c["name"]: c for c in manifest["configs"]}
|
||||
# closure state: the currently-filtered entry list that the gallery reflects
|
||||
current = {"entries": list(entries)}
|
||||
|
||||
def _label(e: dict) -> str:
|
||||
return f'{e["clip"]} | {e["config"]} | {e["n_masks"]}m/{e["n_objects"]}o'
|
||||
|
||||
def gallery_for(cfg_sel: str, clip_sel: str):
|
||||
es = entries
|
||||
if cfg_sel and cfg_sel != "(all)":
|
||||
es = [e for e in es if e["config"] == cfg_sel]
|
||||
if clip_sel and clip_sel != "(all)":
|
||||
es = [e for e in es if e["clip"] == clip_sel]
|
||||
current["entries"] = es
|
||||
return [(str(viz / e["sam"]), _label(e)) for e in es]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = current["entries"][evt.index]
|
||||
tv = str(viz / e["tracks"]) if e.get("tracks") else None
|
||||
info = {**e, "config_params": cfg_params.get(e["config"], {})}
|
||||
return str(viz / e["sam"]), tv, json.dumps(info, indent=2)
|
||||
|
||||
with gr.Blocks(title="FastSAM config sweep") as demo:
|
||||
gr.Markdown("### FastSAM config sweep — masks + chosen per-object points + CoTracker tracks\n"
|
||||
"Filter by **config** and/or **clip**, scroll the gallery, click a tile to inspect. "
|
||||
"Tile label: `clip | config | <#masks>m/<#objects-with-points>o`. "
|
||||
"Fewer, cleaner masks = less over-segmentation.")
|
||||
with gr.Row():
|
||||
cfg_dd = gr.Dropdown(["(all)"] + cfg_names, value="(all)", label="config")
|
||||
clip_dd = gr.Dropdown(["(all)"] + clip_ids, value="(all)", label="clip")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gallery_for("(all)", "(all)"), columns=4, height=620, label="panels (click to view)")
|
||||
with gr.Column():
|
||||
samimg = gr.Image(label="frame-0: SAM masks + chosen points")
|
||||
trkvid = gr.Video(label="CoTracker tracks (colored by object)")
|
||||
meta = gr.Code(label="info", language="json")
|
||||
cfg_dd.change(gallery_for, [cfg_dd, clip_dd], g)
|
||||
clip_dd.change(gallery_for, [cfg_dd, clip_dd], g)
|
||||
g.select(show, None, [samimg, trkvid, meta])
|
||||
# allowed_paths: without this gradio blocks serving the PNGs/mp4s that live
|
||||
# outside its app root -> the three side panels error out on click.
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--model", default="FastSAM-s.pt")
|
||||
r.add_argument("--device", default="cuda")
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--configs", default=None, help="inline JSON list of config dicts (overrides sweep)")
|
||||
r.add_argument("--configs-json", default=None, help="path to a JSON list of config dicts")
|
||||
r.add_argument("--no-tracks", action="store_true", help="skip the (slow) track-overlay mp4s")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7880)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,221 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Interactive dashboard for the SPARSE ~1-per-object + few-background sampler.
|
||||
|
||||
Per Yongqi's advice, our track budget for the sparse recipe is dramatically smaller
|
||||
than MotionStream's 1000-2500: keep 1 track per SAM object, then add ``num_sampled``
|
||||
extra points drawn from the rest (either uniformly at random, or weighted by the
|
||||
precomputed low-rank informativeness). The dashboard lets us eyeball whether the
|
||||
resulting subset actually covers the "action" of each clip before we commit to
|
||||
launching training runs on it.
|
||||
|
||||
CPU only, gradio share (no GPU needed).
|
||||
|
||||
.venv/bin/python data_pipeline/sparse_sampling_dashboard.py \
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/synthetic_toy --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from track_informativeness import _draw_overlay, _read_frames # noqa: E402
|
||||
|
||||
|
||||
def sparse_sample(oid: np.ndarray, tw: np.ndarray | None, vis0: np.ndarray, num_sampled: int, mode: str,
|
||||
seed: int) -> np.ndarray:
|
||||
"""One point per SAM object (weighted-by-`tw` inside each object) + ``num_sampled`` extras
|
||||
drawn from the rest via ``mode`` in {"random", "weighted"}. Returns kept bool [N]."""
|
||||
rng = np.random.default_rng(int(seed))
|
||||
N = oid.shape[0]
|
||||
valid = vis0 > 0.5
|
||||
keep = np.zeros(N, bool)
|
||||
# (a) 1 per object
|
||||
for o in np.unique(oid):
|
||||
if int(o) < 0:
|
||||
continue
|
||||
idx = np.where((oid == o) & valid)[0]
|
||||
if idx.size == 0:
|
||||
continue
|
||||
if tw is not None and tw[idx].sum() > 0:
|
||||
p = tw[idx] / tw[idx].sum()
|
||||
pick = int(rng.choice(idx, p=p))
|
||||
else:
|
||||
pick = int(rng.choice(idx))
|
||||
keep[pick] = True
|
||||
# (b) num_sampled extras from the remaining valid points
|
||||
pool = np.where((~keep) & valid)[0]
|
||||
k = int(min(max(num_sampled, 0), pool.size))
|
||||
if k > 0:
|
||||
if mode == "weighted" and tw is not None and tw[pool].sum() > 0:
|
||||
p = tw[pool] / tw[pool].sum()
|
||||
picks = rng.choice(pool, size=k, replace=False, p=p)
|
||||
else:
|
||||
picks = rng.choice(pool, size=k, replace=False)
|
||||
keep[picks] = True
|
||||
return keep
|
||||
|
||||
|
||||
def _load(data_dir: Path) -> list[dict]:
|
||||
man = json.loads((data_dir / "videos2caption.json").read_text())
|
||||
labels_dir = data_dir / "sam_labels"
|
||||
items = []
|
||||
for it in man:
|
||||
stem = Path(it["path"]).stem
|
||||
vpath = str(data_dir / "videos" / it["path"])
|
||||
npzp = it.get("points_path") or str(data_dir / "tracks" / f"{stem}.npz")
|
||||
d = np.load(npzp)
|
||||
lp = labels_dir / f"{stem}.npy"
|
||||
labels = np.load(lp) if lp.exists() else None # [H,W] int16, -1 = background
|
||||
items.append({
|
||||
"stem": stem,
|
||||
"caption": it["cap"][0] if isinstance(it.get("cap"), list) else "",
|
||||
"fps": int(it.get("fps", 24)),
|
||||
"vpath": vpath,
|
||||
"tracks": d["tracks"].astype(np.float32),
|
||||
"vis": d["visibility"].astype(np.float32),
|
||||
"oid": d["object_ids"].astype(np.int64) if "object_ids" in d else None,
|
||||
"tw": d["track_weights"].astype(np.float32) if "track_weights" in d else None,
|
||||
"labels": labels,
|
||||
})
|
||||
return items
|
||||
|
||||
|
||||
def build_ui(items: list[dict], out_dir: Path):
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
def _render_frame(clip_idx: int, num_sampled: int, mode: str, seed: int, mask_alpha: float = 0.45):
|
||||
it = items[int(clip_idx)]
|
||||
frames = _read_frames(it["vpath"])
|
||||
tracks = it["tracks"][:frames.shape[0]]
|
||||
vis = it["vis"][:frames.shape[0]]
|
||||
oid = it["oid"]
|
||||
tw = it["tw"]
|
||||
labels = it.get("labels")
|
||||
vis0 = vis[0]
|
||||
if oid is None:
|
||||
return None, "No object_ids in npz — run segment_tracks first."
|
||||
keep = sparse_sample(oid, tw, vis0, int(num_sampled), mode, int(seed))
|
||||
n_kept = int(keep.sum())
|
||||
n_obj = int(len(np.unique(oid[oid >= 0])))
|
||||
# Blend SAM segmentation on top of frame 0: color each segment by the same tab20
|
||||
# index we use for the tracked points (mask id == object id).
|
||||
base = frames[0].astype(np.float32)
|
||||
if labels is not None:
|
||||
cmap = plt.get_cmap("tab20")
|
||||
H, W = base.shape[:2]
|
||||
overlay = np.zeros_like(base)
|
||||
uniq = np.unique(labels)
|
||||
for u in uniq:
|
||||
if int(u) < 0:
|
||||
continue
|
||||
col = np.array(cmap(int(u) % 20)[:3]) * 255
|
||||
m = labels == u
|
||||
overlay[m] = col
|
||||
mask = (labels >= 0)[..., None]
|
||||
base = np.where(mask, (1 - mask_alpha) * base + mask_alpha * overlay, base)
|
||||
base = base.clip(0, 255).astype(np.uint8)
|
||||
fig, ax = plt.subplots(1, 1, figsize=(9, 5.5))
|
||||
ax.imshow(base)
|
||||
x, y = tracks[0, :, 0], tracks[0, :, 1]
|
||||
ax.scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.35) # dropped
|
||||
# kept: color per unique object id (SAME palette as the mask overlay above)
|
||||
cmap = plt.get_cmap("tab20")
|
||||
for o in np.unique(oid[keep]):
|
||||
mask_pts = keep & (oid == o)
|
||||
col = "white" if int(o) < 0 else cmap(int(o) % 20)
|
||||
ax.scatter(x[mask_pts], y[mask_pts], c=[col], s=90, edgecolors="black", linewidths=1.2)
|
||||
ax.set_title(
|
||||
f"{it['stem']}: {n_kept} tracks kept ({n_obj} SAM objects, +{num_sampled} background, "
|
||||
f"{mode}, seed {seed})",
|
||||
fontsize=10)
|
||||
ax.axis("off")
|
||||
img_path = str(out_dir / f"preview_{it['stem']}.png")
|
||||
fig.savefig(img_path, dpi=90, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
info = json.dumps(
|
||||
{
|
||||
"caption": it["caption"],
|
||||
"n_kept": n_kept,
|
||||
"n_objects_present": n_obj,
|
||||
"kept_from_objects": int(len(np.unique(oid[keep][oid[keep] >= 0]))),
|
||||
"kept_from_background": int((keep & (oid < 0)).sum()),
|
||||
},
|
||||
indent=2)
|
||||
return img_path, info
|
||||
|
||||
def _render_video(clip_idx: int, num_sampled: int, mode: str, seed: int):
|
||||
it = items[int(clip_idx)]
|
||||
frames = _read_frames(it["vpath"])
|
||||
tracks = it["tracks"][:frames.shape[0]]
|
||||
vis = it["vis"][:frames.shape[0]]
|
||||
oid = it["oid"]
|
||||
tw = it["tw"]
|
||||
keep = sparse_sample(oid, tw, vis[0], int(num_sampled), mode, int(seed))
|
||||
kidx = np.where(keep)[0]
|
||||
if kidx.size == 0:
|
||||
return None
|
||||
cmap = plt.get_cmap("tab20")
|
||||
cols = np.zeros((kidx.size, 3), dtype=np.uint8)
|
||||
for o in np.unique(oid[kidx]):
|
||||
m = (oid[kidx] == o)
|
||||
rgba = (1.0, 1.0, 1.0, 1.0) if int(o) < 0 else cmap(int(o) % 20)
|
||||
cols[m] = (np.array(rgba[:3]) * 255).astype(np.uint8)
|
||||
ov = _draw_overlay(frames, tracks[:, kidx].copy(), vis[:, kidx], cols, 14, 3, 0.5)
|
||||
out = str(out_dir / f"preview_{it['stem']}_overlay.mp4")
|
||||
imageio.mimsave(out, ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
return out
|
||||
|
||||
with gr.Blocks(title="Sparse track sampler dashboard") as demo:
|
||||
gr.Markdown("### Sparse-sampling recipe: 1 track per SAM object + `num_sampled` extras\n"
|
||||
"Frame 0 shows all 2500 grid tracks in gray, and the kept ones colored by object. "
|
||||
"Set `num_sampled` = 0 for pure 1-per-object; raise it to add background context. "
|
||||
"Toggle `weighted` to bias the extras toward high-lowrank-informativeness points. "
|
||||
"The overlay video shows only the kept tracks over the whole clip.")
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
clip_dd = gr.Dropdown([f"{i:02d} — {items[i]['stem']}" for i in range(len(items))],
|
||||
value=f"00 — {items[0]['stem']}",
|
||||
label="Clip",
|
||||
type="index")
|
||||
num_slider = gr.Slider(0, 200, value=20, step=1, label="num_sampled (extras)")
|
||||
mode_radio = gr.Radio(["random", "weighted"], value="weighted", label="extras sampling mode")
|
||||
seed_slider = gr.Slider(0, 100, value=0, step=1, label="seed")
|
||||
info_box = gr.Code(label="info", language="json")
|
||||
vid_btn = gr.Button("Render overlay video", variant="primary")
|
||||
with gr.Column(scale=2):
|
||||
frame_img = gr.Image(label="frame 0 with kept tracks", height=460)
|
||||
overlay_vid = gr.Video(label="kept-tracks overlay video", height=460)
|
||||
|
||||
for src in [clip_dd, num_slider, mode_radio, seed_slider]:
|
||||
src.change(_render_frame, [clip_dd, num_slider, mode_radio, seed_slider], [frame_img, info_box])
|
||||
vid_btn.click(_render_video, [clip_dd, num_slider, mode_radio, seed_slider], overlay_vid)
|
||||
demo.load(_render_frame, [clip_dd, num_slider, mode_radio, seed_slider], [frame_img, info_box])
|
||||
return demo
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--data-dir", required=True)
|
||||
p.add_argument("--out-dir",
|
||||
default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/sparse_sampling_dashboard_out")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7894)
|
||||
p.add_argument("--share", action="store_true")
|
||||
args = p.parse_args()
|
||||
out_dir = Path(args.out_dir)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
items = _load(Path(args.data_dir))
|
||||
demo = build_ui(items, out_dir)
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(out_dir)])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,57 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Split a videos2caption.json manifest into NUM_SHARDS shards for data-parallel preprocess.
|
||||
|
||||
The i2v_track preprocess reads a merge.txt (``<clips_dir>,<json>``) whose JSON is a
|
||||
*list* of dicts (json.load, NOT jsonlines) — see
|
||||
fastvideo/dataset/preprocessing_datasets.py:_load_raw_data. So each shard is itself a
|
||||
JSON array + its own merge.txt, and each shard is fed to one single-GPU v1_preprocess.py
|
||||
process.
|
||||
|
||||
Also injects a ``duration`` field when missing: FrameSamplingStage.should_keep drops any
|
||||
row with ``duration is None`` (preprocessing_datasets.py:183), and OpenVid's
|
||||
videos2caption.json carries only num_frames/fps. duration = num_frames / fps.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--manifest", required=True, help="videos2caption.json (JSON array)")
|
||||
ap.add_argument("--clips-dir", required=True, help="dir with the .mp4 clips (merge.txt col 1)")
|
||||
ap.add_argument("--out-dir", required=True, help="output shards dir (shard_XXXXX/ created inside)")
|
||||
ap.add_argument("--num-shards", type=int, required=True)
|
||||
a = ap.parse_args()
|
||||
|
||||
items = json.load(open(a.manifest))
|
||||
filled = 0
|
||||
for it in items:
|
||||
if it.get("duration") is None:
|
||||
nf, fps = it.get("num_frames"), it.get("fps")
|
||||
if nf and fps:
|
||||
it["duration"] = float(nf) / float(fps)
|
||||
filled += 1
|
||||
|
||||
out = Path(a.out_dir)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
# Contiguous shards keep each shard's clips locally clustered on disk.
|
||||
n = a.num_shards
|
||||
total = len(items)
|
||||
per = (total + n - 1) // n
|
||||
written = 0
|
||||
for k in range(n):
|
||||
chunk = items[k * per:(k + 1) * per]
|
||||
sd = out / f"shard_{k:05d}"
|
||||
sd.mkdir(parents=True, exist_ok=True)
|
||||
mani = sd / "videos2caption.json"
|
||||
mani.write_text(json.dumps(chunk, indent=2))
|
||||
(sd / "merge.txt").write_text(f"{a.clips_dir},{mani}")
|
||||
written += len(chunk)
|
||||
print(f"[split] {total} rows ({filled} got injected duration) -> {n} shards "
|
||||
f"(~{per}/shard) under {out}; wrote {written} rows total")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,289 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Author synthetic point-track motion controls for WanTrack.
|
||||
|
||||
A "track" is just a per-point (x, y) trajectory over time, so we can *author*
|
||||
counterfactual motion directly (no CoTracker needed -- that is only for
|
||||
*extracting* tracks from real video). These authored tracks are the test-time
|
||||
control signal: feed them to the model and check the generated motion follows.
|
||||
|
||||
All coordinates are NORMALIZED to [0, 1] (x/width, y/height) -- the same space
|
||||
the preprocessing stores and the model trains on. Outputs:
|
||||
tracks float32 [T, N, 2] normalized (x, y) per point per frame
|
||||
visibility float32 [T, N] 1 = active/visible, 0 = inactive/occluded
|
||||
|
||||
Modes mirror MotionStream's control surface:
|
||||
- ``pan`` : global translation (camera/scene pan)
|
||||
- ``zoom`` : scale toward/away from a center (dolly in/out)
|
||||
- ``rotate`` : rotate the grid about a center
|
||||
- ``drag`` : move only points inside a region (the sparse "drag" handle),
|
||||
with a smooth spatial falloff; everything else stays put
|
||||
- ``swirl`` : rotation whose angle decays with radius
|
||||
- ``static`` : no motion (tests "does text+first-frame alone reproduce motion?")
|
||||
|
||||
``select_*`` helpers turn a dense field into a SPARSE control by zeroing the
|
||||
visibility of unselected points -- this is how MotionStream lets you drag just a
|
||||
few points. (Requires a model trained with point-subsampling aug to work well.)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Grid + small helpers
|
||||
# ----------------------------------------------------------------------
|
||||
def make_grid(grid_size: int = 50) -> np.ndarray:
|
||||
"""Frame-0 normalized grid positions, returned as [N, 2] (x, y) in [0, 1].
|
||||
|
||||
Row-major (matches CoTracker's grid and our 50x50 preprocessing), so
|
||||
``reshape(grid_size, grid_size, 2)`` is [row(y), col(x)].
|
||||
"""
|
||||
ys, xs = np.meshgrid(
|
||||
np.linspace(0.0, 1.0, grid_size, dtype=np.float32),
|
||||
np.linspace(0.0, 1.0, grid_size, dtype=np.float32),
|
||||
indexing="ij",
|
||||
)
|
||||
return np.stack([xs.reshape(-1), ys.reshape(-1)], axis=-1) # [N, 2]
|
||||
|
||||
|
||||
def _ramp(num_frames: int, ease: bool = True) -> np.ndarray:
|
||||
"""Time ramp in [0, 1] over ``num_frames`` (optionally smooth ease-in-out)."""
|
||||
t = np.linspace(0.0, 1.0, num_frames, dtype=np.float32)
|
||||
if ease:
|
||||
t = t * t * (3.0 - 2.0 * t) # smoothstep
|
||||
return t
|
||||
|
||||
|
||||
def _visibility_from_bounds(tracks: np.ndarray, margin: float = 0.02) -> np.ndarray:
|
||||
"""1 where a point is inside [(-m), (1+m)] in both axes, else 0."""
|
||||
lo, hi = -margin, 1.0 + margin
|
||||
inb = (tracks[..., 0] >= lo) & (tracks[..., 0] <= hi) & (tracks[..., 1] >= lo) & (tracks[..., 1] <= hi)
|
||||
return inb.astype(np.float32)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Motion fields -> (tracks [T,N,2], visibility [T,N])
|
||||
# ----------------------------------------------------------------------
|
||||
def static(grid: np.ndarray, num_frames: int) -> tuple[np.ndarray, np.ndarray]:
|
||||
tracks = np.repeat(grid[None], num_frames, axis=0)
|
||||
return tracks, np.ones((num_frames, grid.shape[0]), np.float32)
|
||||
|
||||
|
||||
def pan(grid: np.ndarray, num_frames: int, dx: float, dy: float,
|
||||
ease: bool = True) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Translate all points by (dx, dy) (normalized units) over the clip."""
|
||||
r = _ramp(num_frames, ease)[:, None, None]
|
||||
tracks = grid[None] + r * np.array([dx, dy], np.float32)[None, None]
|
||||
return tracks.astype(np.float32), _visibility_from_bounds(tracks)
|
||||
|
||||
|
||||
def zoom(grid: np.ndarray, num_frames: int, scale: float,
|
||||
center: tuple[float, float] = (0.5, 0.5), ease: bool = True
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Scale toward (scale<1) or away from (scale>1) ``center`` (dolly)."""
|
||||
c = np.array(center, np.float32)[None, None]
|
||||
r = _ramp(num_frames, ease)[:, None, None]
|
||||
s = 1.0 + (scale - 1.0) * r
|
||||
tracks = c + (grid[None] - c) * s
|
||||
return tracks.astype(np.float32), _visibility_from_bounds(tracks)
|
||||
|
||||
|
||||
def rotate(grid: np.ndarray, num_frames: int, degrees: float,
|
||||
center: tuple[float, float] = (0.5, 0.5), ease: bool = True
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
c = np.array(center, np.float32)[None]
|
||||
r = _ramp(num_frames, ease)
|
||||
out = np.empty((num_frames, grid.shape[0], 2), np.float32)
|
||||
rel = grid - c
|
||||
for t in range(num_frames):
|
||||
a = np.deg2rad(degrees) * r[t]
|
||||
ca, sa = np.cos(a), np.sin(a)
|
||||
rot = np.array([[ca, -sa], [sa, ca]], np.float32)
|
||||
out[t] = rel @ rot.T + c
|
||||
return out, _visibility_from_bounds(out)
|
||||
|
||||
|
||||
def drag(grid: np.ndarray, num_frames: int, *, center: tuple[float, float],
|
||||
dx: float, dy: float, radius: float = 0.15, falloff: str = "smooth",
|
||||
ease: bool = True) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Move only points within ``radius`` of ``center`` by (dx, dy).
|
||||
|
||||
``falloff='smooth'`` weights displacement by a cosine window of the distance
|
||||
(handle-like drag); ``'hard'`` moves all in-radius points equally. Points
|
||||
outside the radius stay put. Visibility stays 1 for all points (dense field
|
||||
with a localized motion) -- use ``select_radius`` to make it *sparse*.
|
||||
"""
|
||||
c = np.array(center, np.float32)[None]
|
||||
d = np.linalg.norm(grid - c, axis=-1) # [N]
|
||||
if falloff == "smooth":
|
||||
w = np.clip(1.0 - d / max(radius, 1e-6), 0.0, 1.0)
|
||||
w = 0.5 - 0.5 * np.cos(np.pi * w) # smooth 0->1
|
||||
else:
|
||||
w = (d <= radius).astype(np.float32)
|
||||
r = _ramp(num_frames, ease)[:, None, None]
|
||||
disp = (w[:, None] * np.array([dx, dy], np.float32)[None])[None] # [1,N,2]
|
||||
tracks = grid[None] + r * disp
|
||||
return tracks.astype(np.float32), _visibility_from_bounds(tracks)
|
||||
|
||||
|
||||
def swirl(grid: np.ndarray, num_frames: int, degrees: float,
|
||||
center: tuple[float, float] = (0.5, 0.5), radius: float = 0.5,
|
||||
ease: bool = True) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Rotation whose angle decays linearly to 0 at ``radius`` from center."""
|
||||
c = np.array(center, np.float32)[None]
|
||||
rel = grid - c
|
||||
dist = np.linalg.norm(rel, axis=-1)
|
||||
decay = np.clip(1.0 - dist / max(radius, 1e-6), 0.0, 1.0)
|
||||
rmp = _ramp(num_frames, ease)
|
||||
out = np.empty((num_frames, grid.shape[0], 2), np.float32)
|
||||
for t in range(num_frames):
|
||||
a = np.deg2rad(degrees) * rmp[t] * decay
|
||||
ca, sa = np.cos(a), np.sin(a)
|
||||
x = rel[:, 0] * ca - rel[:, 1] * sa
|
||||
y = rel[:, 0] * sa + rel[:, 1] * ca
|
||||
out[t] = np.stack([x, y], -1) + c
|
||||
return out, _visibility_from_bounds(out)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Sparsify (sparse "drag a few points" control)
|
||||
# ----------------------------------------------------------------------
|
||||
def select_radius(visibility: np.ndarray, grid: np.ndarray, *,
|
||||
center: tuple[float, float], radius: float) -> np.ndarray:
|
||||
"""Zero visibility outside ``radius`` of ``center`` (keep only a local handle)."""
|
||||
c = np.array(center, np.float32)[None]
|
||||
keep = (np.linalg.norm(grid - c, axis=-1) <= radius).astype(np.float32) # [N]
|
||||
return visibility * keep[None]
|
||||
|
||||
|
||||
def select_random(visibility: np.ndarray, k: int, seed: int = 0) -> np.ndarray:
|
||||
"""Keep only ``k`` random points active (rest visibility 0)."""
|
||||
rng = np.random.default_rng(seed)
|
||||
n = visibility.shape[1]
|
||||
keep = np.zeros(n, np.float32)
|
||||
keep[rng.choice(n, size=min(k, n), replace=False)] = 1.0
|
||||
return visibility * keep[None]
|
||||
|
||||
|
||||
def select_stride(visibility: np.ndarray, grid_size: int, stride: int) -> np.ndarray:
|
||||
"""Keep a coarse sub-grid (every ``stride``-th point in both dims) active."""
|
||||
n = visibility.shape[1]
|
||||
mask2d = np.zeros((grid_size, grid_size), np.float32)
|
||||
mask2d[::stride, ::stride] = 1.0
|
||||
return visibility * mask2d.reshape(1, n)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Motion transfer: reuse tracks extracted from a real/other video
|
||||
# ----------------------------------------------------------------------
|
||||
def from_npz(npz_path: str, num_frames: int | None = None
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Load tracks from an ``extract_tracks.py`` npz, normalized to [0,1]."""
|
||||
data = np.load(npz_path)
|
||||
tr = data["tracks"].astype(np.float32) # [T, N, 2] pixels
|
||||
vis = data["visibility"].astype(np.float32) # [T, N]
|
||||
w = float(data["width"]) if "width" in data else 1.0
|
||||
h = float(data["height"]) if "height" in data else 1.0
|
||||
tr = tr.copy()
|
||||
tr[..., 0] /= max(w, 1e-6)
|
||||
tr[..., 1] /= max(h, 1e-6)
|
||||
if num_frames is not None:
|
||||
tr, vis = tr[:num_frames], vis[:num_frames]
|
||||
return tr, vis
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Convenience: named presets
|
||||
# ----------------------------------------------------------------------
|
||||
def preset(name: str, num_frames: int, grid_size: int = 50, *,
|
||||
strength: float = 0.25) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Author a named motion preset. ``strength`` scales translations/zoom."""
|
||||
g = make_grid(grid_size)
|
||||
name = name.lower()
|
||||
if name == "static":
|
||||
return static(g, num_frames)
|
||||
if name in ("pan_right", "pan_left", "pan_up", "pan_down"):
|
||||
dx = {"pan_right": strength, "pan_left": -strength}.get(name, 0.0)
|
||||
dy = {"pan_down": strength, "pan_up": -strength}.get(name, 0.0)
|
||||
return pan(g, num_frames, dx, dy)
|
||||
if name == "zoom_in":
|
||||
return zoom(g, num_frames, 1.0 + strength)
|
||||
if name == "zoom_out":
|
||||
return zoom(g, num_frames, 1.0 - strength)
|
||||
if name in ("rotate_cw", "rotate_ccw"):
|
||||
return rotate(g, num_frames, 30.0 * (1 if name == "rotate_cw" else -1))
|
||||
if name == "swirl":
|
||||
return swirl(g, num_frames, 60.0)
|
||||
if name == "drag_center_right":
|
||||
return drag(g, num_frames, center=(0.5, 0.5), dx=strength, dy=0.0, radius=0.2)
|
||||
raise ValueError(f"unknown preset {name!r}")
|
||||
|
||||
|
||||
PRESETS = ["static", "pan_right", "pan_left", "pan_up", "pan_down",
|
||||
"zoom_in", "zoom_out", "rotate_cw", "rotate_ccw", "swirl",
|
||||
"drag_center_right"]
|
||||
|
||||
|
||||
def to_pixel(tracks: np.ndarray, height: int, width: int) -> np.ndarray:
|
||||
"""Normalized [0,1] tracks -> pixel coords for overlay/EPE."""
|
||||
out = tracks.copy()
|
||||
out[..., 0] *= width
|
||||
out[..., 1] *= height
|
||||
return out
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# CLI: overlay a preset on a first frame (sanity check, no model/GPU)
|
||||
# ----------------------------------------------------------------------
|
||||
def _overlay_preview(frame0: np.ndarray, tracks_px: np.ndarray, vis: np.ndarray,
|
||||
stride: int = 3) -> np.ndarray:
|
||||
import colorsys
|
||||
|
||||
from PIL import Image, ImageDraw
|
||||
h, w, _ = frame0.shape
|
||||
img = Image.fromarray(frame0).convert("RGB")
|
||||
draw = ImageDraw.Draw(img)
|
||||
n = tracks_px.shape[1]
|
||||
gs = int(round(n ** 0.5))
|
||||
keep = np.zeros(n, bool)
|
||||
keep.reshape(gs, gs)[::stride, ::stride] = True
|
||||
for pi in range(n):
|
||||
if not keep[pi]:
|
||||
continue
|
||||
col = tuple(int(c * 255) for c in colorsys.hsv_to_rgb((pi % gs) / gs, 1.0, 1.0))
|
||||
pts = [(float(tracks_px[t, pi, 0]), float(tracks_px[t, pi, 1]))
|
||||
for t in range(tracks_px.shape[0]) if vis[t, pi] >= 0.5]
|
||||
if len(pts) >= 2:
|
||||
draw.line(pts, fill=col, width=1)
|
||||
if pts:
|
||||
x, y = pts[-1]
|
||||
draw.ellipse([x - 2, y - 2, x + 2, y + 2], fill=col)
|
||||
return np.asarray(img)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description="Preview a synthetic track preset over a first frame.")
|
||||
p.add_argument("--frame", type=Path, required=True, help="First-frame image (png/jpg).")
|
||||
p.add_argument("--preset", type=str, default="pan_right", choices=PRESETS)
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--grid-size", type=int, default=50)
|
||||
p.add_argument("--strength", type=float, default=0.25)
|
||||
p.add_argument("--out", type=Path, required=True)
|
||||
args = p.parse_args()
|
||||
|
||||
import imageio.v2 as imageio
|
||||
frame0 = np.asarray(imageio.imread(args.frame))[..., :3]
|
||||
h, w, _ = frame0.shape
|
||||
tracks, vis = preset(args.preset, args.num_frames, args.grid_size, strength=args.strength)
|
||||
overlay = _overlay_preview(frame0, to_pixel(tracks, h, w), vis)
|
||||
args.out.parent.mkdir(parents=True, exist_ok=True)
|
||||
imageio.imwrite(str(args.out), overlay)
|
||||
print(f"[synthetic-tracks] preset={args.preset} frames={args.num_frames} "
|
||||
f"points={tracks.shape[1]} -> {args.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,142 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Controllability test for an overfitted TrackWan model.
|
||||
|
||||
Disentangles "did it overfit on the video content" from "did it learn to follow
|
||||
the tracks" by holding the content conditioning (first frame + text) fixed and
|
||||
varying ONLY the control tracks:
|
||||
|
||||
- gt : the clip's own tracks (reconstruction; should match GT)
|
||||
- none : no tracks at all (if motion still appears -> content overfit)
|
||||
- pan_*/zoom/drag : authored counterfactuals (if motion follows -> real control)
|
||||
- swap : another clip's tracks (motion transfer)
|
||||
|
||||
For each, it generates a video, overlays the control tracks, saves an mp4, and
|
||||
computes CoTracker EPE (input tracks vs tracks re-extracted from the generation).
|
||||
Low EPE on counterfactuals = genuine track control.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def _overlay(frames_thwc: np.ndarray, tracks_norm: np.ndarray, vis: np.ndarray, stride: int = 3) -> list[np.ndarray]:
|
||||
from fastvideo.train.callbacks.track_validation import (_draw_overlay, _grid_colors, _subsample)
|
||||
T, H, W, _ = frames_thwc.shape
|
||||
tt = min(T, tracks_norm.shape[0])
|
||||
fr, tr, vs = frames_thwc[:tt], tracks_norm[:tt], vis[:tt]
|
||||
grid = int(round(tr.shape[1]**0.5))
|
||||
tr, vs = _subsample(tr, vs, grid, stride)
|
||||
colors = _grid_colors(grid, stride)
|
||||
if colors.shape[0] != tr.shape[1]:
|
||||
colors = _grid_colors(int(round(tr.shape[1]**0.5)) or 1, 1)[:tr.shape[1]]
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
return _draw_overlay(fr, trpx, vs, colors, 12, 2, 0.5)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True, help="dcp_to_diffusers export dir")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data",
|
||||
default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset")
|
||||
p.add_argument("--out", required=True)
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import (compute_epe, load_cotracker)
|
||||
ct = load_cotracker(model.device)
|
||||
|
||||
# latent T -> pixel T
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
Tpx = (num_lat_t - 1) * ratio + 1
|
||||
|
||||
results = []
|
||||
for ci, s in zip(args.clips, samples, strict=False):
|
||||
gt_tracks = s["track_points"][0].numpy()[:Tpx] # [T,N,2] normalized
|
||||
gt_vis = s["track_visibility"][0].numpy()[:Tpx]
|
||||
swap = samples[(args.clips.index(ci) + 1) % len(samples)]
|
||||
swap_tracks = swap["track_points"][0].numpy()[:Tpx]
|
||||
swap_vis = swap["track_visibility"][0].numpy()[:Tpx]
|
||||
|
||||
g = st.make_grid(50)
|
||||
pan_t, pan_v = st.pan(g, Tpx, 0.25, 0.0)
|
||||
zoom_t, zoom_v = st.zoom(g, Tpx, 1.25)
|
||||
drag_t, drag_v = st.drag(g, Tpx, center=(0.5, 0.5), dx=0.3, dy=0.0, radius=0.2)
|
||||
drag_v_sparse = st.select_radius(drag_v, g, center=(0.5, 0.5), radius=0.2)
|
||||
|
||||
controls = {
|
||||
"gt": (gt_tracks, gt_vis),
|
||||
"none": (None, None),
|
||||
"pan_right": (pan_t, pan_v),
|
||||
"zoom_in": (zoom_t, zoom_v),
|
||||
"drag_dense": (drag_t, drag_v),
|
||||
"drag_sparse": (drag_t, drag_v_sparse),
|
||||
"swap": (swap_tracks, swap_vis),
|
||||
}
|
||||
|
||||
# GT reference (decode the real clip)
|
||||
ref = twi.decode_reference(model, s["vae_latent"])
|
||||
imageio.mimsave(os.path.join(args.out, f"clip{ci}_reference.mp4"),
|
||||
_overlay(ref, gt_tracks, gt_vis),
|
||||
fps=args.fps,
|
||||
macro_block_size=1)
|
||||
|
||||
for name, (tr, vs) in controls.items():
|
||||
tp = torch.from_numpy(tr)[None].float() if tr is not None else None
|
||||
tv = torch.from_numpy(vs)[None].float() if vs is not None else None
|
||||
lat = twi.generate(model,
|
||||
first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp,
|
||||
track_visibility=tv,
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=args.steps,
|
||||
seed=args.seed)
|
||||
frames = twi.decode_to_pixels(model, lat)
|
||||
ov_tr = tr if tr is not None else np.zeros((Tpx, 1, 2), np.float32)
|
||||
ov_vs = vs if vs is not None else np.zeros((Tpx, 1), np.float32)
|
||||
imageio.mimsave(os.path.join(args.out, f"clip{ci}_{name}.mp4"),
|
||||
_overlay(frames, ov_tr, ov_vs),
|
||||
fps=args.fps,
|
||||
macro_block_size=1)
|
||||
|
||||
epe = None
|
||||
if tr is not None:
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
tr_px = tr.copy()
|
||||
tr_px[..., 0] *= W
|
||||
tr_px[..., 1] *= H
|
||||
epe = compute_epe(frames, tr_px, vs, ct, model.device)["epe"]
|
||||
results.append((ci, name, epe))
|
||||
print(f"[clip {ci}] {name:12s} EPE={epe if epe is None else round(epe,2)}", flush=True)
|
||||
|
||||
print("\n=== EPE summary (lower = follows control better) ===")
|
||||
for ci, name, epe in results:
|
||||
print(f" clip {ci} {name:12s} {('n/a' if epe is None else f'{epe:7.2f} px')}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Quick motion-CFG test: does guidance>1 improve action adherence on an existing ckpt?"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st
|
||||
import trackwan_infer as twi
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument(
|
||||
"--data",
|
||||
default=
|
||||
"/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200/preprocessed_i2v_track/combined_parquet_dataset"
|
||||
)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_droid_overfit.yaml")
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--scales", type=float, nargs="+", default=[1.0, 2.0, 3.0])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
a = p.parse_args()
|
||||
model, tc = twi.load_trackwan(a.export, a.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(a.data, a.clips, text_len)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import compute_epe, load_cotracker
|
||||
ct = load_cotracker(model.device)
|
||||
nlt = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
Tpx = (nlt - 1) * ratio + 1
|
||||
# grid-snapped dense drag control (moves a radius-0.2 region by +0.3 x)
|
||||
g = st.make_grid(50)
|
||||
dt, dv = st.drag(g, Tpx, center=(0.5, 0.5), dx=0.3, dy=0.0, radius=0.2)
|
||||
disp = np.sqrt(((dt - dt[0:1])**2).sum(-1))
|
||||
moving = disp.max(0) > 8 / 832.0 # normalized thresh ~ px
|
||||
mv = np.where(moving)[0]
|
||||
print(f"{'clip':>4} {'scale':>6} {'EPE_all':>8} {'EPE_moving':>11}")
|
||||
for ci, s in zip(a.clips, samples, strict=False):
|
||||
tp = torch.from_numpy(dt)[None].float()
|
||||
tvv = torch.from_numpy(dv)[None].float()
|
||||
for sc in a.scales:
|
||||
lat = twi.generate(model,
|
||||
first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp,
|
||||
track_visibility=tvv,
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=a.steps,
|
||||
seed=1000,
|
||||
guidance_scale=sc)
|
||||
fr = twi.decode_to_pixels(model, lat)
|
||||
H, W = fr.shape[1], fr.shape[2]
|
||||
dpx = dt.copy()
|
||||
dpx[..., 0] *= W
|
||||
dpx[..., 1] *= H
|
||||
epe_all = compute_epe(fr, dpx, dv, ct, model.device)["epe"]
|
||||
epe_mv = compute_epe(fr, dpx[:, mv], dv[:, mv], ct, model.device)["epe"]
|
||||
print(f"{ci:>4} {sc:>6.1f} {epe_all:>8.2f} {epe_mv:>11.2f}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,477 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Diagnostic bake-off: which point tracks are *informative* vs *generic*?
|
||||
|
||||
A point that moves a lot is not necessarily important. If the camera pans, every
|
||||
background point moves, but that motion is shared by the crowd, so any one background
|
||||
point is redundant. We compare several ways to score / select useful points, from
|
||||
simple statistics to clustering and learned embeddings.
|
||||
|
||||
Per-point scalar scores (heatmaps, bright = informative):
|
||||
abs_disp max_t ||p_t - p_0|| naive raw motion
|
||||
affine_residual residual after a robust per-frame global removes camera pan/zoom/rot
|
||||
affine x0 -> p_t
|
||||
lowrank_residual residual after subtracting mean trajectory unique vs ALL points, but
|
||||
+ top-k shared SVD modes (ONE subspace) assumes a single subspace
|
||||
motion_outlier 1 - HDBSCAN membership prob on trajectory density outlier in motion space
|
||||
embeddings (#1 motion clustering)
|
||||
subspace_resid per-cluster low-rank residual after union-of-subspaces; the
|
||||
Spectral motion segmentation per-object lowrank (#2 SSC-lite)
|
||||
dino_uniqueness 1 - membership prob clustering DINOv2 semantic ⊕ motion (#3)
|
||||
appearance ⊕ trajectory embeddings
|
||||
|
||||
Grouping / subset panels:
|
||||
motion_clusters HDBSCAN labels on trajectory embeddings motion groups (#1)
|
||||
subspace_clusters Spectral motion segmentation labels object subspaces (#2)
|
||||
dpp / fps / greedy diverse K-subset selection what sampling wants (#4)
|
||||
|
||||
``render`` (CPU; DINOv2 on CPU) writes per clip a big multi-panel PNG + an overlay mp4
|
||||
colored by low-rank informativeness + stats. ``serve`` is a gradio gallery (--share).
|
||||
|
||||
.venv/bin/python data_pipeline/track_informativeness.py render --data-dir <root> --out <dir> --limit 8
|
||||
.venv/bin/python data_pipeline/track_informativeness.py serve --viz-dir <dir> --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
K_SELECT = 64 # size of the illustrative "useful subset" for the selection methods
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- overlay / io
|
||||
def _draw_overlay(frames: Any,
|
||||
tracks: Any,
|
||||
vis: Any,
|
||||
colors: Any,
|
||||
tail: int = 12,
|
||||
radius: int = 2,
|
||||
vis_thresh: float = 0.5) -> list:
|
||||
"""Self-contained track overlay (dots + short tails); no fastvideo/triton dependency."""
|
||||
from PIL import Image, ImageDraw
|
||||
T, N, _ = tracks.shape
|
||||
out = []
|
||||
for t in range(T):
|
||||
img = Image.fromarray(np.ascontiguousarray(frames[t]))
|
||||
dr = ImageDraw.Draw(img)
|
||||
t0 = max(0, t - tail)
|
||||
for i in range(N):
|
||||
col = (int(colors[i][0]), int(colors[i][1]), int(colors[i][2]))
|
||||
if t - t0 >= 1:
|
||||
dr.line([(float(p[0]), float(p[1])) for p in tracks[t0:t + 1, i]], fill=col, width=1)
|
||||
if vis[t, i] > vis_thresh:
|
||||
x, y = float(tracks[t, i, 0]), float(tracks[t, i, 1])
|
||||
dr.ellipse([x - radius, y - radius, x + radius, y + radius], fill=col)
|
||||
out.append(np.asarray(img))
|
||||
return out
|
||||
|
||||
|
||||
def _read_frames(path: str) -> np.ndarray:
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001
|
||||
import av
|
||||
c = av.open(path)
|
||||
return np.stack([f.to_ndarray(format="rgb24") for f in c.decode(video=0)])
|
||||
|
||||
|
||||
def _norm(x: np.ndarray, pct: float = 97.0) -> np.ndarray:
|
||||
lo = float(x.min())
|
||||
hi = float(np.percentile(x, pct))
|
||||
return np.zeros_like(x) if hi <= lo else np.clip((x - lo) / (hi - lo), 0.0, 1.0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- statistical scores
|
||||
def abs_disp(tracks: np.ndarray) -> np.ndarray:
|
||||
return np.sqrt(((tracks - tracks[0:1])**2).sum(-1)).max(0)
|
||||
|
||||
|
||||
def affine_residual(tracks: np.ndarray, vis: np.ndarray, iters: int = 3) -> np.ndarray:
|
||||
T, N, _ = tracks.shape
|
||||
X0 = np.concatenate([tracks[0], np.ones((N, 1), np.float32)], 1) # [N,3]
|
||||
resid = np.zeros((T, N), np.float32)
|
||||
for t in range(1, T):
|
||||
pt = tracks[t]
|
||||
w = (vis[t] > 0.5).astype(np.float64)
|
||||
if w.sum() < 6:
|
||||
w = np.ones(N)
|
||||
for _ in range(iters):
|
||||
W = X0.T * w
|
||||
A = np.linalg.lstsq((W @ X0), (W @ pt), rcond=None)[0]
|
||||
r = np.sqrt(((X0 @ A - pt)**2).sum(-1)) + 1e-6
|
||||
delta = 1.5 * np.median(r[w > 0]) + 1e-3
|
||||
w = np.where(r <= delta, 1.0, delta / r) * (vis[t] > 0.5)
|
||||
resid[t] = r
|
||||
return np.median(resid[1:], 0)
|
||||
|
||||
|
||||
def global_affine_motion(tracks: np.ndarray, vis: np.ndarray) -> float:
|
||||
"""Camera-fixedness proxy: median per-frame magnitude of the fitted GLOBAL affine
|
||||
displacement (in px), normalized by frame diagonal. ~0 => fixed camera; large => the
|
||||
whole scene translates/zooms (moving/egocentric camera). Robust (Huber) fit."""
|
||||
T, N, _ = tracks.shape
|
||||
X0 = np.concatenate([tracks[0], np.ones((N, 1), np.float32)], 1)
|
||||
mags = []
|
||||
for t in range(1, T):
|
||||
pt = tracks[t]
|
||||
w = (vis[t] > 0.5).astype(np.float64)
|
||||
if w.sum() < 6:
|
||||
w = np.ones(N)
|
||||
for _ in range(2):
|
||||
Wm = X0.T * w
|
||||
A = np.linalg.lstsq((Wm @ X0), (Wm @ pt), rcond=None)[0]
|
||||
r = np.sqrt(((X0 @ A - pt)**2).sum(-1)) + 1e-6
|
||||
delta = 1.5 * np.median(r[w > 0]) + 1e-3
|
||||
w = np.where(r <= delta, 1.0, delta / r) * (vis[t] > 0.5)
|
||||
mags.append(float(np.median(np.sqrt(((X0 @ A - tracks[0])**2).sum(-1)))))
|
||||
diag = float(np.hypot(tracks[..., 0].max(), tracks[..., 1].max()) + 1e-6)
|
||||
return round(float(np.median(mags)) / diag, 4)
|
||||
|
||||
|
||||
def displacement_matrix(tracks: np.ndarray) -> np.ndarray:
|
||||
"""[N, 2T] per-point displacement relative to frame 0."""
|
||||
T, N, _ = tracks.shape
|
||||
return (tracks - tracks[0:1]).transpose(1, 0, 2).reshape(N, 2 * T)
|
||||
|
||||
|
||||
def lowrank_residual(tracks: np.ndarray, rank: int = 3) -> np.ndarray:
|
||||
D = displacement_matrix(tracks)
|
||||
Dc = D - D.mean(0, keepdims=True)
|
||||
U, S, Vt = np.linalg.svd(Dc, full_matrices=False)
|
||||
r = min(rank, S.shape[0])
|
||||
return np.sqrt(((Dc - (U[:, :r] * S[:r]) @ Vt[:r])**2).sum(-1))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- embeddings + clustering
|
||||
def traj_embedding(tracks: np.ndarray, n_pca: int = 16) -> np.ndarray:
|
||||
from sklearn.decomposition import PCA
|
||||
D = displacement_matrix(tracks)
|
||||
D = (D - D.mean(0)) / (D.std(0) + 1e-6)
|
||||
d = int(min(n_pca, min(D.shape) - 1))
|
||||
return PCA(n_components=d, random_state=0).fit_transform(D).astype(np.float32)
|
||||
|
||||
|
||||
def motion_cluster(emb: np.ndarray):
|
||||
"""#1: HDBSCAN on trajectory embeddings -> labels, outlierness score (1 - membership prob)."""
|
||||
from sklearn.cluster import HDBSCAN
|
||||
mcs = max(10, emb.shape[0] // 100)
|
||||
cl = HDBSCAN(min_cluster_size=mcs, min_samples=5).fit(emb)
|
||||
return cl.labels_, (1.0 - cl.probabilities_).astype(np.float32)
|
||||
|
||||
|
||||
def subspace_cluster(tracks: np.ndarray, emb: np.ndarray, k: int):
|
||||
"""#2: Spectral motion segmentation into k groups; per-cluster low-rank residual."""
|
||||
from sklearn.cluster import SpectralClustering
|
||||
k = int(np.clip(k, 2, 8))
|
||||
lab = SpectralClustering(n_clusters=k,
|
||||
affinity="nearest_neighbors",
|
||||
n_neighbors=15,
|
||||
assign_labels="cluster_qr",
|
||||
random_state=0).fit_predict(emb)
|
||||
D = displacement_matrix(tracks)
|
||||
resid = np.zeros(lab.shape[0], np.float32)
|
||||
for c in np.unique(lab):
|
||||
idx = np.where(lab == c)[0]
|
||||
Z = D[idx] - D[idx].mean(0, keepdims=True)
|
||||
r = min(4, min(Z.shape) - 1)
|
||||
if r > 0:
|
||||
U, S, Vt = np.linalg.svd(Z, full_matrices=False)
|
||||
resid[idx] = np.sqrt(((Z - (U[:, :r] * S[:r]) @ Vt[:r])**2).sum(-1))
|
||||
return lab, resid
|
||||
|
||||
|
||||
def dino_point_features(frame0: np.ndarray, xy: np.ndarray, model, device: str) -> np.ndarray:
|
||||
"""#3: sample DINOv2 patch tokens at each frame-0 point (bilinear). Returns [N,C] L2-normed."""
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
H, W = frame0.shape[0], frame0.shape[1]
|
||||
gw, gh = max(1, round(W / 14)), max(1, round(H / 14))
|
||||
rw, rh = gw * 14, gh * 14
|
||||
img = torch.from_numpy(frame0).float().permute(2, 0, 1)[None] / 255.0
|
||||
img = F.interpolate(img, size=(rh, rw), mode="bilinear", align_corners=False)
|
||||
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
|
||||
std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
|
||||
img = ((img - mean) / std).to(device)
|
||||
with torch.no_grad():
|
||||
tok = model.forward_features(img)["x_norm_patchtokens"][0] # [gh*gw, C]
|
||||
C = tok.shape[-1]
|
||||
grid = tok.reshape(gh, gw, C).permute(2, 0, 1)[None] # [1,C,gh,gw]
|
||||
# normalized sample coords in [-1,1] from pixel xy
|
||||
gx = (xy[:, 0] / W) * 2 - 1
|
||||
gy = (xy[:, 1] / H) * 2 - 1
|
||||
samp = torch.from_numpy(np.stack([gx, gy], -1)).float().view(1, -1, 1, 2).to(device)
|
||||
feat = F.grid_sample(grid, samp, mode="bilinear", align_corners=False)[0, :, :, 0].T # [N,C]
|
||||
feat = feat.cpu().numpy()
|
||||
return (feat / (np.linalg.norm(feat, axis=1, keepdims=True) + 1e-6)).astype(np.float32)
|
||||
|
||||
|
||||
def dino_motion_cluster(dino_feat: np.ndarray, traj_emb: np.ndarray):
|
||||
"""Cluster fused (semantic ⊕ motion) embedding -> labels, uniqueness (1 - membership prob)."""
|
||||
from sklearn.cluster import HDBSCAN
|
||||
from sklearn.decomposition import PCA
|
||||
dz = PCA(n_components=int(min(16, min(dino_feat.shape) - 1)), random_state=0).fit_transform(dino_feat)
|
||||
dz = (dz - dz.mean(0)) / (dz.std(0) + 1e-6)
|
||||
tz = (traj_emb - traj_emb.mean(0)) / (traj_emb.std(0) + 1e-6)
|
||||
fused = np.concatenate([dz, tz], 1)
|
||||
cl = HDBSCAN(min_cluster_size=max(10, fused.shape[0] // 100), min_samples=5).fit(fused)
|
||||
return cl.labels_, (1.0 - cl.probabilities_).astype(np.float32)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- subset selection (#4)
|
||||
def select_fps(emb: np.ndarray, K: int, seed: int) -> np.ndarray:
|
||||
sel = [int(seed)]
|
||||
d = np.linalg.norm(emb - emb[seed], axis=1)
|
||||
for _ in range(K - 1):
|
||||
i = int(np.argmax(d))
|
||||
sel.append(i)
|
||||
d = np.minimum(d, np.linalg.norm(emb - emb[i], axis=1))
|
||||
return np.array(sorted(set(sel)))
|
||||
|
||||
|
||||
def select_dpp(quality: np.ndarray, emb: np.ndarray, K: int) -> np.ndarray:
|
||||
"""Greedy MAP for a DPP with L = diag(q) S diag(q), S a Gaussian similarity kernel."""
|
||||
from scipy.spatial.distance import cdist
|
||||
q = _norm(quality) + 0.05
|
||||
d2 = cdist(emb, emb, "sqeuclidean")
|
||||
S = np.exp(-d2 / (np.median(d2) + 1e-6))
|
||||
L = (q[:, None] * q[None, :]) * S
|
||||
N = L.shape[0]
|
||||
cis = np.zeros((K, N))
|
||||
di2s = np.diag(L).copy()
|
||||
j = int(np.argmax(di2s))
|
||||
sel = [j]
|
||||
for it in range(1, K):
|
||||
k = it - 1
|
||||
ei = (L[j] - cis[:k, j] @ cis[:k]) / np.sqrt(max(di2s[j], 1e-12))
|
||||
cis[k] = ei
|
||||
di2s = di2s - ei**2
|
||||
di2s[sel] = -np.inf
|
||||
j = int(np.argmax(di2s))
|
||||
if di2s[j] <= 1e-10:
|
||||
break
|
||||
sel.append(j)
|
||||
return np.array(sorted(sel))
|
||||
|
||||
|
||||
def select_greedy_recon(xy0: np.ndarray, D: np.ndarray, K: int) -> np.ndarray:
|
||||
"""Add the point whose displacement is least predictable by RBF-interpolating the
|
||||
already-selected points over frame-0 positions -> targets motion discontinuities."""
|
||||
from scipy.spatial.distance import cdist
|
||||
sel = [int(np.argmax(np.linalg.norm(D - D.mean(0), axis=1)))]
|
||||
for _ in range(K - 1):
|
||||
S = np.array(sel)
|
||||
d2 = cdist(xy0, xy0[S], "sqeuclidean")
|
||||
W = np.exp(-d2 / (np.median(d2) + 1e-6))
|
||||
W /= (W.sum(1, keepdims=True) + 1e-9)
|
||||
err = np.linalg.norm(D - W @ D[S], axis=1)
|
||||
err[S] = -1
|
||||
sel.append(int(np.argmax(err)))
|
||||
return np.array(sorted(sel))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- render
|
||||
def _spearman(a, b) -> float:
|
||||
from scipy.stats import spearmanr
|
||||
return round(float(spearmanr(a, b).correlation), 3)
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import imageio.v2 as imageio
|
||||
|
||||
dino_model, dino_dev = None, "cpu"
|
||||
if not args.no_dino:
|
||||
try:
|
||||
import torch
|
||||
dino_dev = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dino_model = torch.hub.load("facebookresearch/dinov2", "dinov2_vits14").to(dino_dev).eval()
|
||||
print(f"[info] DINOv2 loaded on {dino_dev}", flush=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[info] DINOv2 unavailable ({type(e).__name__}: {e}); skipping semantic panels", flush=True)
|
||||
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
turbo = matplotlib.colormaps["turbo"]
|
||||
entries = []
|
||||
for k, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
vpath = str(data / "videos" / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = _read_frames(vpath)
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]]
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]]
|
||||
N = tracks.shape[1]
|
||||
x, y = tracks[0, :, 0], tracks[0, :, 1]
|
||||
D = displacement_matrix(tracks)
|
||||
emb = traj_embedding(tracks)
|
||||
|
||||
# scalar scores
|
||||
scores = {
|
||||
"abs_disp": abs_disp(tracks),
|
||||
"affine_residual": affine_residual(tracks, vis),
|
||||
"lowrank_residual": lowrank_residual(tracks),
|
||||
}
|
||||
m_lab, m_out = motion_cluster(emb)
|
||||
scores["motion_outlier"] = m_out
|
||||
k_est = len(np.unique(m_lab[m_lab >= 0])) or 4
|
||||
s_lab, s_res = subspace_cluster(tracks, emb, k_est)
|
||||
scores["subspace_resid"] = s_res
|
||||
d_lab = None
|
||||
if dino_model is not None:
|
||||
try:
|
||||
dfeat = dino_point_features(frames[0], tracks[0], dino_model, dino_dev)
|
||||
d_lab, d_uniq = dino_motion_cluster(dfeat, emb)
|
||||
scores["dino_uniqueness"] = d_uniq
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[info] dino features failed for {stem}: {e}", flush=True)
|
||||
|
||||
# subset selections (#4)
|
||||
seed = int(np.argmax(scores["lowrank_residual"]))
|
||||
subsets = {
|
||||
"dpp": select_dpp(scores["lowrank_residual"], emb, K_SELECT),
|
||||
"fps": select_fps(emb, K_SELECT, seed),
|
||||
"greedy_recon": select_greedy_recon(tracks[0], D, K_SELECT),
|
||||
}
|
||||
|
||||
# ---- figure: row1 scalar scores, row2 cluster maps + subsets
|
||||
label_panels = [("motion_clusters", m_lab), ("subspace_clusters", s_lab)]
|
||||
if d_lab is not None:
|
||||
label_panels.append(("dino_clusters", d_lab))
|
||||
panels = ([("score", n, s)
|
||||
for n, s in scores.items()] + [("labels", n, lab)
|
||||
for n, lab in label_panels] + [("subset", n, idx)
|
||||
for n, idx in subsets.items()])
|
||||
ncol = 5
|
||||
nrow = int(np.ceil(len(panels) / ncol))
|
||||
fig, axes = plt.subplots(nrow, ncol, figsize=(4.4 * ncol, 3.1 * nrow))
|
||||
axes = np.array(axes).reshape(-1)
|
||||
for ax in axes:
|
||||
ax.axis("off")
|
||||
for ax, (kind, name, val) in zip(axes, panels, strict=False):
|
||||
ax.imshow(frames[0])
|
||||
if kind == "score":
|
||||
ax.scatter(x, y, c=_norm(val), cmap="turbo", s=10, vmin=0, vmax=1)
|
||||
ax.set_title(f"{name}", fontsize=10)
|
||||
elif kind == "labels":
|
||||
noise = val < 0
|
||||
ax.scatter(x[noise], y[noise], c="0.5", s=6)
|
||||
ax.scatter(x[~noise], y[~noise], c=val[~noise], cmap="tab20", s=10)
|
||||
ax.set_title(f"{name} ({len(np.unique(val[val >= 0]))} groups)", fontsize=10)
|
||||
else: # subset
|
||||
ax.scatter(x, y, c="0.4", s=3)
|
||||
ax.scatter(x[val], y[val], c="red", s=26, edgecolors="white", linewidths=0.4)
|
||||
ax.set_title(f"{name} (K={len(val)})", fontsize=10)
|
||||
fig.suptitle(f"{stem} | {(it['cap'][0] if isinstance(it.get('cap'), list) else '')[:90]}", fontsize=11)
|
||||
fig.tight_layout()
|
||||
heat = f"{stem}_scores.png"
|
||||
fig.savefig(str(out / heat), dpi=80, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
|
||||
# ---- overlay colored by lowrank informativeness
|
||||
lr = _norm(scores["lowrank_residual"])
|
||||
pcols = (np.array([turbo(v)[:3] for v in lr]) * 255).astype(np.uint8)
|
||||
G = int(round(N**0.5))
|
||||
sel = (np.arange(N).reshape(G, G)[::max(1, G // 25), ::max(1, G // 25)].reshape(-1) if G *
|
||||
G == N else np.arange(0, N, max(1, N // 625)))
|
||||
ov = _draw_overlay(frames, tracks[:, sel].copy(), vis[:, sel], pcols[sel], 14, 2, 0.5)
|
||||
vid = f"{stem}_lowrank.mp4"
|
||||
imageio.mimsave(str(out / vid), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
|
||||
# ---- stats
|
||||
names = list(scores.keys())
|
||||
rho = {f"{a}|{b}": _spearman(scores[a], scores[b]) for i, a in enumerate(names) for b in names[i + 1:]}
|
||||
lr_raw = scores["lowrank_residual"]
|
||||
lo_info = lr_raw <= np.percentile(lr_raw, 50)
|
||||
|
||||
def _cov_generic(idx: Any, m_lab: Any = m_lab, lo_info: Any = lo_info) -> dict:
|
||||
grp = m_lab[idx]
|
||||
ngrp = len(np.unique(grp[grp >= 0]))
|
||||
tot = max(len(np.unique(m_lab[m_lab >= 0])), 1)
|
||||
return {"coverage_of_motion_groups": f"{ngrp}/{tot}", "frac_generic": round(float(lo_info[idx].mean()), 3)}
|
||||
|
||||
stats = {
|
||||
"id": stem,
|
||||
"n_points": int(N),
|
||||
"camera_motion": global_affine_motion(tracks, vis), # ~0 = fixed cam, high = moving
|
||||
"n_motion_groups": int(len(np.unique(m_lab[m_lab >= 0]))),
|
||||
"noise_frac": round(float((m_lab < 0).mean()), 3),
|
||||
"spearman": rho,
|
||||
"subset_quality": {
|
||||
name: _cov_generic(idx)
|
||||
for name, idx in subsets.items()
|
||||
},
|
||||
}
|
||||
entries.append({"id": stem, "heat": heat, "video": vid, "stats": stats})
|
||||
print(
|
||||
f"[info] [{k}/{len(items)}] {stem}: cam_motion={stats['camera_motion']}, "
|
||||
f"{stats['n_motion_groups']} motion groups, noise={stats['noise_frac']}, "
|
||||
f"rho(abs|lowrank)={rho.get('abs_disp|lowrank_residual')}",
|
||||
flush=True)
|
||||
|
||||
(out / "manifest.json").write_text(json.dumps(entries, indent=2))
|
||||
print(f"[info] rendered {len(entries)} clips -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
entries = json.loads((viz / "manifest.json").read_text())
|
||||
gal = [(str(viz / e["heat"]), e["id"]) for e in entries]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = entries[evt.index]
|
||||
return str(viz / e["heat"]), str(viz / e["video"]), json.dumps(e["stats"], indent=2)
|
||||
|
||||
with gr.Blocks(title="Track informativeness bake-off") as demo:
|
||||
gr.Markdown(
|
||||
"### Which point tracks are useful? Statistics vs clustering vs embeddings vs subset selection\n"
|
||||
"**Row 1 (scores, bright = informative):** `abs_disp` raw motion · `affine_residual` after "
|
||||
"removing camera affine · `lowrank_residual` unique vs all (one subspace) · `motion_outlier` "
|
||||
"HDBSCAN density outlier · `subspace_resid` per-object lowrank · `dino_uniqueness` semantic⊕motion. "
|
||||
"**Row 2:** `motion_clusters` / `subspace_clusters` / `dino_clusters` (colored groups, gray = noise) "
|
||||
"and the diverse K-subsets `dpp` / `fps` / `greedy_recon` (red = chosen). "
|
||||
"Click a clip -> panels + overlay colored by low-rank informativeness. `stats.subset_quality` shows "
|
||||
"how well each subset covers the motion groups and how many chosen points are generic.")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gal, columns=2, height=640, label="clips (click)")
|
||||
with gr.Column():
|
||||
heat = gr.Image(label="method panels")
|
||||
vid = gr.Video(label="overlay colored by low-rank informativeness")
|
||||
meta = gr.Code(label="stats", language="json")
|
||||
g.select(show, None, [heat, vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--no-dino", action="store_true", help="skip the DINOv2 semantic panels")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7883)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/usr/bin/env bash
|
||||
# Overlap tracking with the slow self-healing download: idempotent multi-node passes
|
||||
# over the growing videos/ dir until download done (98 parts) AND ~all tracked.
|
||||
# MUST run inside srun on a compute node (login-node `sleep` gets killed by the harness).
|
||||
# Waits for any in-flight tracking pass first, so it's safe to relaunch (no double-submit).
|
||||
set +e
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
cd "$WORK/FastVideo"
|
||||
DC(){ ls "$WORK"/openvid_1m/_extracted/*.done 2>/dev/null | wc -l; }
|
||||
RUNNING(){ squeue -u vlm-s4duan -h -n cotracker_dp -o %i 2>/dev/null; }
|
||||
NODES=${NODES:-8}
|
||||
while true; do
|
||||
while [ -n "$(RUNNING)" ]; do sleep 90; done # let in-flight pass finish
|
||||
ls "$WORK"/openvid_1m/videos/*.mp4 > "$WORK"/openvid_1m/videos_all.txt 2>/dev/null
|
||||
nv=$(wc -l < "$WORK"/openvid_1m/videos_all.txt 2>/dev/null); nv=${nv:-0}
|
||||
nt=$(ls "$WORK"/openvid_1m/tracks/*.npz 2>/dev/null | wc -l)
|
||||
echo "[$(date +%H:%M)] parts=$(DC)/98 videos=$nv tracks=$nt"
|
||||
if [ "$(DC)" -ge 98 ] && [ "$nt" -ge "$((nv - nv/40 - 200))" ]; then
|
||||
echo "TRACK_LOOP_DONE videos=$nv tracks=$nt"; break; fi
|
||||
if [ "$nv" -le "$((nt + 300))" ]; then echo " little new; wait 300s"; sleep 300; continue; fi
|
||||
jid=$(CLIPS_DIR="$WORK"/openvid_1m/clips VIDEO_LIST="$WORK"/openvid_1m/videos_all.txt \
|
||||
OUT_DIR="$WORK"/openvid_1m/tracks NODES=$NODES PROCS_PER_GPU=2 FPS=24 NUM_FRAMES=121 GRID=50 \
|
||||
bash data_pipeline/run_tracks_slurm.sh 2>&1 | grep -oE "[0-9]+$" | tail -1)
|
||||
echo " submitted track job $jid over $nv videos"
|
||||
[ -z "$jid" ] && sleep 180
|
||||
done
|
||||
@@ -0,0 +1,117 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-checkpoint *swap sensitivity* (one process per checkpoint).
|
||||
|
||||
Linearized, generation-free version of the controllability EPE test: at a fixed
|
||||
noise + timestep, measure how much the predicted velocity ``v`` moves when we swap
|
||||
ONLY one conditioner and hold the rest fixed:
|
||||
|
||||
S_track = ||v(swap tracks) - v(base)|| / ||v(base)||
|
||||
S_ff = ||v(swap first-frame) - v(base)|| / ||v(base)||
|
||||
S_text = ||v(zero text) - v(base)|| / ||v(base)||
|
||||
|
||||
Averaged over a few timesteps and the given clips. This captures the *backbone's*
|
||||
responsiveness to each input (unlike weight-norm or input-projection magnitude,
|
||||
which only see the input side). Expectation for the control50 run: S_track high at
|
||||
step 500, decaying by 6000 (control prior forgotten), while S_ff stays high.
|
||||
|
||||
One checkpoint per process (loads via the framework loader, so weights/DTensor are
|
||||
handled correctly). Launch once per --step across GPUs.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def _v(model: Any, latents: torch.Tensor, ff: torch.Tensor, tp: torch.Tensor, tv: torch.Tensor, txt: torch.Tensor,
|
||||
mask: torch.Tensor, img: torch.Tensor, ts: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
with torch.no_grad(), torch.autocast(model.device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=ts, attn_metadata=None):
|
||||
v = model.transformer(hidden_states=model_in,
|
||||
encoder_hidden_states=txt,
|
||||
encoder_attention_mask=mask,
|
||||
timestep=ts,
|
||||
encoder_hidden_states_image=img,
|
||||
track_points=tp,
|
||||
track_visibility=tv,
|
||||
return_dict=False)
|
||||
return v.float()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True, help="single export dir (export_<step>)")
|
||||
p.add_argument("--step", type=int, default=-1, help="label only")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data",
|
||||
default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset")
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--timestep-idx", type=int, nargs="+", default=[2, 15, 27])
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
args = p.parse_args()
|
||||
|
||||
dtype = torch.bfloat16
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
device = model.device
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(30, device=device)
|
||||
tlist = [sched.timesteps[i].reshape(1).to(device, dtype) for i in args.timestep_idx]
|
||||
|
||||
pairs = []
|
||||
for k, ci in enumerate(args.clips):
|
||||
s = samples[k]
|
||||
sw = samples[(k + 1) % len(samples)]
|
||||
_, _, T, H, W = s["first_frame_latent"].shape
|
||||
g = torch.Generator(device="cpu").manual_seed(args.seed + ci)
|
||||
latents = torch.randn((1, 16, T, H, W), generator=g, dtype=torch.float32).to(device)
|
||||
|
||||
def d(x: torch.Tensor) -> torch.Tensor:
|
||||
return x.to(device, dtype)
|
||||
|
||||
pairs.append(
|
||||
dict(lat=latents,
|
||||
ff=d(s["first_frame_latent"]),
|
||||
ff_sw=d(sw["first_frame_latent"]),
|
||||
tp=d(s["track_points"]),
|
||||
tv=d(s["track_visibility"]),
|
||||
tp_sw=d(sw["track_points"]),
|
||||
tv_sw=d(sw["track_visibility"]),
|
||||
txt=d(s["text_embedding"]),
|
||||
mask=d(s["text_attention_mask"]),
|
||||
img=d(s["clip_feature"])))
|
||||
|
||||
st = sf = sx = 0.0
|
||||
n = 0
|
||||
for pr in pairs:
|
||||
for ts in tlist:
|
||||
vb = _v(model, pr["lat"], pr["ff"], pr["tp"], pr["tv"], pr["txt"], pr["mask"], pr["img"], ts, dtype)
|
||||
nb = vb.norm() + 1e-6
|
||||
vt = _v(model, pr["lat"], pr["ff"], pr["tp_sw"], pr["tv_sw"], pr["txt"], pr["mask"], pr["img"], ts, dtype)
|
||||
vf = _v(model, pr["lat"], pr["ff_sw"], pr["tp"], pr["tv"], pr["txt"], pr["mask"], pr["img"], ts, dtype)
|
||||
vx = _v(model, pr["lat"], pr["ff"], pr["tp"], pr["tv"], torch.zeros_like(pr["txt"]), pr["mask"], pr["img"],
|
||||
ts, dtype)
|
||||
st += float((vt - vb).norm() / nb)
|
||||
sf += float((vf - vb).norm() / nb)
|
||||
sx += float((vx - vb).norm() / nb)
|
||||
n += 1
|
||||
print(f"RESULT step={args.step} S_track={st/n:.4f} S_ff={sf/n:.4f} S_text={sx/n:.4f}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,176 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Standalone TrackWan inference: generate a video from (first frame + text + tracks).
|
||||
|
||||
Loads a checkpoint exported with ``dcp_to_diffusers`` into the *same*
|
||||
``WanTrackModel`` wrapper used for training (so train/inference parity is exact),
|
||||
then runs a flow-matching denoise loop feeding the point-track control. Used by
|
||||
the controllability tests and the Gradio app.
|
||||
|
||||
Backend only -- no track authoring here (see ``synthetic_tracks.py``) and no
|
||||
metrics (see ``fastvideo/eval/metrics/motion/cotracker_epe``).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def _ensure_dist_env() -> None:
|
||||
for k, v in [("RANK", "0"), ("LOCAL_RANK", "0"), ("WORLD_SIZE", "1"), ("MASTER_ADDR", "127.0.0.1"),
|
||||
("MASTER_PORT", "29600")]:
|
||||
os.environ.setdefault(k, v)
|
||||
|
||||
|
||||
def load_trackwan(export_dir: str, yaml_path: str) -> tuple[Any, Any]:
|
||||
"""Build a ``WanTrackModel`` with the exported (trained) weights on 1 GPU."""
|
||||
_ensure_dist_env()
|
||||
from fastvideo.distributed import (
|
||||
maybe_init_distributed_environment_and_model_parallel, )
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
from fastvideo.train.models.wantrack.wantrack import WanTrackModel
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
|
||||
cfg = load_run_config(yaml_path)
|
||||
tc = cfg.training
|
||||
tc.distributed.tp_size = 1
|
||||
tc.distributed.sp_size = 1
|
||||
tc.distributed.num_gpus = 1
|
||||
tc.distributed.hsdp_replicate_dim = 1
|
||||
tc.distributed.hsdp_shard_dim = 1
|
||||
tc.model_path = export_dir # vae + components loaded from the export
|
||||
|
||||
model = WanTrackModel(
|
||||
init_from=export_dir,
|
||||
training_config=tc,
|
||||
trainable=False,
|
||||
flow_shift=float(tc.pipeline_config.flow_shift),
|
||||
enable_gradient_checkpointing_type=None,
|
||||
)
|
||||
model.init_preprocessors(tc)
|
||||
model.transformer.eval()
|
||||
return model, tc
|
||||
|
||||
|
||||
def load_conditioning_from_parquet(data_path: str, indices: list[int], text_len: int) -> list[dict[str, Any]]:
|
||||
"""Pull (first_frame_latent, text, vae_latent, GT tracks, caption) for clips."""
|
||||
import glob
|
||||
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
|
||||
files = sorted(glob.glob(os.path.join(data_path, "**", "*.parquet"), recursive=True))
|
||||
rows: list[dict[str, Any]] = []
|
||||
for f in files:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
sel = [rows[i] for i in indices]
|
||||
batch = collate_rows_from_parquet_schema(sel,
|
||||
pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len),
|
||||
cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in sel]
|
||||
out = []
|
||||
for i in range(len(sel)):
|
||||
out.append({
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(), # [1,SeqLen,Dim] CLIP frame-0
|
||||
"track_points": batch["track_points"][i:i + 1].clone(), # [1,T,N,2] normalized
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(), # [1,T,N]
|
||||
"caption": str(infos[i].get("caption", "")),
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(model: Any,
|
||||
*,
|
||||
first_frame_latent: torch.Tensor,
|
||||
text_embedding: torch.Tensor,
|
||||
text_attention_mask: torch.Tensor,
|
||||
track_points: torch.Tensor | None,
|
||||
track_visibility: torch.Tensor | None,
|
||||
clip_feature: torch.Tensor | None = None,
|
||||
num_steps: int = 30,
|
||||
seed: int = 0,
|
||||
guidance_scale: float = 1.0) -> torch.Tensor:
|
||||
"""Denoise from noise -> normalized latents [1,16,T,H,W]. Tracks may be None.
|
||||
|
||||
``guidance_scale`` > 1 = MOTION classifier-free guidance:
|
||||
v = v(no-tracks) + s*(v(tracks) - v(no-tracks)). The no-tracks branch is the
|
||||
base first-frame+text I2V prediction (track channels zeroed), so it works on any
|
||||
checkpoint (no training-time motion dropout needed)."""
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
ff = first_frame_latent.to(device, dtype) # [1,16,T,H,W] (already normalized)
|
||||
cond20 = model._build_i2v_cond_concat(ff) # [1,20,T,H,W]
|
||||
txt = text_embedding.to(device, dtype)
|
||||
mask = text_attention_mask.to(device, dtype)
|
||||
tp = track_points.to(device, dtype) if track_points is not None else None
|
||||
tv = track_visibility.to(device, dtype) if track_visibility is not None else None
|
||||
img = clip_feature.to(device, dtype) if clip_feature is not None else None
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(int(num_steps), device=device)
|
||||
use_cfg = guidance_scale != 1.0 and tp is not None
|
||||
for t in sched.timesteps:
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1) # [1,36,T,H,W]
|
||||
ts = t.reshape(1).to(device, dtype)
|
||||
|
||||
def _fwd(tpp: torch.Tensor | None,
|
||||
tvv: torch.Tensor | None,
|
||||
mi: torch.Tensor = model_in,
|
||||
tsv: torch.Tensor = ts) -> torch.Tensor:
|
||||
with torch.autocast(device.type, dtype=dtype), set_forward_context(current_timestep=tsv,
|
||||
attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi,
|
||||
encoder_hidden_states=txt,
|
||||
encoder_attention_mask=mask,
|
||||
timestep=tsv,
|
||||
encoder_hidden_states_image=img,
|
||||
track_points=tpp,
|
||||
track_visibility=tvv,
|
||||
return_dict=False)
|
||||
|
||||
if use_cfg:
|
||||
v_cond = _fwd(tp, tv).float()
|
||||
v_uncond = _fwd(None, None).float() # base I2V (track channels -> 0)
|
||||
v = v_uncond + guidance_scale * (v_cond - v_uncond)
|
||||
else:
|
||||
v = _fwd(tp, tv).float()
|
||||
latents = sched.step(v, t, latents.float(), return_dict=False)[0]
|
||||
return latents
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_to_pixels(model: Any, latents: torch.Tensor) -> np.ndarray:
|
||||
"""Normalized latents [1,16,T,H,W] -> uint8 frames [T,H,W,3]."""
|
||||
px = model.decode_latents(latents.permute(0, 2, 1, 3, 4))[0] # [3,T,H,W] in [0,1]
|
||||
video = (px.clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
return np.transpose(video, (1, 2, 3, 0)) # [T,H,W,3]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_reference(model: Any, vae_latent: torch.Tensor) -> np.ndarray:
|
||||
"""Decode the raw GT vae_latent (needs normalize first) -> uint8 [T,H,W,3]."""
|
||||
from fastvideo.training.training_utils import normalize_dit_input
|
||||
raw = vae_latent.to(model.device, torch.bfloat16)
|
||||
norm = normalize_dit_input("wan", raw, model.vae)
|
||||
px = model.decode_latents(norm.permute(0, 2, 1, 3, 4))[0]
|
||||
video = (px.clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
return np.transpose(video, (1, 2, 3, 0))
|
||||
@@ -0,0 +1,325 @@
|
||||
#!/usr/bin/env python3
|
||||
"""2-track validation ablation.
|
||||
|
||||
For each val sample, keep exactly 2 tracks (one per top-motion object) and zero
|
||||
the rest. This is the extreme sparse-inference case: what the user actually gets
|
||||
when they draw 2 traces at inference. Compares checkpoints side-by-side.
|
||||
|
||||
Uses paper CFG (wt=3.0, wm=1.5) unless overridden. Runs `_sample`-style denoise
|
||||
matching track_validation.py exactly.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/two_track_val.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800 \
|
||||
--yaml examples/train/scenario/worldmodel/finetune_wantrack_synth_stage2_paperLR.yaml \
|
||||
--val-parquet /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset \
|
||||
--wandb-run-name twotrack_mb4800
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__)))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
from fastvideo.forward_context import set_forward_context # noqa: E402
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import ( # noqa: E402
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
|
||||
def load_val_samples(parquet_dir: str, text_len: int) -> list:
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
fs = sorted(glob.glob(os.path.join(parquet_dir, "**", "*.parquet"), recursive=True))
|
||||
if not fs:
|
||||
raise FileNotFoundError(parquet_dir)
|
||||
rows = []
|
||||
for f in fs:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
batch = collate_rows_from_parquet_schema(rows, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in rows]
|
||||
return [{
|
||||
"id": str(infos[i].get("id", f"clip{i:03d}")) if i < len(infos) else f"clip{i:03d}",
|
||||
"file_name": rows[i].get("file_name", ""),
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i:i + 1].clone(), # [1,T,N,2] normalised
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(),
|
||||
"object_ids": batch["object_ids"][i:i + 1].clone() if "object_ids" in batch else None,
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
} for i in range(len(rows))]
|
||||
|
||||
|
||||
def pick_two_tracks(sample: dict, seed: int, num_bg: int = 0,
|
||||
num_fg: int = 2) -> tuple[torch.Tensor, torch.Tensor, list]:
|
||||
"""Return (tp_2, tv_2, info) — ``num_fg`` foreground tracks + ``num_bg`` bg anchor tracks.
|
||||
|
||||
Foreground picks require the track to STAY IN FRAME [0,1] for all 121 frames — no
|
||||
point in showing traces that fly off the canvas.
|
||||
"""
|
||||
tp = sample["track_points"].float() # [1,T,N,2]
|
||||
tv = sample["track_visibility"].float() # [1,T,N]
|
||||
oid = sample["object_ids"] # [1,N]
|
||||
if oid is None:
|
||||
raise ValueError("no object_ids in val sample")
|
||||
oid = oid[0].to(torch.int64) # [N]
|
||||
|
||||
tp0 = tp[0] # [T, N, 2]
|
||||
N = tp0.shape[1]
|
||||
motion_per_track = torch.linalg.norm(tp0[-1] - tp0[0], dim=-1) # [N]
|
||||
|
||||
# Which tracks stay fully within [0,1] on both x and y across ALL frames?
|
||||
in_frame = ((tp0 >= 0.0) & (tp0 <= 1.0)).all(dim=-1).all(dim=0) # [N]
|
||||
|
||||
# Rank object_ids by max in-frame within-object track motion.
|
||||
fg_oids = torch.unique(oid[oid >= 0])
|
||||
obj_best = [] # (oid, best_track_idx, best_motion)
|
||||
for o in fg_oids.tolist():
|
||||
idxs = ((oid == o) & in_frame & (tv[0, 0] > 0.5)).nonzero(as_tuple=True)[0]
|
||||
if idxs.numel() == 0:
|
||||
continue
|
||||
m = motion_per_track[idxs]
|
||||
best = idxs[int(torch.argmax(m).item())]
|
||||
obj_best.append((o, int(best.item()), float(m.max().item())))
|
||||
obj_best.sort(key=lambda x: -x[2])
|
||||
picked_tracks = [idx for _, idx, _ in obj_best[:num_fg]]
|
||||
info = [{"oid": int(o), "motion_max": mm, "track_idx": pk}
|
||||
for (o, pk, mm) in obj_best[:num_fg]]
|
||||
rng = np.random.RandomState(seed)
|
||||
vis0 = tv[0, 0] # [N] visibility at frame 0
|
||||
|
||||
# Add num_bg background anchor tracks — near-static (bottom of motion), visible at frame 0.
|
||||
bg_picked = []
|
||||
if num_bg > 0:
|
||||
bg_pool = ((oid == -1) & (vis0 > 0.5)).nonzero(as_tuple=True)[0]
|
||||
if bg_pool.numel() > 0:
|
||||
# Sort bg tracks by motion (ascending), take lowest-motion ones and pick uniformly
|
||||
bg_motion = motion_per_track[bg_pool]
|
||||
n_pool = int(min(bg_pool.numel(), max(num_bg * 20, num_bg)))
|
||||
top_idxs = torch.argsort(bg_motion)[:n_pool]
|
||||
candidate = bg_pool[top_idxs]
|
||||
if candidate.numel() >= num_bg:
|
||||
sel = rng.choice(candidate.numel(), num_bg, replace=False)
|
||||
bg_picked = [int(candidate[k].item()) for k in sel]
|
||||
else:
|
||||
bg_picked = [int(k.item()) for k in candidate]
|
||||
for k in bg_picked:
|
||||
info.append({"oid": -1, "motion_max": float(motion_per_track[k].item()), "track_idx": k})
|
||||
|
||||
# Build zeroed visibility: only picked tracks visible.
|
||||
tv_new = torch.zeros_like(tv)
|
||||
for k in picked_tracks + bg_picked:
|
||||
tv_new[0, :, k] = tv[0, :, k]
|
||||
|
||||
return tp, tv_new, info
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_denoise(model, sample: dict, tp: torch.Tensor, tv: torch.Tensor,
|
||||
*, w_text: float, w_motion: float, num_steps: int, seed: int) -> np.ndarray:
|
||||
"""Reproduce track_validation._sample denoise EXACTLY."""
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
|
||||
ff = sample["first_frame_latent"].to(device, dtype)
|
||||
txt = sample["text_embedding"].to(device, dtype)
|
||||
mask = sample["text_attention_mask"].to(device, dtype)
|
||||
clip = sample["clip_feature"].to(device, dtype)
|
||||
tp_d = tp.to(device, dtype)
|
||||
tv_d = tv.to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(int(num_steps), device=device)
|
||||
|
||||
cfg_on = (w_text != 1.0) or (w_motion != 1.0)
|
||||
txt_null = torch.zeros_like(txt) if cfg_on else None
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, ts_):
|
||||
with torch.autocast(device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=ts_, attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=ts_,
|
||||
encoder_hidden_states_image=clip,
|
||||
track_points=tp_e, track_visibility=tv_e,
|
||||
return_dict=False)
|
||||
|
||||
for tt in sched.timesteps:
|
||||
mi = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
ts_ = tt.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp_d, tv_d, mi, ts_)
|
||||
if not cfg_on:
|
||||
v = v_full
|
||||
elif w_text != 1.0 and w_motion == 1.0:
|
||||
v_no_text = _fwd(txt_null, tp_d, tv_d, mi, ts_)
|
||||
v = v_no_text + w_text * (v_full - v_no_text)
|
||||
elif w_text == 1.0 and w_motion != 1.0:
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
v = v_no_motion + w_motion * (v_full - v_no_motion)
|
||||
else:
|
||||
v_no_text = _fwd(txt_null, tp_d, tv_d, mi, ts_)
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
alpha = w_text / max(w_text + w_motion, 1e-6)
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + w_text * (v_full - v_no_text) + w_motion * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), tt, latents.float(), return_dict=False)[0]
|
||||
|
||||
return twi.decode_to_pixels(model, latents)
|
||||
|
||||
|
||||
def draw_track_overlay(video: np.ndarray, tp: torch.Tensor, tv: torch.Tensor,
|
||||
info: list | None = None) -> np.ndarray:
|
||||
"""Draw kept tracks on the video. Foreground = red/blue, background anchors = gray."""
|
||||
from PIL import Image, ImageDraw
|
||||
T, H, W, _ = video.shape
|
||||
tp_np = tp[0].cpu().numpy() # [T, N, 2]
|
||||
tv_np = tv[0].cpu().numpy() # [T, N]
|
||||
N = tp_np.shape[1]
|
||||
kept = [k for k in range(N) if tv_np[:, k].max() > 0.5]
|
||||
fg_colors = [(255, 60, 60), (60, 200, 255)]
|
||||
bg_color = (180, 180, 180)
|
||||
# figure out which are fg vs bg by looking at info if given
|
||||
is_bg_by_idx = {}
|
||||
if info is not None:
|
||||
for entry in info:
|
||||
is_bg_by_idx[entry["track_idx"]] = (entry["oid"] == -1)
|
||||
out = []
|
||||
fg_seen = 0
|
||||
for t in range(T):
|
||||
img = Image.fromarray(np.ascontiguousarray(video[t]))
|
||||
dr = ImageDraw.Draw(img)
|
||||
fg_seen = 0
|
||||
for k in kept:
|
||||
if tv_np[t, k] <= 0.5:
|
||||
continue
|
||||
x = float(tp_np[t, k, 0] * W)
|
||||
y = float(tp_np[t, k, 1] * H)
|
||||
if is_bg_by_idx.get(k, False):
|
||||
col = bg_color
|
||||
r = 3
|
||||
else:
|
||||
col = fg_colors[fg_seen % len(fg_colors)]
|
||||
fg_seen += 1
|
||||
r = 5
|
||||
dr.ellipse([x - r, y - r, x + r, y + r], fill=col, outline=(255, 255, 255))
|
||||
out.append(np.asarray(img))
|
||||
return np.stack(out)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True)
|
||||
p.add_argument("--yaml", required=True)
|
||||
p.add_argument("--val-parquet", required=True)
|
||||
p.add_argument("--wandb-project", default="wantrack-bidir")
|
||||
p.add_argument("--wandb-run-name", required=True)
|
||||
p.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/two_track_val")
|
||||
p.add_argument("--w-text", type=float, default=3.0)
|
||||
p.add_argument("--w-motion", type=float, default=1.5)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--num-fg-tracks", type=int, default=2,
|
||||
help="Number of foreground (moving object) tracks to keep.")
|
||||
p.add_argument("--num-bg-tracks", type=int, default=0,
|
||||
help="Extra static bg anchor tracks in addition to the fg ones.")
|
||||
p.add_argument("--swap-text-from", type=int, default=None,
|
||||
help="For every sample, replace text_embedding+mask with THIS sample idx's text.")
|
||||
p.add_argument("--swap-clip-from", type=int, default=None,
|
||||
help="For every sample, replace clip_feature with THIS sample idx's clip.")
|
||||
p.add_argument("--swap-ff-from", type=int, default=None,
|
||||
help="For every sample, replace first_frame_latent with THIS sample idx's ff.")
|
||||
args = p.parse_args()
|
||||
|
||||
out_dir = Path(args.out_dir) / args.wandb_run_name
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
import wandb
|
||||
|
||||
print(f"[2trk] loading model {args.model_dir}", flush=True)
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(getattr(tc.data, "text_padding_length", 256))
|
||||
print(f"[2trk] text_len={text_len}", flush=True)
|
||||
|
||||
print(f"[2trk] loading val samples", flush=True)
|
||||
samples = load_val_samples(args.val_parquet, text_len)
|
||||
print(f"[2trk] {len(samples)} val samples", flush=True)
|
||||
|
||||
run = wandb.init(project=args.wandb_project, name=args.wandb_run_name,
|
||||
config={"model_dir": args.model_dir, "w_text": args.w_text,
|
||||
"w_motion": args.w_motion, "steps": args.steps})
|
||||
|
||||
for i, s in enumerate(samples):
|
||||
sid = s["id"] or f"clip{i:03d}"
|
||||
print(f"\n[2trk] === sample {i}: {sid} ({s['file_name'][:40]}) ===", flush=True)
|
||||
# Optional cross-sample field swaps (isolate which condition drives quality)
|
||||
if args.swap_text_from is not None:
|
||||
src = samples[args.swap_text_from]
|
||||
s = {**s, "text_embedding": src["text_embedding"].clone(),
|
||||
"text_attention_mask": src["text_attention_mask"].clone()}
|
||||
print(f"[2trk] SWAPPED text from sample {args.swap_text_from}", flush=True)
|
||||
if args.swap_clip_from is not None:
|
||||
src = samples[args.swap_clip_from]
|
||||
s = {**s, "clip_feature": src["clip_feature"].clone()}
|
||||
print(f"[2trk] SWAPPED clip from sample {args.swap_clip_from}", flush=True)
|
||||
if args.swap_ff_from is not None:
|
||||
src = samples[args.swap_ff_from]
|
||||
s = {**s, "first_frame_latent": src["first_frame_latent"].clone()}
|
||||
print(f"[2trk] SWAPPED first_frame_latent from sample {args.swap_ff_from}", flush=True)
|
||||
tp2, tv2, info = pick_two_tracks(s, seed=args.seed + i,
|
||||
num_fg=args.num_fg_tracks, num_bg=args.num_bg_tracks)
|
||||
print(f"[2trk] picked: {info}", flush=True)
|
||||
|
||||
t0 = time.time()
|
||||
frames = sample_denoise(model, s, tp2, tv2, w_text=args.w_text, w_motion=args.w_motion,
|
||||
num_steps=args.steps, seed=args.seed + i)
|
||||
dt = time.time() - t0
|
||||
print(f"[2trk] denoise {dt:.1f}s -> frames {frames.shape}", flush=True)
|
||||
|
||||
overlay = draw_track_overlay(frames, tp2, tv2, info)
|
||||
fn_raw = out_dir / f"{sid}_gen.mp4"
|
||||
fn_ov = out_dir / f"{sid}_gen_overlay.mp4"
|
||||
imageio.mimsave(str(fn_raw), frames, fps=args.fps, macro_block_size=1)
|
||||
imageio.mimsave(str(fn_ov), overlay, fps=args.fps, macro_block_size=1)
|
||||
|
||||
# GT for reference
|
||||
try:
|
||||
ref = twi.decode_reference(model, s["vae_latent"].to(model.device, torch.bfloat16))
|
||||
ref_ov = draw_track_overlay(ref, tp2, tv2, info)
|
||||
fn_ref = out_dir / f"{sid}_gt_overlay.mp4"
|
||||
imageio.mimsave(str(fn_ref), ref_ov, fps=args.fps, macro_block_size=1)
|
||||
wandb.log({
|
||||
f"sample_{i:02d}_gen": wandb.Video(str(fn_ov), fps=args.fps, format="mp4"),
|
||||
f"sample_{i:02d}_gt": wandb.Video(str(fn_ref), fps=args.fps, format="mp4"),
|
||||
"sample": i, "caption": s["caption"][:200], "picked": str(info),
|
||||
})
|
||||
except Exception as e:
|
||||
print(f"[2trk] gt fail: {e}", flush=True)
|
||||
wandb.log({
|
||||
f"sample_{i:02d}_gen": wandb.Video(str(fn_ov), fps=args.fps, format="mp4"),
|
||||
"sample": i, "caption": s["caption"][:200], "picked": str(info),
|
||||
})
|
||||
|
||||
wandb.finish()
|
||||
print("[2trk] DONE", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,79 @@
|
||||
#!/bin/bash
|
||||
# Publish the processed (parquet) dataset to a HuggingFace dataset repo as a directory tree,
|
||||
# mirroring noctuashap/openvid-wantrack-processed's layout (raw parquet files, not tarred).
|
||||
# Uses `hf upload-large-folder` -- resumable and built for multi-TB uploads: re-running skips
|
||||
# files already on the Hub, so a killed upload just continues.
|
||||
#
|
||||
# Usage:
|
||||
# REPO=FastVideo/openvid-wantrack-processed-v2 bash data_pipeline/upload_parquets.sh
|
||||
# ... DRY_RUN=1 ... # verify completeness + print the command, upload nothing
|
||||
set -uo pipefail
|
||||
|
||||
REPO=${REPO:?set REPO=<owner>/<name>}
|
||||
PARQUET_ROOT=${PARQUET_ROOT:-/home/hal-shared/motionstream/data/openvid-wantrack-parquets}
|
||||
SRC_ROOT=${SRC_ROOT:-/home/hal-shared/motionstream/data/openvid-wantrack}
|
||||
PRIVATE=${PRIVATE:-1}
|
||||
NUM_WORKERS=${NUM_WORKERS:-8}
|
||||
README=${README:-data_pipeline/notes/processed_dataset_README.md}
|
||||
DRY_RUN=${DRY_RUN:-0}
|
||||
|
||||
cd "$(dirname "$0")/.."
|
||||
HF=$(command -v hf || command -v huggingface-cli) || { echo "[pub] ERROR: hf CLI not found" >&2; exit 1; }
|
||||
|
||||
# --- safety: verify the shards PRESENT under PARQUET_ROOT are duplicate-free before publishing.
|
||||
# Validates whatever is present (no hardcoded 260, no source-video cross-check), so partial /
|
||||
# derivative sets like the bf16 copy work; still catches real corruption via duplicate ids.
|
||||
# Set VERIFY=0 to skip verification entirely.
|
||||
if [[ "${VERIFY:-1}" == "1" ]]; then
|
||||
echo "[pub] verifying present shards are duplicate-free before upload ..."
|
||||
python - "$PARQUET_ROOT" <<'PY'
|
||||
import glob, os, sys
|
||||
import pyarrow.parquet as pq
|
||||
pbase = sys.argv[1]
|
||||
shards = sorted(d for d in os.listdir(pbase) if d.startswith("shard"))
|
||||
bad=[]; total=0; nsh=0
|
||||
for s in shards:
|
||||
fs=glob.glob(f"{pbase}/{s}/**/*.parquet", recursive=True)
|
||||
if not fs:
|
||||
continue
|
||||
ids=[i for f in fs for i in pq.read_table(f, columns=["id"]).column("id").to_pylist()]
|
||||
if len(ids)>0 and len(ids)==len(set(ids)):
|
||||
total+=len(ids); nsh+=1
|
||||
else:
|
||||
bad.append((s, len(ids), len(set(ids))))
|
||||
if bad:
|
||||
print(f"[pub] REFUSING: {len(bad)} shard(s) empty/duplicate-id: {bad[:8]}")
|
||||
sys.exit(1)
|
||||
print(f"[pub] OK: {nsh} shard(s) present, {total:,} clips (no duplicate ids)")
|
||||
PY
|
||||
[[ $? -eq 0 ]] || { echo "[pub] aborted -- fix the shards above, then re-run" >&2; exit 1; }
|
||||
else
|
||||
echo "[pub] VERIFY=0 -> skipping shard verification"
|
||||
fi
|
||||
|
||||
n_files=$(find "$PARQUET_ROOT" -name '*.parquet' | wc -l)
|
||||
size=$(du -sh "$PARQUET_ROOT" 2>/dev/null | cut -f1)
|
||||
echo "[pub] repo=$REPO files=$n_files size=$size private=$PRIVATE"
|
||||
|
||||
if [[ "$DRY_RUN" == "1" ]]; then
|
||||
echo "[pub] DRY RUN -- would run:"
|
||||
echo " $HF upload-large-folder $REPO $PARQUET_ROOT --repo-type dataset --include '*.parquet' --num-workers $NUM_WORKERS"
|
||||
echo " (+ README.md upload)"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# --- create repo + upload README once -------------------------------------------------
|
||||
vis=(); [[ "$PRIVATE" == "1" ]] && vis=(--private)
|
||||
"$HF" repo create "$REPO" --repo-type dataset "${vis[@]}" 2>/dev/null \
|
||||
&& echo "[pub] created $REPO" || echo "[pub] repo exists (ok)"
|
||||
[[ -f "$README" ]] && "$HF" upload "$REPO" "$README" README.md --repo-type dataset >/dev/null 2>&1 \
|
||||
&& echo "[pub] README uploaded"
|
||||
|
||||
# --- upload the parquet tree (resumable) ----------------------------------------------
|
||||
# --include '*.parquet' skips any stray files; the shard*/combined_parquet_dataset/... tree
|
||||
# is preserved in the repo. Re-run this exact command to resume after any interruption.
|
||||
echo "[pub] uploading parquet tree (resumable; re-run to continue if interrupted) ..."
|
||||
"$HF" upload-large-folder "$REPO" "$PARQUET_ROOT" \
|
||||
--repo-type dataset --include '*.parquet' --num-workers "$NUM_WORKERS"
|
||||
|
||||
echo "[pub] done -> https://huggingface.co/datasets/$REPO"
|
||||
@@ -0,0 +1,115 @@
|
||||
#!/bin/bash
|
||||
# Package each shard's tracks/ into a tar and upload to a HuggingFace dataset repo,
|
||||
# mirroring noctuashap/openvid-wantrack-tracks layout: one tars-NNNNN.tar per shard,
|
||||
# each holding ~1000 .npz (flat, no directory prefix).
|
||||
#
|
||||
# Only shards whose progress.json marks tracks done are packaged. Resumable: a shard
|
||||
# already present in the repo (checked via the HF API) is skipped.
|
||||
#
|
||||
# Usage:
|
||||
# REPO=<user-or-org>/<name> SHARDS=0-170 bash data_pipeline/upload_tracks.sh
|
||||
# REPO=FastVideo/openvid-wantrack-tracks-v2 SHARDS=0-170 PRIVATE=1 bash data_pipeline/upload_tracks.sh
|
||||
# ... DRY_RUN=1 ... # build tars + report, do NOT create repo or upload
|
||||
set -uo pipefail
|
||||
|
||||
REPO=${REPO:?set REPO=<owner>/<name>}
|
||||
SHARDS=${SHARDS:-0-259}
|
||||
DATA_ROOT_BASE=${DATA_ROOT_BASE:-/home/hal-shared/motionstream/data/openvid-wantrack/shard}
|
||||
STAGING=${STAGING:-/home/hal-shared/motionstream/data/openvid-wantrack/_upload_tars}
|
||||
PRIVATE=${PRIVATE:-1} # create the repo private by default; you flip it public in the UI
|
||||
KEEP_TARS=${KEEP_TARS:-0} # 1 = keep local tar after upload (default: delete to save disk)
|
||||
DRY_RUN=${DRY_RUN:-0}
|
||||
PREFIX=${PREFIX:-tracks} # tar basename: ${PREFIX}-00042.tar
|
||||
README=${README:-data_pipeline/notes/tracks_dataset_README.md} # uploaded as README.md if present
|
||||
|
||||
cd "$(dirname "$0")/.."
|
||||
shard_root() { printf "%s%03d" "$DATA_ROOT_BASE" "$1"; }
|
||||
mkdir -p "$STAGING"
|
||||
|
||||
# expand "0-9,20,30-35"
|
||||
expand() {
|
||||
local tok lo hi out=(); IFS=',' read -ra toks <<< "$1"
|
||||
for tok in "${toks[@]}"; do
|
||||
if [[ "$tok" =~ ^([0-9]+)-([0-9]+)$ ]]; then
|
||||
for ((i=${BASH_REMATCH[1]}; i<=${BASH_REMATCH[2]}; i++)); do out+=("$i"); done
|
||||
elif [[ "$tok" =~ ^[0-9]+$ ]]; then out+=("$tok")
|
||||
else echo "[upload] ERROR: bad SHARDS token '$tok'" >&2; exit 1; fi
|
||||
done
|
||||
printf '%s\n' "${out[@]}"
|
||||
}
|
||||
mapfile -t LIST < <(expand "$SHARDS")
|
||||
|
||||
tracks_done() { # tracks_done <n>
|
||||
python - "$(shard_root "$1")/progress.json" <<'PY'
|
||||
import json,sys
|
||||
from pathlib import Path
|
||||
p=Path(sys.argv[1])
|
||||
sys.exit(0 if p.exists() and json.loads(p.read_text()).get("phases",{}).get("tracks",{}).get("done") else 1)
|
||||
PY
|
||||
}
|
||||
|
||||
# --- ensure repo exists (unless dry run) -------------------------------------------
|
||||
if [[ "$DRY_RUN" != "1" ]]; then
|
||||
vis=(); [[ "$PRIVATE" == "1" ]] && vis=(--private)
|
||||
hf repo create "$REPO" --repo-type dataset "${vis[@]}" 2>/dev/null \
|
||||
&& echo "[upload] created dataset repo $REPO" \
|
||||
|| echo "[upload] repo $REPO already exists (ok)"
|
||||
# names of files already in the repo (via the hub API), to skip re-upload on resume
|
||||
mapfile -t REMOTE < <(python - "$REPO" <<'PY'
|
||||
import sys
|
||||
from huggingface_hub import HfApi
|
||||
try:
|
||||
print("\n".join(HfApi().list_repo_files(sys.argv[1], repo_type="dataset")))
|
||||
except Exception:
|
||||
pass
|
||||
PY
|
||||
)
|
||||
remote_has() { printf '%s\n' "${REMOTE[@]:-}" | grep -qx "$1"; }
|
||||
|
||||
# upload README once (if present and not already there)
|
||||
if [[ -f "$README" ]] && ! remote_has "README.md"; then
|
||||
hf upload "$REPO" "$README" "README.md" --repo-type dataset >/dev/null 2>&1 \
|
||||
&& echo "[upload] README.md uploaded" || echo "[upload] WARN: README upload failed" >&2
|
||||
fi
|
||||
else
|
||||
remote_has() { return 1; }
|
||||
fi
|
||||
|
||||
n_up=0 n_skip=0 n_todo=0
|
||||
T0=$(date +%s)
|
||||
for s in "${LIST[@]}"; do
|
||||
root="$(shard_root "$s")"
|
||||
tar_name=$(printf "%s-%05d.tar" "$PREFIX" "$s")
|
||||
if ! tracks_done "$s"; then continue; fi
|
||||
n_todo=$((n_todo+1))
|
||||
if remote_has "$tar_name"; then n_skip=$((n_skip+1)); echo "[upload] $tar_name already in repo, skip"; continue; fi
|
||||
|
||||
ntracks=$(ls "$root"/tracks/*.npz 2>/dev/null | wc -l || echo 0)
|
||||
[[ "$ntracks" -gt 0 ]] || { echo "[upload] WARN shard $s: no npz despite progress=done, skip" >&2; continue; }
|
||||
|
||||
tar_path="$STAGING/$tar_name"
|
||||
# Clip names start with '---', so a glob/ls would feed tar filenames it reads as options.
|
||||
# find -print0 | tar --null -T - is dash-safe; paths come out './name' (matches the
|
||||
# reference repo's layout). Exclude any leftover *.tmp.npz from an interrupted write.
|
||||
echo "[upload] packing shard $s: $ntracks npz -> $tar_name"
|
||||
( cd "$root/tracks" && find . -maxdepth 1 -type f -name '*.npz' ! -name '*.tmp.npz' -print0 ) \
|
||||
| tar -cf "$tar_path" --null -C "$root/tracks" -T -
|
||||
sz=$(du -h "$tar_path" 2>/dev/null | cut -f1)
|
||||
|
||||
if [[ "$DRY_RUN" == "1" ]]; then
|
||||
echo "[upload] DRY_RUN: built $tar_name ($sz), not uploading"
|
||||
[[ "$KEEP_TARS" == "1" ]] || rm -f "$tar_path"
|
||||
continue
|
||||
fi
|
||||
|
||||
if hf upload "$REPO" "$tar_path" "$tar_name" --repo-type dataset >/dev/null 2>&1; then
|
||||
n_up=$((n_up+1)); echo "[upload] shard $s -> $tar_name ($sz) uploaded"
|
||||
[[ "$KEEP_TARS" == "1" ]] || rm -f "$tar_path"
|
||||
else
|
||||
echo "[upload] ERROR uploading $tar_name (kept at $tar_path)" >&2
|
||||
fi
|
||||
done
|
||||
|
||||
echo "[upload] done in $(( ($(date +%s)-T0)/60 ))m: $n_up uploaded, $n_skip already present, $n_todo eligible"
|
||||
[[ "$DRY_RUN" == "1" ]] && echo "[upload] (dry run: repo not created, nothing uploaded)"
|
||||
echo "[upload] repo: https://huggingface.co/datasets/$REPO"
|
||||
@@ -0,0 +1,60 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Re-run validation as a clean GT-vs-generated side-by-side (no track overlay).
|
||||
|
||||
For each clip: decode the GT clip, generate with its GT tracks, and write a
|
||||
left=GT / right=generated side-by-side video + a first-frame comparison PNG, so
|
||||
the first-frame match and overall reconstruction are easy to eyeball.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track/combined_parquet_dataset")
|
||||
p.add_argument("--out", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/val_compare_4k")
|
||||
p.add_argument("--fig", default="research_log/figures")
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
os.makedirs(args.fig, exist_ok=True)
|
||||
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
|
||||
for ci, s in zip(args.clips, samples):
|
||||
Tpx = (s["first_frame_latent"].shape[2] - 1) * ratio + 1
|
||||
gt = twi.decode_reference(model, s["vae_latent"]) # [T,H,W,3]
|
||||
lat = twi.generate(model, first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"], text_attention_mask=s["text_attention_mask"],
|
||||
track_points=s["track_points"][:, :Tpx], track_visibility=s["track_visibility"][:, :Tpx],
|
||||
num_steps=args.steps, seed=args.seed)
|
||||
gen = twi.decode_to_pixels(model, lat)
|
||||
T = min(len(gt), len(gen))
|
||||
side = np.concatenate([gt[:T], gen[:T]], axis=2) # [T,H,2W,3] (GT|gen)
|
||||
imageio.mimsave(os.path.join(args.out, f"val_clip{ci}_GTleft_GENright.mp4"),
|
||||
list(side), fps=24, macro_block_size=1)
|
||||
ff = np.concatenate([gt[0], gen[0]], axis=1)
|
||||
imageio.imwrite(os.path.join(args.fig, f"val_firstframe_clip{ci}_GTleft_GENright.png"), ff)
|
||||
mse0 = float(((gt[0].astype(np.float32) - gen[0].astype(np.float32)) ** 2).mean())
|
||||
print(f"[clip {ci}] first-frame MSE(GT,gen)={mse0:8.1f} -> {args.out}/val_clip{ci}_GTleft_GENright.mp4", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,160 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Visualize CoTracker point tracks over a generated video.
|
||||
|
||||
Produces, for one (video, tracks) pair:
|
||||
- ``<out>_overlay.mp4``: every frame with visible points drawn as dots plus a short
|
||||
motion tail (last ``--tail`` frames). Colour encodes the point's initial grid position.
|
||||
- ``<out>_trajectories.png``: all full trajectories drawn over frame 0 (static summary).
|
||||
|
||||
Pure CPU; depends only on numpy + PIL + imageio + torchvision (already in the venv), so it
|
||||
does NOT need CoTracker or a GPU.
|
||||
|
||||
Example:
|
||||
.venv/bin/python data_pipeline/visualize_tracks.py \
|
||||
--video /.../smoke_wan21_1.3b_480p/videos/vid_000000.mp4 \
|
||||
--tracks /.../smoke_wan21_1.3b_480p/tracks/vid_000000.npz \
|
||||
--out /.../smoke_wan21_1.3b_480p/viz/vid_000000
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import colorsys
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--video", type=Path, required=True, help="Source .mp4")
|
||||
p.add_argument("--tracks", type=Path, required=True, help=".npz from extract_tracks.py")
|
||||
p.add_argument("--out", type=Path, required=True, help="Output path prefix (no extension).")
|
||||
p.add_argument("--stride", type=int, default=2, help="Sub-sample the NxN grid by this factor for clarity.")
|
||||
p.add_argument("--tail", type=int, default=12, help="Motion-tail length in frames.")
|
||||
p.add_argument("--radius", type=int, default=2, help="Point radius in px.")
|
||||
p.add_argument("--fps", type=int, default=16, help="Output video fps.")
|
||||
p.add_argument("--vis-thresh", type=float, default=0.5, help="Visibility threshold.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def load_frames(path: Path) -> np.ndarray:
|
||||
"""Return frames (T, H, W, 3) uint8.
|
||||
|
||||
Recent torchvision dropped ``torchvision.io.read_video``; use decord, falling
|
||||
back to imageio's ffmpeg reader.
|
||||
"""
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(str(path), ctx=cpu(0))
|
||||
frames = vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001 - fall back to ffmpeg
|
||||
reader = imageio.get_reader(str(path), format="ffmpeg")
|
||||
frames = np.stack([np.asarray(f) for f in reader], axis=0)
|
||||
reader.close()
|
||||
return frames[..., :3].astype(np.uint8)
|
||||
|
||||
|
||||
def grid_colors(grid_size: int, stride: int) -> np.ndarray:
|
||||
"""One RGB colour per (sub-sampled) grid point, encoding its initial position."""
|
||||
idx = np.arange(0, grid_size, stride)
|
||||
gy, gx = np.meshgrid(idx, idx, indexing="ij")
|
||||
nx = gx.reshape(-1) / max(grid_size - 1, 1)
|
||||
ny = gy.reshape(-1) / max(grid_size - 1, 1)
|
||||
cols = np.empty((nx.shape[0], 3), dtype=np.uint8)
|
||||
for i, (x, y) in enumerate(zip(nx, ny)):
|
||||
r, g, b = colorsys.hsv_to_rgb(float(x), 1.0, 0.5 + 0.5 * float(y))
|
||||
cols[i] = (int(r * 255), int(g * 255), int(b * 255))
|
||||
return cols
|
||||
|
||||
|
||||
def subsample(tracks: np.ndarray, vis: np.ndarray, grid_size: int, stride: int):
|
||||
"""tracks (T,N,2), vis (T,N) with N==grid_size**2 -> sub-sampled by stride in both grid dims."""
|
||||
t = tracks.shape[0]
|
||||
if tracks.shape[1] != grid_size * grid_size:
|
||||
return tracks, vis # unknown layout; keep as-is
|
||||
tr = tracks.reshape(t, grid_size, grid_size, 2)[:, ::stride, ::stride, :].reshape(t, -1, 2)
|
||||
vs = vis.reshape(t, grid_size, grid_size)[:, ::stride, ::stride].reshape(t, -1)
|
||||
return tr, vs
|
||||
|
||||
|
||||
def draw_overlay(frames, tracks, vis, colors, tail, radius, vis_thresh) -> list[np.ndarray]:
|
||||
t, h, w, _ = frames.shape
|
||||
n = tracks.shape[1]
|
||||
out = []
|
||||
for fi in range(t):
|
||||
img = Image.fromarray(frames[fi]).convert("RGB")
|
||||
draw = ImageDraw.Draw(img)
|
||||
t0 = max(0, fi - tail)
|
||||
for pi in range(n):
|
||||
col = tuple(int(c) for c in colors[pi])
|
||||
# tail: consecutive visible positions in the window
|
||||
pts = []
|
||||
for tj in range(t0, fi + 1):
|
||||
if vis[tj, pi] >= vis_thresh:
|
||||
x, y = float(tracks[tj, pi, 0]), float(tracks[tj, pi, 1])
|
||||
if 0 <= x < w and 0 <= y < h:
|
||||
pts.append((x, y))
|
||||
else:
|
||||
pts = [] # break the tail on occlusion
|
||||
if len(pts) >= 2:
|
||||
draw.line(pts, fill=col, width=1)
|
||||
if vis[fi, pi] >= vis_thresh:
|
||||
x, y = float(tracks[fi, pi, 0]), float(tracks[fi, pi, 1])
|
||||
if 0 <= x < w and 0 <= y < h:
|
||||
draw.ellipse([x - radius, y - radius, x + radius, y + radius], fill=col)
|
||||
out.append(np.asarray(img))
|
||||
return out
|
||||
|
||||
|
||||
def draw_trajectories(frame0, tracks, vis, colors, vis_thresh) -> np.ndarray:
|
||||
h, w, _ = frame0.shape
|
||||
img = Image.fromarray(frame0).convert("RGB")
|
||||
draw = ImageDraw.Draw(img)
|
||||
n = tracks.shape[1]
|
||||
for pi in range(n):
|
||||
col = tuple(int(c) for c in colors[pi])
|
||||
pts = [(float(tracks[tj, pi, 0]), float(tracks[tj, pi, 1]))
|
||||
for tj in range(tracks.shape[0])
|
||||
if vis[tj, pi] >= vis_thresh and 0 <= tracks[tj, pi, 0] < w and 0 <= tracks[tj, pi, 1] < h]
|
||||
if len(pts) >= 2:
|
||||
draw.line(pts, fill=col, width=1)
|
||||
return np.asarray(img)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
data = np.load(args.tracks)
|
||||
tracks = data["tracks"].astype(np.float32) # (T, N, 2)
|
||||
vis = data["visibility"].astype(np.float32) # (T, N)
|
||||
grid_size = int(data["grid_size"]) if "grid_size" in data else int(round(tracks.shape[1] ** 0.5))
|
||||
|
||||
frames = load_frames(args.video)
|
||||
t = min(frames.shape[0], tracks.shape[0])
|
||||
frames, tracks, vis = frames[:t], tracks[:t], vis[:t]
|
||||
|
||||
tracks, vis = subsample(tracks, vis, grid_size, args.stride)
|
||||
colors = grid_colors(grid_size, args.stride)
|
||||
if colors.shape[0] != tracks.shape[1]: # layout fallback: cycle a rainbow
|
||||
colors = grid_colors(int(round(tracks.shape[1] ** 0.5)) or 1, 1)[:tracks.shape[1]]
|
||||
|
||||
args.out.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
overlay = draw_overlay(frames, tracks, vis, colors, args.tail, args.radius, args.vis_thresh)
|
||||
mp4_path = args.out.with_name(args.out.name + "_overlay.mp4")
|
||||
imageio.mimsave(str(mp4_path), overlay, fps=args.fps, macro_block_size=1)
|
||||
|
||||
traj = draw_trajectories(frames[0], tracks, vis, colors, args.vis_thresh)
|
||||
png_path = args.out.with_name(args.out.name + "_trajectories.png")
|
||||
imageio.imwrite(str(png_path), traj)
|
||||
|
||||
visible_frac = float((vis >= args.vis_thresh).mean())
|
||||
print(f"[viz] frames={t} points_drawn={tracks.shape[1]} (grid {grid_size}x{grid_size}, stride {args.stride}) "
|
||||
f"mean_visible={visible_frac:.2f}")
|
||||
print(f"[viz] wrote {mp4_path}")
|
||||
print(f"[viz] wrote {png_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,3 +1,3 @@
|
||||
#! /bin/bash
|
||||
|
||||
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
|
||||
hf download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
|
||||
|
||||
@@ -11,7 +11,7 @@ def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Interactive track-control demo for TrackWan (MotionStream-style).
|
||||
|
||||
Pick a preprocessed clip (gives the first frame + text + the option of its own
|
||||
tracks), author a motion control, preview the tracks over the first frame, then
|
||||
generate and watch whether the model follows them.
|
||||
|
||||
Control modes (mirroring MotionStream):
|
||||
- Trajectory / drag : click the first frame to set a drag handle, choose a
|
||||
direction + radius; optionally make it SPARSE (only the handle points active).
|
||||
- Camera-like : pan / zoom / rotate / swirl presets.
|
||||
- Motion transfer : reuse another clip's extracted tracks.
|
||||
- GT : the clip's own tracks (reconstruction sanity check).
|
||||
|
||||
Launch (single GPU, on the cluster)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=1 \
|
||||
HF_HUB_CACHE=/.../hf_cache_clean PYTHONPATH=$PWD \
|
||||
.venv/bin/python examples/inference/gradio/trackwan/app.py \
|
||||
--export /.../trackwan_1.3b_overfit4k --port 7860
|
||||
|
||||
Then SSH-forward the port and open http://localhost:7860 .
|
||||
|
||||
NOTE: sparse drag control needs a model trained with point-subsampling aug
|
||||
(WANTRACK_AUG=1). A clean-overfit (aug-off) checkpoint mostly ignores sparse
|
||||
controls -- use dense presets to test it.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
REPO = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", ".."))
|
||||
sys.path.insert(0, os.path.join(REPO, "data_pipeline"))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
STATE: dict = {} # model, tc, samples, Tpx, H, W, out_dir
|
||||
|
||||
|
||||
def _overlay(frames_thwc, tracks_norm, vis, stride=3):
|
||||
from fastvideo.train.callbacks.track_validation import (_draw_overlay, _grid_colors, _subsample)
|
||||
T, H, W, _ = frames_thwc.shape
|
||||
tt = min(T, tracks_norm.shape[0])
|
||||
fr, tr, vs = frames_thwc[:tt], tracks_norm[:tt], vis[:tt]
|
||||
grid = int(round(tr.shape[1] ** 0.5))
|
||||
tr, vs = _subsample(tr, vs, grid, stride)
|
||||
colors = _grid_colors(grid, stride)
|
||||
if colors.shape[0] != tr.shape[1]:
|
||||
colors = _grid_colors(int(round(tr.shape[1] ** 0.5)) or 1, 1)[:tr.shape[1]]
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
return _draw_overlay(fr, trpx, vs, colors, 12, 2, 0.5)
|
||||
|
||||
|
||||
def _first_frame_rgb(clip_idx: int) -> np.ndarray:
|
||||
"""Decode the clip's GT first frame for display/clicking."""
|
||||
s = STATE["samples"][clip_idx]
|
||||
ref = twi.decode_reference(STATE["model"], s["vae_latent"]) # [T,H,W,3]
|
||||
return ref[0]
|
||||
|
||||
|
||||
def _build_tracks(clip_idx, mode, preset_name, strength, cx, cy, radius,
|
||||
sparsity, drag_dx, drag_dy, transfer_idx):
|
||||
"""Return (tracks_norm [T,N,2], vis [T,N]) or (None, None) for 'none'."""
|
||||
Tpx = STATE["Tpx"]
|
||||
g = st.make_grid(50)
|
||||
if mode == "gt":
|
||||
s = STATE["samples"][clip_idx]
|
||||
return s["track_points"][0].numpy()[:Tpx], s["track_visibility"][0].numpy()[:Tpx]
|
||||
if mode == "none":
|
||||
return None, None
|
||||
if mode == "transfer":
|
||||
s = STATE["samples"][int(transfer_idx)]
|
||||
return s["track_points"][0].numpy()[:Tpx], s["track_visibility"][0].numpy()[:Tpx]
|
||||
if mode == "preset":
|
||||
tr, vs = st.preset(preset_name, Tpx, 50, strength=float(strength))
|
||||
elif mode == "drag":
|
||||
tr, vs = st.drag(g, Tpx, center=(float(cx), float(cy)), dx=float(drag_dx),
|
||||
dy=float(drag_dy), radius=float(radius))
|
||||
else:
|
||||
tr, vs = st.static(g, Tpx)
|
||||
# sparsity
|
||||
if sparsity == "radius (handle only)":
|
||||
vs = st.select_radius(vs, g, center=(float(cx), float(cy)), radius=float(radius))
|
||||
elif sparsity == "coarse grid":
|
||||
vs = st.select_stride(vs, 50, 4)
|
||||
return tr, vs
|
||||
|
||||
|
||||
def on_select_clip(clip_idx):
|
||||
img = _first_frame_rgb(int(clip_idx))
|
||||
cap = STATE["samples"][int(clip_idx)]["caption"][:200]
|
||||
return img, cap
|
||||
|
||||
|
||||
def on_image_click(clip_idx, evt: gr.SelectData):
|
||||
"""Set drag center (normalized) from a click on the first frame."""
|
||||
img = _first_frame_rgb(int(clip_idx))
|
||||
H, W, _ = img.shape
|
||||
x, y = evt.index[0] / W, evt.index[1] / H
|
||||
return round(float(x), 3), round(float(y), 3)
|
||||
|
||||
|
||||
def on_preview(clip_idx, mode, preset_name, strength, cx, cy, radius, sparsity,
|
||||
drag_dx, drag_dy, transfer_idx):
|
||||
img = _first_frame_rgb(int(clip_idx))
|
||||
tr, vs = _build_tracks(int(clip_idx), mode, preset_name, strength, cx, cy, radius,
|
||||
sparsity, drag_dx, drag_dy, transfer_idx)
|
||||
if tr is None:
|
||||
return img, "no tracks (mode=none)"
|
||||
frames = np.repeat(img[None], tr.shape[0], 0)
|
||||
ov = _overlay(frames, tr, vs)[0]
|
||||
active = int((vs[0] > 0.5).sum())
|
||||
return ov, f"{active} active points (of {vs.shape[1]})"
|
||||
|
||||
|
||||
def on_generate(clip_idx, mode, preset_name, strength, cx, cy, radius, sparsity,
|
||||
drag_dx, drag_dy, transfer_idx, steps, seed):
|
||||
clip_idx = int(clip_idx)
|
||||
s = STATE["samples"][clip_idx]
|
||||
tr, vs = _build_tracks(clip_idx, mode, preset_name, strength, cx, cy, radius,
|
||||
sparsity, drag_dx, drag_dy, transfer_idx)
|
||||
tp = torch.from_numpy(tr)[None].float() if tr is not None else None
|
||||
tv = torch.from_numpy(vs)[None].float() if vs is not None else None
|
||||
lat = twi.generate(STATE["model"], first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"], text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp, track_visibility=tv, clip_feature=s["clip_feature"],
|
||||
num_steps=int(steps), seed=int(seed))
|
||||
frames = twi.decode_to_pixels(STATE["model"], lat)
|
||||
ov_tr = tr if tr is not None else np.zeros((frames.shape[0], 1, 2), np.float32)
|
||||
ov_vs = vs if vs is not None else np.zeros((frames.shape[0], 1), np.float32)
|
||||
out_frames = _overlay(frames, ov_tr, ov_vs)
|
||||
path = os.path.join(STATE["out_dir"], f"gen_clip{clip_idx}_{mode}_{preset_name}.mp4")
|
||||
imageio.mimsave(path, out_frames, fps=24, macro_block_size=1)
|
||||
|
||||
epe_txt = "EPE: n/a (no control tracks)"
|
||||
if tr is not None:
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import compute_epe
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
res = compute_epe(frames, trpx, vs, STATE["ct"], STATE["model"].device)
|
||||
if res["epe"] is not None:
|
||||
epe_txt = f"EPE: {res['epe']:.2f} px ({res['n_points']} pts) — lower = follows control"
|
||||
return path, epe_txt
|
||||
|
||||
|
||||
def build_ui(num_clips):
|
||||
with gr.Blocks(title="TrackWan — interactive motion control") as demo:
|
||||
gr.Markdown("## TrackWan — interactive point-track control\n"
|
||||
"Pick a clip, author a motion control, **Preview tracks**, then **Generate**. "
|
||||
"Click the first frame to set the drag handle center.")
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
clip = gr.Dropdown(list(range(num_clips)), value=0, label="Clip")
|
||||
caption = gr.Textbox(label="Prompt", interactive=False, lines=2)
|
||||
mode = gr.Radio(["preset", "drag", "transfer", "gt", "none"], value="preset", label="Control mode")
|
||||
preset_name = gr.Dropdown(st.PRESETS, value="pan_right", label="Preset (preset mode)")
|
||||
strength = gr.Slider(0.0, 0.6, value=0.25, step=0.01, label="Strength (pan/zoom)")
|
||||
with gr.Row():
|
||||
cx = gr.Number(value=0.5, label="drag center x (click frame)")
|
||||
cy = gr.Number(value=0.5, label="drag center y")
|
||||
radius = gr.Slider(0.05, 0.6, value=0.2, step=0.01, label="Radius (drag/handle)")
|
||||
with gr.Row():
|
||||
drag_dx = gr.Slider(-0.5, 0.5, value=0.3, step=0.01, label="drag dx")
|
||||
drag_dy = gr.Slider(-0.5, 0.5, value=0.0, step=0.01, label="drag dy")
|
||||
sparsity = gr.Radio(["full", "radius (handle only)", "coarse grid"], value="full",
|
||||
label="Sparsity (sparse needs aug-trained ckpt)")
|
||||
transfer_idx = gr.Dropdown(list(range(num_clips)), value=min(1, num_clips - 1),
|
||||
label="Transfer from clip (transfer mode)")
|
||||
with gr.Row():
|
||||
steps = gr.Slider(10, 60, value=30, step=5, label="Denoise steps")
|
||||
seed = gr.Number(value=1000, label="Seed")
|
||||
with gr.Row():
|
||||
preview_btn = gr.Button("Preview tracks")
|
||||
gen_btn = gr.Button("Generate", variant="primary")
|
||||
with gr.Column(scale=1):
|
||||
frame_img = gr.Image(label="First frame (click to set drag center) / track preview", type="numpy")
|
||||
status = gr.Textbox(label="Status", interactive=False)
|
||||
out_video = gr.Video(label="Generated (tracks overlaid)")
|
||||
epe_box = gr.Textbox(label="Motion fidelity", interactive=False)
|
||||
|
||||
ctrl = [clip, mode, preset_name, strength, cx, cy, radius, sparsity, drag_dx, drag_dy, transfer_idx]
|
||||
clip.change(on_select_clip, [clip], [frame_img, caption])
|
||||
frame_img.select(on_image_click, [clip], [cx, cy])
|
||||
preview_btn.click(on_preview, ctrl, [frame_img, status])
|
||||
gen_btn.click(on_generate, ctrl + [steps, seed], [out_video, epe_box])
|
||||
demo.load(on_select_clip, [clip], [frame_img, caption])
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True, help="dcp_to_diffusers export dir")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset")
|
||||
p.add_argument("--num-clips", type=int, default=10)
|
||||
p.add_argument("--out-dir", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/gradio_out")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7860)
|
||||
p.add_argument("--share", action="store_true", help="Create a public gradio.live link (no SSH forward needed).")
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, list(range(args.num_clips)), text_len)
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import load_cotracker
|
||||
STATE.update(model=model, tc=tc, samples=samples, Tpx=(num_lat_t - 1) * ratio + 1,
|
||||
out_dir=args.out_dir, ct=load_cotracker(model.device))
|
||||
|
||||
demo = build_ui(len(samples))
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[args.out_dir])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,478 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""TrackWan action-recording demo — draw a motion, watch the model follow it.
|
||||
|
||||
Designed to answer one question: *does this checkpoint actually follow the point
|
||||
tracks (action-following), and how does that change across checkpoints?*
|
||||
|
||||
Interaction (the part the old demo lacked):
|
||||
1. Pick a clip (gives the first frame + prompt) and a checkpoint.
|
||||
2. Move the mouse over the first frame and press **Space** -> records your cursor
|
||||
path for 5 seconds (121 frames). The green dot is where the action starts,
|
||||
the red dot where it ends.
|
||||
3. The 50x50 CoTracker training grid is used directly: the grid points within the
|
||||
action radius of your start are dragged along your recorded path, the rest stay
|
||||
static. This is grid-snapped / in-distribution (NOT an off-grid patch added on
|
||||
top of a static grid). Preview shows the full field (moving=green, static=gray).
|
||||
4. Generate -> original + generation side by side, with the full field overlaid,
|
||||
plus the moving-point CoTracker-EPE of how well the arm followed your path.
|
||||
|
||||
Checkpoints hot-swap (only the transformer weights reload, ~3 s) so you can flip
|
||||
between e.g. step-2000 / 3000 / 4000 on the same clip and action.
|
||||
|
||||
Visibility = "dense field" keeps the whole 50x50 grid visible (the trained coverage;
|
||||
recommended for aug-off checkpoints); "sparse" makes only the moved points visible.
|
||||
|
||||
Launch (single GPU on the cluster)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=1 \
|
||||
PYTHONPATH=$PWD .venv/bin/python \
|
||||
examples/inference/gradio/trackwan/app_action.py \
|
||||
--exports-root /.../control_now_exports --share
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import glob
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
REPO = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", ".."))
|
||||
sys.path.insert(0, os.path.join(REPO, "data_pipeline"))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
STATE: dict = {}
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------- utils
|
||||
def _img_b64(arr: np.ndarray) -> str:
|
||||
buf = io.BytesIO()
|
||||
Image.fromarray(arr).save(buf, format="PNG")
|
||||
return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode()
|
||||
|
||||
|
||||
def _first_frame_rgb(clip_idx: int) -> np.ndarray:
|
||||
s = STATE["samples"][clip_idx]
|
||||
return twi.decode_reference(STATE["model"], s["vae_latent"])[0]
|
||||
|
||||
|
||||
def _original_video(clip_idx: int) -> str:
|
||||
"""Decode + cache the GT clip as an mp4 (checkpoint-independent)."""
|
||||
cache = STATE.setdefault("orig_cache", {})
|
||||
if clip_idx in cache:
|
||||
return cache[clip_idx]
|
||||
s = STATE["samples"][clip_idx]
|
||||
frames = twi.decode_reference(STATE["model"], s["vae_latent"]) # [T,H,W,3]
|
||||
path = os.path.join(STATE["out_dir"], f"orig_clip{clip_idx}.mp4")
|
||||
imageio.mimsave(path, frames, fps=24, macro_block_size=1)
|
||||
cache[clip_idx] = path
|
||||
return path
|
||||
|
||||
|
||||
def swap_checkpoint(name: str) -> str:
|
||||
"""Load a checkpoint's transformer weights in place (fast; no text/vae reload)."""
|
||||
if STATE.get("cur_ckpt") == name:
|
||||
return f"checkpoint: {name} (loaded)"
|
||||
from safetensors.torch import load_file
|
||||
try:
|
||||
from torch.distributed.tensor import DTensor
|
||||
except Exception: # older torch
|
||||
from torch.distributed._tensor import DTensor
|
||||
path = STATE["exports"][name]
|
||||
sd = load_file(os.path.join(path, "transformer", "model.safetensors"))
|
||||
tgt = STATE["model"].transformer
|
||||
# Params are FSDP DTensors (even on 1 GPU); copy into each param's local shard
|
||||
# instead of load_state_dict (which errors mixing Tensor and DTensor).
|
||||
params = dict(tgt.named_parameters())
|
||||
n_ok = 0
|
||||
with torch.no_grad():
|
||||
for k, p in params.items():
|
||||
if k not in sd:
|
||||
continue
|
||||
loc = p.to_local() if isinstance(p, DTensor) else p.data
|
||||
src = sd[k].to(dtype=loc.dtype, device=loc.device)
|
||||
if tuple(loc.shape) == tuple(src.shape):
|
||||
loc.copy_(src)
|
||||
n_ok += 1
|
||||
tgt.eval()
|
||||
STATE["cur_ckpt"] = name
|
||||
warn = "" if n_ok == len(params) else f" [WARN copied {n_ok}/{len(params)} params]"
|
||||
return f"checkpoint: {name} (swapped){warn}"
|
||||
|
||||
|
||||
def _grid_action_tracks(traj, Tpx: int, radius: float, falloff: str, sparse: str, num_bg: int, seed: int):
|
||||
"""Build the input tracks for the drawn action, matching the model's training regime.
|
||||
|
||||
Modes (``sparse``):
|
||||
"dense" -- full 50x50 = 2500 tracks; the ones within ``radius`` of the draw-start
|
||||
follow the path, the rest stay static and visible. Matches the OLD
|
||||
training regime (WANTRACK_SPARSE=0, sample K=1000-2500 per step).
|
||||
"sparse" -- ONLY the moving handle + ``num_bg`` random background points from the
|
||||
training grid. Total N ~= handle_size + num_bg. Matches the SPARSE
|
||||
training regime (WANTRACK_SPARSE=1, 1-per-object + ~20 background).
|
||||
``num_bg`` should be small (5-30) to match training.
|
||||
|
||||
Returns (tracks[Tpx,N,2], vis[Tpx,N], moving[N]) or (None, None, None).
|
||||
"""
|
||||
traj = np.asarray(traj, np.float32)
|
||||
if traj.ndim != 2 or len(traj) < 2:
|
||||
return None, None, None
|
||||
idx = np.linspace(0, len(traj) - 1, Tpx)
|
||||
path = np.stack([
|
||||
np.interp(idx, np.arange(len(traj)), traj[:, 0]),
|
||||
np.interp(idx, np.arange(len(traj)), traj[:, 1]),
|
||||
], 1).astype(np.float32) # [Tpx,2]
|
||||
g = st.make_grid(50) # [2500,2] frame-0 grid (x,y) in [0,1]
|
||||
start = path[0]
|
||||
d = np.linalg.norm(g - start[None], axis=-1) # [2500]
|
||||
if falloff == "smooth":
|
||||
w = np.clip(1.0 - d / max(radius, 1e-6), 0.0, 1.0)
|
||||
w = (0.5 - 0.5 * np.cos(np.pi * w)).astype(np.float32) # handle-like weighting
|
||||
else: # hard
|
||||
w = (d <= radius).astype(np.float32)
|
||||
handle_mask = w > 1e-3 # [2500] which grid points are the moving handle
|
||||
disp = path - start[None] # [Tpx,2] displacement from start each frame
|
||||
mode = str(sparse)
|
||||
if mode == "sparse":
|
||||
# Take the moving handle + `num_bg` random background points (indices from the
|
||||
# non-handle pool), matching the sparse training sampler shape (~num_objects + extras).
|
||||
rng = np.random.default_rng(int(seed))
|
||||
bg_pool = np.where(~handle_mask)[0]
|
||||
n_bg = int(min(max(num_bg, 0), bg_pool.size))
|
||||
bg_idx = rng.choice(bg_pool, size=n_bg, replace=False) if n_bg > 0 else np.zeros(0, np.int64)
|
||||
handle_idx = np.where(handle_mask)[0]
|
||||
keep = np.concatenate([handle_idx, bg_idx])
|
||||
g_k = g[keep] # [N,2]
|
||||
w_k = w[keep] # [N] (background weights are ~0 -> stay static)
|
||||
tracks = (g_k[None] + w_k[None, :, None] * disp[:, None, :]).astype(np.float32) # [Tpx,N,2]
|
||||
tracks = np.clip(tracks, 0.0, 1.0)
|
||||
moving = np.zeros(keep.size, bool)
|
||||
moving[:handle_idx.size] = True
|
||||
vis = np.ones(tracks.shape[:2], np.float32) # all sparse points visible from start
|
||||
return tracks, vis, moving
|
||||
# "dense" (original behavior): 2500 tracks, handle moves, rest static + visible.
|
||||
tracks = (g[None] + w[None, :, None] * disp[:, None, :]).astype(np.float32) # [Tpx,2500,2]
|
||||
tracks = np.clip(tracks, 0.0, 1.0)
|
||||
vis = np.ones(tracks.shape[:2], np.float32)
|
||||
return tracks, vis, handle_mask
|
||||
|
||||
|
||||
def _overlay_grid(frames: np.ndarray, tracks: np.ndarray, vis: np.ndarray, moving, stride: int = 3) -> np.ndarray:
|
||||
"""Draw the input tracks: moving handle (green + tails) + static background (gray dots).
|
||||
|
||||
For dense mode (N = 50*50 = 2500) we grid-subsample by ``stride`` for legibility. For
|
||||
sparse mode (N ~= handle + num_bg, much smaller) we draw every point, no subsample.
|
||||
"""
|
||||
from fastvideo.train.callbacks.track_validation import _draw_overlay, _subsample
|
||||
T, H, W, _ = frames.shape
|
||||
tt = min(T, tracks.shape[0])
|
||||
tr, vs, mv = tracks[:tt], vis[:tt], moving
|
||||
if tr.shape[1] == 2500:
|
||||
tr, vs = _subsample(tr, vs, 50, stride)
|
||||
mv_full = np.broadcast_to(moving[None, :, None].astype(np.float32),
|
||||
(tt, moving.shape[0], 2)).copy()
|
||||
mv, _ = _subsample(mv_full, vis[:tt], 50, stride)
|
||||
is_mv = mv[0, :, 0] > 0.5
|
||||
else: # sparse: draw all points
|
||||
is_mv = mv.astype(bool)
|
||||
colors = np.where(is_mv[:, None], np.array([0, 255, 100], np.uint8),
|
||||
np.array([120, 120, 120], np.uint8)).astype(np.uint8)
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
return _draw_overlay(frames[:tt], trpx, vs, colors, 12, 2, 0.5)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------- handlers
|
||||
def on_clip(clip_idx):
|
||||
clip_idx = int(clip_idx)
|
||||
img = _first_frame_rgb(clip_idx)
|
||||
cap = STATE["samples"][clip_idx]["caption"][:240]
|
||||
return cap, _img_b64(img), _original_video(clip_idx)
|
||||
|
||||
|
||||
def on_ckpt(name):
|
||||
return swap_checkpoint(name)
|
||||
|
||||
|
||||
def _build_tracks(traj_json, radius, falloff, sparse, num_bg, seed):
|
||||
"""Parse the drawn trajectory -> (tracks[Tpx,N,2], vis, moving[N]) or (None, msg).
|
||||
|
||||
``sparse``: "dense" (full 2500 grid, background static) matches the OLD training
|
||||
regime, "sparse" (handle + ``num_bg`` random background) matches the SPARSE
|
||||
training regime our overfit runs used (WANTRACK_SPARSE=1, ~20 background extras).
|
||||
"""
|
||||
import json
|
||||
if not traj_json or traj_json.strip() in ("", "[]"):
|
||||
return None, "Record an action first: move the mouse over the frame and press Space (5 s)."
|
||||
try:
|
||||
traj = json.loads(traj_json)
|
||||
except Exception as e: # noqa: BLE001
|
||||
return None, f"bad trajectory json: {e}"
|
||||
tracks, vis, moving = _grid_action_tracks(traj, STATE["Tpx"], float(radius), str(falloff),
|
||||
str(sparse), int(num_bg), int(seed))
|
||||
if tracks is None:
|
||||
return None, "trajectory too short — hold Space and move the mouse for the full 5 s."
|
||||
return (tracks, vis, moving), None
|
||||
|
||||
|
||||
def _first_frame_cached(clip_idx):
|
||||
cache = STATE.setdefault("_ff_cache", {})
|
||||
if clip_idx not in cache:
|
||||
cache[clip_idx] = _first_frame_rgb(clip_idx)
|
||||
return cache[clip_idx]
|
||||
|
||||
|
||||
def on_preview(clip_idx, traj_json, radius, falloff, sparse, num_bg, seed):
|
||||
"""Show the input tracks overlaid on the (frozen) first frame -- BEFORE generating."""
|
||||
clip_idx = int(clip_idx)
|
||||
built, err = _build_tracks(traj_json, radius, falloff, sparse, num_bg, seed)
|
||||
if built is None:
|
||||
return None, err
|
||||
tracks, vis, moving = built
|
||||
ff = _first_frame_cached(clip_idx) # [H,W,3]
|
||||
frames = np.repeat(ff[None], tracks.shape[0], axis=0) # freeze frame 0, T copies
|
||||
out = _overlay_grid(frames, tracks, vis, moving)
|
||||
path = os.path.join(STATE["out_dir"], f"preview_clip{clip_idx}.mp4")
|
||||
imageio.mimsave(path, out, fps=24, macro_block_size=1)
|
||||
kind = "SPARSE" if str(sparse) == "sparse" else "DENSE 50x50"
|
||||
return path, (f"PREVIEW ({kind}, N={tracks.shape[1]}): {int(moving.sum())} moving handle points (green, tails) "
|
||||
f"and {int((~moving).sum())} background points (gray). "
|
||||
f"This matches the training track format. Press Generate.")
|
||||
|
||||
|
||||
def on_generate(clip_idx, ckpt_name, traj_json, radius, falloff, sparse, steps, seed,
|
||||
w_text: float = 3.0, w_motion: float = 1.5, mode: str = "joint", num_bg: int = 20):
|
||||
"""Generator: streams per-step denoise progress to the log, yields the video at the end."""
|
||||
import time
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
clip_idx = int(clip_idx)
|
||||
# Sparse mode uses the same "seed" slider as the tracks background seed AND the noise seed
|
||||
# (fine for interactive use; the two are independent knobs technically).
|
||||
built, err = _build_tracks(traj_json, radius, falloff, sparse, num_bg, seed)
|
||||
if built is None:
|
||||
yield gr.update(), f"[error] {err}", gr.update()
|
||||
return
|
||||
tracks, vis, moving = built
|
||||
steps = int(steps)
|
||||
t0 = time.time()
|
||||
|
||||
yield gr.update(), f"[1/4] swapping to {ckpt_name} ...", gr.update()
|
||||
swap_checkpoint(ckpt_name)
|
||||
|
||||
model = STATE["model"]
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
s = STATE["samples"][clip_idx]
|
||||
ff = s["first_frame_latent"].to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
txt = s["text_embedding"].to(device, dtype)
|
||||
mask = s["text_attention_mask"].to(device, dtype)
|
||||
img = s["clip_feature"].to(device, dtype)
|
||||
tp = torch.from_numpy(tracks)[None].to(device, dtype)
|
||||
tv = torch.from_numpy(vis)[None].to(device, dtype)
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(steps, device=device)
|
||||
|
||||
# Match TrackValidationCallback._sample: MotionStream joint text+motion CFG (Eq. 2).
|
||||
# mode="joint" -> full formula (3 NFE / step). Default. Matches training-time val.
|
||||
# mode="text_only" -> text CFG only, tracks in forward (2 NFE / step).
|
||||
# mode="no_track" -> tracks=None, text CFG only (2 NFE / step).
|
||||
tp_eff = None if mode == "no_track" else tp
|
||||
tv_eff = None if mode == "no_track" else tv
|
||||
wt, wm = float(w_text), float(w_motion)
|
||||
cfg_on = (wt != 1.0) or (wm != 1.0)
|
||||
txt_null = torch.zeros_like(txt) if cfg_on else None
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, tsv):
|
||||
with torch.no_grad(), torch.autocast(device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=tsv, attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=tsv,
|
||||
encoder_hidden_states_image=img, track_points=tp_e,
|
||||
track_visibility=tv_e, return_dict=False)
|
||||
|
||||
for i, t in enumerate(sched.timesteps):
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
ts = t.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp_eff, tv_eff, model_in, ts)
|
||||
if not cfg_on:
|
||||
v = v_full
|
||||
elif tp_eff is None or mode == "text_only":
|
||||
v_no_text = _fwd(txt_null, tp_eff, tv_eff, model_in, ts)
|
||||
v = v_no_text + wt * (v_full - v_no_text)
|
||||
else: # mode="joint" (Eq. 2)
|
||||
v_no_text = _fwd(txt_null, tp_eff, tv_eff, model_in, ts)
|
||||
v_no_motion = _fwd(txt, None, None, model_in, ts)
|
||||
alpha = wt / (wt + wm) if (wt + wm) > 0 else 0.5
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + wt * (v_full - v_no_text) + wm * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), t, latents.float(), return_dict=False)[0]
|
||||
yield gr.update(), f"[2/4] denoising {i + 1}/{steps} ({time.time() - t0:.1f}s, cfg={mode})", gr.update()
|
||||
|
||||
yield gr.update(), f"[3/4] decoding latents ... ({time.time() - t0:.1f}s)", gr.update()
|
||||
frames = twi.decode_to_pixels(model, latents)
|
||||
out = _overlay_grid(frames, tracks, vis, moving)
|
||||
tag = ckpt_name.replace(" ", "")
|
||||
path = os.path.join(STATE["out_dir"], f"action_clip{clip_idx}_{tag}.mp4")
|
||||
imageio.mimsave(path, out, fps=24, macro_block_size=1)
|
||||
|
||||
yield gr.update(), f"[4/4] computing CoTracker EPE ... ({time.time() - t0:.1f}s)", gr.update()
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import compute_epe
|
||||
# EPE on the MOVING points only (did the arm actually follow your path)
|
||||
mv_idx = np.where(moving)[0]
|
||||
ppx = tracks[:, mv_idx].copy()
|
||||
ppx[..., 0] *= W
|
||||
ppx[..., 1] *= H
|
||||
res = compute_epe(frames, ppx, vis[:, mv_idx], STATE["ct"], device)
|
||||
epe = res["epe"]
|
||||
epe_msg = (f"moving-point EPE = {epe:.2f}px ({res['n_points']} moved pts) — "
|
||||
f"lower = the arm region followed your drawn path."
|
||||
if epe is not None else "EPE n/a (no points re-tracked)")
|
||||
yield path, f"[done] generated in {time.time() - t0:.1f}s", epe_msg
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------- ui
|
||||
def _canvas_js(Tpx: int) -> str:
|
||||
return ("() => {\n"
|
||||
" const cvs = document.getElementById('cvs'); if(!cvs||cvs._init) return; cvs._init=true;\n"
|
||||
" const ctx = cvs.getContext('2d');\n"
|
||||
f" const st = {{img:new Image(), mouse:[0.5,0.5], rec:false, traj:[], N:{Tpx}}};\n"
|
||||
" cvs.width=512; cvs.height=288;\n"
|
||||
" function draw(){ ctx.clearRect(0,0,cvs.width,cvs.height);\n"
|
||||
" if(st.img.complete && st.img.width) ctx.drawImage(st.img,0,0,cvs.width,cvs.height);\n"
|
||||
" if(st.traj.length){ ctx.strokeStyle='#00ff66'; ctx.lineWidth=3; ctx.beginPath();\n"
|
||||
" st.traj.forEach((p,i)=>{const x=p[0]*cvs.width,y=p[1]*cvs.height; i?ctx.lineTo(x,y):ctx.moveTo(x,y);}); ctx.stroke();\n"
|
||||
" const s=st.traj[0], e=st.traj[st.traj.length-1];\n"
|
||||
" ctx.fillStyle='#00ff66'; ctx.beginPath(); ctx.arc(s[0]*cvs.width,s[1]*cvs.height,6,0,7); ctx.fill();\n"
|
||||
" ctx.fillStyle='#ff3333'; ctx.beginPath(); ctx.arc(e[0]*cvs.width,e[1]*cvs.height,6,0,7); ctx.fill(); } }\n"
|
||||
" cvs.addEventListener('mousemove', e=>{const r=cvs.getBoundingClientRect(); st.mouse=[(e.clientX-r.left)/r.width,(e.clientY-r.top)/r.height];});\n"
|
||||
" window._setbg=(b64)=>{ if(!b64) return; const im=new Image(); im.onload=()=>{ const ar=im.width/im.height; cvs.width=512; cvs.height=Math.round(512/ar); st.img=im; st.traj=[]; draw(); }; im.src=b64; };\n"
|
||||
" window._clear=()=>{ st.traj=[]; const tb=document.querySelector('#traj_json textarea'); if(tb){tb.value=''; tb.dispatchEvent(new Event('input',{bubbles:true}));} draw(); };\n"
|
||||
" window._startRec=()=>{ if(st.rec) return; st.rec=true; st.traj=[]; let n=0; const stt=document.getElementById('recstatus');\n"
|
||||
" const iv=setInterval(()=>{ st.traj.push([st.mouse[0],st.mouse[1]]); n++; draw(); if(stt) stt.textContent='\\u25CF REC '+n+'/'+st.N;\n"
|
||||
" if(n>=st.N){ clearInterval(iv); st.rec=false; if(stt) stt.textContent='recorded '+st.N+' frames'; \n"
|
||||
" const tb=document.querySelector('#traj_json textarea'); tb.value=JSON.stringify(st.traj); tb.dispatchEvent(new Event('input',{bubbles:true})); } }, 1000/24); };\n"
|
||||
" document.addEventListener('keydown', e=>{ if(e.code==='Space'){ e.preventDefault(); window._startRec(); } });\n"
|
||||
" draw();\n"
|
||||
"}")
|
||||
|
||||
|
||||
def build_ui(num_clips: int, ckpt_names: list[str], Tpx: int):
|
||||
canvas_html = ("<div style='user-select:none'>"
|
||||
"<canvas id='cvs' style='border:1px solid #888;cursor:crosshair;max-width:100%'></canvas>"
|
||||
"<div id='recstatus' style='font-weight:bold;color:#0a0;height:1.4em'></div>"
|
||||
"<div style='font-size:0.85em;color:#888'>Move the mouse over the frame, press "
|
||||
"<b>Space</b> to record your action for 5 s. Green=start, red=end.</div></div>")
|
||||
with gr.Blocks(title="TrackWan — action recorder") as demo:
|
||||
gr.Markdown("## TrackWan — record an action, test action-following across checkpoints\n"
|
||||
"Pick a clip + checkpoint, **press Space** and move the mouse over the frame to draw a "
|
||||
"5 second motion, then **Generate**. The 50x50 grid points near your start are dragged along your path (rest stay static).")
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
ckpt = gr.Dropdown(ckpt_names, value=ckpt_names[-1], label="Checkpoint (hot-swap)")
|
||||
clip = gr.Dropdown(list(range(num_clips)), value=0, label="Clip (first frame + prompt)")
|
||||
caption = gr.Textbox(label="Prompt", interactive=False, lines=2)
|
||||
gr.HTML(canvas_html)
|
||||
with gr.Row():
|
||||
rec_btn = gr.Button("● Record (Space)")
|
||||
clr_btn = gr.Button("Clear")
|
||||
with gr.Row():
|
||||
radius = gr.Slider(0.02, 0.5, value=0.15, step=0.01,
|
||||
label="Action radius (grid pts within this move)")
|
||||
falloff = gr.Radio(["hard", "smooth"], value="hard", label="Falloff")
|
||||
with gr.Row():
|
||||
background = gr.Radio(["dense", "sparse"],
|
||||
value="sparse",
|
||||
label="Track budget (matches training: WANTRACK_SPARSE)")
|
||||
num_bg = gr.Slider(0, 60, value=20, step=1,
|
||||
label="Sparse: # background points (WANTRACK_EXTRA_RANDOM)")
|
||||
with gr.Row():
|
||||
steps = gr.Slider(10, 60, value=30, step=5, label="Denoise steps")
|
||||
seed = gr.Number(value=1000, label="Seed (also seeds sparse background pool)")
|
||||
with gr.Row():
|
||||
w_text = gr.Slider(1.0, 8.0, value=3.0, step=0.5, label="Text CFG (w_t)")
|
||||
w_motion = gr.Slider(1.0, 5.0, value=1.5, step=0.25, label="Motion CFG (w_m)")
|
||||
mode = gr.Radio(["joint", "text_only", "no_track"],
|
||||
value="joint",
|
||||
label="CFG mode (joint = Eq. 2; text_only = drop motion arm; no_track = tracks=None)")
|
||||
with gr.Row():
|
||||
preview_btn = gr.Button("Preview traces")
|
||||
gen_btn = gr.Button("Generate", variant="primary")
|
||||
status = gr.Textbox(label="Status / generation log", interactive=False, lines=2)
|
||||
with gr.Column(scale=1):
|
||||
orig_vid = gr.Video(label="Original clip")
|
||||
gen_vid = gr.Video(label="Preview traces / Generated (your action overlaid)")
|
||||
epe_box = gr.Textbox(label="Action-following (moving-point CoTracker EPE)", interactive=False)
|
||||
|
||||
ff_b64 = gr.Textbox(elem_id="ff_b64", visible=False)
|
||||
traj_json = gr.Textbox(elem_id="traj_json", visible=False)
|
||||
|
||||
clip.change(on_clip, [clip], [caption, ff_b64, orig_vid])
|
||||
ckpt.change(on_ckpt, [ckpt], [status])
|
||||
ff_b64.change(None, [ff_b64], None, js="(b64)=>{ window._setbg(b64); }")
|
||||
rec_btn.click(None, None, None, js="()=>{ window._startRec(); }")
|
||||
clr_btn.click(None, None, None, js="()=>{ window._clear(); }")
|
||||
preview_btn.click(on_preview,
|
||||
[clip, traj_json, radius, falloff, background, num_bg, seed],
|
||||
[gen_vid, status])
|
||||
gen_btn.click(on_generate,
|
||||
[clip, ckpt, traj_json, radius, falloff, background, steps, seed,
|
||||
w_text, w_motion, mode, num_bg],
|
||||
[gen_vid, status, epe_box])
|
||||
demo.load(on_clip, [clip], [caption, ff_b64, orig_vid])
|
||||
demo.load(None, None, None, js=_canvas_js(Tpx))
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--exports-root", required=True,
|
||||
help="dir containing export_* checkpoint dirs (each a dcp_to_diffusers export)")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp_now/combined_parquet_dataset")
|
||||
p.add_argument("--num-clips", type=int, default=10)
|
||||
p.add_argument("--out-dir", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/gradio_action_out")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7861)
|
||||
p.add_argument("--share", action="store_true")
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
roots = sorted(glob.glob(os.path.join(args.exports_root, "export_*")))
|
||||
if not roots:
|
||||
raise SystemExit(f"no export_* dirs under {args.exports_root}")
|
||||
exports = {f"step {os.path.basename(r).split('_')[-1]}": r for r in roots}
|
||||
names = list(exports.keys())
|
||||
|
||||
model, tc = twi.load_trackwan(exports[names[-1]], args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, list(range(args.num_clips)), text_len)
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import load_cotracker
|
||||
STATE.update(model=model, tc=tc, samples=samples, Tpx=(num_lat_t - 1) * ratio + 1,
|
||||
out_dir=args.out_dir, ct=load_cotracker(model.device), exports=exports,
|
||||
cur_ckpt=names[-1])
|
||||
|
||||
demo = build_ui(len(samples), names, STATE["Tpx"])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[args.out_dir])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,543 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""TrackWan interactive rollout demo.
|
||||
|
||||
Two point-input modes on the same canvas:
|
||||
- **Trace mode** — hover, press *Space* to record a 5-second (121-frame) drag.
|
||||
- **Anchor mode** — click once to drop a static point (a track whose (x,y) is
|
||||
constant across all 121 frames).
|
||||
|
||||
Design (matches training-time validation exactly):
|
||||
- Loads conditioning (text_embedding, clip_feature, first_frame_latent) DIRECTLY from
|
||||
the preprocessed openvid parquet. Same path as
|
||||
``fastvideo/train/callbacks/track_validation.py``.
|
||||
- First frame image shown to the user is decoded from the VAE latent (``decode_reference``).
|
||||
- Both traces and anchors are packed into a single ``track_points[T,N,2]`` tensor.
|
||||
Anchors get their (x,y) broadcast across T=121; traces are linearly interpolated.
|
||||
- Generation: the SAME motion-CFG denoise loop as ``TrackValidationCallback._sample``
|
||||
(MotionStream Eq. 2, wt=3.0, wm=1.5 by default).
|
||||
|
||||
Launch::
|
||||
|
||||
srun --overlap --jobid=<job> --ntasks=1 -w <node> --chdir=$PWD bash -lc "
|
||||
source .venv/bin/activate
|
||||
export ... TRACKWAN_TRACK_BIAS=1 CUDA_VISIBLE_DEVICES=0
|
||||
python examples/inference/gradio/trackwan/app_multi_trace.py \\
|
||||
--model-dir /path/to/diffusers_export --yaml <yaml> \\
|
||||
--data-path /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
"
|
||||
|
||||
share=True doesn't work on aarch64 (no frpc); use SSH port-forward on the compute node.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
REPO = Path(__file__).resolve().parents[4]
|
||||
NUM_FRAMES = 121
|
||||
FPS = 24
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# frame b64 + track construction
|
||||
# =============================================================================
|
||||
def img_to_b64(arr: np.ndarray) -> str:
|
||||
buf = io.BytesIO()
|
||||
Image.fromarray(arr).save(buf, format="PNG")
|
||||
return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode()
|
||||
|
||||
|
||||
def build_track_points(traces_json: str, anchors_json: str) -> tuple[np.ndarray, np.ndarray, int, int]:
|
||||
"""(traces, anchors) → (track_points[T,N,2], track_visibility[T,N], n_tr, n_an).
|
||||
|
||||
Each trace is a list of normalized (x,y) samples at ~24 Hz over 5 s → interpolated to T=121.
|
||||
Each anchor is a single normalized (x,y) → broadcast across all T frames.
|
||||
"""
|
||||
T = NUM_FRAMES
|
||||
try:
|
||||
traces = json.loads(traces_json) if traces_json else []
|
||||
except Exception:
|
||||
traces = []
|
||||
try:
|
||||
anchors = json.loads(anchors_json) if anchors_json else []
|
||||
except Exception:
|
||||
anchors = []
|
||||
n_tr, n_an = len(traces), len(anchors)
|
||||
N = n_tr + n_an
|
||||
if N == 0:
|
||||
return np.zeros((T, 0, 2), np.float32), np.zeros((T, 0), np.float32), 0, 0
|
||||
tp = np.zeros((T, N, 2), np.float32)
|
||||
tv = np.ones((T, N), np.float32)
|
||||
for i, traj in enumerate(traces):
|
||||
if not traj:
|
||||
continue
|
||||
arr = np.asarray(traj, np.float32)
|
||||
if arr.shape[0] < 2:
|
||||
tp[:, i] = arr[0]
|
||||
else:
|
||||
src = np.arange(arr.shape[0])
|
||||
tgt = np.linspace(0, arr.shape[0] - 1, T)
|
||||
tp[:, i, 0] = np.interp(tgt, src, arr[:, 0])
|
||||
tp[:, i, 1] = np.interp(tgt, src, arr[:, 1])
|
||||
for j, (ax, ay) in enumerate(anchors):
|
||||
tp[:, n_tr + j, 0] = ax
|
||||
tp[:, n_tr + j, 1] = ay
|
||||
return tp, tv, n_tr, n_an
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# canvas JS (traces via Space + anchors via click, mode radio decides)
|
||||
# =============================================================================
|
||||
def canvas_js(num_frames: int, fps: int) -> str:
|
||||
palette = ["#ff3c3c", "#3cd23c", "#3c78ff", "#ffc83c", "#c83cc8", "#3cdcdc", "#ff821e", "#9696ff"]
|
||||
return ("() => {\n"
|
||||
" const cvs = document.getElementById('cvs'); if (!cvs || cvs._init) return; cvs._init=true;\n"
|
||||
" const ctx = cvs.getContext('2d');\n"
|
||||
f" const PAL = {json.dumps(palette)};\n"
|
||||
f" const N = {num_frames}; const DT = {int(1000 / fps)};\n"
|
||||
" const st = { img:new Image(), mouse:[0.5,0.5], rec:false, current:[], "
|
||||
"traces:[], anchors:[], mode:'trace' };\n"
|
||||
" window._st = st;\n"
|
||||
" cvs.width = 832; cvs.height = 480;\n"
|
||||
" function commit_traces(){\n"
|
||||
" const tb = document.querySelector('#traces_json textarea');\n"
|
||||
" if (tb) { tb.value = JSON.stringify(st.traces); tb.dispatchEvent(new Event('input',{bubbles:true})); }\n"
|
||||
" }\n"
|
||||
" function commit_anchors(){\n"
|
||||
" const tb = document.querySelector('#anchors_json textarea');\n"
|
||||
" if (tb) { tb.value = JSON.stringify(st.anchors); tb.dispatchEvent(new Event('input',{bubbles:true})); }\n"
|
||||
" }\n"
|
||||
" function draw(){\n"
|
||||
" ctx.clearRect(0,0,cvs.width,cvs.height);\n"
|
||||
" if (st.img.complete && st.img.width) ctx.drawImage(st.img, 0, 0, cvs.width, cvs.height);\n"
|
||||
" // anchors: solid gray dots with dark ring\n"
|
||||
" st.anchors.forEach((p, i) => { const x=p[0]*cvs.width, y=p[1]*cvs.height;\n"
|
||||
" ctx.beginPath(); ctx.arc(x,y,7,0,7); ctx.fillStyle='#888'; ctx.fill();\n"
|
||||
" ctx.strokeStyle='#000'; ctx.lineWidth=2; ctx.stroke();\n"
|
||||
" ctx.fillStyle='#fff'; ctx.font='bold 10px sans-serif'; ctx.textAlign='center'; ctx.textBaseline='middle';\n"
|
||||
" ctx.fillText('A'+(i+1), x, y);\n"
|
||||
" });\n"
|
||||
" // traces: colored polyline with start/end markers\n"
|
||||
" st.traces.forEach((traj, i) => {\n"
|
||||
" const color = PAL[i % PAL.length];\n"
|
||||
" ctx.strokeStyle=color; ctx.lineWidth=3; ctx.beginPath();\n"
|
||||
" traj.forEach((p, j) => { const x=p[0]*cvs.width, y=p[1]*cvs.height; j?ctx.lineTo(x,y):ctx.moveTo(x,y); });\n"
|
||||
" ctx.stroke();\n"
|
||||
" const s = traj[0], e = traj[traj.length-1];\n"
|
||||
" ctx.fillStyle=color; ctx.beginPath(); ctx.arc(s[0]*cvs.width, s[1]*cvs.height, 6, 0, 7); ctx.fill(); ctx.strokeStyle='#000'; ctx.lineWidth=1; ctx.stroke();\n"
|
||||
" ctx.fillStyle='#ff3333'; ctx.beginPath(); ctx.arc(e[0]*cvs.width, e[1]*cvs.height, 6, 0, 7); ctx.fill(); ctx.strokeStyle='#000'; ctx.stroke();\n"
|
||||
" ctx.fillStyle='#000'; ctx.font='13px sans-serif'; ctx.textAlign='left'; ctx.textBaseline='alphabetic'; ctx.fillText('T'+(i+1), s[0]*cvs.width+8, s[1]*cvs.height-8);\n"
|
||||
" });\n"
|
||||
" // live-drawing current trace\n"
|
||||
" if (st.rec && st.current.length) {\n"
|
||||
" const color = PAL[st.traces.length % PAL.length];\n"
|
||||
" ctx.strokeStyle=color; ctx.lineWidth=3; ctx.beginPath();\n"
|
||||
" st.current.forEach((p, j) => { const x=p[0]*cvs.width, y=p[1]*cvs.height; j?ctx.lineTo(x,y):ctx.moveTo(x,y); });\n"
|
||||
" ctx.stroke();\n"
|
||||
" }\n"
|
||||
" }\n"
|
||||
" cvs.addEventListener('mousemove', e => {\n"
|
||||
" const r = cvs.getBoundingClientRect();\n"
|
||||
" st.mouse = [(e.clientX - r.left) / r.width, (e.clientY - r.top) / r.height];\n"
|
||||
" });\n"
|
||||
" cvs.addEventListener('click', e => {\n"
|
||||
" if (st.rec) return;\n"
|
||||
" const r = cvs.getBoundingClientRect();\n"
|
||||
" const cx = (e.clientX - r.left) * (cvs.width / r.width);\n"
|
||||
" const cy = (e.clientY - r.top) * (cvs.height / r.height);\n"
|
||||
" if (st.mode === 'anchor') {\n"
|
||||
" // toggle delete-if-near, else place\n"
|
||||
" let idx = -1, best = 14*14;\n"
|
||||
" st.anchors.forEach((p, i) => { const dx = p[0]*cvs.width - cx, dy = p[1]*cvs.height - cy;\n"
|
||||
" const d2 = dx*dx + dy*dy; if (d2 < best) { best = d2; idx = i; } });\n"
|
||||
" if (idx >= 0) st.anchors.splice(idx, 1);\n"
|
||||
" else st.anchors.push([cx / cvs.width, cy / cvs.height]);\n"
|
||||
" commit_anchors(); draw();\n"
|
||||
" }\n"
|
||||
" });\n"
|
||||
" window._setmode = (m) => { st.mode = m || 'trace'; };\n"
|
||||
" window._setbg = (b64) => {\n"
|
||||
" if (!b64) return;\n"
|
||||
" const im = new Image();\n"
|
||||
" im.onload = () => { const ar = im.width / im.height; cvs.width = 832; cvs.height = Math.round(832 / ar); st.img = im; draw(); };\n"
|
||||
" im.src = b64;\n"
|
||||
" };\n"
|
||||
" window._clear_traces = () => {\n"
|
||||
" st.traces = []; st.current = []; st.rec = false; commit_traces(); draw();\n"
|
||||
" const stt = document.getElementById('recstatus'); if (stt) stt.textContent = 'traces cleared. Space for trace #1';\n"
|
||||
" };\n"
|
||||
" window._undo_trace = () => {\n"
|
||||
" if (st.traces.length) st.traces.pop();\n"
|
||||
" commit_traces(); draw();\n"
|
||||
" const stt = document.getElementById('recstatus'); if (stt) stt.textContent = st.traces.length + ' trace(s). Space for another';\n"
|
||||
" };\n"
|
||||
" window._clear_anchors = () => { st.anchors = []; commit_anchors(); draw(); };\n"
|
||||
" window._undo_anchor = () => { if (st.anchors.length) st.anchors.pop(); commit_anchors(); draw(); };\n"
|
||||
" window._startRec = () => {\n"
|
||||
" if (st.rec) return;\n"
|
||||
" st.rec = true; st.current = []; let n = 0;\n"
|
||||
" const stt = document.getElementById('recstatus');\n"
|
||||
" const iv = setInterval(() => {\n"
|
||||
" st.current.push([st.mouse[0], st.mouse[1]]); n++; draw();\n"
|
||||
" if (stt) stt.textContent = 'REC trace #' + (st.traces.length + 1) + ' — ' + n + '/' + N;\n"
|
||||
" if (n >= N) {\n"
|
||||
" clearInterval(iv); st.rec = false;\n"
|
||||
" st.traces.push(st.current); st.current = [];\n"
|
||||
" if (stt) stt.textContent = 'committed trace #' + st.traces.length + ' (Space for another, or Generate)';\n"
|
||||
" commit_traces(); draw();\n"
|
||||
" }\n"
|
||||
" }, DT);\n"
|
||||
" };\n"
|
||||
" document.addEventListener('keydown', e => {\n"
|
||||
" if (e.code !== 'Space' || e.repeat) return;\n"
|
||||
" const tag = document.activeElement && document.activeElement.tagName;\n"
|
||||
" if (tag === 'INPUT' || tag === 'TEXTAREA') return;\n"
|
||||
" if (st.mode !== 'trace') return; // Space only records in trace mode\n"
|
||||
" e.preventDefault(); window._startRec();\n"
|
||||
" });\n"
|
||||
" draw();\n"
|
||||
"}")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Gradio UI
|
||||
# =============================================================================
|
||||
def build_ui(state):
|
||||
canvas_html = ("<div style='user-select:none'>"
|
||||
"<canvas id='cvs' style='border:1px solid #888;cursor:crosshair;max-width:100%'></canvas>"
|
||||
"<div id='recstatus' style='font-weight:bold;color:#0a0;height:1.4em;margin-top:6px'>"
|
||||
"pick a clip, click <b>Load frame</b>, then choose a mode below</div>"
|
||||
"</div>")
|
||||
labels = state["labels"]
|
||||
with gr.Blocks(title="TrackWan interactive rollout") as demo:
|
||||
gr.Markdown(
|
||||
"## TrackWan — interactive rollout demo\n"
|
||||
"1. Pick a preprocessed clip → **Load frame**.\n"
|
||||
"2. Choose a mode:\n"
|
||||
" - **Trace** — hover on the frame and press **Space** to record a 5 s drag.\n"
|
||||
" - **Anchor** — click on the frame to drop a *static* point. Click an existing anchor to remove it.\n"
|
||||
"3. **Generate**. Denoise formula matches training's validation exactly (MotionStream Eq. 2).")
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
clip = gr.Dropdown(labels, value=labels[0] if labels else None, label="Preprocessed clip")
|
||||
caption = gr.Textbox(label="Caption (from parquet)", interactive=False, lines=2)
|
||||
load_btn = gr.Button("Load frame", variant="primary")
|
||||
|
||||
gr.Markdown("**Input mode**")
|
||||
mode = gr.Radio(["trace", "anchor"], value="trace",
|
||||
label="Mode",
|
||||
info="Trace = Space to record a moving drag. Anchor = single click for a static point.")
|
||||
|
||||
gr.Markdown("**Traces** (moving)")
|
||||
with gr.Row():
|
||||
rec_btn = gr.Button("● Record (Space)")
|
||||
undo_trace_btn = gr.Button("Undo trace")
|
||||
clear_trace_btn = gr.Button("Clear traces")
|
||||
|
||||
gr.Markdown("**Anchors** (static)")
|
||||
with gr.Row():
|
||||
undo_anchor_btn = gr.Button("Undo anchor")
|
||||
clear_anchor_btn = gr.Button("Clear anchors")
|
||||
|
||||
info_md = gr.Markdown("_no traces or anchors yet._")
|
||||
|
||||
gr.Markdown("**Generation** (denoise formula matches training's validation exactly)")
|
||||
seed_in = gr.Slider(0, 9999, 1000, step=1, label="Seed")
|
||||
steps_in = gr.Slider(10, 60, 30, step=1, label="Denoise steps")
|
||||
w_text = gr.Slider(1.0, 8.0, 3.0, step=0.25, label="Text CFG (w_t)")
|
||||
w_motion = gr.Slider(1.0, 5.0, 1.5, step=0.25, label="Motion CFG (w_m)")
|
||||
gen_btn = gr.Button("Generate video", variant="primary")
|
||||
with gr.Column(scale=2):
|
||||
gr.HTML(canvas_html)
|
||||
out_video = gr.Video(label="Generated video", height=380)
|
||||
status = gr.Markdown("_ready_")
|
||||
log_box = gr.Textbox(label="Generation log", interactive=False, lines=8,
|
||||
max_lines=16, value="_no runs yet._")
|
||||
|
||||
# hidden bridges
|
||||
ff_b64 = gr.Textbox(elem_id="ff_b64", visible=False)
|
||||
traces_json = gr.Textbox(elem_id="traces_json", visible=False, value="[]")
|
||||
anchors_json = gr.Textbox(elem_id="anchors_json", visible=False, value="[]")
|
||||
|
||||
# server-side per-page state
|
||||
st_clip_idx = gr.State(0)
|
||||
st_frame = gr.State(None)
|
||||
|
||||
# ---------------- caption on clip change ----------------
|
||||
def _cap(clip_name):
|
||||
i = state["label_to_idx"].get(clip_name, 0)
|
||||
return i, state["samples"][i]["caption"][:400]
|
||||
clip.change(_cap, [clip], [st_clip_idx, caption])
|
||||
|
||||
# ---------------- load frame (decode ref frame from VAE latent) ----------------
|
||||
first_frame_cache: dict[int, np.ndarray] = {}
|
||||
|
||||
def _load(clip_name):
|
||||
i = state["label_to_idx"].get(clip_name, 0)
|
||||
if i not in first_frame_cache:
|
||||
first_frame_cache[i] = state["decode_first_frame"](i)
|
||||
frame = first_frame_cache[i]
|
||||
msg = f"_clip loaded ({state['samples'][i]['caption'][:60]}...)_"
|
||||
return i, frame, img_to_b64(frame), msg
|
||||
|
||||
load_btn.click(_load, [clip], [st_clip_idx, st_frame, ff_b64, status])
|
||||
|
||||
# preload first clip on startup
|
||||
def _preload():
|
||||
i = 0
|
||||
if i not in first_frame_cache:
|
||||
first_frame_cache[i] = state["decode_first_frame"](i)
|
||||
frame = first_frame_cache[i]
|
||||
return (i, state["samples"][i]["caption"][:400], frame, img_to_b64(frame),
|
||||
"_default clip preloaded. Draw a trace (Space) or place anchors (click)._")
|
||||
demo.load(_preload, None, [st_clip_idx, caption, st_frame, ff_b64, status])
|
||||
|
||||
# ---------------- JS bridges ----------------
|
||||
ff_b64.change(None, ff_b64, None, js="(b) => window._setbg && window._setbg(b)")
|
||||
mode.change(None, mode, None, js="(m) => window._setmode && window._setmode(m)")
|
||||
rec_btn.click(None, None, None, js="() => window._startRec && window._startRec()")
|
||||
undo_trace_btn.click(None, None, None, js="() => window._undo_trace && window._undo_trace()")
|
||||
clear_trace_btn.click(None, None, None, js="() => window._clear_traces && window._clear_traces()")
|
||||
undo_anchor_btn.click(None, None, None, js="() => window._undo_anchor && window._undo_anchor()")
|
||||
clear_anchor_btn.click(None, None, None, js="() => window._clear_anchors && window._clear_anchors()")
|
||||
|
||||
def _tinfo(t, a):
|
||||
try:
|
||||
n_t = len(json.loads(t) if t else [])
|
||||
except Exception:
|
||||
n_t = 0
|
||||
try:
|
||||
n_a = len(json.loads(a) if a else [])
|
||||
except Exception:
|
||||
n_a = 0
|
||||
if n_t == 0 and n_a == 0:
|
||||
return "_no traces or anchors yet._"
|
||||
return f"_{n_t} trace(s) + {n_a} anchor(s) → {n_t + n_a} total track(s)_"
|
||||
traces_json.change(_tinfo, [traces_json, anchors_json], info_md)
|
||||
anchors_json.change(_tinfo, [traces_json, anchors_json], info_md)
|
||||
|
||||
# ---------------- generate (MATCHES TrackValidationCallback._sample exactly) ----------------
|
||||
def _generate(clip_idx, traces_str, anchors_str, seed, steps, w_t, w_m):
|
||||
def logln(log, msg):
|
||||
line = time.strftime("%H:%M:%S ") + msg
|
||||
print(line, flush=True)
|
||||
return (log + "\n" + line) if log else line
|
||||
|
||||
log = ""
|
||||
tp_np, tv_np, n_tr, n_an = build_track_points(traces_str, anchors_str)
|
||||
log = logln(log, f"clip idx={clip_idx}: {n_tr} trace(s) + {n_an} anchor(s) → N={tp_np.shape[1]}")
|
||||
if tp_np.shape[1] == 0:
|
||||
return None, "_record at least one trace or place at least one anchor_", log
|
||||
|
||||
ts = time.strftime("%Y%m%d_%H%M%S")
|
||||
req_dir = Path(state["out_dir"]) / f"req_{ts}"
|
||||
req_dir.mkdir(parents=True, exist_ok=True)
|
||||
np.savez(req_dir / "tracks.npz", track_points=tp_np, track_visibility=tv_np)
|
||||
(req_dir / "meta.json").write_text(json.dumps({
|
||||
"clip_idx": int(clip_idx), "seed": int(seed), "steps": int(steps),
|
||||
"w_text": float(w_t), "w_motion": float(w_m),
|
||||
"num_traces": n_tr, "num_anchors": n_an,
|
||||
}, indent=2))
|
||||
log = logln(log, f"req -> {req_dir.name}, denoising ...")
|
||||
|
||||
try:
|
||||
t0 = time.time()
|
||||
mp4 = state["run_generate"](int(clip_idx), tp_np, tv_np,
|
||||
int(seed), int(steps), float(w_t), float(w_m),
|
||||
req_dir)
|
||||
elapsed = time.time() - t0
|
||||
log = logln(log, f"done in {elapsed:.1f}s → {mp4}")
|
||||
return str(mp4), f"_generated in {elapsed:.1f}s_", log
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return None, f"_generation failed: {e}_", logln(log, f"EXC: {e}")
|
||||
|
||||
gen_btn.click(_generate,
|
||||
[st_clip_idx, traces_json, anchors_json, seed_in, steps_in, w_text, w_motion],
|
||||
[out_video, status, log_box])
|
||||
|
||||
demo.load(None, None, None, js=canvas_js(NUM_FRAMES, FPS))
|
||||
|
||||
return demo
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Generation (MATCHES TrackValidationCallback._sample)
|
||||
# =============================================================================
|
||||
def make_generator(model, tc, samples):
|
||||
"""Return a callable (clip_idx, tp_np, tv_np, seed, steps, w_t, w_m, out_dir) -> mp4 path.
|
||||
|
||||
The denoise formula is copied verbatim from
|
||||
``fastvideo/train/callbacks/track_validation.py::_sample`` (MotionStream Eq. 2).
|
||||
"""
|
||||
import torch
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
flow_shift = float(model.timestep_shift)
|
||||
transformer = model.transformer
|
||||
|
||||
@torch.no_grad()
|
||||
def _run(clip_idx: int, tp_np: np.ndarray, tv_np: np.ndarray,
|
||||
seed: int, steps: int, w_t: float, w_m: float, out_dir: Path) -> str:
|
||||
s = samples[clip_idx]
|
||||
ff = s["first_frame_latent"].to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
txt = s["text_embedding"].to(device, dtype)
|
||||
mask = s["text_attention_mask"].to(device, dtype)
|
||||
img = s["clip_feature"].to(device, dtype)
|
||||
# user tracks (override the sample's tracks)
|
||||
tp = torch.from_numpy(tp_np)[None].to(device, dtype) # [1,T,N,2]
|
||||
tv = torch.from_numpy(tv_np)[None].to(device, dtype) # [1,T,N]
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=flow_shift)
|
||||
sched.set_timesteps(int(steps), device=device)
|
||||
cfg_on = (w_t != 1.0) or (w_m != 1.0)
|
||||
txt_null = torch.zeros_like(txt) if cfg_on else None
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, tsv):
|
||||
with torch.autocast(device.type, dtype=dtype), set_forward_context(current_timestep=tsv,
|
||||
attn_metadata=None):
|
||||
return transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=tsv,
|
||||
encoder_hidden_states_image=img, track_points=tp_e,
|
||||
track_visibility=tv_e, return_dict=False)
|
||||
|
||||
for tt in sched.timesteps:
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
tss = tt.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp, tv, model_in, tss)
|
||||
if not cfg_on:
|
||||
v = v_full
|
||||
else:
|
||||
v_no_text = _fwd(txt_null, tp, tv, model_in, tss)
|
||||
v_no_motion = _fwd(txt, None, None, model_in, tss)
|
||||
alpha = w_t / (w_t + w_m) if (w_t + w_m) > 0 else 0.5
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + w_t * (v_full - v_no_text) + w_m * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), tt, latents.float(), return_dict=False)[0]
|
||||
|
||||
px = model.decode_latents(latents.permute(0, 2, 1, 3, 4))[0] # [3,T,H,W] in [0,1]
|
||||
video = (px.clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
frames = np.transpose(video, (1, 2, 3, 0))
|
||||
mp4 = out_dir / "generation.mp4"
|
||||
imageio.mimsave(mp4, frames, fps=FPS, macro_block_size=1)
|
||||
return str(mp4)
|
||||
|
||||
return _run
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Efficient parquet loader (reads only enough files to get N rows, NOT all 4494)
|
||||
# =============================================================================
|
||||
def _load_first_n_from_parquet(data_path: str, n: int, text_len: int) -> list:
|
||||
"""Grab the first N rows from a preprocessed parquet dataset."""
|
||||
import glob
|
||||
import pyarrow.parquet as pq
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
|
||||
files = sorted(glob.glob(os.path.join(data_path, "**", "*.parquet"), recursive=True))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"no *.parquet under {data_path}")
|
||||
rows: list = []
|
||||
for f in files:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
if len(rows) >= n:
|
||||
break
|
||||
sel = rows[:n]
|
||||
print(f"[app] loaded {len(sel)} rows from parquet", flush=True)
|
||||
batch = collate_rows_from_parquet_schema(sel, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in sel]
|
||||
out = []
|
||||
for i in range(len(sel)):
|
||||
out.append({
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i:i + 1].clone(),
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(),
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# main
|
||||
# =============================================================================
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--model-dir", required=True, help="diffusers-exported ckpt dir")
|
||||
ap.add_argument("--yaml", required=True, help="training yaml")
|
||||
ap.add_argument("--data-path", required=True,
|
||||
help="preprocessed parquet root (combined_parquet_dataset)")
|
||||
ap.add_argument("--num-clips", type=int, default=30, help="how many parquet clips to preload")
|
||||
ap.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/multi_trace_out")
|
||||
ap.add_argument("--host", default="0.0.0.0")
|
||||
ap.add_argument("--port", type=int, default=7864)
|
||||
args = ap.parse_args()
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
state: dict = {"out_dir": args.out_dir}
|
||||
|
||||
# Load model
|
||||
print(f"[app] loading model from {args.model_dir} ...", flush=True)
|
||||
sys.path.insert(0, str(REPO / "data_pipeline"))
|
||||
import trackwan_infer as twi
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
state["model"] = model
|
||||
state["tc"] = tc
|
||||
|
||||
# Load parquet conditioning efficiently
|
||||
print(f"[app] loading {args.num_clips} parquet clips ...", flush=True)
|
||||
samples = _load_first_n_from_parquet(args.data_path, args.num_clips, text_len)
|
||||
labels = []
|
||||
label_to_idx = {}
|
||||
for i, s in enumerate(samples):
|
||||
cap = (s["caption"] or f"clip_{i}").strip().replace("\n", " ")
|
||||
lab = f"{i:03d} {cap[:60]}"
|
||||
labels.append(lab)
|
||||
label_to_idx[lab] = i
|
||||
state["samples"] = samples
|
||||
state["labels"] = labels
|
||||
state["label_to_idx"] = label_to_idx
|
||||
|
||||
# First-frame decoder (from vae_latent) — used as the drawing background
|
||||
def _decode_first_frame(i: int) -> np.ndarray:
|
||||
return twi.decode_reference(model, samples[i]["vae_latent"])[0] # [H,W,3] uint8
|
||||
state["decode_first_frame"] = _decode_first_frame
|
||||
|
||||
# Generator (validation callback's exact denoise loop)
|
||||
state["run_generate"] = make_generator(model, tc, samples)
|
||||
|
||||
print(f"[app] ready. {len(samples)} clips loaded.", flush=True)
|
||||
demo = build_ui(state)
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port,
|
||||
share=False, allowed_paths=[args.out_dir])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,127 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""TrackWan eval-artifact VIEWER — browse pre-generated controllability videos.
|
||||
|
||||
No GPU / no model: just serves the artifacts written by
|
||||
``data_pipeline/gen_eval_artifacts.py``. Pick a checkpoint, a clip, and a
|
||||
counterfactual control; see, side by side:
|
||||
|
||||
- the original clip (and the original with its GT tracks),
|
||||
- the generation (raw),
|
||||
- the generation with the INPUT control tracks (cyan = what we asked for),
|
||||
- the generation with the tracks RE-EXTRACTED from it (yellow = what it did),
|
||||
- the EPE heatmap (per-point error: green = followed, red = ignored),
|
||||
|
||||
plus the CoTracker EPE for that control.
|
||||
|
||||
Launch (no GPU needed)::
|
||||
|
||||
.venv/bin/python examples/inference/gradio/trackwan/app_viewer.py \
|
||||
--artifacts-root /.../eval_artifacts --share
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
|
||||
import gradio as gr
|
||||
|
||||
CONTROL_ORDER = ["gt", "none", "pan_right", "zoom_in", "drag_dense", "drag_sparse", "swap"]
|
||||
STEPS: dict = {} # name -> {"dir": path, "m": manifest}
|
||||
|
||||
|
||||
def _p(ckpt, fname):
|
||||
if not fname:
|
||||
return None
|
||||
path = os.path.join(STEPS[ckpt]["dir"], fname)
|
||||
return path if os.path.exists(path) else None
|
||||
|
||||
|
||||
def update(ckpt, clip, control):
|
||||
if ckpt not in STEPS:
|
||||
return [None] * 7 + ["", ""]
|
||||
clips = STEPS[ckpt]["m"]["clips"]
|
||||
e = clips.get(str(clip))
|
||||
if e is None:
|
||||
return [None] * 7 + ["(clip not in this checkpoint)", ""]
|
||||
c = e["controls"].get(control, {})
|
||||
parts = []
|
||||
if c.get("epe") is not None:
|
||||
parts.append(f"EPE_all={c['epe']:.1f}px")
|
||||
if c.get("epe_moving") is not None:
|
||||
parts.append(f"EPE_moving={c['epe_moving']:.1f}px ({c.get('n_moving', 0)} moving pts)")
|
||||
if c.get("sensitivity_roi") is not None:
|
||||
parts.append(f"sensitivity_roi={c['sensitivity_roi']:.3f}")
|
||||
if c.get("localization_iou") is not None:
|
||||
parts.append(f"localization_IoU={c['localization_iou']:.3f}")
|
||||
if c.get("bg_leakage") is not None:
|
||||
parts.append(f"bg_leakage={c['bg_leakage']:.3f}")
|
||||
txt = " | ".join(parts) if parts else "n/a (GT baseline / no control tracks)"
|
||||
return (_p(ckpt, e.get("original")), _p(ckpt, e.get("original_tracks")),
|
||||
_p(ckpt, c.get("gen")), _p(ckpt, c.get("input")),
|
||||
_p(ckpt, c.get("tracked")), _p(ckpt, c.get("heat")), _p(ckpt, c.get("diff")),
|
||||
e.get("caption", ""), txt)
|
||||
|
||||
|
||||
def build_ui(ckpt_names, clip_ids):
|
||||
with gr.Blocks(title="TrackWan — eval viewer") as demo:
|
||||
gr.Markdown("## TrackWan — controllability eval viewer\n"
|
||||
"Browse pre-generated videos per **checkpoint × clip × control**. "
|
||||
"Cyan = input tracks (asked), yellow = re-tracked from the generation (did), "
|
||||
"EPE heatmap green→red = per-point EPE (followed→ignored). "
|
||||
"**Diff** = |this control − GT-track generation|: bright where the control changed the "
|
||||
"output (the real test — dark everywhere = the model ignored the control).")
|
||||
with gr.Row():
|
||||
ckpt = gr.Dropdown(ckpt_names, value=ckpt_names[0], label="Checkpoint")
|
||||
clip = gr.Dropdown(clip_ids, value=clip_ids[0], label="Clip")
|
||||
control = gr.Dropdown(CONTROL_ORDER, value="drag_dense", label="Control")
|
||||
caption = gr.Textbox(label="Prompt", interactive=False, lines=2)
|
||||
epe_box = gr.Textbox(label="Metrics (EPE_moving + counterfactual sensitivity/IoU = the honest ones)",
|
||||
interactive=False)
|
||||
with gr.Row():
|
||||
v_orig = gr.Video(label="Original clip")
|
||||
v_orig_tr = gr.Video(label="Original + GT tracks")
|
||||
with gr.Row():
|
||||
v_gen = gr.Video(label="Generated (raw)")
|
||||
v_input = gr.Video(label="Generated + INPUT tracks (cyan)")
|
||||
with gr.Row():
|
||||
v_tracked = gr.Video(label="Generated + RE-TRACKED (yellow)")
|
||||
v_heat = gr.Video(label="EPE heatmap (green=followed, red=ignored)")
|
||||
with gr.Row():
|
||||
v_diff = gr.Video(label="DIFF vs GT-track gen (bright = control changed the output)")
|
||||
gr.Markdown("")
|
||||
|
||||
outs = [v_orig, v_orig_tr, v_gen, v_input, v_tracked, v_heat, v_diff, caption, epe_box]
|
||||
for comp in (ckpt, clip, control):
|
||||
comp.change(update, [ckpt, clip, control], outs)
|
||||
demo.load(update, [ckpt, clip, control], outs)
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--artifacts-root", required=True, help="dir containing step*/manifest.json")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7870)
|
||||
p.add_argument("--share", action="store_true")
|
||||
args = p.parse_args()
|
||||
|
||||
for d in sorted(glob.glob(os.path.join(args.artifacts_root, "*"))):
|
||||
mf = os.path.join(d, "manifest.json")
|
||||
if os.path.isfile(mf):
|
||||
with open(mf) as f:
|
||||
m = json.load(f)
|
||||
STEPS[m["name"]] = {"dir": d, "m": m}
|
||||
if not STEPS:
|
||||
raise SystemExit(f"no */manifest.json under {args.artifacts_root}")
|
||||
|
||||
names = list(STEPS.keys())
|
||||
clip_ids = sorted(int(k) for k in STEPS[names[0]]["m"]["clips"])
|
||||
demo = build_ui(names, clip_ids)
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[args.artifacts_root])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,80 @@
|
||||
#!/usr/bin/env bash
|
||||
# Acquire a single gang allocation of N nodes on the Slinky (k8s-backed) pool.
|
||||
#
|
||||
# Why this is not just `sbatch -N<N>`: the k8s autoscaler destroys idle nodes, so the pool only
|
||||
# contains as many nodes as there is current demand. Slurm REJECTS (does not queue) a job asking
|
||||
# for more nodes than physically exist in the partition:
|
||||
# "Batch job submission failed: Requested node configuration is not available"
|
||||
# So we have to manufacture demand first, then grab the gang allocation:
|
||||
# 1. keep N+SLACK short 1-node warmups queued -> autoscaler grows the pool
|
||||
# 2. probe with `sbatch --test-only -N<N>` until Slurm says it *could* schedule it
|
||||
# 3. submit the real -N<N> hold (now accepted; pends behind the warmups)
|
||||
# 4. drop the warmups so the hold starts immediately
|
||||
#
|
||||
# Note: other people's jobs (e.g. the demo hold) occupy nodes too, so the pool has to grow to
|
||||
# roughly N + (nodes held by others) before a -N<N> request becomes satisfiable.
|
||||
#
|
||||
# Usage: NODES=12 JOB=wan14b_hold bash examples/train/acquire_nodes.sh
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
NODES="${NODES:-12}"
|
||||
JOB="${JOB:-wan14b_hold}"
|
||||
TIME="${TIME:-120:00:00}"
|
||||
SLACK="${SLACK:-3}" # extra warmups beyond NODES, to out-pace other users' holds
|
||||
WARM_SECS="${WARM_SECS:-900}"
|
||||
MAX_ROUNDS="${MAX_ROUNDS:-60}"
|
||||
mkdir -p "$WORK/logs"
|
||||
say() { echo "[acquire $(date +%H:%M:%S)] $*"; }
|
||||
|
||||
EXIST=$(squeue -h -u "$USER" -n "${JOB}" -o '%i' 2>/dev/null | head -1)
|
||||
if [ -n "$EXIST" ]; then say "reusing existing hold $EXIST"; echo "ALLOC=$EXIST"; exit 0; fi
|
||||
|
||||
WANT_WARM=$(( NODES + SLACK ))
|
||||
|
||||
for round in $(seq 1 "$MAX_ROUNDS"); do
|
||||
pool=$(sinfo -h -p all -N -o '%N' 2>/dev/null | sort -u | wc -l)
|
||||
|
||||
# Probe: can Slurm schedule -N$NODES at all? (--test-only never actually submits)
|
||||
probe=$(sbatch --test-only -N"$NODES" --gres=gpu:4 --ntasks-per-node=1 --exclusive \
|
||||
-t "$TIME" -p all --wrap='hostname' 2>&1 | head -1)
|
||||
|
||||
if echo "$probe" | grep -q "to start at"; then
|
||||
say "pool=$pool — probe OK, submitting real -N$NODES hold"
|
||||
OUT=$(sbatch -N"$NODES" --gres=gpu:4 --ntasks-per-node=1 --exclusive -t "$TIME" \
|
||||
-p all -J "$JOB" --requeue --chdir="$WORK" -o "$WORK/logs/${JOB}_%j.out" \
|
||||
--wrap='srun sleep infinity' 2>&1)
|
||||
if echo "$OUT" | grep -q "Submitted batch"; then
|
||||
ALLOC=$(echo "$OUT" | grep -oE '[0-9]+' | head -1)
|
||||
say "SUBMITTED hold jobid=$ALLOC — clearing warmups so it can start"
|
||||
scancel -u "$USER" -n warmup 2>/dev/null
|
||||
for i in $(seq 1 180); do
|
||||
st=$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)
|
||||
if [ "$st" = R ]; then
|
||||
say "RUNNING on: $(squeue -h -j "$ALLOC" -o '%N')"
|
||||
echo "ALLOC=$ALLOC"; exit 0
|
||||
fi
|
||||
[ -z "$st" ] && { say "hold $ALLOC vanished"; break; }
|
||||
sleep 10
|
||||
done
|
||||
say "hold $ALLOC queued but not started yet (jobid kept)"; echo "ALLOC=$ALLOC"; exit 0
|
||||
fi
|
||||
say "real submit rejected despite OK probe: $(echo "$OUT" | tail -1)"
|
||||
else
|
||||
say "pool=$pool — probe says not schedulable yet"
|
||||
fi
|
||||
|
||||
# Keep warmup pressure up so the autoscaler keeps growing the pool.
|
||||
have=$(squeue -h -u "$USER" -n warmup -o '%i' 2>/dev/null | wc -l)
|
||||
need=$(( WANT_WARM - have ))
|
||||
if [ "$need" -gt 0 ]; then
|
||||
say "submitting $need warmup(s) (have $have, want $WANT_WARM) to drive scale-up"
|
||||
for _ in $(seq 1 "$need"); do
|
||||
sbatch -N1 --gres=gpu:4 --exclusive -p all -t 00:20:00 -J warmup \
|
||||
-o /dev/null --chdir="$WORK" --wrap="srun sleep $WARM_SECS" >/dev/null 2>&1
|
||||
done
|
||||
fi
|
||||
sleep 30
|
||||
done
|
||||
|
||||
say "FAILED to acquire $NODES nodes after $MAX_ROUNDS rounds"
|
||||
exit 1
|
||||
@@ -0,0 +1,44 @@
|
||||
#!/usr/bin/env bash
|
||||
# Wait for step A to finish, then seed and launch step B — so the held allocation never sits
|
||||
# idle between the two overfit stages.
|
||||
#
|
||||
# "A is done" means BOTH: no step-A launcher process is alive, AND a COMPLETE checkpoint-1000
|
||||
# exists (dcp/.metadata present). Requiring both avoids launching B off a half-written
|
||||
# checkpoint if A died mid-save.
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
ALLOC="${ALLOC:-728}"
|
||||
A_JOB="${A_JOB:-wan14b_stepA_v2}"
|
||||
A_CKPT="$WORK/wantrack_14b_synth_sparse_fixed_out/checkpoint-1000"
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY}"
|
||||
say() { echo "[chain $(date +%H:%M:%S)] $*"; }
|
||||
|
||||
say "waiting for step A to reach a complete $A_CKPT ..."
|
||||
while :; do
|
||||
if [ -f "$A_CKPT/dcp/.metadata" ]; then
|
||||
# Checkpoint is complete; wait for the launcher to actually exit before reusing the nodes.
|
||||
# Use the launcher's PID file, NOT `pgrep -f run_wan14b_held.sh` — that pattern also matches
|
||||
# any watcher/shell whose command line contains the string, so the check never goes false
|
||||
# and the chain waits forever on nodes that are already idle.
|
||||
A_RUNFILE="$WORK/logs/${A_JOB}.running"
|
||||
if [ ! -f "$A_RUNFILE" ] || ! kill -0 "$(cat "$A_RUNFILE" 2>/dev/null)" 2>/dev/null; then
|
||||
say "step A finished and $A_CKPT is complete"
|
||||
break
|
||||
fi
|
||||
say "checkpoint-1000 complete, waiting for step A launcher (pid $(cat "$A_RUNFILE")) to exit ..."
|
||||
fi
|
||||
[ -n "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" ] || { say "ALLOC $ALLOC vanished — aborting chain"; exit 1; }
|
||||
sleep 60
|
||||
done
|
||||
|
||||
say "seeding step B output dir from A's checkpoint-1000"
|
||||
bash "$REPO/examples/train/run_stepB_seed.sh" || { say "seed failed"; exit 1; }
|
||||
|
||||
say "launching step B (random-track overfit, WANTRACK_FIXED_SAMPLE=0)"
|
||||
cd "$REPO"
|
||||
exec env ALLOC="$ALLOC" NODES=4 JOB=wan14b_stepB_v2 PORT=30918 \
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_synth_sparse_random_14b.yaml \
|
||||
WANTRACK_FIXED_SAMPLE=0 WANTRACK_FREEZE_HEAD=0 TRACKWAN_TRACK_BIAS=1 \
|
||||
WANDB_API_KEY="$WANDB_API_KEY" \
|
||||
bash examples/train/run_wan14b_held.sh
|
||||
@@ -0,0 +1,98 @@
|
||||
#!/usr/bin/env bash
|
||||
# Launch a bidir teacher training run on a HELD node allocation, so it survives this
|
||||
# cluster's ~80-min node-cordon recycling (a held allocation keeps its nodes because a
|
||||
# drain can't complete while a job holds them — same reason the shao_wm2 sleep-infinity
|
||||
# node lived a day). We hold N nodes with sleep infinity, then run torchrun training
|
||||
# INSIDE the allocation via `srun --overlap`, auto-restarting on any exit (resume from
|
||||
# the latest checkpoint). Verified: one srun --overlap fans out across all held nodes.
|
||||
#
|
||||
# Usage:
|
||||
# WANDB_API_KEY=<key> CFG=<config.yaml> NODES=<n> JOB=<name> \
|
||||
# bash examples/train/run_openvid_bidir_held.sh
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
CFG="${CFG:-examples/train/scenario/worldmodel/finetune_wantrack_openvid_sparse_1p3b.yaml}"
|
||||
NODES="${NODES:-4}"; GPUS=4
|
||||
JOB="${JOB:-openvid_bidir_1p3b}"
|
||||
PORT="${PORT:-29500}"
|
||||
# Stable, per-model W&B run id so crash-relaunches RESUME a single run instead of minting a
|
||||
# fresh run every attempt (that's what fragmented the project into dozens of runs). wandb.init
|
||||
# honors WANDB_RUN_ID + WANDB_RESUME from the env, so no code change is needed. Distinct per
|
||||
# JOB (1p3b vs 14b); override with WANDB_RUN_ID=... to point restarts at an existing run.
|
||||
WANDB_RUN_ID="${WANDB_RUN_ID:-${JOB}}"
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
TOTAL_GPUS=$(( NODES * GPUS ))
|
||||
mkdir -p "$WORK/logs"
|
||||
# output_dir (for checkpoint cleanup) — read from the config
|
||||
OUTPUT_DIR=$(grep -oE 'output_dir:[^#]*' "$REPO/$CFG" | head -1 | sed 's/.*output_dir:[[:space:]]*//; s/"//g' | xargs)
|
||||
|
||||
# Drop a half-written LATEST checkpoint (a crash mid-save leaves checkpoint-N/dcp with no
|
||||
# .metadata) so resume_from_checkpoint=latest falls back to the previous VALID checkpoint (or
|
||||
# scratch) instead of crash-looping on "metadata is None". Keeps only complete checkpoints.
|
||||
clean_bad_ckpt() {
|
||||
[ -n "$OUTPUT_DIR" ] || return 0
|
||||
local latest
|
||||
latest=$(ls -d "$OUTPUT_DIR"/checkpoint-* 2>/dev/null | sed 's/.*checkpoint-//' | sort -n | tail -1)
|
||||
[ -n "$latest" ] || return 0
|
||||
local d="$OUTPUT_DIR/checkpoint-$latest"
|
||||
if [ ! -f "$d/dcp/.metadata" ]; then
|
||||
echo "[held] checkpoint-$latest is incomplete (no dcp/.metadata) — removing so resume uses the last good one"
|
||||
rm -rf "$d"
|
||||
fi
|
||||
}
|
||||
|
||||
# --- 1) hold the nodes (survives cordons; sleep infinity never releases them) --------
|
||||
# Reuse an already-held allocation by passing ALLOC=<jobid> in the env — lets us swap
|
||||
# env vars / config without losing the nodes to the cold pool.
|
||||
if [ -z "${ALLOC:-}" ]; then
|
||||
echo "[held] requesting $NODES nodes ($TOTAL_GPUS GPUs) ..."
|
||||
ALLOC=$(sbatch -N"$NODES" --gres=gpu:$GPUS --ntasks-per-node=1 --exclusive -t 120:00:00 \
|
||||
-p all -J "${JOB}_hold" --chdir="$WORK" -o "$WORK/logs/${JOB}_hold_%j.out" \
|
||||
--wrap='srun sleep infinity' | grep -oE '[0-9]+' | head -1)
|
||||
[ -z "$ALLOC" ] && { echo "[held] sbatch failed (cold pool? submit a warmup + retry)"; exit 1; }
|
||||
else
|
||||
echo "[held] reusing existing allocation JobID=$ALLOC"
|
||||
fi
|
||||
echo "[held] allocation JobID=$ALLOC ; waiting for it to start ..."
|
||||
for i in $(seq 1 60); do
|
||||
[ "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" = R ] && break; sleep 5
|
||||
done
|
||||
[ "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" = R ] || { echo "[held] $ALLOC not running"; exit 1; }
|
||||
NODELIST=$(squeue -h -j "$ALLOC" -o '%N')
|
||||
MASTER=$(scontrol show hostnames "$NODELIST" | head -1)
|
||||
echo "[held] nodes=$NODELIST master=$MASTER (cancel with: scancel $ALLOC)"
|
||||
|
||||
# --- 2) run training inside the held allocation, auto-restart on exit -----------------
|
||||
attempt=0
|
||||
while :; do
|
||||
# bail out if the held allocation itself is gone (should be rare)
|
||||
[ -n "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" ] || { echo "[held] allocation $ALLOC vanished — re-run this script"; exit 1; }
|
||||
attempt=$((attempt + 1))
|
||||
clean_bad_ckpt # remove any half-written checkpoint before (re)starting
|
||||
echo "=== [train] attempt $attempt on alloc $ALLOC (resume from latest valid checkpoint) ==="
|
||||
srun --overlap --jobid="$ALLOC" --nodes="$NODES" --ntasks="$NODES" --ntasks-per-node=1 \
|
||||
--chdir="$REPO" bash -lc "
|
||||
source .venv/bin/activate
|
||||
export HOME=$WORK HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch MPLCONFIGDIR=$WORK/.mpl \
|
||||
TRITON_CACHE_DIR=$WORK/.cache/triton_${JOB} TORCHINDUCTOR_CACHE_DIR=$WORK/.cache/inductor_${JOB} \
|
||||
TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0 PYTHONPATH=$REPO \
|
||||
WANDB_API_KEY=$WANDB_API_KEY WANDB_MODE=online \
|
||||
WANDB_RUN_ID=$WANDB_RUN_ID WANDB_RESUME=allow \
|
||||
WANTRACK_AUG=1 WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 WANTRACK_EXTRA_MODE=random \
|
||||
WANTRACK_PMASK=0 WANTRACK_FIXED_SAMPLE=0 WANTRACK_MOTION_DROP=0 WANTRACK_TEXT_DROP=0 \
|
||||
WANTRACK_DEBUG=1 TRACKWAN_TRACK_BIAS=${TRACKWAN_TRACK_BIAS:-0} \
|
||||
TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=1800 TORCH_NCCL_TRACE_BUFFER_SIZE=2000 \
|
||||
NCCL_SOCKET_NTHREADS=4 NCCL_NSOCKS_PERTHREAD=8
|
||||
torchrun --nnodes=$NODES --nproc-per-node=$GPUS --node-rank=\$SLURM_PROCID \
|
||||
--rdzv-backend=c10d --rdzv-endpoint=$MASTER:$PORT \
|
||||
fastvideo/train/entrypoint/train.py --config $CFG \
|
||||
--training.distributed.num_gpus $TOTAL_GPUS
|
||||
"
|
||||
rc=$?
|
||||
if [ $rc -eq 0 ]; then echo "[train] finished cleanly (rc=0)"; break; fi
|
||||
echo "[train] exited rc=$rc — nodes still held, relaunching from checkpoint in 20s ..."
|
||||
sleep 20
|
||||
done
|
||||
echo "[held] training complete. Freeing allocation $ALLOC."
|
||||
scancel "$ALLOC" 2>/dev/null || true
|
||||
@@ -0,0 +1,44 @@
|
||||
#!/usr/bin/env bash
|
||||
# Launch the OpenVid bidir teacher run (Wan2.1 1.3B Fun-I2V) with the SPARSE point
|
||||
# conditioning recipe + MotionStream stage-1 hparams. Wraps examples/train/run_slurm.sh
|
||||
# with this cluster's settings (partition=all, 4 GPU/node) and the WANTRACK_* env that
|
||||
# selects sparse mode with NO track masking and NO fixed/overfit sampling.
|
||||
#
|
||||
# WANDB_API_KEY must be exported in the environment before calling (never hard-coded).
|
||||
#
|
||||
# Usage: WANDB_API_KEY=<key> bash examples/train/run_openvid_bidir_slurm.sh [num_nodes] [extra --dotted overrides]
|
||||
set -euo pipefail
|
||||
# CFG overridable for the 14B run: CFG=.../finetune_wantrack_openvid_sparse_14b.yaml
|
||||
CFG="${CFG:-examples/train/scenario/worldmodel/finetune_wantrack_openvid_sparse_1p3b.yaml}"
|
||||
NODES="${1:-8}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
# ---- sparse point conditioning (1-per-SAM-object + EXTRA_RANDOM extras), no masking ----
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_PMASK=0 # no stochastic track masking (initial stage)
|
||||
export WANTRACK_FIXED_SAMPLE=0 # no deterministic/overfit sampling
|
||||
export WANTRACK_MOTION_DROP=0
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}" # log sampling/coverage stats (throttled)
|
||||
|
||||
export WANDB_MODE=online
|
||||
|
||||
# ---- compute nodes: /home is NOT mounted; point every ~-based cache at writable Lustre/tmp ----
|
||||
export HOME=/mnt/lustre/vlm-s4duan
|
||||
export HF_HOME=/mnt/lustre/vlm-s4duan/.hf
|
||||
export TORCH_HOME=/mnt/lustre/vlm-s4duan/.torch
|
||||
export MPLCONFIGDIR=/mnt/lustre/vlm-s4duan/.mpl
|
||||
export TORCHINDUCTOR_CACHE_DIR=/tmp/inductor_cache # node-local; triton cache set in run_slurm.sh
|
||||
|
||||
# ---- this cluster ----
|
||||
export PARTITION=all
|
||||
export NUM_GPUS=4 # GPUs per node
|
||||
export MEM=0 # all node memory
|
||||
export CPUS_PER_TASK=128
|
||||
export JOB_NAME="${JOB_NAME:-openvid_bidir_1p3b}"
|
||||
export OUTPUT_DIR=/mnt/lustre/vlm-s4duan/logs/slurm
|
||||
|
||||
bash examples/train/run_slurm.sh "$CFG" "$NODES" "$@"
|
||||
Executable
+53
@@ -0,0 +1,53 @@
|
||||
#!/usr/bin/env bash
|
||||
# Launch a stage-2 finetuning run (MotionStream-style track masking, lr 1e-6,
|
||||
# 800 steps) on a HELD node allocation, like ``run_openvid_bidir_held.sh``.
|
||||
#
|
||||
# Two variants (pass VARIANT=frozen or VARIANT=unfrozen):
|
||||
# VARIANT=frozen → freezes track_encoder (matches MotionStream paper)
|
||||
# VARIANT=unfrozen → keeps track_encoder trainable (ablation)
|
||||
#
|
||||
# WANDB_API_KEY must be exported before calling.
|
||||
#
|
||||
# Usage:
|
||||
# WANDB_API_KEY=<key> VARIANT=frozen bash examples/train/run_openvid_stage2_slurm.sh [num_nodes]
|
||||
set -euo pipefail
|
||||
VARIANT="${VARIANT:-frozen}"
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
case "$VARIANT" in
|
||||
frozen)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage2_frozen.yaml
|
||||
JOB=openvid_stage2_frozen
|
||||
PORT=30100
|
||||
FREEZE=1
|
||||
;;
|
||||
unfrozen)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage2_unfrozen.yaml
|
||||
JOB=openvid_stage2_unfrozen
|
||||
PORT=30200
|
||||
FREEZE=0
|
||||
;;
|
||||
*) echo "unknown VARIANT=$VARIANT (frozen|unfrozen)"; exit 1 ;;
|
||||
esac
|
||||
|
||||
# Stage-2 recipe: stochastic mid-frame track masking (MotionStream Sec. 3.1)
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_PMASK=0.2 # 20% chance to zero contiguous frame chunks
|
||||
export WANTRACK_MASK_CHUNK=8 # 8-frame chunks (~1/3 sec at 24 fps)
|
||||
export WANTRACK_FIXED_SAMPLE=0
|
||||
export WANTRACK_MOTION_DROP=0
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}"
|
||||
export WANTRACK_FREEZE_HEAD="$FREEZE" # 1=freeze track_encoder (MotionStream), 0=trainable
|
||||
export TRACKWAN_TRACK_BIAS=1 # merged-bias init used bias=True convs
|
||||
|
||||
export WANDB_MODE=online
|
||||
export CFG JOB PORT
|
||||
|
||||
# call the same held-alloc launcher used for stage 1
|
||||
CFG="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
Executable
+34
@@ -0,0 +1,34 @@
|
||||
#!/usr/bin/env bash
|
||||
# Stage-2 REPLACEMENT (hi-LR): the empirical eval showed stages 2/3 at lr=1e-6
|
||||
# barely moved weights (~0.02%/600 steps), so behavior barely changed. Try lr=5e-5
|
||||
# to actually shift the model. Aggressive drop rates so it truly sees the sparse
|
||||
# regime the user hits at inference time.
|
||||
#
|
||||
# Init: merged_bias_ckpt4800 (stage-1 end, clean base).
|
||||
# Recipe: TRACK_DROP=0.5 + MOTION_DROP=0.3 + PMASK=0.2 + freeze head.
|
||||
set -euo pipefail
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_TRACK_DROP=0.5
|
||||
export WANTRACK_MOTION_DROP=0.3
|
||||
export WANTRACK_PMASK=0.2
|
||||
export WANTRACK_MASK_CHUNK=8
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_FIXED_SAMPLE=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}"
|
||||
export WANTRACK_FREEZE_HEAD=1
|
||||
export TRACKWAN_TRACK_BIAS=1
|
||||
|
||||
export WANDB_MODE=online
|
||||
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage2v2_hilr.yaml
|
||||
JOB=openvid_stage2v2_hilr
|
||||
PORT=30600
|
||||
|
||||
CFG="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
Executable
+55
@@ -0,0 +1,55 @@
|
||||
#!/usr/bin/env bash
|
||||
# Stage-3 experiments: per-track dropout (WANTRACK_TRACK_DROP=0.5) to fix the
|
||||
# sparse-conditioning distribution mismatch at inference. Three variants:
|
||||
#
|
||||
# VARIANT=A_chunkmask TRACK_DROP=0.5 + PMASK=0.2 (from stage-2 frozen 800)
|
||||
# VARIANT=B_motiondrop TRACK_DROP=0.5 + MOTION_DROP=0.10 (from stage-2 frozen 800)
|
||||
# VARIANT=C_from3700 TRACK_DROP=0.5 + PMASK=0.2 (from merged-bias 3700 — earlier ckpt sanity)
|
||||
#
|
||||
# All three: lr 1e-6, freeze head, 600 steps.
|
||||
set -euo pipefail
|
||||
VARIANT="${VARIANT:-A_chunkmask}"
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
case "$VARIANT" in
|
||||
A_chunkmask)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage3A_chunkmask.yaml
|
||||
JOB=openvid_stage3A_chunkmask
|
||||
PORT=30300
|
||||
PMASK=0.2; MOTION_DROP=0.0
|
||||
;;
|
||||
B_motiondrop)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage3B_motiondrop.yaml
|
||||
JOB=openvid_stage3B_motiondrop
|
||||
PORT=30400
|
||||
PMASK=0.0; MOTION_DROP=0.10
|
||||
;;
|
||||
C_from3700)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage3C_from3700.yaml
|
||||
JOB=openvid_stage3C_from3700
|
||||
PORT=30500
|
||||
PMASK=0.2; MOTION_DROP=0.0
|
||||
;;
|
||||
*) echo "unknown VARIANT=$VARIANT"; exit 1 ;;
|
||||
esac
|
||||
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_TRACK_DROP=0.5 # per-track dropout (the fix)
|
||||
export WANTRACK_PMASK="$PMASK" # chunked mid-frame mask (0 for B)
|
||||
export WANTRACK_MASK_CHUNK=8
|
||||
export WANTRACK_MOTION_DROP="$MOTION_DROP" # full-track drop (0 for A/C)
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_FIXED_SAMPLE=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}"
|
||||
export WANTRACK_FREEZE_HEAD=1 # head is converged
|
||||
export TRACKWAN_TRACK_BIAS=1
|
||||
|
||||
export WANDB_MODE=online
|
||||
export CFG JOB PORT
|
||||
|
||||
CFG="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user