Compare commits

...
Author SHA1 Message Date
kevin314 63a02d55f4 Add clip fix 2026-08-08 12:39:46 +00:00
kevin314 3f7d495fbd Update stage a/b script 2026-08-02 10:26:16 +00:00
kevin314 808bb8e697 Fix ckpt resuming 2026-08-01 09:28:57 +00:00
kevin314 661fa51d9e Update stage 1 2026-07-31 01:35:05 +00:00
kevin314 59843ee2d1 Merge branch 'trackwan_bidir' into klin/trackwan 2026-07-28 09:13:22 +00:00
kevin314 cb80364aaa Update preprocessing 2026-07-24 23:47:31 +00:00
shaoxiongduan 92cc893cfc trackwan: 14B teacher pipeline, layer-wise track LR, and overfit tooling
Training:
- finetune.py: optional track/base param groups (WANTRACK_TRACK_GROUP) so the
  track pathway can take a different LR (WANTRACK_TRACK_LR_MULT); patch_embedding
  gradient masked to the track slot [:, 36:]. Raises if nothing matches.
- track_warmup.py callback: hold the pretrained DiT at lr 0 while the track
  pathway warms up, then ramp it back in.

14B configs + held-alloc launchers:
- finetune_wantrack_{synth_sparse_fixed,synth_sparse_random,openvid_sparse_merged_bias,
  synth_stage2}_14b + openvid_14b_{warmup,lwlr} + 1.3B stage2 configs.
- run_wan14b_held.sh (held alloc + auto-restart + rank cleanup, per-stage WANTRACK
  env), acquire_nodes.sh, chain_stepAB.sh, run_stepB_seed.sh, run_stepC_merge.sh,
  watch_gate.sh, run_synth_stage2_slurm.sh.

Init / diagnostics:
- convert_trackwan_init_v2.py: --pe-src to lift patch_embedding[:, 36:] and
  --track-src to lift track_encoder together (co-adapted merge); sharded-source support.
- check_gate_growth.py, gradient_flow_analysis.py, gradient_flow_summarize.py.

Data pipeline: segment_tracks.py sharding; decode/segment/tracks slurm scripts;
precompute_null_text, two_track_val, cfg_ablation, lwlr_experiment.

Inference: app_multi_trace.py updates.
2026-07-23 03:52:56 +00:00
kevin314 6ef5b6b28f CoTracker compile flag + decode prefetch + benchmark knobs 2026-07-17 11:22:15 +00:00
kevin314 948e9d7610 Fuse preprocessing operations 2026-07-17 03:41:06 +00:00
kevin314 ec5a7c7e73 Update scripts 2026-07-16 10:10:35 +00:00
kevin314 c31dd6853e Optimize preproc 2026-07-16 07:43:56 +00:00
shaoxiongduan ac919b1c87 [feat] WanTrack: OpenVid preprocess + Wan2.2 synthetic-gen pipeline, training configs, first-frame saturation fix
- data_pipeline: SAM+CoTracker preprocessing, sharded Slurm launchers, OpenVid download/filter,
  and a data-parallel Wan2.2-T2V-A14B synthetic video generator with the first-frame
  over-brightness fix (over-generate 129 -> drop first 8 -> keep 121).
- examples/train: OpenVid bidir/stage2/stage3 launchers + finetune YAMLs (sparse, merged,
  merged_bias, stage2v2 hi-lr, stage3A/B/C).
- research_log: first-frame saturation root-cause writeup (DiT position-0 latent, not VAE decode).
- fastvideo: trackwan model / track-encoder / preprocessing updates, dcp_to_diffusers export, FA4.
2026-07-14 08:58:22 +00:00
kevin314 d4eab03809 Fix tracking 2026-07-10 07:13:01 +00:00
kevin314 c1332cb4fb Init 2026-07-08 07:29:59 +00:00
shaoxiongduan 8bc4e538d0 [feat] WanTrack tooling: sparse interactive app + warmstart YAML
examples/inference/gradio/trackwan/app_action.py:
  Bring the interactive drawing app in line with the sparse training
  regime our overfit runs used (WANTRACK_SPARSE=1). New 'Track budget'
  radio: 'dense' keeps the old full 50x50 grid; 'sparse' (new default)
  builds a much smaller [Tpx, num_moving + num_bg, 2] tensor -- the
  moving handle points from the drawn path + num_bg random background
  points from the training grid (seeded by the seed slider). New
  'num_bg' slider (0-60, default 20) matches WANTRACK_EXTRA_RANDOM=20.
  Overlay logic updated to draw every point in sparse mode instead of
  grid-strided subsampling that assumed N==2500.

examples/train/scenario/worldmodel/finetune_wantrack_synth_sparse_random_warmstart.yaml:
  Test config for the curriculum idea: start from sparse-fixed@1000
  and let random-sampling sparse training take over. Compared cold-
  start sparse-random (~0.8%) vs warmstart (~3.2%) vs fixed (~8.8%)
  divergence.
2026-07-07 03:24:11 +00:00
shaoxiongduan 495dd27d33 [feat] WanTrack: golf/synth overfit YAMLs for A/B ablations
Add the training configs used for the golf-overfit and 20-clip synthetic
overfit series exercising every combination of the recent changes:

  golf_nonzero / golf_overfit_nonzero     -- non-zero-init proj ablation
  synth_generic / synth_varied             -- prompt-conditioning ablation
                                              (generic vs per-clip caption)
  synth_varied_d64                         -- id_dim=64 (paper spec)
  synth_varied_d64_nobias                  -- + bias=False on track head
  synth_varied_d64_nobias_fixed            -- + WANTRACK_FIXED_SAMPLE=1
  synth_curriculum                         -- random-sample warmstart from
                                              a fixed-sample checkpoint
  synth_varied_full_val                    -- max_train_steps=1 harness that
                                              regenerates all 20 val clips
                                              in 3 modes from a checkpoint
  synth_sparse_fixed / synth_sparse_random -- 1-per-object + 20 background
                                              extras (Yongqi's recipe)

All resume from the trackwan_1.3b_i2v_control_init family (or the
d64_nobias variant) and enable the new joint text+motion CFG (w_t=3,
w_m=1.5) in track_validation.
2026-07-06 04:11:32 +00:00
shaoxiongduan 4d0571eabf [feat] WanTrack tooling: I2V gen flag, joint-CFG gradio app, sparse dashboard
data_pipeline/generate_videos.py:
  Add --image arg -> when set, forwards to VideoGenerator.generate_video
  via image_path=... so the same script can drive either T2V or I2V. Used
  to generate the shared-first-frame synthetic-toy overfit dataset from a
  single seed image.

examples/inference/gradio/trackwan/app_action.py:
  Bring the interactive drawing app in line with the training-time val
  recipe: joint text+motion CFG (MotionStream Eq. 2) with w_text and
  w_motion sliders + a mode radio (joint / text_only / no_track).
  `on_generate` now branches on mode instead of the old bare single
  forward, and text_only / no_track let the user compare against the
  same counterfactuals the val callback logs.

data_pipeline/sparse_sampling_dashboard.py:
  New interactive dashboard for eyeballing the sparse recipe
  (1 track per SAM object + num_sampled background extras). Overlays
  the FastSAM segmentation on frame 0 in the same tab20 palette used
  for the tracked points, so mask/point correspondence is visible at a
  glance. Sliders for num_sampled, random-vs-weighted extras mode, and
  seed; a button to render the full kept-tracks overlay video.
2026-07-06 04:11:04 +00:00
shaoxiongduan 7f0c06e7f2 [feat] WanTrack: paper-aligned track head + fixed-sample + sparse sampler
Changes to the track pathway inspired by (and now matching or extending)
the MotionStream paper:

track_encoder.py:
  * bias=False on temporal_conv and proj (paper Eq. 1 requires zero at
    empty cells; nn.Conv3d's default bias=True was leaving a per-channel
    constant term across every spatial cell, drowning the sparse signal).
  * zero_init default now False (paper explicitly rejects ControlNet
    zero-conv; matches Yongqi's recipe).
  * WANTRACK_FIXED_SAMPLE=1 -> deterministic sinusoidal IDs (arange(N))
    instead of a fresh random draw every forward. Same video always
    scatters the same phi_n at the same cells across steps.

wantrack.py:
  * WANTRACK_FIXED_SAMPLE plumbing: per-sample seeds derived from
    info_list[i]['id'] so the same clip -> same subset every step.
  * WANTRACK_SPARSE=1 sampler: 1 track per SAM object + WANTRACK_EXTRA_RANDOM
    background extras (uniform or lowrank-weighted via WANTRACK_EXTRA_MODE).
    Bypasses the dense 1000-2500 code path.
  * on_train_start hook: WANTRACK_NONZERO_INIT=1 re-inits track_encoder.proj
    with Kaiming after checkpoint load, avoiding the load/save round-trip
    that corrupted our earlier attempts to force non-zero.

track_validation.py:
  * paired 'no_motion_cfg' third variant per sample (with-tracks + text CFG
    only, drops the motion arm) alongside the existing 'generated' (joint
    text+motion CFG) and 'no_track' branches -- isolates how much of gen's
    quality comes from motion CFG amplification vs the base pathway.
2026-07-06 04:10:18 +00:00
shaoxiongduan aea5cadb5f [feat] WanTrack val: MotionStream-aligned joint text+motion CFG
Wire compositional text + motion CFG into TrackValidationCallback._sample,
following MotionStream Eq. 2 exactly:

  v_no_text   = v(∅, c_m)           # drop text, keep tracks
  v_no_motion = v(c_t, ∅)           # keep text, drop tracks
  v_base      = α·v_no_text + (1-α)·v_no_motion,  α = wt/(wt+wm)
  v̂           = v_base + wt·(v_full - v_no_text) + wm·(v_full - v_no_motion)

3 NFE/step when either wt or wm != 1. At wt=wm=1 collapses to single v_full
(no CFG). No-track counterfactual branch keeps plain text CFG only (motion
CFG undefined without tracks). "Unconditional" text = zero embedding, matching
WANTRACK_TEXT_DROP training.

Configs updated to MotionStream's recommended wt=3.0, wm=1.5.
2026-07-04 03:22:52 +00:00
shaoxiongduan c5444621f9 [feat] WanTrack val: paired no-track counterfactual generation
TrackValidation now also generates a NO-TRACK video per sample (same first-frame +
prompt, track_points=None -> zero track map), logged as track_val/no_track alongside
track_val/generated. If the two match, the model is ignoring the tracks; their
divergence quantifies track influence. Sanity-checked: identical at step 0 (proj=0),
should separate as the track pathway trains. Toggle via paired_no_track (default on).
2026-07-03 08:54:31 +00:00
shaoxiongduan c6d0233eab [bugfix] WanTrack: track subsampling was a no-op when max_points == N
_augment_tracks gated the object-coverage + low-rank-weighted + uniform draw on
`0 < max_points < N`, so with the default WANTRACK_MAX_POINTS=2500 and the full
2500-track grid (max_points == N) the entire subsample block was skipped -- the
model saw all 2500 tracks every step and only chunk-masking ran. Gate on
`min_points < N` instead so K~U[1000,2500] actually samples (verified: kept
1179/1437/1573 across steps, was a constant 2500).

Also: TrackValidation now applies the same sampler (loads object_ids/track_weights,
runs _augment_tracks deterministically) so validation is conditioned on -- and
overlays -- the sampled subset, not the full grid. Env-gated WANTRACK_DEBUG=1 logs
per-step sampling stats (kept count, object coverage, coord ranges). Adds a
single-video (golf) overfit config for debugging.
2026-07-03 08:18:44 +00:00
shaoxiongduan de8e9e20b4 [feat] WanTrack: EPIC-200 / static-EPIC-200 / Wan-200-sampler training configs
Sampler-on finetune configs (coverage + low-rank-weighted + uniform track
sampling, p_mask chunk masking) across three data regimes: unfiltered
egocentric EPIC-200, camera-static-filtered EPIC-200 (single kitchen), and the
Wan-200 synthetic set re-run with the sampler on (vs the AUG=0 overfit baseline).
2026-07-03 07:08:33 +00:00
shaoxiongduan 4fe7638101 [feat] WanTrack: track-informativeness, sampler-viz, static-window dataset tooling
- track_informativeness.py: compare abs-motion vs affine/low-rank residual vs
  motion/subspace clustering vs DINOv2 semantic embeddings for scoring which
  point tracks are informative; per-clip camera-motion metric; gradio viz.
- sampling_viz.py: render the exact _augment_tracks sampler (coverage + low-rank
  weighted + uniform) on a preprocessed dataset to inspect which traces it keeps.
- find_static_windows.py: phase-correlation per-frame camera-motion signal (recovers
  after turns) + greedy non-overlapping static-window extraction to build a
  relatively-still clip dataset from long untrimmed videos.
- segment_viz.py: FastSAM config sweep + fixes (config/clip filters, allowed_paths).
2026-07-03 07:08:19 +00:00
shaoxiongduan 6c23b4552a [feat] WanTrack: low-rank informativeness track sampler + weights plumbing
segment_tracks.py precomputes a per-track low-rank motion-residual weight
(track_weights) alongside object_ids; schema + i2v_track preprocess carry it
into the parquet. _augment_tracks now samples the 1000-2500 track subset as
object-coverage (>=1 pt/SAM segment) + informativeness-weighted draw
(track_weights, abs-motion fallback) + a uniform draw (WANTRACK_UNIFORM_FRAC)
so static/background context survives.
2026-07-03 07:08:03 +00:00
shaoxiongduan 7c6f8f3f23 [feat] WanTrack: SAM object segmentation for track coverage + viz
- segment_tracks.py: FastSAM on frame-0 -> per-grid-point object_ids ([N], -1=bg)
  added to each tracks .npz (smallest containing mask wins).
- Plumb object_ids through i2v_track preprocess + schema (optional/empty on
  pre-segmentation datasets; collate already guards missing keys) -> prepare_batch
  -> _augment_tracks object-coverage sampling.
- segment_viz.py: gradio viz of SAM masks + chosen per-object points + CoTracker
  tracks (colored by object) on the source videos.
Deps: ultralytics (FastSAM) + opencv-python-headless.
2026-07-02 07:14:23 +00:00
shaoxiongduan 21f25791e2 [feat] WanTrack: MotionStream-aligned track sampling + motion-CFG + Wan-200 cfg
- _augment_tracks: default random 1000-2500 track sampling (was 1-200), contiguous
  chunk masking (was per-frame), diversity weighting (favor moving/unique tracks,
  keep some static), and an object-coverage hook (>=1 track per SAM segment via
  optional object_ids). Sampling ON by default; CFG dropout now default off.
- trackwan_infer.generate: motion classifier-free guidance (guidance_scale) using
  the base no-tracks I2V branch as the unconditional.
- test_motion_cfg.py: quick guidance-scale EPE sweep.
- finetune_wantrack_wan200.yaml: 50->200 data-scale run (same params).
2026-07-02 07:05:14 +00:00
shaoxiongduan d203b5f557 [feat] WanTrack: dataset-inspection dashboard + interactive draw-trace upgrades
- droid_dataset_dashboard.py: render per-clip CoTracker overlays + motion-coverage
  stats (frac_moving, mean/max motion) + near-duplicate grouping, and a gradio
  gallery sorted least-motion-first to surface weak/duplicate clips.
- app_action.py: DTensor-safe checkpoint hot-swap (copy into each param's local
  shard instead of load_state_dict, which errors on FSDP DTensors); a Preview-traces
  button that overlays the synthetic control on the frozen first frame before
  generating; and a streaming per-step denoise log.
2026-07-01 06:36:44 +00:00
shaoxiongduan b588be05d1 [feat] WanTrack: DROID track-conditioning dataset + overfit config
convert_droid.py builds a set where the first frame is similar but motion is
diverse, the design that forces trace reliance: it scans DROID (LeRobot mirror),
clusters frame-0 embeddings, and writes the tightest cluster as 121-frame libx264
mp4s + manifest with ONE generic caption, dropping straight into extract_tracks ->
i2v_track preprocess. Adds finetune_wantrack_droid_overfit.yaml (Fun-Control init,
WANTRACK_AUG=0 clean overfit, fresh output_dir) and gitignores .gradio/.
2026-07-01 06:36:43 +00:00
shaoxiongduan 42d4f524a0 [feat] WanTrack: controllability eval + diagnostics tooling
Add the instruments used to diagnose whether the model follows the tracks:
- track_sensitivity.py: per-checkpoint swap-sensitivity (relative velocity shift
  when only the tracks / first-frame / text change) -- intrinsic input attribution.
- control_sensitivity metric: intervention-diff (generate twice, same conditioning,
  different control) scored in the moving ROI, so a static background can't dilute it.
- gen_eval_artifacts.py + app_viewer.py: pre-generate counterfactual videos
  (input vs re-extracted tracks, EPE heatmap) and browse them with no GPU.
- app_action.py: mouse-draw a motion and hot-swap checkpoints.
- test_controllability.py / cotracker_epe / app.py: clip_feature + minor updates.
2026-07-01 06:36:43 +00:00
shaoxiongduan 7349e07cec [feat] WanTrack: Fun-Control init to break the track-encoder deadlock
Initialize the 52ch TrackWan from Wan2.1-Fun-1.3B-Control so the track-slot
patch-embed channels start from Fun-Control's PRETRAINED control channels
(non-zero). Combined with the zero-init track head this gives step-0 == teacher
while still letting gradient reach the track encoder -- avoiding the double
zero-init deadlock of the from-scratch Fun-InP path (where both the track
patch-embed channels and the track-head proj are zero, so the track pathway
never trains). Adds the converter, the matching track_config channel handling,
and the InP init updates.
2026-07-01 06:36:42 +00:00
shaoxiongduan 8af0c43a65 [fix] WanTrack: enforce I2V via Fun-InP base + CLIP cross-attention
The overfit teacher ignored the first frame (validation looked T2V): the
init came from FastWan2.1-T2V-1.3B, whose image-conditioning channels are
zero/untrained and which has no CLIP pathway, so first-frame and track
conditioning received ~no gradient.

Re-base on Wan2.1-Fun-1.3B-InP (pretrained concat-I2V + CLIP):
- convert_trackwan_init: widen patch-embed base_in(36)->52, keep the 36
  pretrained image-cond channels, zero-init only the 16 track channels;
  inherit image_dim from the base config; merge sharded safetensors.
- trackwan config: image_dim=1280 builds the DiT image_embedder.
- preprocess: compute + store clip_feature (CLIP frame-0 embedding) and
  add image_encoder/image_processor to required modules.
- schema: clip_feature_{bytes,shape,dtype} in pyarrow_schema_i2v_track.
- wantrack / track_validation / trackwan_infer: feed
  encoder_hidden_states_image through train, validation, and inference.
- yaml: point at the Fun-InP init + re-preprocessed data.

Verified: step-0 validation reproduces the GT first frame (PSNR ~24-26 dB
vs ~9-11 before), i.e. I2V is enforced before any training.
2026-07-01 06:36:41 +00:00
shaoxiongduan 34a0a81be1 [feat] WanTrack: bidirectional I2V + point-track (MotionStream) finetune + eval
TrackWan = Wan2.2 + MotionStream point-track motion control. Stage-1 bidirectional
I2V finetune (text kept), built on the new modular trainer.

- DiT: TrackWanTransformer3DModel (52ch in = 16 noisy + 20 I2V + 16 track), TrackEncoder
  (sinusoidal track-ID embed, scatter, 4x temporal conv, zero-init head); config + registry.
- Trainer: WanTrackModel wrapper, i2v_track preprocess pipeline + pyarrow schema,
  TrackValidationCallback, finetune_wantrack_i2v.yaml.
- data_pipeline/: video gen (Wan2.2-A14B), CoTracker3 track extraction, 52ch init
  conversion, synthetic track authoring, controllability (CoTracker EPE) eval, standalone
  inference, Gradio demo.
- eval: motion.cotracker_epe metric. research_log/ lab notebook.
- fixes: FA3 flash_attn_func tuple unwrap; decord fallback for removed torchvision.io.read_video;
  preprocess flush write_remainder.
2026-07-01 06:36:40 +00:00
191 changed files with 20850 additions and 55 deletions
+3
View File
@@ -135,3 +135,6 @@ fastvideo/tests/ssim/reference_videos/**
*.nvimlog
.nvimlog
.python-version
# Gradio runtime/flagging artifacts
.gradio/
+70
View File
@@ -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)"
+38
View File
@@ -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
View File
@@ -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"
+30
View File
@@ -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))
"
+56
View File
@@ -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
+41
View File
@@ -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)))."
+38
View File
@@ -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"
+29
View File
@@ -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
View File
@@ -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
+66
View File
@@ -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
View File
@@ -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
View File
@@ -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
+83
View File
@@ -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
+111
View File
@@ -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()
+148
View File
@@ -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"
+53
View File
@@ -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()
+241
View File
@@ -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()
+94
View File
@@ -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()
+237
View File
@@ -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()
+171
View File
@@ -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()
+159
View File
@@ -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()
+245
View File
@@ -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()
+132
View File
@@ -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()
+162
View File
@@ -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()
+113
View File
@@ -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()
+47
View File
@@ -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)"
+29
View File
@@ -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"
+224
View File
@@ -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()
+416
View File
@@ -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()
+536
View File
@@ -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()
+145
View File
@@ -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()
+109
View File
@@ -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()
+353
View File
@@ -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()
+239
View File
@@ -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()
+123
View File
@@ -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()
+122
View File
@@ -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()
+235
View File
@@ -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()
+246
View File
@@ -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()
+35
View File
@@ -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()
+65
View File
@@ -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.
+40
View File
@@ -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.
+70
View File
@@ -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.
+148
View File
@@ -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.
+65
View File
@@ -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()
+100
View File
@@ -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()
+61
View File
@@ -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()
+81
View File
@@ -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()
+137
View File
@@ -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()
+43
View File
@@ -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"
+25
View File
@@ -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'"
+21
View File
@@ -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'"
+32
View File
@@ -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"
+405
View 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"
+215
View File
@@ -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
+88
View File
@@ -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\"'"
+29
View File
@@ -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'"
+25
View File
@@ -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'"
+26
View File
@@ -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"
+62
View 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"
+119
View File
@@ -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
+32
View File
@@ -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'"
+25
View File
@@ -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'"
+213
View File
@@ -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()
+377
View File
@@ -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()
+78
View File
@@ -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()
+363
View File
@@ -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()
+133
View File
@@ -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()
+346
View File
@@ -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()
+221
View File
@@ -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()
+57
View File
@@ -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()
+289
View File
@@ -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()
+142
View File
@@ -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()
+67
View File
@@ -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()
+477
View File
@@ -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()
+26
View File
@@ -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
+117
View File
@@ -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()
+176
View File
@@ -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))
+325
View File
@@ -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()
+79
View File
@@ -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"
+115
View File
@@ -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"
+60
View File
@@ -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()
+160
View File
@@ -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 -1
View File
@@ -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
+1 -1
View File
@@ -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,
+230
View File
@@ -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&nbsp;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&nbsp;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()
+80
View File
@@ -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
+44
View File
@@ -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
+98
View File
@@ -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
+44
View File
@@ -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" "$@"
+53
View File
@@ -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" "$@"
+34
View File
@@ -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" "$@"
+55
View File
@@ -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