Compare commits
22
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
efca231efc | ||
|
|
36d6f8785e | ||
|
|
92cc893cfc | ||
|
|
ac919b1c87 | ||
|
|
8bc4e538d0 | ||
|
|
495dd27d33 | ||
|
|
4d0571eabf | ||
|
|
7f0c06e7f2 | ||
|
|
aea5cadb5f | ||
|
|
c5444621f9 | ||
|
|
c6d0233eab | ||
|
|
de8e9e20b4 | ||
|
|
4fe7638101 | ||
|
|
6c23b4552a | ||
|
|
7c6f8f3f23 | ||
|
|
21f25791e2 | ||
|
|
d203b5f557 | ||
|
|
b588be05d1 | ||
|
|
42d4f524a0 | ||
|
|
7349e07cec | ||
|
|
8af0c43a65 | ||
|
|
34a0a81be1 |
@@ -135,3 +135,6 @@ fastvideo/tests/ssim/reference_videos/**
|
||||
*.nvimlog
|
||||
.nvimlog
|
||||
.python-version
|
||||
|
||||
# Gradio runtime/flagging artifacts
|
||||
.gradio/
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
#!/usr/bin/env bash
|
||||
# FastVideo (trackwan_bidir) env — CORRECTED build. CUDA 12.9 container, aarch64 + GB200 (sm_100a).
|
||||
# Ordering fix: install repo-pinned torch (2.11.0 cu128) FIRST, build the local Blackwell kernel
|
||||
# against it, then core deps (relaxed kernel pin so the local 0.3.0 satisfies), then FA4.
|
||||
set +e
|
||||
say(){ echo; echo "==================== $* ===================="; }
|
||||
|
||||
export DEBIAN_FRONTEND=noninteractive
|
||||
export CUDA_HOME=${CUDA_HOME:-/usr/local/cuda}
|
||||
export PATH=$CUDA_HOME/bin:/mnt/lustre/vlm-s4duan/bin:$PATH
|
||||
export LD_LIBRARY_PATH=$CUDA_HOME/lib64:$LD_LIBRARY_PATH
|
||||
export UV_CACHE_DIR=/mnt/lustre/vlm-s4duan/.uv/cache
|
||||
export UV_PYTHON_INSTALL_DIR=/mnt/lustre/vlm-s4duan/.uv/python
|
||||
export TORCH_CUDA_ARCH_LIST=10.0a
|
||||
export MAX_JOBS=${MAX_JOBS:-64}
|
||||
cd /mnt/lustre/vlm-s4duan/FastVideo || { echo NO_REPO; exit 1; }
|
||||
|
||||
say "ENSURE toolchain"
|
||||
if ! command -v g++ >/dev/null || ! command -v cmake >/dev/null || ! command -v ninja >/dev/null; then
|
||||
apt-get update -y >/tmp/apt.log 2>&1 && apt-get install -y --no-install-recommends g++ gcc make cmake ninja-build git ca-certificates >>/tmp/apt.log 2>&1 && echo "apt ok" || { echo APT_FAIL; tail -20 /tmp/apt.log; }
|
||||
fi
|
||||
nvcc --version | tail -2
|
||||
|
||||
say "FRESH venv (python 3.12 on Lustre)"
|
||||
rm -rf .venv
|
||||
uv venv --python 3.12 --seed .venv || exit 1
|
||||
source .venv/bin/activate
|
||||
python --version
|
||||
|
||||
say "STAGE 1: torch 2.11.0 cu128 (repo pin, Blackwell-capable)"
|
||||
uv pip install --torch-backend=cu128 torch==2.11.0 torchvision torchaudio 2>&1 | tail -15
|
||||
python -c "import torch;print('TORCH',torch.__version__,torch.version.cuda,'avail',torch.cuda.is_available())" || echo TORCH_FAIL
|
||||
|
||||
say "STAGE 2: submodules + build local fastvideo-kernel (sm_100a) against torch 2.11"
|
||||
git submodule update --init --recursive fastvideo-kernel/include/cutlass fastvideo-kernel/include/tk 2>&1 | tail -5
|
||||
( cd fastvideo-kernel && TORCH_CUDA_ARCH_LIST=10.0a bash ./build.sh ) 2>&1 | tail -20
|
||||
echo "KERNEL_EXIT=$?"
|
||||
python -c "import fastvideo_kernel;print('kernel import OK')" || echo KERNEL_IMPORT_FAIL
|
||||
|
||||
say "STAGE 3: core deps uv pip install -e .[dev] (local kernel already satisfies relaxed pin)"
|
||||
uv pip install -e ".[dev]" 2>&1 | tail -40
|
||||
echo "CORE_EXIT=$?"
|
||||
python -c "import fastvideo;print('fastvideo import OK')" || echo FASTVIDEO_IMPORT_FAIL
|
||||
|
||||
say "STAGE 4: FA4 — OFFICIAL Dao-AILab flash_attn/cute (NOT the XOR-op fork; fork is stale/broken)"
|
||||
# Package name is flash-attn-4; pass the bare git URL (name-qualified spec is rejected).
|
||||
# --prerelease allow lets it pull the matching nvidia-cutlass-dsl==4.6.0.dev0 + quack>=0.5.3.
|
||||
uv pip install --prerelease allow \
|
||||
"git+https://github.com/Dao-AILab/flash-attention.git#subdirectory=flash_attn/cute" 2>&1 | tail -25
|
||||
python -c "from flash_attn.cute.interface import _flash_attn_fwd,_flash_attn_bwd; print('FA4_CUTE_OK')" 2>&1
|
||||
|
||||
say "VERIFY"
|
||||
python - <<'PY'
|
||||
import importlib, torch
|
||||
print("torch", torch.__version__, torch.version.cuda, "cuda_avail", torch.cuda.is_available())
|
||||
if torch.cuda.is_available():
|
||||
print("device0", torch.cuda.get_device_name(0), "cc", torch.cuda.get_device_capability(0))
|
||||
x = torch.randn(2048, 2048, device="cuda", dtype=torch.bfloat16)
|
||||
y = (x @ x); torch.cuda.synchronize()
|
||||
print("bf16 matmul on GPU OK", tuple(y.shape))
|
||||
for m in ["fastvideo", "fastvideo_kernel", "flash_attn"]:
|
||||
try: importlib.import_module(m); print("import OK:", m)
|
||||
except Exception as e: print("import FAIL:", m, repr(e)[:180])
|
||||
try:
|
||||
from flash_attn.cute.interface import _flash_attn_fwd; print("FA4 flash_attn.cute.interface OK")
|
||||
except Exception as e:
|
||||
print("FA4 cute interface:", repr(e)[:180])
|
||||
PY
|
||||
|
||||
say "DONE build_env.sh (v2)"
|
||||
@@ -0,0 +1,38 @@
|
||||
#!/usr/bin/env bash
|
||||
set +e
|
||||
export HF_HOME=/mnt/lustre/vlm-s4duan/.hf
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
cd /mnt/lustre/vlm-s4duan/FastVideo || exit 1
|
||||
source .venv/bin/activate
|
||||
|
||||
BASE=/mnt/lustre/vlm-s4duan/models/Wan2.1-Fun-1.3B-InP-Diffusers
|
||||
OUT=/mnt/lustre/vlm-s4duan/models/trackwan_1.3b_init
|
||||
mkdir -p /mnt/lustre/vlm-s4duan/models
|
||||
|
||||
echo "==================== download Fun-InP diffusers base ===================="
|
||||
python - <<PY
|
||||
from huggingface_hub import snapshot_download
|
||||
p = snapshot_download("weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers", local_dir="$BASE")
|
||||
print("downloaded to", p)
|
||||
PY
|
||||
echo "DL_EXIT=$?"
|
||||
|
||||
echo "==================== base transformer sanity (I2V => in_channels 36) ===================="
|
||||
python - <<PY
|
||||
import json
|
||||
c = json.load(open("$BASE/transformer/config.json"))
|
||||
print("in_channels:", c.get("in_channels"), "image_dim:", c.get("image_dim"), "num_layers:", c.get("num_layers"), "hidden:", c.get("hidden_size") or c.get("dim"))
|
||||
PY
|
||||
|
||||
echo "==================== convert_trackwan_init.py ===================="
|
||||
python data_pipeline/convert_trackwan_init.py --base "$BASE" --out "$OUT"
|
||||
echo "CONVERT_EXIT=$?"
|
||||
|
||||
echo "==================== trackwan init result ===================="
|
||||
ls -la "$OUT" 2>&1 | head -25
|
||||
python - <<PY
|
||||
import json
|
||||
c = json.load(open("$OUT/transformer/config.json"))
|
||||
print("trackwan in_channels:", c.get("in_channels"), "| has track_config:", "track_config" in c)
|
||||
PY
|
||||
echo "==================== DONE build_trackwan_init ===================="
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Build videos2caption.json + merge.txt for the final track dataset.
|
||||
|
||||
Pairs each 121f/720p clip (--clips-dir) with its CoTracker npz (--tracks-dir) and
|
||||
caption (--captions JSON {video.mp4: caption}). Output is exactly the manifest
|
||||
`fastvideo/pipelines/preprocess/v1_preprocess.py --preprocess_task i2v_track` consumes
|
||||
(same format used for the overfit run), so the dataset is training-ready.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--clips-dir", required=True)
|
||||
ap.add_argument("--tracks-dir", required=True)
|
||||
ap.add_argument("--captions", required=True, help="JSON {video.mp4: caption}")
|
||||
ap.add_argument("--out-dir", required=True)
|
||||
ap.add_argument("--fps", type=int, default=24)
|
||||
ap.add_argument("--num-frames", type=int, default=121)
|
||||
ap.add_argument("--height", type=int, default=720)
|
||||
ap.add_argument("--width", type=int, default=1280)
|
||||
a = ap.parse_args()
|
||||
|
||||
caps = json.load(open(a.captions))
|
||||
tracks = {p.stem: p for p in Path(a.tracks_dir).glob("*.npz")}
|
||||
items, missing_track, missing_cap = [], 0, 0
|
||||
for clip in sorted(Path(a.clips_dir).glob("*.mp4")):
|
||||
tr = tracks.get(clip.stem)
|
||||
if tr is None:
|
||||
missing_track += 1; continue
|
||||
cap = caps.get(clip.name, "")
|
||||
if not cap:
|
||||
missing_cap += 1
|
||||
items.append({
|
||||
"path": clip.name,
|
||||
"cap": [cap],
|
||||
"points_path": str(tr.resolve()),
|
||||
"fps": float(a.fps),
|
||||
"num_frames": a.num_frames,
|
||||
"resolution": {"width": a.width, "height": a.height},
|
||||
})
|
||||
outd = Path(a.out_dir); outd.mkdir(parents=True, exist_ok=True)
|
||||
(outd / "videos2caption.json").write_text(json.dumps(items, indent=2))
|
||||
(outd / "merge.txt").write_text(f"{a.clips_dir},{outd / 'videos2caption.json'}")
|
||||
print(f"[manifest] {len(items)} clips paired (clip+track+caption); "
|
||||
f"skipped {missing_track} clips w/o track, {missing_cap} w/o caption")
|
||||
print(f"[manifest] wrote {outd}/videos2caption.json + merge.txt")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,241 @@
|
||||
#!/usr/bin/env python3
|
||||
"""CFG-ablation eval on a trackwan checkpoint.
|
||||
|
||||
Reproduces ``fastvideo/train/callbacks/track_validation.py::_sample`` denoise
|
||||
formula EXACTLY, but with configurable ``w_text`` and ``w_motion`` so we can
|
||||
isolate which CFG branch is causing artifacts.
|
||||
|
||||
Runs each val sample under N configs; uploads all videos to one wandb run for
|
||||
side-by-side visual comparison.
|
||||
|
||||
Usage (single-GPU, on any held alloc):
|
||||
python data_pipeline/cfg_ablation.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/synth_stage2_paperLR_ckpt400 \
|
||||
--yaml examples/train/scenario/worldmodel/finetune_wantrack_synth_stage2_paperLR.yaml \
|
||||
--val-parquet /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset \
|
||||
--wandb-run-name cfg_abl_paperLR_ckpt400
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__)))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
from fastvideo.forward_context import set_forward_context # noqa: E402
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import ( # noqa: E402
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
|
||||
def load_val_samples(parquet_dir: str, text_len: int) -> list:
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
fs = sorted(glob.glob(os.path.join(parquet_dir, "**", "*.parquet"), recursive=True))
|
||||
if not fs:
|
||||
raise FileNotFoundError(parquet_dir)
|
||||
rows = []
|
||||
for f in fs:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
batch = collate_rows_from_parquet_schema(rows, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in rows]
|
||||
return [{
|
||||
"id": str(infos[i].get("id", f"clip{i:03d}")) if i < len(infos) else f"clip{i:03d}",
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i:i + 1].clone(),
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(),
|
||||
"object_ids": batch["object_ids"][i:i + 1].clone() if "object_ids" in batch else None,
|
||||
"track_weights": batch["track_weights"][i:i + 1].clone() if "track_weights" in batch else None,
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
} for i in range(len(rows))]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def generate_with_cfg(model, sample: dict, *, w_text: float, w_motion: float,
|
||||
num_steps: int, seed: int,
|
||||
null_text: torch.Tensor | None = None,
|
||||
null_text_mask: torch.Tensor | None = None,
|
||||
no_track: bool = False) -> np.ndarray:
|
||||
"""Reproduce track_validation._sample denoise EXACTLY, w/ configurable CFG.
|
||||
|
||||
If ``null_text`` is provided, use it as the ∅ text (canonical Wan/T5 convention
|
||||
= UMT5-encoded empty string). Otherwise fall back to zeros_like(txt).
|
||||
If ``no_track`` is True, pass tp=None/tv=None to every forward (matches the
|
||||
val callback's `track_val/no_track` panel).
|
||||
"""
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
|
||||
ff = sample["first_frame_latent"].to(device, dtype)
|
||||
txt = sample["text_embedding"].to(device, dtype)
|
||||
mask = sample["text_attention_mask"].to(device, dtype)
|
||||
clip = sample["clip_feature"].to(device, dtype)
|
||||
if no_track:
|
||||
tp = None
|
||||
tv = None
|
||||
else:
|
||||
# CRITICAL: sparse-sample tracks the same way training + track_validation do.
|
||||
# Feeding the raw 2500 tracks is out-of-distribution and produces melted output.
|
||||
tp_raw = sample["track_points"].float()
|
||||
tv_raw = sample["track_visibility"].float()
|
||||
oid = sample.get("object_ids")
|
||||
tw = sample.get("track_weights")
|
||||
aug_gen = torch.Generator(device=device).manual_seed(int(seed))
|
||||
tp_s, tv_s = model._augment_tracks(tp_raw, tv_raw, aug_gen, object_ids=oid, track_weights=tw)
|
||||
if tp_s is None or tv_s is None: # motion_drop fired -> fall back to raw
|
||||
tp_s, tv_s = tp_raw, tv_raw
|
||||
tp = tp_s.to(device, dtype)
|
||||
tv = tv_s.to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(int(num_steps), device=device)
|
||||
|
||||
do_text_cfg = (w_text != 1.0)
|
||||
do_motion_cfg = (w_motion != 1.0) and not no_track # no tracks -> motion CFG is a no-op
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, ts_):
|
||||
with torch.autocast(device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=ts_, attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=ts_,
|
||||
encoder_hidden_states_image=clip,
|
||||
track_points=tp_e, track_visibility=tv_e,
|
||||
return_dict=False)
|
||||
|
||||
if do_text_cfg:
|
||||
if null_text is not None:
|
||||
# Broadcast the [seq_len, dim] null embedding to match txt shape [B, seq_len, dim]
|
||||
nl = null_text.to(device, dtype)
|
||||
if nl.dim() == 2:
|
||||
nl = nl.unsqueeze(0).expand_as(txt).contiguous()
|
||||
txt_null = nl
|
||||
else:
|
||||
txt_null = torch.zeros_like(txt)
|
||||
else:
|
||||
txt_null = None
|
||||
|
||||
for tt in sched.timesteps:
|
||||
mi = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
ts_ = tt.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp, tv, mi, ts_)
|
||||
if not do_text_cfg and not do_motion_cfg:
|
||||
v = v_full
|
||||
elif do_text_cfg and not do_motion_cfg:
|
||||
v_no_text = _fwd(txt_null, tp, tv, mi, ts_)
|
||||
v = v_no_text + w_text * (v_full - v_no_text)
|
||||
elif do_motion_cfg and not do_text_cfg:
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
v = v_no_motion + w_motion * (v_full - v_no_motion)
|
||||
else:
|
||||
# joint text+motion CFG (matches track_validation._sample)
|
||||
v_no_text = _fwd(txt_null, tp, tv, mi, ts_)
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
alpha = w_text / max(w_text + w_motion, 1e-6)
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + w_text * (v_full - v_no_text) + w_motion * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), tt, latents.float(), return_dict=False)[0]
|
||||
|
||||
return twi.decode_to_pixels(model, latents)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True)
|
||||
p.add_argument("--yaml", required=True)
|
||||
p.add_argument("--val-parquet", required=True)
|
||||
p.add_argument("--wandb-project", default="wantrack-bidir")
|
||||
p.add_argument("--wandb-run-name", required=True)
|
||||
p.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/cfg_ablation")
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--configs", default="3.0,1.5;1.0,1.5;1.0,1.0",
|
||||
help="';'-sep list of 'w_text,w_motion' pairs")
|
||||
p.add_argument("--null-text-pt",
|
||||
help="Path to precomputed UMT5(``'') embedding .pt file. If given, this is "
|
||||
"used as the ∅ text branch instead of zeros_like(txt).")
|
||||
p.add_argument("--no-track", action="store_true",
|
||||
help="Force track_points=None throughout denoise. Reproduces the val "
|
||||
"callback's `track_val/no_track` panel.")
|
||||
args = p.parse_args()
|
||||
|
||||
null_text = None
|
||||
null_text_mask = None
|
||||
if args.null_text_pt:
|
||||
blob = torch.load(args.null_text_pt, map_location="cpu")
|
||||
null_text = blob["embedding"]
|
||||
null_text_mask = blob.get("attention_mask")
|
||||
print(f"[abl] loaded null-text embedding: shape={list(null_text.shape)} "
|
||||
f"norm={null_text.norm().item():.3f}", flush=True)
|
||||
|
||||
import wandb
|
||||
Path(args.out_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
print(f"[abl] loading model {args.model_dir}", flush=True)
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(getattr(tc.data, "text_padding_length", 256))
|
||||
|
||||
print(f"[abl] loading val samples from {args.val_parquet}", flush=True)
|
||||
samples = load_val_samples(args.val_parquet, text_len)
|
||||
print(f"[abl] {len(samples)} val samples", flush=True)
|
||||
|
||||
configs = [tuple(float(x) for x in pair.split(","))
|
||||
for pair in args.configs.split(";")]
|
||||
print(f"[abl] configs: {configs}", flush=True)
|
||||
|
||||
run = wandb.init(project=args.wandb_project, name=args.wandb_run_name,
|
||||
config={"model_dir": args.model_dir, "steps": args.steps,
|
||||
"configs": [list(c) for c in configs]})
|
||||
|
||||
for i, s in enumerate(samples):
|
||||
sid = s["id"] or f"clip{i:03d}"
|
||||
print(f"[abl] === sample {i}: {sid} ===", flush=True)
|
||||
wandb_log = {"sample": i, "caption": s["caption"][:120]}
|
||||
for wt, wm in configs:
|
||||
t0 = time.time()
|
||||
frames = generate_with_cfg(model, s, w_text=wt, w_motion=wm,
|
||||
num_steps=args.steps, seed=args.seed + i,
|
||||
null_text=null_text, null_text_mask=null_text_mask,
|
||||
no_track=args.no_track)
|
||||
dt = time.time() - t0
|
||||
tag = f"wt{wt}_wm{wm}".replace(".", "p")
|
||||
fn = Path(args.out_dir) / f"{sid}_{tag}.mp4"
|
||||
imageio.mimsave(str(fn), frames, fps=args.fps, macro_block_size=1)
|
||||
wandb_log[f"gen_{tag}"] = wandb.Video(str(fn), fps=args.fps, format="mp4")
|
||||
print(f"[abl] {tag}: {dt:.1f}s -> {fn.name}", flush=True)
|
||||
wandb.log(wandb_log)
|
||||
|
||||
# also render GT reference once
|
||||
for i, s in enumerate(samples):
|
||||
try:
|
||||
ref = twi.decode_reference(model, s["vae_latent"].to(model.device, torch.bfloat16))
|
||||
fn = Path(args.out_dir) / f"{s['id']}_gt.mp4"
|
||||
imageio.mimsave(str(fn), ref, fps=args.fps, macro_block_size=1)
|
||||
wandb.log({"sample": i, "reference_gt": wandb.Video(str(fn), fps=args.fps, format="mp4")})
|
||||
except Exception as e:
|
||||
print(f"[abl] gt render failed for sample {i}: {e}", flush=True)
|
||||
|
||||
wandb.finish()
|
||||
print("[abl] DONE", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Verify a preprocessed i2v_track parquet row by decoding its VAE latent back to pixels.
|
||||
|
||||
The preprocess stores the RAW VAE.encode() mean (diffusers' AutoencoderKLWan.decode
|
||||
consumes raw latents directly; the *training* path is what applies latents_mean/std).
|
||||
So the check is: raw latent -> vae.decode -> compare to the source clip (PSNR + montage).
|
||||
|
||||
Usage:
|
||||
python data_pipeline/check_720p_latents.py --parquet-glob '<dir>/**/*.parquet' \
|
||||
--clips-dir <clips> --model <model_path> --out <png> [--num 1]
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
|
||||
import imageio.v3 as iio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
|
||||
def _psnr(a: np.ndarray, b: np.ndarray) -> float:
|
||||
mse = float(((a.astype(np.float32) - b.astype(np.float32)) ** 2).mean())
|
||||
return 99.0 if mse == 0 else 10.0 * float(np.log10(255.0 * 255.0 / mse))
|
||||
|
||||
|
||||
def _read_clip(path: str, num_frames: int) -> np.ndarray:
|
||||
"""Source clip -> uint8 [T,H,W,3] (first num_frames)."""
|
||||
frames = []
|
||||
for i, f in enumerate(iio.imiter(path, plugin="pyav")):
|
||||
if i >= num_frames:
|
||||
break
|
||||
frames.append(f)
|
||||
return np.stack(frames)
|
||||
|
||||
|
||||
def _resize_to(frames: np.ndarray, h: int, w: int) -> np.ndarray:
|
||||
"""uint8 [T,H,W,3] -> bilinear-resized uint8, matching the preprocess transform."""
|
||||
if frames.shape[1] == h and frames.shape[2] == w:
|
||||
return frames
|
||||
t = torch.from_numpy(frames).permute(0, 3, 1, 2).float()
|
||||
t = torch.nn.functional.interpolate(t, size=(h, w), mode="bilinear", align_corners=False, antialias=True)
|
||||
return t.permute(0, 2, 3, 1).clamp(0, 255).to(torch.uint8).numpy()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__)
|
||||
ap.add_argument("--parquet-glob", required=True)
|
||||
ap.add_argument("--clips-dir", required=True)
|
||||
ap.add_argument("--model", default="/mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_nobias_init")
|
||||
ap.add_argument("--out", default="/mnt/lustre/vlm-s4duan/openvid_1m/_sanity720/decode_check.png")
|
||||
ap.add_argument("--num", type=int, default=1, help="rows to check")
|
||||
a = ap.parse_args()
|
||||
|
||||
files = sorted(glob.glob(a.parquet_glob, recursive=True))
|
||||
assert files, f"no parquet matched {a.parquet_glob}"
|
||||
|
||||
from diffusers import AutoencoderKLWan
|
||||
vae = AutoencoderKLWan.from_pretrained(a.model, subfolder="vae", torch_dtype=torch.float32).to("cuda").eval()
|
||||
|
||||
rows = []
|
||||
for f in files:
|
||||
for r in pq.read_table(f).to_pylist():
|
||||
rows.append(r)
|
||||
if len(rows) >= a.num:
|
||||
break
|
||||
if len(rows) >= a.num:
|
||||
break
|
||||
|
||||
panels = []
|
||||
for r in rows:
|
||||
lat = np.frombuffer(r["vae_latent_bytes"], np.float32).reshape(r["vae_latent_shape"]).copy()
|
||||
z = torch.from_numpy(lat)[None].to("cuda", torch.float32) # [1,16,T,h,w]
|
||||
with torch.no_grad():
|
||||
px = vae.decode(z, return_dict=False)[0] # [1,3,T,H,W] in [-1,1]
|
||||
dec = ((px[0].permute(1, 2, 3, 0).clamp(-1, 1) + 1) * 127.5).to(torch.uint8).cpu().numpy() # [T,H,W,3]
|
||||
src = _read_clip(f"{a.clips_dir}/{r['file_name']}.mp4", dec.shape[0])
|
||||
# record width/height are (H, W) — the base pipeline stores shape[-2], shape[-1]
|
||||
tgt = _resize_to(src, int(r["width"]), int(r["height"]))
|
||||
n = min(len(dec), len(tgt))
|
||||
p_all = _psnr(dec[:n], tgt[:n])
|
||||
p_f0 = _psnr(dec[0], tgt[0])
|
||||
print(f"{r['file_name']}: latent {tuple(r['vae_latent_shape'])} -> decoded {dec.shape}, "
|
||||
f"src(resized) {tgt.shape} | PSNR all-frames {p_all:.2f} dB, frame0 {p_f0:.2f} dB")
|
||||
# montage: [src | decoded] for frames 0, T/2, T-1
|
||||
idxs = [0, n // 2, n - 1]
|
||||
panels.append(np.concatenate([np.concatenate([tgt[i], dec[i]], axis=1) for i in idxs], axis=0))
|
||||
|
||||
iio.imwrite(a.out, np.concatenate(panels, axis=0))
|
||||
print(f"wrote montage (left=source, right=VAE round-trip; rows = frame 0 / mid / last) -> {a.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,94 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Report the WanTrack "gate" — patch_embedding.weight[:, 36:] — from a DCP checkpoint.
|
||||
|
||||
The overfit stages (A/B) exist to grow this slot from EXACTLY ZERO up to roughly the scale the
|
||||
successful 1.3B merged init reached (std ~0.011), co-adapting it with track_encoder. If it is
|
||||
still ~0 there is no bootstrap happening and step C would merge nothing; if it explodes past the
|
||||
pretrained channel scale (~0.0124) the track pathway is overwhelming the image prior.
|
||||
|
||||
Loads ONLY the tensors of interest (DCP resharding is per-tensor), so it costs ~MBs, not the
|
||||
~200GB a full checkpoint load would.
|
||||
|
||||
Usage: python data_pipeline/check_gate_growth.py <checkpoint-dir> [more dirs...]
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.distributed.checkpoint as dcp
|
||||
from torch.distributed.checkpoint import FileSystemReader
|
||||
|
||||
BASE_IN = 36
|
||||
# NOTE: the live module nests the conv, so a TRAINED checkpoint stores
|
||||
# 'patch_embedding.proj.weight', while the converter writes 'patch_embedding.weight' into the
|
||||
# init safetensors. Try both — silently matching neither is how the earlier LWLR experiment
|
||||
# ended up measuring nothing.
|
||||
PE_KEYS = ["patch_embedding.proj.weight", "patch_embedding.weight"]
|
||||
KEYS = PE_KEYS + [
|
||||
"track_encoder.proj.weight",
|
||||
"track_encoder.temporal_conv.weight",
|
||||
]
|
||||
|
||||
|
||||
def _find(md_keys, want: str) -> str | None:
|
||||
"""DCP nests model tensors under a role prefix (e.g. 'student.'); match on the suffix."""
|
||||
exact = [k for k in md_keys if k == want]
|
||||
if exact:
|
||||
return exact[0]
|
||||
cands = [k for k in md_keys if k.endswith(want)]
|
||||
return cands[0] if cands else None
|
||||
|
||||
|
||||
def report(ckpt: Path) -> None:
|
||||
dcp_dir = ckpt / "dcp" if (ckpt / "dcp").is_dir() else ckpt
|
||||
reader = FileSystemReader(str(dcp_dir))
|
||||
md = reader.read_metadata()
|
||||
planned = md.state_dict_metadata
|
||||
sd = {}
|
||||
resolved = {}
|
||||
for want in KEYS:
|
||||
k = _find(list(planned.keys()), want)
|
||||
if k is None:
|
||||
continue
|
||||
meta = planned[k]
|
||||
sd[k] = torch.empty(tuple(meta.size), dtype=meta.properties.dtype)
|
||||
resolved[want] = k
|
||||
if not sd:
|
||||
print(f"{ckpt.name}: none of {KEYS} found in checkpoint metadata")
|
||||
return
|
||||
dcp.load(sd, checkpoint_id=str(dcp_dir))
|
||||
|
||||
out = [f"{ckpt.name:>18s}"]
|
||||
pe_k = next((resolved[k] for k in PE_KEYS if k in resolved), None)
|
||||
if pe_k is None:
|
||||
out.append("!! patch_embedding NOT FOUND — cannot report gate")
|
||||
else:
|
||||
pe = sd[pe_k].float()
|
||||
out.append(f"gate[:, {BASE_IN}:] std={pe[:, BASE_IN:].std():.6f}")
|
||||
out.append(f"pretrained std={pe[:, :BASE_IN].std():.6f}")
|
||||
for w in ("track_encoder.proj.weight", "track_encoder.temporal_conv.weight"):
|
||||
if w in resolved:
|
||||
out.append(f"{w.split('.')[1]}_norm={sd[resolved[w]].float().norm():.4f}")
|
||||
print(" ".join(out))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if len(sys.argv) < 2:
|
||||
print(__doc__)
|
||||
sys.exit(1)
|
||||
for key, default in [("RANK", "0"), ("LOCAL_RANK", "0"), ("WORLD_SIZE", "1"),
|
||||
("MASTER_ADDR", "127.0.0.1"), ("MASTER_PORT", "29555")]:
|
||||
os.environ.setdefault(key, default)
|
||||
print(f"reference: 1.3B merged-from-overfit gate std = 0.011418 (grown from 0.000000)")
|
||||
for p in sys.argv[1:]:
|
||||
try:
|
||||
report(Path(p))
|
||||
except Exception as e: # keep going across checkpoints
|
||||
print(f"{Path(p).name}: ERROR {type(e).__name__}: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 0b (DROID): build a track-conditioning training set from DROID.
|
||||
|
||||
Goal (per the project decision): a set where the **first frame is similar** across
|
||||
clips but the **motion is diverse**, with a single generic prompt, so the model is
|
||||
forced to read the point tracks to reduce loss (content cannot determine the output).
|
||||
|
||||
DROID has no scene grouping in its LeRobot metadata, so "similar first frame" is
|
||||
obtained by *clustering frame-0*: we download a candidate pool, embed each first
|
||||
frame (downscaled RGB), and pick the ``--num-clips`` episodes whose first frames are
|
||||
mutually most similar (the tightest dense cluster -- the centroid that minimises the
|
||||
radius to its k-th nearest neighbour, then its k nearest).
|
||||
|
||||
Output mirrors ``generate_videos.py`` exactly so the existing downstream
|
||||
(``extract_tracks.py`` -> ``v1_preprocess.py --preprocess_task i2v_track``) ingests it
|
||||
unchanged:
|
||||
|
||||
<out>/videos/vid_000000.mp4 ... libx264, ``--num-frames`` frames @ ``--fps``
|
||||
<out>/videos2caption.json FastVideo manifest (one generic caption)
|
||||
<out>/merge.txt "<videos_dir>,<json_path>"
|
||||
|
||||
Source: HuggingFace LeRobot mirror ``IPEC-COMMUNITY/droid_lerobot`` (v2.0): per-episode
|
||||
av1 mp4 at ``videos/chunk-{c:03d}/{video_key}/episode_{i:06d}.mp4`` (av1 -> decoded
|
||||
with PyAV; re-encoded to libx264 so downstream decord reads it). 180x320, 15 fps.
|
||||
|
||||
Run on a compute node (network + a little CPU; no GPU)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 .venv/bin/python data_pipeline/convert_droid.py \
|
||||
--out-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200 \
|
||||
--num-clips 200 --candidate-scan 1500
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
REPO = "IPEC-COMMUNITY/droid_lerobot"
|
||||
CHUNK_SIZE = 1000 # meta/info.json chunks_size
|
||||
DEFAULT_VIDEO_KEY = "observation.images.exterior_image_1_left"
|
||||
DEFAULT_PROMPT = "a robot arm manipulating objects on a tabletop"
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--out-dir", type=Path, required=True)
|
||||
p.add_argument("--num-clips", type=int, default=200, help="How many episodes to keep (the cluster size).")
|
||||
p.add_argument("--candidate-scan",
|
||||
type=int,
|
||||
default=1500,
|
||||
help="How many episodes to download+embed before clustering. Larger => tighter cluster.")
|
||||
p.add_argument("--video-key", type=str, default=DEFAULT_VIDEO_KEY)
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--fps", type=int, default=24, help="Written fps == train_fps so preprocess does NOT resample.")
|
||||
p.add_argument("--frame-start", type=int, default=0, help="First source-frame index of the kept window.")
|
||||
p.add_argument("--prompt", type=str, default=DEFAULT_PROMPT)
|
||||
p.add_argument("--embed-size", type=int, default=64, help="Frame-0 is resized to NxN for the similarity embedding.")
|
||||
p.add_argument("--no-cluster", action="store_true", help="Skip clustering; keep the first --num-clips candidates.")
|
||||
p.add_argument("--download-workers", type=int, default=16)
|
||||
p.add_argument("--seed", type=int, default=0)
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def list_candidate_episodes(num_frames: int, scan: int) -> list[int]:
|
||||
"""Episode indices (length >= num_frames), first ``scan`` in index order."""
|
||||
from huggingface_hub import hf_hub_download
|
||||
ep_path = hf_hub_download(REPO, "meta/episodes.jsonl", repo_type="dataset")
|
||||
out: list[int] = []
|
||||
with open(ep_path) as f:
|
||||
for line in f:
|
||||
d = json.loads(line)
|
||||
if int(d.get("length", 0)) >= num_frames:
|
||||
out.append(int(d["episode_index"]))
|
||||
if len(out) >= scan:
|
||||
break
|
||||
return out
|
||||
|
||||
|
||||
def episode_video_path(ep_idx: int, video_key: str) -> str:
|
||||
chunk = ep_idx // CHUNK_SIZE
|
||||
return f"videos/chunk-{chunk:03d}/{video_key}/episode_{ep_idx:06d}.mp4"
|
||||
|
||||
|
||||
def download_one(ep_idx: int, video_key: str) -> tuple[int, str | None]:
|
||||
from huggingface_hub import hf_hub_download
|
||||
try:
|
||||
p = hf_hub_download(REPO, episode_video_path(ep_idx, video_key), repo_type="dataset")
|
||||
return ep_idx, p
|
||||
except Exception: # noqa: BLE001
|
||||
return ep_idx, None
|
||||
|
||||
|
||||
def decode_frames(path: str, start: int, count: int) -> np.ndarray | None:
|
||||
"""Decode ``count`` frames starting at ``start`` (RGB uint8 [T,H,W,3]) via PyAV (handles av1)."""
|
||||
import av
|
||||
try:
|
||||
container = av.open(path)
|
||||
frames = []
|
||||
for idx, frame in enumerate(container.decode(video=0)):
|
||||
if idx >= start:
|
||||
frames.append(frame.to_ndarray(format="rgb24"))
|
||||
if len(frames) >= count:
|
||||
break
|
||||
container.close()
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if len(frames) < count:
|
||||
return None
|
||||
return np.stack(frames, axis=0)
|
||||
|
||||
|
||||
def embed_frame0(path: str, size: int) -> np.ndarray | None:
|
||||
from PIL import Image
|
||||
fr = decode_frames(path, 0, 1)
|
||||
if fr is None:
|
||||
return None
|
||||
img = Image.fromarray(fr[0]).resize((size, size), Image.BILINEAR)
|
||||
v = np.asarray(img, dtype=np.float32).reshape(-1)
|
||||
n = np.linalg.norm(v) + 1e-8
|
||||
return v / n
|
||||
|
||||
|
||||
def pick_tightest_cluster(embs: np.ndarray, k: int) -> np.ndarray:
|
||||
"""Return indices of the k mutually-most-similar embeddings.
|
||||
|
||||
Centroid = the point minimising the distance to its k-th nearest neighbour
|
||||
(densest region); then take that point's k nearest neighbours.
|
||||
"""
|
||||
# cosine distance (embs are L2-normalised) -> 1 - sim
|
||||
sim = embs @ embs.T
|
||||
dist = 1.0 - sim
|
||||
np.fill_diagonal(dist, 0.0)
|
||||
part = np.partition(dist, kth=min(k, dist.shape[0] - 1), axis=1)
|
||||
kth = part[:, min(k, dist.shape[0] - 1)] # radius to k-th NN per row
|
||||
centroid = int(np.argmin(kth))
|
||||
order = np.argsort(dist[centroid])[:k]
|
||||
return order
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
out_dir = args.out_dir
|
||||
videos_dir = out_dir / "videos"
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
json_path = out_dir / "videos2caption.json"
|
||||
merge_path = out_dir / "merge.txt"
|
||||
|
||||
print(f"[droid] listing candidates (length>={args.num_frames}, scan={args.candidate_scan})...", flush=True)
|
||||
candidates = list_candidate_episodes(args.num_frames, args.candidate_scan)
|
||||
print(f"[droid] {len(candidates)} candidate episodes", flush=True)
|
||||
|
||||
# Download (cached) + embed frame-0 in parallel.
|
||||
print(f"[droid] downloading + embedding frame-0 ({args.download_workers} workers)...", flush=True)
|
||||
ep_to_path: dict[int, str] = {}
|
||||
with ThreadPoolExecutor(max_workers=args.download_workers) as ex:
|
||||
futs = {ex.submit(download_one, ep, args.video_key): ep for ep in candidates}
|
||||
for i, fut in enumerate(as_completed(futs), 1):
|
||||
ep, path = fut.result()
|
||||
if path is not None:
|
||||
ep_to_path[ep] = path
|
||||
if i % 100 == 0:
|
||||
print(f"[droid] downloaded {i}/{len(candidates)}", flush=True)
|
||||
print(f"[droid] downloaded {len(ep_to_path)} mp4s; embedding...", flush=True)
|
||||
|
||||
eps_ok: list[int] = []
|
||||
embs: list[np.ndarray] = []
|
||||
for ep in candidates:
|
||||
if ep not in ep_to_path:
|
||||
continue
|
||||
e = embed_frame0(ep_to_path[ep], args.embed_size)
|
||||
if e is not None:
|
||||
eps_ok.append(ep)
|
||||
embs.append(e)
|
||||
embs_arr = np.stack(embs, axis=0)
|
||||
print(f"[droid] embedded {len(eps_ok)} frame-0s", flush=True)
|
||||
|
||||
if args.no_cluster or len(eps_ok) <= args.num_clips:
|
||||
sel_local = list(range(min(args.num_clips, len(eps_ok))))
|
||||
print(f"[droid] no clustering -> first {len(sel_local)} candidates", flush=True)
|
||||
else:
|
||||
sel_local = pick_tightest_cluster(embs_arr, args.num_clips).tolist()
|
||||
# report cluster tightness
|
||||
sub = embs_arr[sel_local]
|
||||
c = sub.mean(0, keepdims=True)
|
||||
c /= np.linalg.norm(c) + 1e-8
|
||||
mean_sim = float((sub @ c.T).mean())
|
||||
print(f"[droid] tightest-cluster of {len(sel_local)} clips; mean cos-sim to centroid={mean_sim:.3f}",
|
||||
flush=True)
|
||||
|
||||
selected_eps = [eps_ok[i] for i in sel_local]
|
||||
|
||||
# Decode + write libx264 mp4s + manifest.
|
||||
import imageio.v2 as imageio
|
||||
records = []
|
||||
seq = 0
|
||||
for ep in selected_eps:
|
||||
frames = decode_frames(ep_to_path[ep], args.frame_start, args.num_frames)
|
||||
if frames is None:
|
||||
print(f"[droid] skip ep{ep}: decode<{args.num_frames} frames", flush=True)
|
||||
continue
|
||||
H, W = int(frames.shape[1]), int(frames.shape[2])
|
||||
name = f"vid_{seq:06d}.mp4"
|
||||
imageio.mimsave(str(videos_dir / name),
|
||||
list(frames),
|
||||
fps=args.fps,
|
||||
codec="libx264",
|
||||
macro_block_size=1,
|
||||
output_params=["-pix_fmt", "yuv420p"])
|
||||
records.append({
|
||||
"idx": seq,
|
||||
"path": name,
|
||||
"cap": [args.prompt],
|
||||
"fps": float(args.fps),
|
||||
"duration": float(args.num_frames) / float(args.fps),
|
||||
"num_frames": int(args.num_frames),
|
||||
"resolution": {
|
||||
"width": W,
|
||||
"height": H
|
||||
},
|
||||
"source_episode": int(ep),
|
||||
})
|
||||
seq += 1
|
||||
if seq % 25 == 0:
|
||||
print(f"[droid] wrote {seq}/{len(selected_eps)}", flush=True)
|
||||
|
||||
json_path.write_text(json.dumps(records, indent=2))
|
||||
merge_path.write_text(f"{videos_dir.resolve()},{json_path.resolve()}\n")
|
||||
print(f"[droid] DONE: wrote {len(records)} clips -> {videos_dir}", flush=True)
|
||||
print(f"[droid] manifest -> {json_path}", flush=True)
|
||||
print(f"[droid] next: extract_tracks.py --data-dir {out_dir} ; then v1_preprocess i2v_track", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,156 @@
|
||||
# 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."""
|
||||
# NOTE(local): TrackEncoder builds BOTH convs with bias=False (load-bearing in the
|
||||
# model — see dits/trackwan/track_encoder.py). Match that exactly: bias-free convs,
|
||||
# emit only the .weight keys (no spurious .bias params that the model has no home for).
|
||||
temporal_conv = nn.Conv3d(ID_DIM, TRACK_CHANNELS, kernel_size=(VAE_T_COMP, 1, 1), stride=(VAE_T_COMP, 1, 1), bias=False)
|
||||
proj = nn.Conv3d(TRACK_CHANNELS, TRACK_CHANNELS, kernel_size=1, bias=False) # default init (NOT zero) -> breaks the deadlock
|
||||
return {
|
||||
"track_encoder.temporal_conv.weight": temporal_conv.weight.detach().clone(),
|
||||
"track_encoder.proj.weight": proj.weight.detach().clone(),
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
global ID_DIM
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--base",
|
||||
required=True,
|
||||
help="Base Wan I2V (e.g. Wan2.1-Fun-1.3B-InP) or T2V diffusers model dir (with transformer/).")
|
||||
p.add_argument("--out", required=True, help="Output model dir for the WanTrack init.")
|
||||
p.add_argument("--id-dim", type=int, default=ID_DIM,
|
||||
help="sinusoidal track-id posemb dim (MotionStream d; 64 for d64 init, default 128)")
|
||||
args = p.parse_args()
|
||||
|
||||
ID_DIM = args.id_dim
|
||||
TRACK_CONFIG["id_dim"] = args.id_dim
|
||||
|
||||
base = Path(args.base)
|
||||
out = Path(args.out)
|
||||
if not (base / "transformer" / "config.json").exists():
|
||||
raise FileNotFoundError(f"{base}/transformer/config.json not found")
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 1) Symlink every top-level entry except transformer/ (vae, text_encoder, ...).
|
||||
for entry in os.listdir(base):
|
||||
if entry == "transformer":
|
||||
continue
|
||||
src = (base / entry).resolve()
|
||||
dst = out / entry
|
||||
if dst.is_symlink() or dst.exists():
|
||||
if dst.is_dir() and not dst.is_symlink():
|
||||
shutil.rmtree(dst)
|
||||
else:
|
||||
dst.unlink()
|
||||
os.symlink(src, dst)
|
||||
|
||||
# 2) Convert transformer/.
|
||||
tdir = out / "transformer"
|
||||
tdir.mkdir(exist_ok=True)
|
||||
cfg = json.loads((base / "transformer" / "config.json").read_text())
|
||||
base_in = int(cfg["in_channels"])
|
||||
cfg["in_channels"] = NEW_IN_CHANNELS
|
||||
cfg["track_config"] = TRACK_CONFIG
|
||||
(tdir / "config.json").write_text(json.dumps(cfg, indent=2))
|
||||
|
||||
sf_files = sorted((base / "transformer").glob("*.safetensors"))
|
||||
if not sf_files:
|
||||
raise FileNotFoundError(f"No transformer safetensors under {base}/transformer")
|
||||
state: dict[str, torch.Tensor] = {}
|
||||
for sf in sf_files:
|
||||
state.update(load_file(str(sf))) # merge shards if the base is sharded
|
||||
|
||||
pe_key = "patch_embedding.weight"
|
||||
if pe_key not in state:
|
||||
cands = [k for k in state if "patch_embed" in k and k.endswith(".weight")]
|
||||
if len(cands) != 1:
|
||||
raise KeyError(f"Could not find patch_embedding weight; candidates={cands}")
|
||||
pe_key = cands[0]
|
||||
w = state[pe_key] # [out, base_in, 1, 2, 2]
|
||||
if w.shape[1] != base_in:
|
||||
raise ValueError(f"{pe_key} in_ch {w.shape[1]} != config in_channels {base_in}")
|
||||
new_w = torch.zeros((w.shape[0], NEW_IN_CHANNELS, *w.shape[2:]), dtype=w.dtype)
|
||||
new_w[:, :base_in] = w # pretrained channels first; new channels zero
|
||||
state[pe_key] = new_w
|
||||
print(f"[convert] padded {pe_key}: {tuple(w.shape)} -> {tuple(new_w.shape)}")
|
||||
|
||||
te_state = build_track_encoder_state()
|
||||
for k, v in te_state.items():
|
||||
state[k] = v.to(w.dtype)
|
||||
print(f"[convert] added {len(te_state)} track_encoder tensors")
|
||||
|
||||
save_file(state, str(tdir / "diffusion_pytorch_model.safetensors"), metadata={"format": "pt"})
|
||||
print(f"[convert] wrote {tdir/'diffusion_pytorch_model.safetensors'} ({len(state)} keys)")
|
||||
print(f"[convert] done -> {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,245 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Build a WanTrack init checkpoint from a base Wan diffusers model, with several
|
||||
strategies for the patch_embedding track-channel slot and the track_encoder.
|
||||
|
||||
Modes for patch_embedding[:, base_in:] (the new 16 in-channels):
|
||||
--pe-init zero : zero-pad (original stage-1 recipe)
|
||||
--pe-init random : normal-init with std matched to base first-N channels
|
||||
|
||||
Modes for track_encoder.{proj, temporal_conv}:
|
||||
--track-src default : default Conv init (used by original convert)
|
||||
--track-src <ckpt.safetensors> : lift both convs from an existing WanTrack safetensors
|
||||
(e.g. a trained 1.3B ckpt). Bias is copied when present
|
||||
and its shape matches.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base /mnt/lustre/vlm-s4duan/models/Wan2.1-I2V-14B-720P-Diffusers \
|
||||
--out /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_random_init \
|
||||
--id-dim 64 --pe-init random
|
||||
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base /mnt/lustre/vlm-s4duan/models/Wan2.1-I2V-14B-720P-Diffusers \
|
||||
--out /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_partial_merged \
|
||||
--id-dim 64 --pe-init random \
|
||||
--track-src /mnt/lustre/vlm-s4duan/exports/synth_stage2_paperLR_ckpt600/transformer/model.safetensors
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file, save_file
|
||||
from torch import nn
|
||||
|
||||
NEW_IN_CHANNELS = 52
|
||||
TRACK_CHANNELS = 16
|
||||
VAE_T_COMP = 4
|
||||
DEFAULT_ID_DIM = 128
|
||||
|
||||
|
||||
def build_track_encoder_default(id_dim: int, dtype: torch.dtype, use_bias: bool = False) -> dict[str, torch.Tensor]:
|
||||
"""Default (non-zero) Conv init for track_encoder.proj + temporal_conv."""
|
||||
temporal_conv = nn.Conv3d(id_dim, TRACK_CHANNELS,
|
||||
kernel_size=(VAE_T_COMP, 1, 1), stride=(VAE_T_COMP, 1, 1),
|
||||
bias=use_bias)
|
||||
proj = nn.Conv3d(TRACK_CHANNELS, TRACK_CHANNELS, kernel_size=1, bias=use_bias)
|
||||
d = {
|
||||
"track_encoder.temporal_conv.weight": temporal_conv.weight.detach().to(dtype).clone(),
|
||||
"track_encoder.proj.weight": proj.weight.detach().to(dtype).clone(),
|
||||
}
|
||||
if use_bias:
|
||||
d["track_encoder.temporal_conv.bias"] = temporal_conv.bias.detach().to(dtype).clone()
|
||||
d["track_encoder.proj.bias"] = proj.bias.detach().to(dtype).clone()
|
||||
return d
|
||||
|
||||
|
||||
def load_src_state(src_path: str) -> dict[str, torch.Tensor]:
|
||||
"""Load a source state dict from a .safetensors FILE or a diffusers model DIR.
|
||||
|
||||
A 14B export is sharded across several *.safetensors, so accept a directory (either the
|
||||
model root or its transformer/ subdir) and merge every shard.
|
||||
"""
|
||||
p = Path(src_path)
|
||||
if p.is_file():
|
||||
return load_file(str(p))
|
||||
tdir = p / "transformer" if (p / "transformer").is_dir() else p
|
||||
shards = sorted(tdir.glob("*.safetensors"))
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"no *.safetensors under {tdir}")
|
||||
state: dict[str, torch.Tensor] = {}
|
||||
for s in shards:
|
||||
state.update(load_file(str(s)))
|
||||
print(f"[src] loaded {len(state)} tensors from {len(shards)} shard(s) in {tdir}")
|
||||
return state
|
||||
|
||||
|
||||
def lift_track_encoder_from_src(src_path: str, id_dim: int, dtype: torch.dtype) -> dict[str, torch.Tensor]:
|
||||
"""Load track_encoder.{proj, temporal_conv}.{weight, bias} from an existing WanTrack ckpt.
|
||||
|
||||
Verifies shape compatibility. If bias is missing in the source we skip it; if the
|
||||
source shapes don't match the expected shapes we fall back to default init for that
|
||||
key and print a warning.
|
||||
"""
|
||||
src_state = load_src_state(src_path)
|
||||
expected_shapes = {
|
||||
"track_encoder.temporal_conv.weight": (TRACK_CHANNELS, id_dim, VAE_T_COMP, 1, 1),
|
||||
"track_encoder.proj.weight": (TRACK_CHANNELS, TRACK_CHANNELS, 1, 1, 1),
|
||||
"track_encoder.temporal_conv.bias": (TRACK_CHANNELS,),
|
||||
"track_encoder.proj.bias": (TRACK_CHANNELS,),
|
||||
}
|
||||
picked: dict[str, torch.Tensor] = {}
|
||||
for k, want in expected_shapes.items():
|
||||
if k not in src_state:
|
||||
continue
|
||||
got = tuple(src_state[k].shape)
|
||||
if got != want:
|
||||
print(f"[track-src] SHAPE MISMATCH on {k}: src={got} vs expected={want}; SKIPPING (will need fallback)")
|
||||
continue
|
||||
picked[k] = src_state[k].detach().to(dtype).clone()
|
||||
print(f"[track-src] lifted {k} shape={got}")
|
||||
# If either weight is missing, we still need to keep the model loadable.
|
||||
fallback = build_track_encoder_default(id_dim, dtype, use_bias=False)
|
||||
for k, v in fallback.items():
|
||||
if k not in picked:
|
||||
print(f"[track-src] {k} not in source -> default init")
|
||||
picked[k] = v
|
||||
return picked
|
||||
|
||||
|
||||
def std_of_first_n(w: torch.Tensor, n: int) -> float:
|
||||
"""Std of the pretrained first-N in-channels slice of patch_embedding.weight."""
|
||||
with torch.no_grad():
|
||||
s = w[:, :n].float().std().item()
|
||||
return float(s)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--base", required=True, help="Base Wan diffusers model dir.")
|
||||
p.add_argument("--out", required=True, help="Output diffusers model dir.")
|
||||
p.add_argument("--id-dim", type=int, default=DEFAULT_ID_DIM)
|
||||
p.add_argument("--pe-init", choices=("zero", "random"), default="zero",
|
||||
help="patch_embedding init strategy for the ADDED track channels [:, base_in:]. "
|
||||
"Base pretrained channels [:, :base_in] are always preserved.")
|
||||
p.add_argument("--pe-random-seed", type=int, default=1234)
|
||||
p.add_argument("--track-src", default=None,
|
||||
help="Path to a .safetensors (or diffusers model dir) with track_encoder.* to lift.")
|
||||
p.add_argument("--pe-src", default=None,
|
||||
help="Path to a .safetensors (or diffusers model dir) to lift patch_embedding.weight[:, base_in:] "
|
||||
"from — i.e. the TRACK SLOT only; the pretrained [:, :base_in] channels always come from "
|
||||
"--base. Use together with --track-src pointing at the SAME source so the encoder and the "
|
||||
"track slot stay CO-ADAPTED (lifting the encoder alone is worse than random init).")
|
||||
p.add_argument("--use-bias-defaults", action="store_true",
|
||||
help="When --track-src is not set, build track_encoder convs WITH bias (matches TRACKWAN_TRACK_BIAS=1 training).")
|
||||
args = p.parse_args()
|
||||
|
||||
id_dim = int(args.id_dim)
|
||||
TRACK_CONFIG = {
|
||||
"id_dim": id_dim,
|
||||
"track_channels": TRACK_CHANNELS,
|
||||
"vae_spatial_compression": 8,
|
||||
"vae_temporal_compression": VAE_T_COMP,
|
||||
"max_track_id": 100_000,
|
||||
"zero_init_head": False,
|
||||
}
|
||||
|
||||
base = Path(args.base)
|
||||
out = Path(args.out)
|
||||
if not (base / "transformer" / "config.json").exists():
|
||||
raise FileNotFoundError(f"{base}/transformer/config.json not found")
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 1) Symlink every top-level entry except transformer/
|
||||
for entry in os.listdir(base):
|
||||
if entry == "transformer":
|
||||
continue
|
||||
src_p = (base / entry).resolve()
|
||||
dst = out / entry
|
||||
if dst.is_symlink() or dst.exists():
|
||||
if dst.is_dir() and not dst.is_symlink():
|
||||
shutil.rmtree(dst)
|
||||
else:
|
||||
dst.unlink()
|
||||
os.symlink(src_p, dst)
|
||||
|
||||
# 2) Rewrite transformer/config.json
|
||||
tdir = out / "transformer"
|
||||
tdir.mkdir(exist_ok=True)
|
||||
cfg = json.loads((base / "transformer" / "config.json").read_text())
|
||||
base_in = int(cfg["in_channels"])
|
||||
cfg["in_channels"] = NEW_IN_CHANNELS
|
||||
cfg["track_config"] = TRACK_CONFIG
|
||||
(tdir / "config.json").write_text(json.dumps(cfg, indent=2))
|
||||
|
||||
# 3) Load full base transformer state
|
||||
sf_files = sorted((base / "transformer").glob("*.safetensors"))
|
||||
if not sf_files:
|
||||
raise FileNotFoundError(f"No safetensors under {base}/transformer")
|
||||
state: dict[str, torch.Tensor] = {}
|
||||
for sf in sf_files:
|
||||
state.update(load_file(str(sf)))
|
||||
dtype = next(iter(state.values())).dtype
|
||||
|
||||
# 4) Grow patch_embedding.weight from base_in -> 52 input channels
|
||||
pe_key = "patch_embedding.weight"
|
||||
if pe_key not in state:
|
||||
raise KeyError(f"{pe_key} not found in base transformer")
|
||||
w = state[pe_key] # [hidden, base_in, 1, 2, 2]
|
||||
if w.shape[1] != base_in:
|
||||
raise ValueError(f"{pe_key} in_ch {w.shape[1]} != config in_channels {base_in}")
|
||||
new_w = torch.empty((w.shape[0], NEW_IN_CHANNELS, *w.shape[2:]), dtype=w.dtype)
|
||||
new_w[:, :base_in] = w # pretrained channels first
|
||||
if args.pe_init == "zero":
|
||||
new_w[:, base_in:] = 0
|
||||
print(f"[convert] pe-init=zero -> new channels are zero")
|
||||
else:
|
||||
std = std_of_first_n(w, base_in)
|
||||
g = torch.Generator().manual_seed(int(args.pe_random_seed))
|
||||
noise = torch.empty_like(new_w[:, base_in:], dtype=torch.float32)
|
||||
noise.normal_(mean=0.0, std=std, generator=g)
|
||||
new_w[:, base_in:] = noise.to(w.dtype)
|
||||
print(f"[convert] pe-init=random N(0, {std:.5f}) matching first-{base_in} std")
|
||||
# 4b) Optionally overwrite ONLY the track slot [:, base_in:] from a trained source.
|
||||
# This is the other half of the "merged" recipe: the track slot and the track_encoder were
|
||||
# co-adapted during the overfit, so they must be lifted TOGETHER (--pe-src + --track-src from
|
||||
# the same ckpt). Lifting the encoder alone leaves it talking to a random projection.
|
||||
if args.pe_src:
|
||||
src_pe_state = load_src_state(args.pe_src)
|
||||
if pe_key not in src_pe_state:
|
||||
raise KeyError(f"{pe_key} not found in --pe-src {args.pe_src}")
|
||||
src_pe = src_pe_state[pe_key]
|
||||
want = (w.shape[0], NEW_IN_CHANNELS, *w.shape[2:])
|
||||
if tuple(src_pe.shape) != want:
|
||||
raise ValueError(f"--pe-src {pe_key} shape {tuple(src_pe.shape)} != expected {want}")
|
||||
slot = src_pe[:, base_in:].to(w.dtype)
|
||||
new_w[:, base_in:] = slot
|
||||
print(f"[pe-src] lifted {pe_key}[:, {base_in}:] from source "
|
||||
f"(std={slot.float().std().item():.5f}, absmax={slot.float().abs().max().item():.5f})")
|
||||
|
||||
state[pe_key] = new_w
|
||||
print(f"[convert] patch_embedding.weight: {tuple(w.shape)} -> {tuple(new_w.shape)}")
|
||||
with torch.no_grad():
|
||||
print(f"[convert] pe[:, :{base_in}] std={new_w[:, :base_in].float().std().item():.5f} (pretrained) | "
|
||||
f"pe[:, {base_in}:] std={new_w[:, base_in:].float().std().item():.5f} (track slot)")
|
||||
|
||||
# 5) Add track_encoder
|
||||
if args.track_src:
|
||||
te = lift_track_encoder_from_src(args.track_src, id_dim, dtype)
|
||||
else:
|
||||
te = build_track_encoder_default(id_dim, dtype, use_bias=args.use_bias_defaults)
|
||||
for k, v in te.items():
|
||||
state[k] = v
|
||||
print(f"[convert] added {len(te)} track_encoder tensors")
|
||||
|
||||
save_file(state, str(tdir / "diffusion_pytorch_model.safetensors"), metadata={"format": "pt"})
|
||||
print(f"[convert] wrote {tdir/'diffusion_pytorch_model.safetensors'} ({len(state)} keys)")
|
||||
print(f"[convert] done -> {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,162 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Decode FastVideo Wan2.2-Syn-121x704x1280_32k parquet latents to 480x832 mp4s.
|
||||
|
||||
The HF dataset stores Wan 2.2 5B VAE latents (48-channel, 16x spatial, 4x
|
||||
temporal). We can't feed these to our Wan 2.1 WanTrack model directly, so we
|
||||
decode → pixels → downsize to our training resolution 480x832 → save mp4.
|
||||
Downstream (SAM + CoTracker + i2v_track preprocess) then treats these as if they
|
||||
were raw synth mp4s.
|
||||
|
||||
Idempotent + resumable: skips samples whose mp4 already exists.
|
||||
|
||||
Slurm sharded via ``--num-shards``/``--shard``. Each worker processes
|
||||
``parquets[shard::num_shards]``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def _resize_frames(frames: np.ndarray, out_h: int, out_w: int) -> np.ndarray:
|
||||
"""Center-crop-then-resize a (T,H,W,3) uint8 stack to (T,out_h,out_w,3)."""
|
||||
T, H, W, _ = frames.shape
|
||||
src_ar = W / H
|
||||
dst_ar = out_w / out_h
|
||||
if src_ar > dst_ar:
|
||||
new_w = int(round(H * dst_ar))
|
||||
x0 = (W - new_w) // 2
|
||||
cropped = frames[:, :, x0:x0 + new_w, :]
|
||||
else:
|
||||
new_h = int(round(W / dst_ar))
|
||||
y0 = (H - new_h) // 2
|
||||
cropped = frames[:, y0:y0 + new_h, :, :]
|
||||
out = np.empty((T, out_h, out_w, 3), dtype=np.uint8)
|
||||
for t in range(T):
|
||||
out[t] = np.asarray(Image.fromarray(cropped[t]).resize((out_w, out_h), Image.BILINEAR))
|
||||
return out
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_row(vae: AutoencoderKLWan, latent_bytes: bytes, latent_shape: list[int],
|
||||
device: torch.device, dtype: torch.dtype) -> np.ndarray:
|
||||
"""Decode one row's Wan 2.2 5B latent to a [T,704,1280,3] uint8 pixel stack."""
|
||||
lat = np.frombuffer(latent_bytes, dtype=np.float32).reshape(latent_shape)
|
||||
lat = torch.from_numpy(lat).to(device=device, dtype=dtype).unsqueeze(0) # [1,C,T,H,W]
|
||||
# unnormalize per-channel
|
||||
m = torch.tensor(vae.config.latents_mean, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
|
||||
s = torch.tensor(vae.config.latents_std, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
|
||||
lat = lat * s + m
|
||||
pix = vae.decode(lat, return_dict=False)[0] # [1,3,T_out,H*16,W*16] in [-1,1]
|
||||
pix = pix.clamp(-1, 1).float()
|
||||
pix = ((pix + 1.0) * 127.5).round().to(torch.uint8)
|
||||
pix = pix[0].permute(1, 2, 3, 0).cpu().numpy() # [T,H,W,3]
|
||||
return pix
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--parquet-dir", required=True, help="dir containing train/Part_*/latents_chunk_*.parquet")
|
||||
p.add_argument("--vae-dir", required=True, help="Wan 2.2 5B VAE dir")
|
||||
p.add_argument("--out-dir", required=True, help="output root; writes videos/vid_*.mp4 + meta/*.json")
|
||||
p.add_argument("--out-h", type=int, default=480)
|
||||
p.add_argument("--out-w", type=int, default=832)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--shard", type=int, default=0)
|
||||
p.add_argument("--num-shards", type=int, default=1)
|
||||
p.add_argument("--dtype", default="bfloat16", choices=["float32", "bfloat16", "float16"])
|
||||
p.add_argument("--limit-rows", type=int, default=None, help="cap rows for smoke test")
|
||||
args = p.parse_args()
|
||||
|
||||
device = torch.device("cuda")
|
||||
dtype = getattr(torch, args.dtype)
|
||||
|
||||
vids_dir = Path(args.out_dir) / "videos"
|
||||
meta_dir = Path(args.out_dir) / "meta"
|
||||
vids_dir.mkdir(parents=True, exist_ok=True)
|
||||
meta_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
files = sorted(glob.glob(os.path.join(args.parquet_dir, "**", "*.parquet"), recursive=True))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"no *.parquet under {args.parquet_dir}")
|
||||
my_files = files[args.shard::args.num_shards]
|
||||
print(f"[dec {args.shard}/{args.num_shards}] {len(my_files)} parquets", flush=True)
|
||||
|
||||
vae = AutoencoderKLWan.from_pretrained(args.vae_dir, torch_dtype=dtype).to(device).eval()
|
||||
|
||||
n_ok = 0
|
||||
n_skip = 0
|
||||
n_err = 0
|
||||
t0 = time.time()
|
||||
row_budget = args.limit_rows or 10**12
|
||||
processed_rows = 0
|
||||
for fk, f in enumerate(my_files):
|
||||
try:
|
||||
table = pq.read_table(f)
|
||||
except Exception as e: # skip unreadable parquets
|
||||
print(f"[dec {args.shard}] parquet read fail {f}: {e}", flush=True)
|
||||
continue
|
||||
# Dataset's `id` field only encodes chunk index (not Part), so the same id
|
||||
# appears in many Parts. Prefix Part directory so vid_ids are unique.
|
||||
part_name = Path(f).parent.name # e.g. "Part_57"
|
||||
rows = table.to_pylist()
|
||||
for r in rows:
|
||||
if processed_rows >= row_budget:
|
||||
break
|
||||
processed_rows += 1
|
||||
row_key = str(r.get("id") or r.get("file_name") or f"{fk:05d}_{n_ok:04d}")
|
||||
vid_id = f"{part_name}_{row_key}"
|
||||
vid_path = vids_dir / f"{vid_id}.mp4"
|
||||
meta_path = meta_dir / f"{vid_id}.json"
|
||||
if vid_path.exists() and meta_path.exists():
|
||||
n_skip += 1
|
||||
continue
|
||||
try:
|
||||
pix = decode_row(vae, r["vae_latent_bytes"], list(r["vae_latent_shape"]),
|
||||
device=device, dtype=dtype)
|
||||
resized = _resize_frames(pix, args.out_h, args.out_w) # [T, out_h, out_w, 3]
|
||||
# atomic mp4 write (imageio infers format from extension → keep .mp4)
|
||||
tmp = vid_path.parent / f".{vid_path.name}.tmp.mp4"
|
||||
imageio.mimsave(str(tmp), resized, fps=args.fps, codec="libx264",
|
||||
quality=7, macro_block_size=1)
|
||||
tmp.rename(vid_path)
|
||||
meta = {
|
||||
"path": vid_path.name,
|
||||
"id": vid_id,
|
||||
"cap": [str(r.get("caption") or "")],
|
||||
"fps": float(args.fps),
|
||||
"num_frames": int(resized.shape[0]),
|
||||
"resolution": [int(args.out_h), int(args.out_w)],
|
||||
"src_resolution": [int(r.get("height", 704)), int(r.get("width", 1280))],
|
||||
"src_dataset": "FastVideo/Wan2.2-Syn-121x704x1280_32k",
|
||||
}
|
||||
meta_path.write_text(json.dumps(meta))
|
||||
n_ok += 1
|
||||
if n_ok % 20 == 0:
|
||||
dt = time.time() - t0
|
||||
rate = n_ok / max(dt, 1e-9)
|
||||
print(f"[dec {args.shard}] ok={n_ok} skip={n_skip} err={n_err} "
|
||||
f"rate={rate:.2f}/s", flush=True)
|
||||
except Exception as e: # skip a bad row, keep going
|
||||
n_err += 1
|
||||
print(f"[dec {args.shard}] row {vid_id} FAILED: {e}", flush=True)
|
||||
if processed_rows >= row_budget:
|
||||
break
|
||||
|
||||
print(f"[dec {args.shard}] DONE ok={n_ok} skip={n_skip} err={n_err} "
|
||||
f"in {(time.time()-t0)/60:.1f}min", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Diagnose the first-frame I2V conditioning + visualize synthetic tracks.
|
||||
|
||||
(1) First-frame conditioning bug: the preprocessing stored ``first_frame_latent``
|
||||
using vae.scaling_factor/shift_factor (absent for Wan -> effectively RAW),
|
||||
while the model normalizes latents with per-channel latents_mean/std. This
|
||||
script decodes the stored conditioning under both interpretations and compares
|
||||
to the GT first frame to localize the space mismatch, and shows the current
|
||||
model's generated first frame.
|
||||
|
||||
(2) Synthetic-track previews: overlays authored controls (pan/zoom/drag/...) on the
|
||||
real first frame so you can see what the control signals look like.
|
||||
|
||||
Outputs PNGs to --out (default research_log/figures/).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def _mse(a, b):
|
||||
a = a[:min(len(a), len(b))].astype(np.float32)
|
||||
b = b[:len(a)].astype(np.float32)
|
||||
return float(((a - b) ** 2).mean())
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _decode_frame0(model, latent_bcthw, *, already_normalized: bool) -> np.ndarray:
|
||||
"""latent [1,16,T,H,W] -> frame0 uint8 (H,W,3). If raw, normalize first."""
|
||||
from fastvideo.training.training_utils import normalize_dit_input
|
||||
x = latent_bcthw.to(model.device, torch.bfloat16)
|
||||
if not already_normalized:
|
||||
x = normalize_dit_input("wan", x, model.vae)
|
||||
px = model.decode_latents(x.permute(0, 2, 1, 3, 4))[0] # [3,T,H,W] in [0,1]
|
||||
f0 = (px[:, 0].clamp(0, 1).float().cpu().numpy() * 255).astype(np.uint8)
|
||||
return np.transpose(f0, (1, 2, 0)) # H,W,3
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track/combined_parquet_dataset")
|
||||
p.add_argument("--out", default="research_log/figures")
|
||||
p.add_argument("--clip", type=int, default=0)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
s = twi.load_conditioning_from_parquet(args.data, [args.clip], text_len)[0]
|
||||
|
||||
vae_latent = s["vae_latent"] # raw GT latents
|
||||
ff = s["first_frame_latent"] # stored conditioning latent
|
||||
|
||||
# GT first frame: vae_latent is raw -> normalize -> decode.
|
||||
gt0 = _decode_frame0(model, vae_latent, already_normalized=False)
|
||||
|
||||
# Stored conditioning decoded TWO ways to find its true space:
|
||||
# (a) "model's view": treat stored as already-normalized (decode_latents denorms it).
|
||||
# This is effectively what the model was conditioned on.
|
||||
cond_as_norm0 = _decode_frame0(model, ff, already_normalized=True)
|
||||
# (b) "as raw": treat stored as raw -> normalize -> decode.
|
||||
cond_as_raw0 = _decode_frame0(model, ff, already_normalized=False)
|
||||
|
||||
# Current model's generated first frame (GT tracks).
|
||||
Tpx = (vae_latent.shape[2] - 1) * int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio) + 1
|
||||
lat = twi.generate(model, first_frame_latent=ff, text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=s["track_points"][:, :Tpx], track_visibility=s["track_visibility"][:, :Tpx],
|
||||
num_steps=args.steps, seed=1000)
|
||||
gen0 = twi.decode_to_pixels(model, lat)[0]
|
||||
|
||||
# Save individual + side-by-side.
|
||||
panels = {"1_gt_first_frame": gt0,
|
||||
"2_cond_as_normalized_MODELVIEW": cond_as_norm0,
|
||||
"3_cond_as_raw_then_normalized": cond_as_raw0,
|
||||
"4_generated_first_frame": gen0}
|
||||
for name, img in panels.items():
|
||||
imageio.imwrite(os.path.join(args.out, f"firstframe_clip{args.clip}_{name}.png"), img)
|
||||
strip = np.concatenate([gt0, cond_as_norm0, cond_as_raw0, gen0], axis=1)
|
||||
imageio.imwrite(os.path.join(args.out, f"firstframe_clip{args.clip}_compare.png"), strip)
|
||||
|
||||
print("=== first-frame diagnosis (clip %d) ===" % args.clip)
|
||||
print(f" MSE(GT, cond-as-normalized / MODEL'S VIEW) = {_mse(gt0, cond_as_norm0):8.1f}")
|
||||
print(f" MSE(GT, cond-as-raw-then-normalized) = {_mse(gt0, cond_as_raw0):8.1f}")
|
||||
print(f" MSE(GT, generated first frame) = {_mse(gt0, gen0):8.1f}")
|
||||
print(" -> if 'as-raw' MSE << 'model's view' MSE, the stored conditioning is RAW")
|
||||
print(" and the model was fed a mis-normalized first frame (the bug).")
|
||||
|
||||
# ---- synthetic track previews on the real first frame ----
|
||||
for name in ["pan_right", "zoom_in", "rotate_cw", "drag_center_right", "swirl"]:
|
||||
tr, vis = st.preset(name, Tpx, 50, strength=0.25)
|
||||
H, W, _ = gt0.shape
|
||||
ov = st._overlay_preview(gt0, st.to_pixel(tr, H, W), vis, stride=3)
|
||||
imageio.imwrite(os.path.join(args.out, f"trackpreview_{name}.png"), ov)
|
||||
print(f"\n[viz] wrote first-frame panels + synthetic track previews to {args.out}/")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,47 @@
|
||||
#!/usr/bin/env bash
|
||||
# Self-healing OpenVidHD download: HF unauthenticated throttles ~per-200GB then
|
||||
# resets after a cooldown. Run N idempotent download shards; watchdog detects a
|
||||
# stall (no zip growth) and kills+cooldowns+relaunches (resume) until all 98 parts
|
||||
# done. Runs as ONE srun on the node so local kill works. Honors $HF_TOKEN if set.
|
||||
set +e
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
export HF_HOME=$WORK/.hf
|
||||
cd "$WORK/FastVideo"; source .venv/bin/activate
|
||||
NSHARD=${NSHARD:-4}
|
||||
DC(){ ls "$WORK"/openvid_1m/_extracted/*.done 2>/dev/null | wc -l; }
|
||||
ZB(){ find "$WORK"/openvid_1m/_zips -type f -printf '%s\n' 2>/dev/null | awk '{s+=$1}END{print s+0}'; }
|
||||
|
||||
pkill -f openvid_download_hd.py 2>/dev/null; sleep 3
|
||||
round=0
|
||||
while [ "$(DC)" -lt 98 ]; do
|
||||
round=$((round+1))
|
||||
echo "[$(date +%H:%M)] ROUND $round start parts=$(DC)/98 ${HF_TOKEN:+(token set)}"
|
||||
pids=()
|
||||
for s in $(seq 0 $((NSHARD-1))); do
|
||||
python data_pipeline/openvid_download_hd.py --shard "$s" --num-shards "$NSHARD" \
|
||||
--videos-dir "$WORK"/openvid_1m/videos --zip-dir "$WORK"/openvid_1m/_zips/"$s" \
|
||||
--only-list "$WORK"/openvid/OpenVidHD_filtered.txt &
|
||||
pids+=($!)
|
||||
done
|
||||
last=$(ZB); stall=0
|
||||
while :; do
|
||||
sleep 120
|
||||
[ "$(DC)" -ge 98 ] && break
|
||||
alive=0; for p in "${pids[@]}"; do kill -0 "$p" 2>/dev/null && alive=1; done
|
||||
[ "$alive" -eq 0 ] && { echo "[$(date +%H:%M)] round done naturally parts=$(DC)/98"; break; }
|
||||
now=$(ZB)
|
||||
if [ "$now" -le "$last" ]; then stall=$((stall+1)); else stall=0; fi
|
||||
last=$now
|
||||
echo "[$(date +%H:%M)] parts=$(DC)/98 zip=$((now/1000000000))GB stall=$stall"
|
||||
if [ "$stall" -ge 2 ]; then
|
||||
echo "[$(date +%H:%M)] STALL -> kill + cooldown"
|
||||
kill "${pids[@]}" 2>/dev/null; sleep 5; kill -9 "${pids[@]}" 2>/dev/null
|
||||
pkill -9 -f openvid_download_hd.py 2>/dev/null
|
||||
break
|
||||
fi
|
||||
done
|
||||
wait 2>/dev/null
|
||||
[ "$(DC)" -ge 98 ] && break
|
||||
echo "[$(date +%H:%M)] cooldown 180s (parts $(DC)/98)"; sleep 180
|
||||
done
|
||||
echo "ALL_PARTS_DONE $(DC)/98 videos=$(ls "$WORK"/openvid_1m/videos/*.mp4 2>/dev/null | wc -l)"
|
||||
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# DP concurrency sweep on ONE GPU: K workers sharing GPU 3, aggregate throughput.
|
||||
set +e
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
source "$WORK/FastVideo/.venv/bin/activate"
|
||||
export CUDA_HOME=/usr/local/cuda
|
||||
export PATH=$CUDA_HOME/bin:$PATH
|
||||
export TORCH_HOME=$WORK/.torch
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
V=$WORK/dp_test/videos; T=$WORK/dp_test/tracks; L=$WORK/dp_test/videos.txt
|
||||
mkdir -p "$V" "$T"
|
||||
if [ ! -f "$V/clip_24.mp4" ]; then
|
||||
for i in $(seq -w 1 24); do cp -f "$WORK/wan_overfit/data/videos/vid_000000.mp4" "$V/clip_$i.mp4"; done
|
||||
fi
|
||||
ls "$V"/*.mp4 > "$L"
|
||||
echo "test set: $(wc -l < "$L") videos on GPU 3 (121f 720p, grid 50)"
|
||||
for K in 1 2 3 4; do
|
||||
rm -f "$T"/*.npz
|
||||
echo "===================== K=$K procs/GPU ====================="
|
||||
pids=()
|
||||
for s in $(seq 0 $((K-1))); do
|
||||
CUDA_VISIBLE_DEVICES=3 python data_pipeline/extract_tracks_mp.py \
|
||||
--video-list "$L" --out-dir "$T" --gpus-per-node 1 \
|
||||
--shard "$s" --num-shards "$K" --fps 24 --num-frames 121 --grid-size 50 &
|
||||
pids+=($!)
|
||||
done
|
||||
wait "${pids[@]}"
|
||||
done
|
||||
echo "SWEEP_DONE"
|
||||
@@ -0,0 +1,224 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Inspect a track-conditioning dataset: are the clips diverse, and do the
|
||||
CoTracker tracks actually capture the motion?
|
||||
|
||||
Two concerns this surfaces:
|
||||
1. Near-duplicate clips (the frame-0 cluster can over-select the same scene).
|
||||
2. Weak control signal: CoTracker uses a frame-0 grid, so anything that ENTERS
|
||||
the frame later (or that the tracker loses) is never queried -> its motion is
|
||||
invisible to the tracks. A clip can look dynamic yet have near-static tracks.
|
||||
|
||||
``render`` (CPU; run on a node) writes, per clip:
|
||||
- <id>__overlay.mp4 the video with the CoTracker tracks drawn (dots + tails),
|
||||
- <id>__thumb.png first frame,
|
||||
and a ``stats.json`` with per-clip motion-coverage + a near-duplicate grouping.
|
||||
|
||||
``serve`` (no GPU) is a gradio gallery: clips sorted by motion (least first, so
|
||||
dead/duplicate clips float to the top), each with its overlay + stats.
|
||||
|
||||
# 1) render artifacts (node)
|
||||
srun --jobid=<job> --overlap --ntasks=1 .venv/bin/python data_pipeline/droid_dataset_dashboard.py render \
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200 --out <dash_dir>
|
||||
# 2) serve (login node ok)
|
||||
.venv/bin/python data_pipeline/droid_dataset_dashboard.py serve --dash-dir <dash_dir> --share
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
MOVE_THRESH_PX = 8.0 # a point "moves" if it travels > this from frame 0 (input px)
|
||||
VIS_THRESH = 0.5
|
||||
|
||||
|
||||
def _load_frames(path: str) -> np.ndarray:
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001
|
||||
import imageio.v2 as imageio
|
||||
rd = imageio.get_reader(path, format="ffmpeg")
|
||||
fr = np.stack([np.asarray(f) for f in rd], axis=0)
|
||||
rd.close()
|
||||
return fr
|
||||
|
||||
|
||||
def _clip_stats(tracks: np.ndarray, vis: np.ndarray) -> dict[str, Any]:
|
||||
"""Motion coverage from tracks [T,N,2] (px) + vis [T,N]."""
|
||||
T, N, _ = tracks.shape
|
||||
v = vis > VIS_THRESH
|
||||
# displacement of each point from its frame-0 position, only where visible
|
||||
disp = np.sqrt(((tracks - tracks[0:1])**2).sum(-1)) # [T,N]
|
||||
max_disp = np.where(v, disp, 0.0).max(0) # [N] per-point peak travel
|
||||
moving = max_disp > MOVE_THRESH_PX
|
||||
return {
|
||||
"n_points": int(N),
|
||||
"frac_visible": float(v.mean()),
|
||||
"frac_moving": float(moving.mean()), # fraction of grid points that ever move
|
||||
"n_moving": int(moving.sum()),
|
||||
"mean_motion_px": float(max_disp.mean()),
|
||||
"p95_motion_px": float(np.percentile(max_disp, 95)),
|
||||
"max_motion_px": float(max_disp.max()),
|
||||
}
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
from fastvideo.train.callbacks.track_validation import (_draw_overlay, _grid_colors, _subsample)
|
||||
data_dir = Path(args.data_dir)
|
||||
manifest = json.loads((data_dir / "videos2caption.json").read_text())
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
import imageio.v2 as imageio
|
||||
from PIL import Image
|
||||
|
||||
embs: list[np.ndarray] = []
|
||||
entries: list[dict[str, Any]] = []
|
||||
for rec in manifest:
|
||||
cid = Path(rec["path"]).stem
|
||||
vpath = str(data_dir / "videos" / rec["path"])
|
||||
ppath = rec.get("points_path")
|
||||
frames = _load_frames(vpath)
|
||||
d = np.load(ppath)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]] # [T,N,2] px
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]]
|
||||
H, W = int(frames.shape[1]), int(frames.shape[2])
|
||||
|
||||
stats = _clip_stats(tracks, vis)
|
||||
|
||||
# overlay (subsampled grid for clarity)
|
||||
grid = int(round(tracks.shape[1]**0.5))
|
||||
tn = tracks / np.array([W, H], np.float32) # normalize for _subsample
|
||||
tr_s, vs_s = _subsample(tn, vis, grid, args.stride)
|
||||
colors = _grid_colors(grid, args.stride)[:tr_s.shape[1]]
|
||||
trpx = tr_s.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
ov = _draw_overlay(frames, trpx, vs_s, colors, args.tail, 2, VIS_THRESH)
|
||||
imageio.mimsave(str(out / f"{cid}__overlay.mp4"), ov, fps=args.fps, macro_block_size=1)
|
||||
Image.fromarray(frames[0]).save(str(out / f"{cid}__thumb.png"))
|
||||
|
||||
# frame-0 embedding for near-duplicate grouping (32x32 gray, L2-norm)
|
||||
g = np.asarray(Image.fromarray(frames[0]).convert("L").resize((32, 32)), np.float32).reshape(-1)
|
||||
embs.append(g / (np.linalg.norm(g) + 1e-8))
|
||||
|
||||
entries.append({"id": cid, "episode": rec.get("source_episode"), "caption": rec["cap"][0], **stats})
|
||||
print(
|
||||
f"[dash] {cid} move%={stats['frac_moving']*100:4.1f} mean={stats['mean_motion_px']:5.1f}px "
|
||||
f"max={stats['max_motion_px']:5.1f}px",
|
||||
flush=True)
|
||||
|
||||
# near-duplicate grouping: union-find on cosine-sim >= threshold
|
||||
E = np.stack(embs, 0)
|
||||
sim = E @ E.T
|
||||
n = len(entries)
|
||||
parent = list(range(n))
|
||||
|
||||
def find(i: int) -> int:
|
||||
while parent[i] != i:
|
||||
parent[i] = parent[parent[i]]
|
||||
i = parent[i]
|
||||
return i
|
||||
|
||||
for i in range(n):
|
||||
for j in range(i + 1, n):
|
||||
if sim[i, j] >= args.dup_thresh:
|
||||
parent[find(i)] = find(j)
|
||||
groups: dict[int, list[int]] = {}
|
||||
for i in range(n):
|
||||
groups.setdefault(find(i), []).append(i)
|
||||
group_id = {}
|
||||
for gi, (_, members) in enumerate(sorted(groups.items(), key=lambda kv: -len(kv[1]))):
|
||||
for m in members:
|
||||
group_id[m] = gi
|
||||
for i, e in enumerate(entries):
|
||||
e["dup_group"] = int(group_id[i])
|
||||
|
||||
n_groups = len(groups)
|
||||
n_dead = sum(1 for e in entries if e["frac_moving"] < 0.02)
|
||||
summary = {
|
||||
"n_clips": n,
|
||||
"n_dup_groups": n_groups,
|
||||
"largest_dup_group": max(len(v) for v in groups.values()),
|
||||
"n_low_motion_clips(<2%)": n_dead,
|
||||
"median_frac_moving": float(np.median([e["frac_moving"] for e in entries])),
|
||||
"median_mean_motion_px": float(np.median([e["mean_motion_px"] for e in entries])),
|
||||
"dup_thresh": args.dup_thresh,
|
||||
"move_thresh_px": MOVE_THRESH_PX,
|
||||
}
|
||||
(out / "stats.json").write_text(json.dumps({"summary": summary, "clips": entries}, indent=2))
|
||||
print(f"[dash] SUMMARY: {json.dumps(summary, indent=2)}", flush=True)
|
||||
print(f"[dash] wrote dashboard artifacts -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
dash = Path(args.dash_dir)
|
||||
data = json.loads((dash / "stats.json").read_text())
|
||||
summary, clips = data["summary"], data["clips"]
|
||||
# sort: least motion first (dead/duplicate clips surface), then by dup group
|
||||
order = sorted(range(len(clips)), key=lambda i: (clips[i]["frac_moving"], clips[i]["dup_group"]))
|
||||
|
||||
def label(i: int) -> str:
|
||||
c = clips[i]
|
||||
return (f"#{i} ep{c['episode']} | move {c['frac_moving']*100:.0f}% | "
|
||||
f"mean {c['mean_motion_px']:.0f}px | dupgrp {c['dup_group']}")
|
||||
|
||||
gallery_items = [(str(dash / f"{clips[i]['id']}__thumb.png"), label(i)) for i in order]
|
||||
|
||||
def show(evt): # gr.SelectData; unannotated so gradio's get_type_hints doesn't eval it
|
||||
i = order[evt.index]
|
||||
c = clips[i]
|
||||
return (str(dash / f"{c['id']}__overlay.mp4"), json.dumps(c, indent=2))
|
||||
|
||||
head = (f"### DROID dataset inspection — {summary['n_clips']} clips\n"
|
||||
f"- near-duplicate groups: **{summary['n_dup_groups']}** "
|
||||
f"(largest group **{summary['largest_dup_group']}** clips)\n"
|
||||
f"- low-motion clips (<2% points move): **{summary['n_low_motion_clips(<2%)']}**\n"
|
||||
f"- median fraction of points moving: **{summary['median_frac_moving']*100:.1f}%**, "
|
||||
f"median mean-motion **{summary['median_mean_motion_px']:.1f}px**\n\n"
|
||||
f"Sorted least-motion first. Overlay = CoTracker tracks (dots + tails); "
|
||||
f"watch for moving objects with NO dots on them (off-frame entries / lost tracks).")
|
||||
|
||||
with gr.Blocks(title="DROID dataset dashboard") as demo:
|
||||
gr.Markdown(head)
|
||||
with gr.Row():
|
||||
gal = gr.Gallery(value=gallery_items, columns=6, height=560, label="clips (least motion first)")
|
||||
with gr.Column():
|
||||
vid = gr.Video(label="track overlay")
|
||||
meta = gr.Code(label="clip stats", language="json")
|
||||
gal.select(show, None, [vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True, help="dataset root (videos/ + videos2caption.json + tracks/)")
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--stride", type=int, default=3, help="subsample the NxN grid for overlay clarity")
|
||||
r.add_argument("--tail", type=int, default=12)
|
||||
r.add_argument("--fps", type=int, default=12)
|
||||
r.add_argument("--dup-thresh", type=float, default=0.985, help="cosine-sim on frame-0 to call clips duplicates")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--dash-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7872)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
args = p.parse_args()
|
||||
args.func(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,416 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Evaluate super-sparse track adherence via CoTracker EPE.
|
||||
|
||||
For each val clip, we build the exact adversarial case the user described:
|
||||
* Pick ONE active foreground trace from the GT parquet — the highest-motion
|
||||
foreground track (object_ids >= 0, visible in frame 0).
|
||||
* Pick N static background anchors — random background tracks (object_ids == -1)
|
||||
with position frozen at frame-0 (visibility=1 across all frames).
|
||||
* Generate video conditioned on (active + anchors) — this is what training
|
||||
was supposed to teach the model to handle.
|
||||
* Also generate an ablation with track_points=None (unconditional motion).
|
||||
* Run CoTracker on both generated videos, initialised at frame 0 with the
|
||||
conditioning points. Compare CoTracker-extracted tracks against the GT
|
||||
trace to compute End-Point-Error in pixel space.
|
||||
|
||||
Metrics logged per clip and aggregated to wandb:
|
||||
epe_active_full EPE(GT_active, CoTracker(gen_full))
|
||||
epe_active_notrack EPE(GT_active, CoTracker(gen_no_track))
|
||||
epe_bg EPE(static_anchor, CoTracker(gen_full)) — should be small
|
||||
DELTA = epe_active_notrack - epe_active_full — positive means training helped.
|
||||
|
||||
Usage:
|
||||
srun --overlap --jobid=529 --ntasks=1 -w hpc-rack-2-13 --chdir=$PWD bash -lc '
|
||||
source .venv/bin/activate
|
||||
export HOME=/mnt/lustre/vlm-s4duan HF_HOME=/mnt/lustre/vlm-s4duan/.hf \
|
||||
TORCH_HOME=/mnt/lustre/vlm-s4duan/.torch MPLCONFIGDIR=/mnt/lustre/vlm-s4duan/.mpl \
|
||||
TRITON_CACHE_DIR=/tmp/triton_eval TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0 \
|
||||
PYTHONPATH=$PWD TRACKWAN_TRACK_BIAS=1 CUDA_VISIBLE_DEVICES=2 \
|
||||
WANDB_API_KEY=<key> WANDB_MODE=online
|
||||
python data_pipeline/eval_sparse_track_epe.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800 \
|
||||
--yaml examples/train/scenario/worldmodel/finetune_wantrack_openvid_sparse_1p3b_merged_bias.yaml \
|
||||
--data-path /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset \
|
||||
--num-clips 5 --steps 30 --w-motion 1.5 \
|
||||
--wandb-run-name eval_mb4800_sparse
|
||||
'
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
# make trackwan_infer importable
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__)))
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# CoTracker
|
||||
# ==========================================================================
|
||||
def load_cotracker(device: str) -> torch.nn.Module:
|
||||
"""Load CoTracker v3 offline. Uses the warm hub cache under $TORCH_HOME/hub."""
|
||||
hub_dir = Path(torch.hub.get_dir())
|
||||
local = hub_dir / "facebookresearch_co-tracker_main"
|
||||
if local.exists():
|
||||
model = torch.hub.load(str(local), "cotracker3_offline", source="local", trust_repo=True)
|
||||
else:
|
||||
model = torch.hub.load("facebookresearch/co-tracker", "cotracker3_offline", trust_repo=True)
|
||||
return model.to(device).eval()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def cotracker_query(cotracker: torch.nn.Module, video: np.ndarray, queries_xy: np.ndarray,
|
||||
device: str) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Track `queries_xy` (N, 2 pixel coords) forward through `video` [T,H,W,3] uint8.
|
||||
|
||||
Returns tracks [T,N,2] (pixel coords) and visibility [T,N] (bool).
|
||||
"""
|
||||
T, H, W, C = video.shape
|
||||
v = torch.from_numpy(video).permute(0, 3, 1, 2).float()[None].to(device) # [1,T,3,H,W]
|
||||
N = queries_xy.shape[0]
|
||||
q = np.zeros((1, N, 3), dtype=np.float32)
|
||||
q[0, :, 0] = 0 # start-frame index
|
||||
q[0, :, 1] = queries_xy[:, 0] # x
|
||||
q[0, :, 2] = queries_xy[:, 1] # y
|
||||
tracks, visibility = cotracker(v, queries=torch.from_numpy(q).to(device))
|
||||
return tracks[0].cpu().numpy(), visibility[0].cpu().numpy()
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Sparse-conditioning sampler
|
||||
# ==========================================================================
|
||||
def pick_sparse_conditioning(row: dict, *, num_active: int, num_bg: int,
|
||||
seed: int) -> dict:
|
||||
"""From a preprocessed parquet row, pick the sparse case: `num_active` moving fg
|
||||
traces (highest-displacement fg tracks visible in frame 0) + `num_bg` static bg
|
||||
anchors (random bg tracks, position frozen at frame-0).
|
||||
|
||||
Everything is normalised to [0,1] to match training. Pixel-space is only used at
|
||||
the metric-computation step.
|
||||
"""
|
||||
tp = row["track_points"] # [T,N,2] normalised
|
||||
tv = row["track_visibility"] # [T,N]
|
||||
oid = row["object_ids"] # [N] int, -1 for bg
|
||||
T, N, _ = tp.shape
|
||||
rng = np.random.RandomState(seed)
|
||||
|
||||
fg = np.where((oid >= 0) & (tv[0] > 0.5))[0]
|
||||
bg = np.where((oid == -1) & (tv[0] > 0.5))[0]
|
||||
if fg.size == 0:
|
||||
raise RuntimeError("no visible foreground tracks in frame 0")
|
||||
# displacement of each fg over the full clip (only over visible frames)
|
||||
disp = np.linalg.norm(tp[-1, fg] - tp[0, fg], axis=-1)
|
||||
active_idx = fg[np.argsort(-disp)[:num_active]] # top-`num_active` moving fg
|
||||
n_bg = min(num_bg, bg.size)
|
||||
bg_idx = rng.choice(bg, size=n_bg, replace=False) if n_bg > 0 else np.zeros((0, ), dtype=int)
|
||||
|
||||
active_gt = tp[:, active_idx] # [T, num_active, 2] — GT motion
|
||||
active_vis = tv[:, active_idx]
|
||||
bg_static_pos = np.broadcast_to(tp[0:1, bg_idx], (T, n_bg, 2)).copy() # position frozen at f0
|
||||
bg_vis = np.ones((T, n_bg), dtype=np.float32)
|
||||
|
||||
cond_tracks = np.concatenate([active_gt, bg_static_pos], axis=1) # [T, num_active+n_bg, 2]
|
||||
cond_vis = np.concatenate([active_vis, bg_vis], axis=1)
|
||||
return {
|
||||
"active_idx": active_idx,
|
||||
"active_gt": active_gt, # [T,num_active,2] normalised
|
||||
"active_vis": active_vis,
|
||||
"bg_static_pos": bg_static_pos, # [T,n_bg,2] normalised (const in t)
|
||||
"cond_tracks": cond_tracks, # [T, num_active+n_bg, 2] normalised
|
||||
"cond_vis": cond_vis,
|
||||
"n_active": len(active_idx),
|
||||
"n_bg": n_bg,
|
||||
}
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Metrics
|
||||
# ==========================================================================
|
||||
def epe_norm_to_px(pred_norm: np.ndarray, gt_norm: np.ndarray, vis: np.ndarray,
|
||||
res_h: int, res_w: int) -> float:
|
||||
"""L2 distance between predicted/GT normalised tracks, evaluated in pixel space,
|
||||
averaged only over visible frames × visible points."""
|
||||
diff = (pred_norm - gt_norm) * np.array([res_w, res_h])[None, None, :]
|
||||
l2 = np.linalg.norm(diff, axis=-1) # [T, N]
|
||||
mask = vis > 0.5
|
||||
if not mask.any():
|
||||
return float("nan")
|
||||
return float(l2[mask].mean())
|
||||
|
||||
|
||||
def epe_px(pred_px: np.ndarray, gt_px: np.ndarray, vis: np.ndarray) -> float:
|
||||
"""L2 in pixel space — both inputs already in pixels."""
|
||||
l2 = np.linalg.norm(pred_px - gt_px, axis=-1)
|
||||
mask = vis > 0.5
|
||||
if not mask.any():
|
||||
return float("nan")
|
||||
return float(l2[mask].mean())
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Trace overlay for wandb visualisation
|
||||
# ==========================================================================
|
||||
def overlay_tracks(video: np.ndarray, tracks_norm: np.ndarray, vis: np.ndarray,
|
||||
colors: list[tuple[int, int, int]], radius: int = 3,
|
||||
tail: int = 12) -> np.ndarray:
|
||||
"""Draw normalised tracks on a video [T,H,W,3] uint8. Small tail + current dot."""
|
||||
from PIL import Image, ImageDraw
|
||||
T, H, W, C = video.shape
|
||||
N = tracks_norm.shape[1]
|
||||
out = []
|
||||
for t in range(T):
|
||||
img = Image.fromarray(np.ascontiguousarray(video[t]))
|
||||
dr = ImageDraw.Draw(img)
|
||||
t0 = max(0, t - tail)
|
||||
for i in range(N):
|
||||
col = colors[i % len(colors)]
|
||||
# tail
|
||||
if t - t0 >= 1:
|
||||
pts = [(float(tracks_norm[k, i, 0] * W), float(tracks_norm[k, i, 1] * H))
|
||||
for k in range(t0, t + 1)]
|
||||
dr.line(pts, fill=col, width=1)
|
||||
if vis[t, i] > 0.5:
|
||||
x = float(tracks_norm[t, i, 0] * W)
|
||||
y = float(tracks_norm[t, i, 1] * H)
|
||||
dr.ellipse([x - radius, y - radius, x + radius, y + radius], fill=col)
|
||||
out.append(np.asarray(img))
|
||||
return np.stack(out)
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Parquet loader (small subset — avoid loading the whole 259k row set)
|
||||
# ==========================================================================
|
||||
def load_first_n_from_parquet(data_path: str, n: int, text_len: int) -> list:
|
||||
import pyarrow.parquet as pq
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
files = sorted(glob.glob(os.path.join(data_path, "**", "*.parquet"), recursive=True))
|
||||
if not files:
|
||||
raise FileNotFoundError(data_path)
|
||||
rows: list = []
|
||||
for f in files:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
if len(rows) >= n:
|
||||
break
|
||||
sel = rows[:n]
|
||||
batch = collate_rows_from_parquet_schema(sel, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in sel]
|
||||
out = []
|
||||
for i in range(len(sel)):
|
||||
out.append({
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i].numpy(), # [T,N,2] normalised
|
||||
"track_visibility": batch["track_visibility"][i].numpy(), # [T,N]
|
||||
"object_ids": batch["object_ids"][i].numpy(), # [N]
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
# ==========================================================================
|
||||
# Main
|
||||
# ==========================================================================
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True)
|
||||
p.add_argument("--yaml", required=True)
|
||||
p.add_argument("--data-path", required=True)
|
||||
p.add_argument("--num-clips", type=int, default=5)
|
||||
p.add_argument("--num-active", type=int, default=1)
|
||||
p.add_argument("--num-bg", type=int, default=20)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--w-motion", type=float, default=1.5)
|
||||
p.add_argument("--seed", type=int, default=1234)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--wandb-project", default="wantrack-bidir")
|
||||
p.add_argument("--wandb-run-name", required=True)
|
||||
p.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/eval_sparse_epe")
|
||||
args = p.parse_args()
|
||||
|
||||
Path(args.out_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
import wandb
|
||||
import imageio.v2 as imageio
|
||||
import trackwan_infer as twi
|
||||
|
||||
print(f"[eval] loading model {args.model_dir}", flush=True)
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(tc.data.text_padding_length) if hasattr(tc.data, "text_padding_length") else 256
|
||||
|
||||
print(f"[eval] loading {args.num_clips} parquet clips", flush=True)
|
||||
samples = load_first_n_from_parquet(args.data_path, args.num_clips, text_len)
|
||||
|
||||
print("[eval] loading CoTracker", flush=True)
|
||||
device = model.device
|
||||
cotracker = load_cotracker(device.type + ":" + str(device.index) if device.index is not None else device.type)
|
||||
|
||||
res_h = int(tc.data.num_height)
|
||||
res_w = int(tc.data.num_width)
|
||||
palette = [(255, 60, 60), (60, 255, 60), (60, 90, 255), (255, 200, 60), (200, 60, 255), (60, 255, 220)]
|
||||
|
||||
run = wandb.init(project=args.wandb_project, name=args.wandb_run_name,
|
||||
config={"model_dir": args.model_dir, "num_clips": args.num_clips,
|
||||
"num_active": args.num_active, "num_bg": args.num_bg,
|
||||
"steps": args.steps, "w_motion": args.w_motion, "res": [res_h, res_w]})
|
||||
|
||||
per_clip: list = []
|
||||
|
||||
for idx, s in enumerate(samples):
|
||||
try:
|
||||
sc = pick_sparse_conditioning(s, num_active=args.num_active, num_bg=args.num_bg,
|
||||
seed=args.seed + idx)
|
||||
except RuntimeError as exc:
|
||||
print(f"[eval] clip {idx}: {exc}; skipping", flush=True)
|
||||
continue
|
||||
|
||||
# ── Generate two videos: with sparse tracks (motion CFG) + without any tracks ──
|
||||
tp_t = torch.from_numpy(sc["cond_tracks"])[None].float() # [1,T,N,2]
|
||||
tv_t = torch.from_numpy(sc["cond_vis"])[None].float() # [1,T,N]
|
||||
|
||||
seed = args.seed + idx
|
||||
common = dict(first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=args.steps, seed=seed)
|
||||
|
||||
print(f"[eval] clip {idx}: generating full ({sc['n_active']} active + {sc['n_bg']} anchors)", flush=True)
|
||||
t0 = time.time()
|
||||
lat_full = twi.generate(model, track_points=tp_t, track_visibility=tv_t,
|
||||
guidance_scale=args.w_motion, **common)
|
||||
gen_full = twi.decode_to_pixels(model, lat_full) # [T,H,W,3] uint8
|
||||
t_full = time.time() - t0
|
||||
print(f"[eval] clip {idx}: full took {t_full:.1f}s", flush=True)
|
||||
|
||||
print(f"[eval] clip {idx}: generating no-track", flush=True)
|
||||
t0 = time.time()
|
||||
lat_no = twi.generate(model, track_points=None, track_visibility=None,
|
||||
guidance_scale=1.0, **common)
|
||||
gen_no = twi.decode_to_pixels(model, lat_no)
|
||||
t_no = time.time() - t0
|
||||
print(f"[eval] clip {idx}: no-track took {t_no:.1f}s", flush=True)
|
||||
|
||||
T, H, W, _ = gen_full.shape
|
||||
|
||||
# ── CoTracker on both generations, queried at frame-0 with same points ──
|
||||
query_norm = sc["cond_tracks"][0] # [N_total, 2] in [0,1]
|
||||
queries_px = query_norm * np.array([W, H])[None] # pixel space
|
||||
print(f"[eval] clip {idx}: CoTracker on full", flush=True)
|
||||
tk_full_px, tk_full_vis = cotracker_query(cotracker, gen_full, queries_px, device=device.type)
|
||||
print(f"[eval] clip {idx}: CoTracker on no-track", flush=True)
|
||||
tk_no_px, tk_no_vis = cotracker_query(cotracker, gen_no, queries_px, device=device.type)
|
||||
|
||||
# GT in pixel space, at the SAME resolution as the generation
|
||||
gt_full_px = sc["cond_tracks"] * np.array([W, H])[None, None, :] # [T,N_total,2]
|
||||
active_slice = slice(0, sc["n_active"])
|
||||
bg_slice = slice(sc["n_active"], sc["n_active"] + sc["n_bg"])
|
||||
|
||||
# Active EPE: CoTracker(gen) vs the GT trace we specified.
|
||||
# Filter by GT visibility (active_vis) AND CoTracker's own visibility.
|
||||
active_vis_gt = sc["active_vis"] # [T,n_active]
|
||||
both_vis_full = (active_vis_gt > 0.5) & (tk_full_vis[:, active_slice] > 0.5)
|
||||
both_vis_no = (active_vis_gt > 0.5) & (tk_no_vis[:, active_slice] > 0.5)
|
||||
|
||||
epe_active_full = epe_px(tk_full_px[:, active_slice],
|
||||
gt_full_px[:, active_slice],
|
||||
both_vis_full.astype(np.float32))
|
||||
epe_active_no = epe_px(tk_no_px[:, active_slice],
|
||||
gt_full_px[:, active_slice],
|
||||
both_vis_no.astype(np.float32))
|
||||
|
||||
# BG EPE (only for full): CoTracker(gen_full) vs static anchor position
|
||||
bg_static_px = sc["bg_static_pos"] * np.array([W, H])[None, None, :]
|
||||
bg_vis_mask = (tk_full_vis[:, bg_slice] > 0.5).astype(np.float32)
|
||||
epe_bg = epe_px(tk_full_px[:, bg_slice], bg_static_px, bg_vis_mask) if sc["n_bg"] > 0 else float("nan")
|
||||
|
||||
delta = epe_active_no - epe_active_full # positive => tracks helped
|
||||
|
||||
# ── Save + log videos ──
|
||||
overlay_full = overlay_tracks(gen_full, sc["cond_tracks"], sc["cond_vis"], palette)
|
||||
overlay_no = overlay_tracks(gen_no, sc["cond_tracks"], sc["cond_vis"], palette)
|
||||
|
||||
# Also show the CoTracker-extracted tracks over the full generation — this is
|
||||
# what the model *actually* produced motion-wise.
|
||||
cotr_tracks_norm = tk_full_px / np.array([W, H])[None, None, :]
|
||||
overlay_cotr = overlay_tracks(gen_full, cotr_tracks_norm, tk_full_vis.astype(np.float32), palette)
|
||||
|
||||
clip_dir = Path(args.out_dir) / f"clip_{idx:03d}"
|
||||
clip_dir.mkdir(parents=True, exist_ok=True)
|
||||
p_full = clip_dir / "gen_full_overlay.mp4"
|
||||
p_no = clip_dir / "gen_notrack_overlay.mp4"
|
||||
p_cotr = clip_dir / "gen_full_cotracker.mp4"
|
||||
imageio.mimsave(p_full, overlay_full, fps=args.fps, macro_block_size=1)
|
||||
imageio.mimsave(p_no, overlay_no, fps=args.fps, macro_block_size=1)
|
||||
imageio.mimsave(p_cotr, overlay_cotr, fps=args.fps, macro_block_size=1)
|
||||
|
||||
print(f"[eval] clip {idx}: EPE_active_full={epe_active_full:.2f} "
|
||||
f"EPE_active_notrack={epe_active_no:.2f} DELTA={delta:.2f} "
|
||||
f"EPE_bg={epe_bg:.2f}", flush=True)
|
||||
|
||||
# per-clip logs
|
||||
entry = {
|
||||
"clip": idx,
|
||||
"epe_active_full_px": epe_active_full,
|
||||
"epe_active_notrack_px": epe_active_no,
|
||||
"delta_px": delta,
|
||||
"epe_bg_px": epe_bg,
|
||||
"n_active": sc["n_active"],
|
||||
"n_bg": sc["n_bg"],
|
||||
"sec_gen_full": t_full,
|
||||
"sec_gen_notrack": t_no,
|
||||
"caption": s["caption"][:120],
|
||||
}
|
||||
per_clip.append(entry)
|
||||
wandb.log({
|
||||
f"clip_{idx:03d}/gen_full": wandb.Video(str(p_full), fps=args.fps, format="mp4"),
|
||||
f"clip_{idx:03d}/gen_notrack": wandb.Video(str(p_no), fps=args.fps, format="mp4"),
|
||||
f"clip_{idx:03d}/gen_full_cotracker": wandb.Video(str(p_cotr), fps=args.fps, format="mp4"),
|
||||
"clip": idx,
|
||||
"epe_active_full_px": epe_active_full,
|
||||
"epe_active_notrack_px": epe_active_no,
|
||||
"delta_px": delta,
|
||||
"epe_bg_px": epe_bg,
|
||||
})
|
||||
|
||||
# ── Aggregate + summary ──
|
||||
if per_clip:
|
||||
arr = lambda k: np.array([r[k] for r in per_clip if not np.isnan(r[k])])
|
||||
summary = {
|
||||
"n_clips": len(per_clip),
|
||||
"mean_epe_active_full_px": float(np.mean(arr("epe_active_full_px"))),
|
||||
"median_epe_active_full_px": float(np.median(arr("epe_active_full_px"))),
|
||||
"mean_epe_active_notrack_px": float(np.mean(arr("epe_active_notrack_px"))),
|
||||
"median_epe_active_notrack_px": float(np.median(arr("epe_active_notrack_px"))),
|
||||
"mean_delta_px": float(np.mean(arr("delta_px"))),
|
||||
"median_delta_px": float(np.median(arr("delta_px"))),
|
||||
"mean_epe_bg_px": float(np.mean(arr("epe_bg_px"))),
|
||||
}
|
||||
(Path(args.out_dir) / f"{args.wandb_run_name}_summary.json").write_text(json.dumps({
|
||||
"summary": summary, "per_clip": per_clip}, indent=2))
|
||||
wandb.log({f"summary/{k}": v for k, v in summary.items()})
|
||||
wandb.summary.update(summary)
|
||||
print("=" * 70)
|
||||
print(f"[eval] summary ({len(per_clip)} clips):")
|
||||
for k, v in summary.items():
|
||||
print(f" {k}: {v:.3f}" if isinstance(v, float) else f" {k}: {v}")
|
||||
print("=" * 70)
|
||||
|
||||
wandb.finish()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,171 @@
|
||||
# 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.
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
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).")
|
||||
return p.parse_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).
|
||||
|
||||
Recent torchvision dropped ``torchvision.io.read_video``; use decord (fast,
|
||||
reliable), 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() # (T, H, W, C) uint8
|
||||
except Exception: # noqa: BLE001 - fall back to ffmpeg
|
||||
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
|
||||
) -> 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
|
||||
|
||||
pred_tracks, pred_vis = model(track_video.to(device), grid_size=grid_size)
|
||||
tracks = pred_tracks[0].float().cpu().numpy() # (T, N, 2) in tracking-res px
|
||||
vis = pred_vis[0].cpu().numpy() # (T, N)
|
||||
|
||||
if downscale != 1.0: # rescale coords back to original resolution
|
||||
tracks[..., 0] *= w / float(sw)
|
||||
tracks[..., 1] *= h / float(sh)
|
||||
return tracks.astype(np.float32), vis
|
||||
|
||||
|
||||
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(".json.tmp")
|
||||
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.limit is not None:
|
||||
videos = videos[:args.limit]
|
||||
if not videos:
|
||||
print(f"[track] no videos found in {videos_dir}", flush=True)
|
||||
return
|
||||
|
||||
model = load_cotracker(args.model, args.device)
|
||||
print(f"[track] loaded {args.model}; {len(videos)} videos, grid={args.grid_size}x{args.grid_size}", flush=True)
|
||||
|
||||
stem_to_points: dict[str, Path] = {}
|
||||
for k, vpath in enumerate(videos, 1):
|
||||
out_path = out_dir / f"{vpath.stem}.npz"
|
||||
stem_to_points[vpath.stem] = out_path
|
||||
if out_path.exists():
|
||||
continue
|
||||
video, h, w = read_video(vpath)
|
||||
tracks, vis = track_one(model, video, args.grid_size, args.downscale, args.device)
|
||||
# np.savez appends ".npz" unless the name already ends in it, so keep the
|
||||
# tmp name ".npz"-terminated to make the atomic replace below match.
|
||||
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],
|
||||
)
|
||||
tmp.replace(out_path)
|
||||
print(f"[track] [{k}/{len(videos)}] {vpath.name} -> {out_path.name} "
|
||||
f"tracks={tracks.shape} vis={vis.shape}", flush=True)
|
||||
|
||||
n = patch_manifest(args.data_dir / args.manifest, stem_to_points)
|
||||
print(f"[track] done; patched points_path into {n} manifest entries", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,145 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Data-parallel CoTracker v3 track extraction for large video sets (OpenVid-1M).
|
||||
|
||||
CoTracker batching does NOT help (s/video is flat, memory linear), so we scale by
|
||||
DATA PARALLELISM: launch many single-video workers, several per GPU, across nodes.
|
||||
Each worker processes video_list[shard::num_shards] on one pinned GPU and is idempotent.
|
||||
|
||||
Sharding + GPU pinning come from Slurm env by default:
|
||||
shard = SLURM_PROCID (0..num_shards-1)
|
||||
num_shards = SLURM_NTASKS
|
||||
gpu = SLURM_LOCALID % gpus_per_node
|
||||
so `srun --ntasks-per-node=(gpus*procs_per_gpu)` gives procs_per_gpu workers per GPU.
|
||||
|
||||
Resamples each clip to --fps / --num-frames and resizes to --height x --width (720p)
|
||||
before tracking, matching the training spec.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, os, time
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
HUB = "facebookresearch/co-tracker"
|
||||
|
||||
|
||||
def load_cotracker(device: str):
|
||||
local = Path(torch.hub.get_dir()) / (HUB.replace("/", "_") + "_main")
|
||||
if local.exists():
|
||||
m = torch.hub.load(str(local), "cotracker3_offline", source="local", trust_repo=True)
|
||||
else:
|
||||
m = torch.hub.load(HUB, "cotracker3_offline", trust_repo=True)
|
||||
return m.to(device).eval()
|
||||
|
||||
|
||||
def read_clip(path: str, fps: int, num_frames: int, H: int, W: int,
|
||||
min_height: int = 0, aspect: float | None = None, aspect_tol: float = 0.1):
|
||||
"""Decode (PyAV — decord has no aarch64 wheels, opencv needs GUI libs), resample
|
||||
to `fps`, take `num_frames`, resize to HxW. Sequential decode so it is codec-robust
|
||||
for OpenVid's varied encodings. Enforces resolution/aspect by probing stream metadata
|
||||
(fast, no decode). Returns (1,T,C,H,W) on success, or a reason string
|
||||
('lowres' | 'aspect' | 'short') on reject."""
|
||||
import av
|
||||
container = av.open(str(path))
|
||||
try:
|
||||
stream = container.streams.video[0]
|
||||
nh, nw = int(stream.height or 0), int(stream.width or 0)
|
||||
if min_height and nh and nh < min_height:
|
||||
return "lowres"
|
||||
if aspect is not None and nw and nh and abs((nw / nh) - aspect) > aspect_tol:
|
||||
return "aspect"
|
||||
native = float(stream.average_rate) if stream.average_rate else float(fps)
|
||||
step = native / float(fps)
|
||||
idxs = [int(round(i * step)) for i in range(num_frames)]
|
||||
picked, wi, cur = [], 0, 0
|
||||
for frame in container.decode(video=0):
|
||||
img = None
|
||||
while wi < len(idxs) and idxs[wi] == cur: # handles repeated idxs (fps up-sample)
|
||||
if img is None:
|
||||
img = frame.to_ndarray(format="rgb24") # (H,W,3) uint8
|
||||
picked.append(img); wi += 1
|
||||
cur += 1
|
||||
if wi >= len(idxs):
|
||||
break
|
||||
finally:
|
||||
container.close()
|
||||
if wi < len(idxs):
|
||||
return "short" # video too short at target fps
|
||||
arr = np.ascontiguousarray(np.stack(picked)) # (T,H,W,C) RGB uint8
|
||||
vid = torch.from_numpy(arr).permute(0, 3, 1, 2).unsqueeze(0).float() # 1,T,C,h,w
|
||||
vid = F.interpolate(vid[0], size=(H, W), mode="bilinear", align_corners=False).unsqueeze(0)
|
||||
return vid
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def track(model, video, grid, device):
|
||||
tr, vis = model(video.to(device), grid_size=grid)
|
||||
return tr[0].float().cpu().numpy(), vis[0].cpu().numpy()
|
||||
|
||||
|
||||
def save_clip(video, path, fps):
|
||||
"""Write the resampled 121f/720p clip so the training video IS the tracked frames."""
|
||||
import imageio
|
||||
arr = video[0].permute(0, 2, 3, 1).clamp(0, 255).byte().cpu().numpy() # (T,C,H,W)->(T,H,W,C)
|
||||
imageio.mimwrite(path, arr, fps=fps, codec="libx264", quality=7, macro_block_size=1)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--video-list", required=True, help="Text file: one video path per line.")
|
||||
ap.add_argument("--out-dir", required=True, help="Where <stem>.npz tracks are written.")
|
||||
ap.add_argument("--grid-size", type=int, default=50)
|
||||
ap.add_argument("--num-shards", type=int, default=int(os.environ.get("SLURM_NTASKS", "1")))
|
||||
ap.add_argument("--shard", type=int, default=int(os.environ.get("SLURM_PROCID", "0")))
|
||||
ap.add_argument("--gpus-per-node", type=int, default=4)
|
||||
ap.add_argument("--fps", type=int, default=24)
|
||||
ap.add_argument("--num-frames", type=int, default=121)
|
||||
ap.add_argument("--height", type=int, default=720)
|
||||
ap.add_argument("--width", type=int, default=1280)
|
||||
ap.add_argument("--min-height", type=int, default=720, help="skip clips whose native height < this (0=off)")
|
||||
ap.add_argument("--aspect", type=float, default=16.0 / 9.0, help="target native W/H aspect")
|
||||
ap.add_argument("--aspect-tol", type=float, default=0.12, help="allowed |native_aspect - target| (large=off)")
|
||||
ap.add_argument("--clips-dir", default=None, help="also save the resampled 121f/720p clip here (aligned to tracks)")
|
||||
a = ap.parse_args()
|
||||
|
||||
local = int(os.environ.get("SLURM_LOCALID", str(a.shard)))
|
||||
gpu = local % a.gpus_per_node
|
||||
torch.cuda.set_device(gpu)
|
||||
device = f"cuda:{gpu}"
|
||||
Path(a.out_dir).mkdir(parents=True, exist_ok=True)
|
||||
if a.clips_dir:
|
||||
Path(a.clips_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
vids = [l.strip() for l in open(a.video_list) if l.strip()]
|
||||
mine = vids[a.shard::a.num_shards]
|
||||
model = load_cotracker(device)
|
||||
print(f"[shard {a.shard}/{a.num_shards}] gpu={gpu} localid={local} videos={len(mine)}", flush=True)
|
||||
|
||||
t0 = time.time(); done = 0; err = 0; rej = {"short": 0, "lowres": 0, "aspect": 0}
|
||||
for vp in mine:
|
||||
out = Path(a.out_dir) / f"{Path(vp).stem}.npz"
|
||||
if out.exists():
|
||||
continue
|
||||
try:
|
||||
res = read_clip(vp, a.fps, a.num_frames, a.height, a.width,
|
||||
a.min_height, a.aspect, a.aspect_tol)
|
||||
if isinstance(res, str):
|
||||
rej[res] = rej.get(res, 0) + 1; continue
|
||||
tr, vis = track(model, res, a.grid_size, device)
|
||||
if a.clips_dir:
|
||||
save_clip(res, str(Path(a.clips_dir) / f"{Path(vp).stem}.mp4"), a.fps)
|
||||
tmp = out.with_suffix(".tmp.npz")
|
||||
np.savez(tmp, tracks=tr.astype(np.float32), visibility=vis, grid_size=a.grid_size,
|
||||
height=a.height, width=a.width, num_frames=tr.shape[0], fps=a.fps)
|
||||
tmp.replace(out); done += 1
|
||||
except Exception as e: # noqa: BLE001
|
||||
err += 1
|
||||
print(f"[shard {a.shard}] ERR {Path(vp).name}: {repr(e)[:110]}", flush=True)
|
||||
dt = time.time() - t0
|
||||
rate = done / dt if dt > 0 else 0.0
|
||||
print(f"[shard {a.shard}] DONE done={done} rej={rej} err={err} in {dt:.1f}s -> {rate:.3f} vid/s", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,353 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Find relatively-static WINDOWS inside (long) videos, before tracking.
|
||||
|
||||
The camera in egocentric video turns and the whole view sweeps off-frame; a filter that
|
||||
seeds points once and tracks the whole clip would die at the first turn and mis-score
|
||||
everything after. Instead we compute a PER-FRAME global-motion signal that RE-SEEDS every
|
||||
frame (so it recovers after a turn: motion spikes during the turn, drops when static), then
|
||||
slide a window and locate the calm stretches.
|
||||
|
||||
Per-frame motion m[t] = median Lucas-Kanade optical-flow magnitude between frames t-1 and t
|
||||
(fresh goodFeaturesToTrack each step), normalized by the frame diagonal. Camera pan/turn ->
|
||||
large; static -> small. A window is static iff m[t] stays low across ALL its frames (we score
|
||||
by a high percentile so one calm-but-not-perfect frame is fine but a turn is not). We also
|
||||
report per-window SURVIVAL: seed a grid at the window START, LK-track to the window END, and
|
||||
measure the fraction still tracked & in-frame -> "would CoTracker keep traces here?".
|
||||
|
||||
.venv/bin/python data_pipeline/find_static_windows.py scan --videos-dir <dir> --out <dir> --window-sec 5 --limit 12
|
||||
.venv/bin/python data_pipeline/find_static_windows.py serve --viz-dir <dir> --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
|
||||
def read_gray_lowres(path: str, max_frames: int = 900, target_w: int = 320) -> tuple[list, float, int]:
|
||||
"""Decode DIRECTLY at low res (fast + tiny memory) -> gray frames + fps + total_frames.
|
||||
Only used for the motion analysis; full-res frames are read on demand for extraction."""
|
||||
import cv2
|
||||
from decord import VideoReader, cpu
|
||||
cap = cv2.VideoCapture(path) # fast metadata read (no full decode)
|
||||
W0 = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) or 640
|
||||
H0 = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) or 360
|
||||
fps = float(cap.get(cv2.CAP_PROP_FPS)) or 30.0
|
||||
cap.release()
|
||||
tw = min(target_w, W0)
|
||||
th = int(round(H0 * tw / max(W0, 1)))
|
||||
try:
|
||||
vr = VideoReader(path, ctx=cpu(0), width=tw, height=th)
|
||||
except Exception: # noqa: BLE001 (older decord without width/height)
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
n_total = len(vr)
|
||||
n = min(n_total, max_frames)
|
||||
batch = vr.get_batch(list(range(n))).asnumpy()
|
||||
gray = [cv2.cvtColor(f, cv2.COLOR_RGB2GRAY) for f in batch]
|
||||
return gray, fps, n_total
|
||||
|
||||
|
||||
def read_rgb_frames(path: str, indices: Any) -> np.ndarray:
|
||||
"""Full-res RGB for a specific set of frame indices (decord random access)."""
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(indices)).asnumpy()
|
||||
|
||||
|
||||
def motion_series(gray: list) -> tuple[np.ndarray, np.ndarray, float]:
|
||||
"""Per-frame normalized global camera motion via phase correlation (FFT-based global
|
||||
shift between consecutive frames, ~1ms/frame). Recovers after turns: spikes during a
|
||||
turn, drops when static. Returns (m, shifts, diag): m[t] = |shift|/diag, shifts[t] =
|
||||
signed (dx,dy)/diag (used to accumulate net view drift over a window)."""
|
||||
import cv2
|
||||
H, W = gray[0].shape
|
||||
diag = float(np.hypot(H, W))
|
||||
hann = cv2.createHanningWindow((W, H), cv2.CV_32F)
|
||||
m = np.zeros(len(gray), np.float32)
|
||||
shifts = np.zeros((len(gray), 2), np.float32)
|
||||
prev = gray[0].astype(np.float32)
|
||||
for t in range(1, len(gray)):
|
||||
cur = gray[t].astype(np.float32)
|
||||
(dx, dy), _ = cv2.phaseCorrelate(prev, cur, hann)
|
||||
shifts[t] = (dx / diag, dy / diag)
|
||||
m[t] = float(np.hypot(dx, dy)) / diag
|
||||
prev = cur
|
||||
m[0] = m[1] if len(m) > 1 else 0.0
|
||||
return m, shifts, diag
|
||||
|
||||
|
||||
def window_drift(shifts: np.ndarray, s: int, L: int) -> float:
|
||||
"""Net view drift over a window: max cumulative excursion of the accumulated per-frame
|
||||
shifts, normalized by diagonal. Large => the camera slowly panned the view off-screen
|
||||
even if no single frame moved much. Free (from the phase-correlation shifts)."""
|
||||
c = np.cumsum(shifts[s:s + L], axis=0)
|
||||
return float(np.max(np.hypot(c[:, 0], c[:, 1]))) if len(c) else 0.0
|
||||
|
||||
|
||||
def window_survival(gray: list, s: int, L: int, grid: int = 24) -> float:
|
||||
"""Seed a grid at frame s, LK-track to s+L-1; fraction still tracked & in-frame."""
|
||||
import cv2
|
||||
H, W = gray[0].shape
|
||||
xs = np.linspace(W * 0.05, W * 0.95, grid)
|
||||
ys = np.linspace(H * 0.05, H * 0.95, grid)
|
||||
gx, gy = np.meshgrid(xs, ys)
|
||||
pts = np.stack([gx.ravel(), gy.ravel()], 1).astype(np.float32).reshape(-1, 1, 2)
|
||||
alive = np.ones(pts.shape[0], bool)
|
||||
cur = pts.copy()
|
||||
lk = dict(winSize=(21, 21), maxLevel=3)
|
||||
for t in range(s + 1, min(s + L, len(gray))):
|
||||
nxt, stt, _ = cv2.calcOpticalFlowPyrLK(gray[t - 1], gray[t], cur, None, **lk)
|
||||
stt = stt.ravel().astype(bool)
|
||||
x, y = nxt[:, 0, 0], nxt[:, 0, 1]
|
||||
inb = (x >= 0) & (x < W) & (y >= 0) & (y < H)
|
||||
alive &= stt & inb
|
||||
cur = nxt
|
||||
return float(alive.mean())
|
||||
|
||||
|
||||
def best_windows(m: np.ndarray, L: int, pct: float = 95, top: int = 1, min_gap: int | None = None) -> list:
|
||||
"""Return up to `top` low-motion windows (start index) minimizing the pct-percentile of
|
||||
per-frame motion, non-overlapping by min_gap (default L)."""
|
||||
if len(m) < L:
|
||||
return []
|
||||
min_gap = min_gap or L
|
||||
scores = np.array([np.percentile(m[s:s + L], pct) for s in range(len(m) - L + 1)])
|
||||
order = np.argsort(scores)
|
||||
picks: list[int] = []
|
||||
for s in order:
|
||||
if all(abs(s - p) >= min_gap for p in picks):
|
||||
picks.append(int(s))
|
||||
if len(picks) >= top:
|
||||
break
|
||||
return picks
|
||||
|
||||
|
||||
def greedy_valid_windows(m: np.ndarray,
|
||||
shifts: np.ndarray,
|
||||
L: int,
|
||||
pct: float,
|
||||
p95_thresh: float,
|
||||
drift_thresh: float,
|
||||
cand_step: int | None = None) -> list:
|
||||
"""Greedily take ALL non-overlapping windows passing p95_motion<=p95_thresh AND
|
||||
net-drift<=drift_thresh, calmest-first. Returns list of (start, p95, drift). All metrics
|
||||
come from the phase-correlation signal (no LK), so this is cheap."""
|
||||
if len(m) < L:
|
||||
return []
|
||||
cand_step = cand_step or max(1, L // 6)
|
||||
cand = [(s, float(np.percentile(m[s:s + L], pct))) for s in range(0, len(m) - L + 1, cand_step)]
|
||||
cand.sort(key=lambda x: x[1]) # calmest first
|
||||
accepted: list[tuple[int, float, float]] = []
|
||||
for s, p in cand:
|
||||
if p > p95_thresh:
|
||||
break # sorted -> everything after is worse
|
||||
if any(abs(s - a[0]) < L for a in accepted):
|
||||
continue # overlaps an accepted window
|
||||
drift = window_drift(shifts, s, L)
|
||||
if drift > drift_thresh:
|
||||
continue
|
||||
accepted.append((int(s), round(p, 4), round(drift, 3)))
|
||||
return sorted(accepted)
|
||||
|
||||
|
||||
def cmd_scan(args: argparse.Namespace) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import imageio.v2 as imageio
|
||||
|
||||
vids = sorted(Path(args.videos_dir).glob("**/*.mp4"))[:args.limit]
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
entries = []
|
||||
for i, vp in enumerate(vids, 1):
|
||||
gray, fps, _ = read_gray_lowres(str(vp), args.max_frames)
|
||||
L = int(round(args.window_sec * fps))
|
||||
if len(gray) < L:
|
||||
print(f"[win] [{i}/{len(vids)}] {vp.name}: only {len(gray)}f < {L}, skip", flush=True)
|
||||
continue
|
||||
m, shifts, diag = motion_series(gray)
|
||||
picks = best_windows(m, L, args.pct, top=args.top)
|
||||
stem = vp.stem
|
||||
# motion curve with chosen windows shaded
|
||||
fig, ax = plt.subplots(figsize=(12, 3.2))
|
||||
ax.plot(np.arange(len(m)) / fps, m, lw=0.8, color="steelblue")
|
||||
ax.axhline(args.calm_thresh, color="red", ls="--", lw=0.8, label=f"calm thresh {args.calm_thresh}")
|
||||
wstats = []
|
||||
for j, s in enumerate(picks):
|
||||
drift = window_drift(shifts, s, L)
|
||||
p95 = float(np.percentile(m[s:s + L], args.pct))
|
||||
ax.axvspan(s / fps, (s + L) / fps, color="limegreen", alpha=0.3)
|
||||
ax.text((s + L / 2) / fps, m.max() * 0.9, f"drift {drift:.3f}", ha="center", fontsize=8)
|
||||
wstats.append({
|
||||
"start_frame": s,
|
||||
"start_sec": round(s / fps, 2),
|
||||
"len_frames": L,
|
||||
"p95_motion": round(p95, 4),
|
||||
"drift": round(drift, 3)
|
||||
})
|
||||
ax.set_xlabel("time (s)")
|
||||
ax.set_ylabel("per-frame motion (norm)")
|
||||
ax.set_title(f"{stem} ({len(m)}f @ {fps:.0f}fps) green = static {args.window_sec}s window")
|
||||
ax.legend(fontsize=8)
|
||||
fig.tight_layout()
|
||||
curve = f"{stem}_motion.png"
|
||||
fig.savefig(str(out / curve), dpi=80)
|
||||
plt.close(fig)
|
||||
# preview: the best window as a short clip (full-res, read on demand)
|
||||
prev = ""
|
||||
if picks:
|
||||
s = picks[0]
|
||||
rgbw = read_rgb_frames(str(vp), range(s, s + L))
|
||||
imageio.mimsave(str(out / f"{stem}_win.mp4"), rgbw, fps=int(round(fps)), macro_block_size=1)
|
||||
prev = f"{stem}_win.mp4"
|
||||
entries.append({
|
||||
"id": stem,
|
||||
"src": str(vp),
|
||||
"fps": round(fps, 2),
|
||||
"n_frames": len(m),
|
||||
"curve": curve,
|
||||
"preview": prev,
|
||||
"windows": wstats,
|
||||
"median_motion": round(float(np.median(m)), 4)
|
||||
})
|
||||
best = wstats[0] if wstats else {}
|
||||
print(
|
||||
f"[win] [{i}/{len(vids)}] {stem}: {len(m)}f, best window @{best.get('start_sec')}s "
|
||||
f"p95={best.get('p95_motion')} surv={best.get('survival')}",
|
||||
flush=True)
|
||||
(out / "manifest.json").write_text(json.dumps(entries, indent=2))
|
||||
print(f"[win] scanned {len(entries)} videos -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_build(args: argparse.Namespace) -> None:
|
||||
"""Scan untrimmed source clips, greedily extract all non-overlapping static windows, and
|
||||
write them as target_frames-frame clips (subsampled to span the whole window) + manifest."""
|
||||
import imageio.v2 as imageio
|
||||
vids = sorted(Path(args.videos_dir).glob("**/*.mp4"))
|
||||
out = Path(args.out)
|
||||
(out / "videos").mkdir(parents=True, exist_ok=True)
|
||||
GENERIC = "a first-person egocentric view of a person performing a kitchen task with their hands"
|
||||
man = []
|
||||
idx = 0
|
||||
for vi, vp in enumerate(vids, 1):
|
||||
if idx >= args.target_count:
|
||||
break
|
||||
try:
|
||||
gray, fps, _ = read_gray_lowres(str(vp), args.max_frames)
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[build] {vp.name}: read fail {e}", flush=True)
|
||||
continue
|
||||
L = int(round(args.window_sec * fps))
|
||||
if len(gray) < L:
|
||||
continue
|
||||
m, shifts, _ = motion_series(gray)
|
||||
wins = greedy_valid_windows(m, shifts, L, args.pct, args.p95_thresh, args.drift_thresh)
|
||||
stride = max(1, round(L / args.target_frames))
|
||||
out_fps = round(fps / stride, 3)
|
||||
part = vp.stem.split("_")[0]
|
||||
n_from_clip = 0
|
||||
for (s, p95, drift) in wins:
|
||||
if idx >= args.target_count:
|
||||
break
|
||||
sel = np.arange(s, s + L, stride)[:args.target_frames]
|
||||
if sel.size < args.target_frames:
|
||||
continue
|
||||
f = read_rgb_frames(str(vp), sel) # full-res, only the window frames
|
||||
f = f[:, :f.shape[1] // 2 * 2, :f.shape[2] // 2 * 2] # even dims for libx264
|
||||
name = f"vid_{idx:06d}.mp4"
|
||||
imageio.mimsave(str(out / "videos" / name), f, fps=int(round(out_fps)), macro_block_size=1)
|
||||
man.append({
|
||||
"idx": idx,
|
||||
"path": name,
|
||||
"cap": [GENERIC],
|
||||
"fps": out_fps,
|
||||
"num_frames": int(args.target_frames),
|
||||
"duration": round(args.target_frames / out_fps, 3),
|
||||
"resolution": [int(f.shape[1]), int(f.shape[2])],
|
||||
"participant": part,
|
||||
"orig": vp.stem,
|
||||
"win_start_sec": round(s / fps, 2),
|
||||
"p95_motion": p95,
|
||||
"drift": drift
|
||||
})
|
||||
idx += 1
|
||||
n_from_clip += 1
|
||||
if n_from_clip:
|
||||
print(f"[build] [{vi}/{len(vids)}] {vp.stem}: +{n_from_clip} windows (total {idx})", flush=True)
|
||||
(out / "videos2caption.json").write_text(json.dumps(man, indent=2))
|
||||
(out / "merge.txt").write_text(f"{out}/videos,{out}/videos2caption.json\n")
|
||||
from collections import Counter
|
||||
print(
|
||||
f"[build] wrote {len(man)} static clips -> {out} | per-participant {dict(Counter(x['participant'] for x in man))}",
|
||||
flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
entries = json.loads((viz / "manifest.json").read_text())
|
||||
gal = [(str(viz / e["curve"]), e["id"]) for e in entries]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = entries[evt.index]
|
||||
vp = str(viz / e["preview"]) if e.get("preview") else None
|
||||
return str(viz / e["curve"]), vp, json.dumps(e["windows"], indent=2)
|
||||
|
||||
with gr.Blocks(title="Static-window finder") as demo:
|
||||
gr.Markdown("### Per-frame camera motion + static-window finder\n"
|
||||
"The curve is per-frame global motion (re-seeded each frame, so it **recovers after a turn** "
|
||||
"— spikes = turns, valleys = static). Green span = the chosen static window; `surv` = grid "
|
||||
"survival seeded at the window start. Click a clip -> motion curve + a preview of the best "
|
||||
"window + per-window stats. We keep windows with low `p95_motion` and high `survival`.")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gal, columns=2, height=560, label="videos (click)")
|
||||
with gr.Column():
|
||||
curve = gr.Image(label="per-frame motion (green = static window)")
|
||||
vid = gr.Video(label="preview of best static window")
|
||||
meta = gr.Code(label="windows", language="json")
|
||||
g.select(show, None, [curve, vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("scan")
|
||||
r.add_argument("--videos-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--window-sec", type=float, default=5.0)
|
||||
r.add_argument("--pct", type=float, default=95.0, help="percentile of per-frame motion used to score a window")
|
||||
r.add_argument("--calm-thresh", type=float, default=0.01, help="reference 'calm' motion line for the plot")
|
||||
r.add_argument("--top", type=int, default=2, help="static windows to find per video")
|
||||
r.add_argument("--max-frames", type=int, default=900, help="cap frames scanned per video (bounds cost)")
|
||||
r.add_argument("--limit", type=int, default=12)
|
||||
r.set_defaults(func=cmd_scan)
|
||||
b = sub.add_parser("build")
|
||||
b.add_argument("--videos-dir", required=True)
|
||||
b.add_argument("--out", required=True)
|
||||
b.add_argument("--window-sec", type=float, default=5.0)
|
||||
b.add_argument("--target-frames", type=int, default=121, help="frames per output clip (window subsampled to span)")
|
||||
b.add_argument("--target-count", type=int, default=200)
|
||||
b.add_argument("--pct", type=float, default=95.0)
|
||||
b.add_argument("--p95-thresh", type=float, default=0.004, help="max p95 per-frame motion for a valid window")
|
||||
b.add_argument("--drift-thresh", type=float, default=0.35, help="max net view drift (frac of diagonal)")
|
||||
b.add_argument("--max-frames", type=int, default=1500, help="cap frames scanned per source clip")
|
||||
b.set_defaults(func=cmd_build)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7890)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,239 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pre-generate controllability-eval artifacts for the viewer app.
|
||||
|
||||
For one checkpoint, over clips x counterfactuals, generate the video and save:
|
||||
- <clip>_<ctrl>__gen.mp4 raw generation
|
||||
- <clip>_<ctrl>__input.mp4 gen + the INPUT control tracks (what we asked for)
|
||||
- <clip>_<ctrl>__tracked.mp4 gen + tracks RE-EXTRACTED from the gen (CoTracker3)
|
||||
- <clip>_<ctrl>__heat.mp4 gen + points colored by per-point EPE (green=followed,
|
||||
red=ignored) -- the EPE heatmap (when tracks apply)
|
||||
plus per clip: <clip>__original.mp4 (real clip) and <clip>__original_tracks.mp4.
|
||||
|
||||
A manifest.json records EPE / n_points / coverage and the file names so the viewer
|
||||
(examples/inference/gradio/trackwan/app_viewer.py) can browse without a GPU.
|
||||
|
||||
Run one process per checkpoint (parallel across GPUs)::
|
||||
|
||||
srun ... env CUDA_VISIBLE_DEVICES=5 .venv/bin/python data_pipeline/gen_eval_artifacts.py \
|
||||
--export <export_3000> --name "step 3000" --out <artifacts_root>/step3000 \
|
||||
--data <funinp_now parquet> --clips 0 1 2 3
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
from fastvideo.eval.metrics.motion.control_sensitivity.metric import ( # noqa: E402
|
||||
compute_control_sensitivity, render_diff_heatmap)
|
||||
|
||||
HEAT_SCALE = 40.0 # px error mapped green(0)->red(>=40)
|
||||
|
||||
|
||||
def _draw(frames: Any, trpx: Any, vis: Any, colors: Any) -> Any:
|
||||
from fastvideo.train.callbacks.track_validation import _draw_overlay
|
||||
return _draw_overlay(frames, trpx, vis, colors, 9, 2, 0.65)
|
||||
|
||||
|
||||
def _solid(m: Any, rgb: Any) -> np.ndarray:
|
||||
return np.tile(np.array([rgb], np.uint8), (m, 1))
|
||||
|
||||
|
||||
def _heat_colors(perpt: Any) -> np.ndarray:
|
||||
e = np.clip(np.nan_to_num(perpt, nan=HEAT_SCALE) / HEAT_SCALE, 0.0, 1.0)
|
||||
r = (255 * e).astype(np.uint8)
|
||||
g = (255 * (1 - e)).astype(np.uint8)
|
||||
b = np.full_like(r, 40)
|
||||
return np.stack([r, g, b], 1)
|
||||
|
||||
|
||||
def _retrack_core(frames: Any,
|
||||
tr_px: Any,
|
||||
vs: Any,
|
||||
ct: Any,
|
||||
device: Any,
|
||||
max_points: int = 600,
|
||||
seed: int = 0,
|
||||
qf: int = 0,
|
||||
vt: float = 0.5) -> dict[str, Any] | None:
|
||||
"""Mirror compute_epe but also return the extracted tracks + per-point error."""
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import _retrack
|
||||
T = min(frames.shape[0], tr_px.shape[0])
|
||||
frames, tr_px, vs = frames[:T], tr_px[:T], vs[:T]
|
||||
idx = np.nonzero(vs[qf] > vt)[0]
|
||||
if idx.size == 0:
|
||||
return None
|
||||
if max_points and idx.size > max_points:
|
||||
rng = np.random.default_rng(seed)
|
||||
idx = np.sort(rng.choice(idx, size=max_points, replace=False))
|
||||
q_xy = tr_px[qf, idx]
|
||||
queries = np.concatenate([np.full((idx.size, 1), qf, np.float32), q_xy], axis=1)
|
||||
rt, _ = _retrack(ct, frames, queries, device) # [T,M,2] px
|
||||
Tr = min(T, rt.shape[0])
|
||||
tgt = tr_px[:Tr][:, idx]
|
||||
pred = rt[:Tr]
|
||||
mask = vs[:Tr][:, idx] > vt
|
||||
d = np.sqrt(((pred - tgt)**2).sum(-1)) # [Tr,M]
|
||||
valid = d[mask]
|
||||
epe = float(valid.mean()) if valid.size else None
|
||||
perpt = np.array([d[:, m][mask[:, m]].mean() if mask[:, m].any() else np.nan for m in range(idx.size)])
|
||||
disp = np.sqrt(((tgt - tgt[0:1])**2).sum(-1)) # input travel from frame 0
|
||||
moving = disp.max(0) > 8.0
|
||||
mm = mask & moving[None, :]
|
||||
epe_moving = float(d[mm].mean()) if d[mm].size else None
|
||||
return dict(tgt=tgt,
|
||||
pred=pred,
|
||||
mask=mask,
|
||||
epe=epe,
|
||||
epe_moving=epe_moving,
|
||||
perpt=perpt,
|
||||
n=int(idx.size),
|
||||
n_moving=int(moving.sum()),
|
||||
cov=float(mask.mean()))
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument("--name", required=True, help="checkpoint label e.g. 'step 3000'")
|
||||
p.add_argument("--out", required=True)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", required=True)
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1, 2, 3])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import load_cotracker
|
||||
ct = load_cotracker(model.device)
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
Tpx = (num_lat_t - 1) * ratio + 1
|
||||
dev = model.device
|
||||
|
||||
def save(frames: Any, name: str) -> str:
|
||||
path = os.path.join(args.out, name)
|
||||
imageio.mimsave(path, frames, fps=args.fps, macro_block_size=1)
|
||||
return name
|
||||
|
||||
manifest = {"name": args.name, "clips": {}}
|
||||
for ci, s in zip(args.clips, samples, strict=False):
|
||||
gt_t = s["track_points"][0].numpy()[:Tpx]
|
||||
gt_v = s["track_visibility"][0].numpy()[:Tpx]
|
||||
swap = samples[(args.clips.index(ci) + 1) % len(samples)]
|
||||
g = st.make_grid(50)
|
||||
pan_t, pan_v = st.pan(g, Tpx, 0.25, 0.0)
|
||||
zoom_t, zoom_v = st.zoom(g, Tpx, 1.25)
|
||||
drag_t, drag_v = st.drag(g, Tpx, center=(0.5, 0.5), dx=0.3, dy=0.0, radius=0.2)
|
||||
drag_vs = st.select_radius(drag_v, g, center=(0.5, 0.5), radius=0.2)
|
||||
controls = {
|
||||
"gt": (gt_t, gt_v),
|
||||
"none": (None, None),
|
||||
"pan_right": (pan_t, pan_v),
|
||||
"zoom_in": (zoom_t, zoom_v),
|
||||
"drag_dense": (drag_t, drag_v),
|
||||
"drag_sparse": (drag_t, drag_vs),
|
||||
"swap": (swap["track_points"][0].numpy()[:Tpx], swap["track_visibility"][0].numpy()[:Tpx]),
|
||||
}
|
||||
|
||||
ref = twi.decode_reference(model, s["vae_latent"])
|
||||
H, W = ref.shape[1], ref.shape[2]
|
||||
gtpx = gt_t.copy()
|
||||
gtpx[..., 0] *= W
|
||||
gtpx[..., 1] *= H
|
||||
entry = {
|
||||
"caption":
|
||||
s["caption"][:240],
|
||||
"H":
|
||||
int(H),
|
||||
"W":
|
||||
int(W),
|
||||
"original":
|
||||
save(ref, f"clip{ci}__original.mp4"),
|
||||
"original_tracks":
|
||||
save(_draw(ref, gtpx, gt_v, _solid(gt_t.shape[1], [0, 220, 255])), f"clip{ci}__original_tracks.mp4"),
|
||||
"controls": {}
|
||||
}
|
||||
|
||||
frames_gt = None
|
||||
for cname, (tr, vs) in controls.items():
|
||||
tp = torch.from_numpy(tr)[None].float() if tr is not None else None
|
||||
tv = torch.from_numpy(vs)[None].float() if vs is not None else None
|
||||
lat = twi.generate(model,
|
||||
first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp,
|
||||
track_visibility=tv,
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=args.steps,
|
||||
seed=args.seed)
|
||||
frames = twi.decode_to_pixels(model, lat)
|
||||
if cname == "gt":
|
||||
frames_gt = frames
|
||||
trpx = None
|
||||
if tr is not None:
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
rec = {
|
||||
"gen": save(frames, f"clip{ci}_{cname}__gen.mp4"),
|
||||
"epe": None,
|
||||
"epe_moving": None,
|
||||
"n_points": 0,
|
||||
"coverage": 0.0
|
||||
}
|
||||
if trpx is not None:
|
||||
core = _retrack_core(frames, trpx, vs, ct, dev)
|
||||
if core is not None:
|
||||
m = core["tgt"].shape[1]
|
||||
rec["input"] = save(_draw(frames, core["tgt"], core["mask"], _solid(m, [0, 220, 255])),
|
||||
f"clip{ci}_{cname}__input.mp4")
|
||||
rec["tracked"] = save(_draw(frames, core["pred"], core["mask"], _solid(m, [255, 220, 0])),
|
||||
f"clip{ci}_{cname}__tracked.mp4")
|
||||
rec["heat"] = save(_draw(frames, core["pred"], core["mask"], _heat_colors(core["perpt"])),
|
||||
f"clip{ci}_{cname}__heat.mp4")
|
||||
rec.update(epe=core["epe"],
|
||||
epe_moving=core["epe_moving"],
|
||||
n_points=core["n"],
|
||||
n_moving=core["n_moving"],
|
||||
coverage=core["cov"])
|
||||
# intervention diff vs the GT-track generation (counterfactual control sensitivity)
|
||||
if cname != "gt" and frames_gt is not None:
|
||||
sens = compute_control_sensitivity(frames_gt, frames, trpx, vs)
|
||||
rec["diff"] = save(render_diff_heatmap(frames_gt, frames), f"clip{ci}_{cname}__diff.mp4")
|
||||
sr = sens["sensitivity_roi"]
|
||||
rec.update(sensitivity_roi=(round(sr, 4) if sr == sr else None),
|
||||
bg_leakage=round(sens["bg_leakage"], 4),
|
||||
localization_iou=round(sens["localization_iou"], 4),
|
||||
mean_diff=round(sens["mean_diff"], 4))
|
||||
entry["controls"][cname] = rec
|
||||
print(
|
||||
f"[{args.name}] clip{ci} {cname:12s} "
|
||||
f"EPE={rec['epe'] if rec['epe'] is None else round(float(rec['epe']),2)} "
|
||||
f"EPEmv={rec['epe_moving'] if rec['epe_moving'] is None else round(float(rec['epe_moving']),2)} "
|
||||
f"sens={rec.get('sensitivity_roi')} iou={rec.get('localization_iou')}",
|
||||
flush=True)
|
||||
manifest["clips"][str(ci)] = entry
|
||||
|
||||
with open(os.path.join(args.out, "manifest.json"), "w") as f:
|
||||
json.dump(manifest, f, indent=2)
|
||||
print("ARTIFACTS_DONE", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-GPU data-parallel worker for large-scale Wan2.2-T2V-A14B synthetic gen.
|
||||
|
||||
One VideoGenerator per GPU (num_gpus=1, no FSDP/SP/TP). Each worker owns a stride
|
||||
slice of the prompt list: prompts[worker_id :: num_workers]. Idempotent/resumable:
|
||||
a video whose final mp4 exists is skipped, so a requeue after a cordon just continues.
|
||||
|
||||
First-frame saturation fix baked in: generate (num_frames + drop) frames, drop the
|
||||
first `drop` decoded frames, keep `num_frames`. (See research_log/08-first-frame-saturation.md.)
|
||||
|
||||
Per-worker manifest shard (manifest_shards/worker_<id>.jsonl) avoids 48 procs racing on
|
||||
one file; merge_manifests.py compiles them into the FastVideo videos2caption.json.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, json, os, time, traceback
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
|
||||
MODEL = "Wan-AI/Wan2.2-T2V-A14B-Diffusers"
|
||||
|
||||
def parse_args():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--prompts", type=Path, required=True)
|
||||
p.add_argument("--output-dir", type=Path, required=True)
|
||||
p.add_argument("--worker-id", type=int, required=True)
|
||||
p.add_argument("--num-workers", type=int, required=True)
|
||||
p.add_argument("--max-videos", type=int, default=None, help="global cap on #prompts consumed (across all workers)")
|
||||
p.add_argument("--height", type=int, default=720)
|
||||
p.add_argument("--width", type=int, default=1280)
|
||||
p.add_argument("--num-frames", type=int, default=121, help="frames to KEEP")
|
||||
p.add_argument("--drop", type=int, default=8, help="leading frames dropped (gen = keep+drop, must be 4k+1)")
|
||||
p.add_argument("--fps", type=int, default=16)
|
||||
p.add_argument("--steps", type=int, default=40)
|
||||
p.add_argument("--guidance-scale", type=float, default=4.0)
|
||||
p.add_argument("--guidance-scale-2", type=float, default=3.0)
|
||||
p.add_argument("--seed-base", type=int, default=1024, help="per-video seed = seed_base + global_prompt_idx")
|
||||
p.add_argument("--shuffle-seed", type=int, default=1234, help="deterministic prompt shuffle (same across workers)")
|
||||
return p.parse_args()
|
||||
|
||||
def load_prompts(path: Path):
|
||||
lines = [ln.strip() for ln in path.read_text(encoding="utf-8").splitlines() if ln.strip()]
|
||||
return lines
|
||||
|
||||
def main():
|
||||
a = parse_args()
|
||||
gen_frames = a.num_frames + a.drop
|
||||
assert (gen_frames - 1) % 4 == 0, f"gen frames {gen_frames} must be 4k+1"
|
||||
|
||||
videos_dir = a.output_dir / "videos"
|
||||
meta_dir = a.output_dir / "meta"
|
||||
shard_dir = a.output_dir / "manifest_shards"
|
||||
for d in (videos_dir, meta_dir, shard_dir):
|
||||
d.mkdir(parents=True, exist_ok=True)
|
||||
shard_manifest = shard_dir / f"worker_{a.worker_id:04d}.jsonl"
|
||||
fail_log = a.output_dir / f"failures_worker_{a.worker_id:04d}.log"
|
||||
|
||||
prompts = load_prompts(a.prompts)
|
||||
# Deterministic global shuffle so prompt order is stable across restarts and workers.
|
||||
order = np.random.RandomState(a.shuffle_seed).permutation(len(prompts))
|
||||
if a.max_videos is not None:
|
||||
order = order[:a.max_videos]
|
||||
# This worker's stride slice.
|
||||
my_positions = list(range(a.worker_id, len(order), a.num_workers))
|
||||
my_global_idx = [int(order[pos]) for pos in my_positions]
|
||||
print(f"[w{a.worker_id}/{a.num_workers}] assigned {len(my_global_idx)} prompts "
|
||||
f"(gen {gen_frames}->keep {a.num_frames}, {a.width}x{a.height}@{a.fps}fps)", flush=True)
|
||||
|
||||
# already-done set from existing mp4s
|
||||
def final_path(idx): return videos_dir / f"vid_{idx:06d}.mp4"
|
||||
|
||||
import imageio.v2 as imageio
|
||||
from fastvideo import VideoGenerator
|
||||
t0 = time.time()
|
||||
g = VideoGenerator.from_pretrained(
|
||||
MODEL, num_gpus=1, use_fsdp_inference=False,
|
||||
dit_cpu_offload=False, vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True, pin_cpu_memory=True,
|
||||
)
|
||||
print(f"[w{a.worker_id}] model ready in {time.time()-t0:.1f}s", flush=True)
|
||||
|
||||
n_ok = 0
|
||||
for gi in my_global_idx:
|
||||
fp = final_path(gi)
|
||||
if fp.exists():
|
||||
continue
|
||||
prompt = prompts[gi]
|
||||
tmp = videos_dir / f".tmp_w{a.worker_id}_{gi:06d}.mp4"
|
||||
t = time.time()
|
||||
try:
|
||||
res = g.generate_video(
|
||||
prompt, save_video=False, return_frames=True,
|
||||
height=a.height, width=a.width, num_frames=gen_frames, fps=a.fps,
|
||||
seed=a.seed_base + gi, num_inference_steps=a.steps,
|
||||
guidance_scale=a.guidance_scale, guidance_scale_2=a.guidance_scale_2,
|
||||
)
|
||||
if isinstance(res, list):
|
||||
res = res[0]
|
||||
frames = np.asarray(res["frames"])[a.drop:a.drop + a.num_frames]
|
||||
if frames.shape[0] != a.num_frames:
|
||||
raise RuntimeError(f"got {frames.shape[0]} frames after drop, want {a.num_frames}")
|
||||
imageio.mimsave(tmp, list(frames), fps=a.fps, format="mp4")
|
||||
os.replace(tmp, fp) # atomic publish
|
||||
except Exception as e: # keep the worker alive
|
||||
with fail_log.open("a") as f:
|
||||
f.write(json.dumps({"idx": gi, "err": repr(e)[:500]}) + "\n")
|
||||
print(f"[w{a.worker_id}] FAIL idx={gi}: {e!r}", flush=True)
|
||||
if tmp.exists(): tmp.unlink()
|
||||
continue
|
||||
dt = time.time() - t
|
||||
rec = {"idx": gi, "path": fp.name, "cap": [prompt], "fps": float(a.fps),
|
||||
"num_frames": a.num_frames, "resolution": {"width": a.width, "height": a.height},
|
||||
"gen_seconds": round(dt, 1)}
|
||||
with shard_manifest.open("a") as f:
|
||||
f.write(json.dumps(rec) + "\n")
|
||||
(meta_dir / f"vid_{gi:06d}.json").write_text(json.dumps(rec))
|
||||
n_ok += 1
|
||||
if n_ok % 10 == 0:
|
||||
print(f"[w{a.worker_id}] {n_ok} done, last {dt:.0f}s idx={gi}", flush=True)
|
||||
print(f"[w{a.worker_id}] DONE_WORKER made {n_ok} new videos", flush=True)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,215 @@
|
||||
# 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("--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
|
||||
|
||||
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=args.num_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")
|
||||
shutil.move(str(produced[0]), str(final_path))
|
||||
except Exception as e: # noqa: BLE001 - keep the batch alive, log and move on
|
||||
with fail_log.open("a") as f:
|
||||
f.write(json.dumps({"idx": idx, "error": repr(e), "prompt": prompt}) + "\n")
|
||||
print(f"[gen] FAILED idx={idx}: {e!r}", flush=True)
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
continue
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
|
||||
record = {
|
||||
"idx": idx,
|
||||
"path": final_path.name, # basename; folder in merge.txt is videos_dir
|
||||
"cap": [prompt],
|
||||
"fps": float(args.fps),
|
||||
"duration": duration,
|
||||
"num_frames": int(args.num_frames),
|
||||
"resolution": {
|
||||
"width": args.width,
|
||||
"height": args.height
|
||||
},
|
||||
}
|
||||
with manifest_jsonl.open("a") as f:
|
||||
f.write(json.dumps(record) + "\n")
|
||||
print(f"[gen] [{n_done}/{len(todo)}] idx={idx} {time.time()-t0:.1f}s -> {final_path.name}", flush=True)
|
||||
|
||||
n = rebuild_manifest(manifest_jsonl, videos_dir, json_path, merge_path)
|
||||
print(f"[gen] done; manifest has {n} entries -> {json_path}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,246 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Layer-wise learning-rate (LWLR) bootstrap experiment.
|
||||
|
||||
Question: on a 14B random-init WanTrack model, does giving the track pathway
|
||||
(track_encoder.* + patch_embedding.weight[:, 36:]) a 100x higher LR than the
|
||||
rest actually let it grow into functional gradient magnitudes within a small
|
||||
number of steps? Or is the encoder stuck regardless?
|
||||
|
||||
Runs a short training loop on real OpenVid samples, tracking weight+grad norms
|
||||
of the track pathway at each step. Compares against the "known good" magnitude
|
||||
observed in 1.3B_merged (per-sample track_encoder.proj grad ~0.5).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from safetensors.torch import load_file as st_load
|
||||
|
||||
# Distributed init BEFORE importing fastvideo (same pattern as gradient_flow_analysis)
|
||||
def _init_distributed_1gpu() -> None:
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
)
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
|
||||
|
||||
_init_distributed_1gpu()
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
from gradient_flow_analysis import ( # noqa: E402
|
||||
build_model_from_diffusers,
|
||||
load_samples,
|
||||
build_i2v_cond_concat,
|
||||
flow_matching_noisy,
|
||||
Sample,
|
||||
)
|
||||
|
||||
|
||||
TRACK_PATH_PARAMS = (
|
||||
"track_encoder.proj.weight",
|
||||
"track_encoder.temporal_conv.weight",
|
||||
"track_encoder.proj.bias", # may not exist depending on TRACKWAN_TRACK_BIAS
|
||||
"track_encoder.temporal_conv.bias",
|
||||
)
|
||||
|
||||
|
||||
def split_param_groups(model, high_lr: float, base_lr: float) -> list[dict]:
|
||||
"""Give track pathway (track_encoder.* + patch_embedding[:, 36:]) a high LR.
|
||||
|
||||
patch_embedding.weight is a SINGLE tensor of shape [hidden, 52, 1, 2, 2] — we can't
|
||||
give sub-slice a different LR through standard param groups. Workaround: put the
|
||||
whole patch_embedding in a "hybrid" group with high LR (the base 36 channels will
|
||||
also move faster, but by only 100x for a short warmup that's small — 100 steps of
|
||||
high LR is still << stage-1 4800 steps of normal LR).
|
||||
"""
|
||||
high_group, base_group = [], []
|
||||
named = dict(model.named_parameters())
|
||||
for n, p in named.items():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
# High-LR pathway: encoder convs (+ bias if present) + the whole patch_embedding.weight
|
||||
is_encoder = n.startswith("track_encoder.")
|
||||
is_patch_embed = (n == "patch_embedding.weight")
|
||||
if is_encoder or is_patch_embed:
|
||||
high_group.append(p)
|
||||
else:
|
||||
base_group.append(p)
|
||||
return [
|
||||
{"params": high_group, "lr": high_lr, "name": "track_pathway"},
|
||||
{"params": base_group, "lr": base_lr, "name": "dit_body"},
|
||||
]
|
||||
|
||||
|
||||
def snapshot_norms(model, sample_grad_i: int) -> dict:
|
||||
"""Return weight/grad norms for the diagnostic params."""
|
||||
out = {}
|
||||
for n, p in model.named_parameters():
|
||||
if not p.requires_grad:
|
||||
continue
|
||||
# Track params of interest
|
||||
keep = (
|
||||
n.startswith("track_encoder.")
|
||||
or n == "patch_embedding.weight"
|
||||
)
|
||||
if not keep:
|
||||
continue
|
||||
w_norm = p.detach().float().norm().item()
|
||||
g_norm = p.grad.detach().float().norm().item() if p.grad is not None else 0.0
|
||||
# For patch_embedding, split into base/track slices
|
||||
if n == "patch_embedding.weight":
|
||||
w_base = p.detach()[:, :36].float()
|
||||
w_track = p.detach()[:, 36:].float()
|
||||
out["patch_embedding.weight[:, :36]"] = {
|
||||
"w_norm": w_base.norm().item(), "g_norm": 0.0,
|
||||
}
|
||||
out["patch_embedding.weight[:, 36:]"] = {
|
||||
"w_norm": w_track.norm().item(), "g_norm": 0.0,
|
||||
}
|
||||
if p.grad is not None:
|
||||
g_base = p.grad.detach()[:, :36].float()
|
||||
g_track = p.grad.detach()[:, 36:].float()
|
||||
out["patch_embedding.weight[:, :36]"]["g_norm"] = g_base.norm().item()
|
||||
out["patch_embedding.weight[:, 36:]"]["g_norm"] = g_track.norm().item()
|
||||
else:
|
||||
out[n] = {"w_norm": w_norm, "g_norm": g_norm}
|
||||
return out
|
||||
|
||||
|
||||
def one_step(model, s: Sample, device: torch.device, flow_shift: float,
|
||||
from_ctx_ns) -> float:
|
||||
"""Forward+backward on one sample; grads accumulate."""
|
||||
vae = s.vae_latent.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
ff = s.first_frame_latent.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
clip = s.clip_feature.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
text = s.text_embedding.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
mask = torch.ones(1, s.text_embedding.shape[0], device=device, dtype=torch.bfloat16)
|
||||
tp = s.track_points.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
tv = s.track_visibility.unsqueeze(0).to(device, dtype=torch.bfloat16)
|
||||
|
||||
num_latent_t = vae.shape[2]
|
||||
expected_T = (num_latent_t - 1) * 4 + 1
|
||||
tp = tp[:, :expected_T]
|
||||
tv = tv[:, :expected_T]
|
||||
|
||||
sigma_u = float(torch.rand(1).item())
|
||||
noise = torch.randn_like(vae, dtype=torch.float32)
|
||||
clean_f = vae.float()
|
||||
noisy_f, _, ts_val = flow_matching_noisy(clean_f, noise, sigma_u,
|
||||
flow_shift=flow_shift)
|
||||
noisy = noisy_f.to(torch.bfloat16)
|
||||
cond20 = build_i2v_cond_concat(ff, vae_temporal_compression=4)
|
||||
hs = torch.cat([noisy, cond20], dim=1)
|
||||
ts = torch.tensor([ts_val], device=device, dtype=torch.bfloat16)
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16), from_ctx_ns(current_timestep=ts, attn_metadata=None):
|
||||
pred = model(
|
||||
hidden_states=hs, encoder_hidden_states=text, encoder_attention_mask=mask,
|
||||
timestep=ts, encoder_hidden_states_image=clip,
|
||||
track_points=tp, track_visibility=tv, return_dict=False,
|
||||
)
|
||||
target = noise - clean_f
|
||||
loss = F.mse_loss(pred.float(), target.float())
|
||||
loss.backward()
|
||||
return float(loss.item())
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--model", default="/mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_random_init")
|
||||
ap.add_argument("--n-samples", type=int, default=16)
|
||||
ap.add_argument("--n-steps", type=int, default=60)
|
||||
ap.add_argument("--batch-accum", type=int, default=4, help="gradient accumulation steps per optim step")
|
||||
ap.add_argument("--high-lr", type=float, default=1e-3, help="LR for track pathway")
|
||||
ap.add_argument("--base-lr", type=float, default=1e-5, help="LR for the rest of DiT")
|
||||
ap.add_argument("--flow-shift", type=float, default=6.0)
|
||||
ap.add_argument("--parquet-glob", type=str,
|
||||
default="/mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset/shard_*/**/*.parquet")
|
||||
ap.add_argument("--out", type=str,
|
||||
default="/mnt/lustre/vlm-s4duan/gradient_analysis_14b/lwlr_trace.json")
|
||||
ap.add_argument("--seed", type=int, default=42)
|
||||
args = ap.parse_args()
|
||||
|
||||
random.seed(args.seed)
|
||||
np.random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
device = torch.device("cuda")
|
||||
|
||||
samples = load_samples(args.parquet_glob, args.n_samples, seed=args.seed)
|
||||
print(f"loaded {len(samples)} samples")
|
||||
assert len(samples) >= args.batch_accum, "need enough samples for accumulation"
|
||||
|
||||
print(f"loading {args.model} ...")
|
||||
model = build_model_from_diffusers(args.model, device)
|
||||
|
||||
# Set up param groups
|
||||
param_groups = split_param_groups(model, high_lr=args.high_lr, base_lr=args.base_lr)
|
||||
n_high = sum(p.numel() for p in param_groups[0]["params"])
|
||||
n_base = sum(p.numel() for p in param_groups[1]["params"])
|
||||
print(f"track_pathway params: {n_high/1e6:.2f}M (LR={args.high_lr:g})")
|
||||
print(f"dit_body params: {n_base/1e6:.2f}M (LR={args.base_lr:g})")
|
||||
optim = torch.optim.AdamW(param_groups, weight_decay=0.0) # no wd so we see raw drift
|
||||
|
||||
trace: list[dict] = []
|
||||
# step 0 (before any update)
|
||||
print("\n=== step 0 (init) ===")
|
||||
norms0 = snapshot_norms(model, -1)
|
||||
for k, v in norms0.items():
|
||||
print(f" {k:<45s} w={v['w_norm']:.4f} g={v['g_norm']:.4f}")
|
||||
trace.append({"step": 0, "loss": None, "norms": norms0})
|
||||
|
||||
step = 0
|
||||
accum = 0
|
||||
losses_accum = []
|
||||
for outer in range(args.n_steps * args.batch_accum):
|
||||
s = samples[outer % len(samples)]
|
||||
loss = one_step(model, s, device, args.flow_shift, set_forward_context)
|
||||
losses_accum.append(loss)
|
||||
accum += 1
|
||||
if accum < args.batch_accum:
|
||||
continue
|
||||
step += 1
|
||||
# snapshot BEFORE step (grads are accumulated over the batch)
|
||||
gs = snapshot_norms(model, -1)
|
||||
optim.step()
|
||||
optim.zero_grad(set_to_none=True)
|
||||
avg_loss = float(np.mean(losses_accum))
|
||||
losses_accum = []
|
||||
accum = 0
|
||||
if step % 5 == 0 or step <= 5:
|
||||
print(f"\n=== step {step} (loss={avg_loss:.4f}) ===")
|
||||
for k, v in gs.items():
|
||||
marker = ""
|
||||
if k == "track_encoder.proj.weight":
|
||||
marker = " <-- 1.3B_merged saw per_sample_norm ~0.50"
|
||||
elif k == "patch_embedding.weight[:, 36:]":
|
||||
marker = " <-- 1.3B_merged saw per_sample_norm ~0.73"
|
||||
print(f" {k:<45s} w={v['w_norm']:.4f} g={v['g_norm']:.4f}{marker}")
|
||||
trace.append({"step": step, "loss": avg_loss, "norms": gs})
|
||||
|
||||
Path(args.out).parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(args.out, "w") as f:
|
||||
json.dump({"config": vars(args), "trace": trace}, f, indent=2)
|
||||
print(f"\nwrote trace -> {args.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Compile per-worker manifest shards -> FastVideo videos2caption.json + merge.txt.
|
||||
Idempotent; run anytime (progress check) or at the end. Dedups by idx, verifies mp4 exists."""
|
||||
from __future__ import annotations
|
||||
import argparse, json
|
||||
from pathlib import Path
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--output-dir", type=Path, required=True)
|
||||
a = p.parse_args()
|
||||
videos_dir = a.output_dir / "videos"
|
||||
shard_dir = a.output_dir / "manifest_shards"
|
||||
recs: dict[int, dict] = {}
|
||||
for sh in sorted(shard_dir.glob("worker_*.jsonl")):
|
||||
for ln in sh.read_text().splitlines():
|
||||
ln = ln.strip()
|
||||
if not ln:
|
||||
continue
|
||||
try:
|
||||
r = json.loads(ln)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if (videos_dir / r["path"]).exists():
|
||||
recs[int(r["idx"])] = r
|
||||
ordered = [recs[k] for k in sorted(recs)]
|
||||
j = a.output_dir / "videos2caption.json"
|
||||
j.write_text(json.dumps(ordered, indent=2))
|
||||
(a.output_dir / "merge.txt").write_text(f"{videos_dir.resolve()},{j.resolve()}\n")
|
||||
# also count mp4s on disk (may exceed manifest if a worker died mid-write)
|
||||
on_disk = len(list(videos_dir.glob("vid_*.mp4")))
|
||||
print(f"manifest: {len(ordered)} entries; {on_disk} mp4s on disk -> {j}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,65 @@
|
||||
# Data Pipeline — Decisions & Rationale
|
||||
|
||||
## Why generation + tracking live OUTSIDE `fastvideo/pipelines/preprocess/`
|
||||
FastVideo's preprocess pipeline is structurally a *consumer of existing .mp4 files*
|
||||
(`VideoCaptionMergedDataset` reads a `merge.txt`→`annotation.json` over a folder and
|
||||
asserts each video path exists). There is no seam to *produce* a video inside it. So
|
||||
generation (Wan2.2) and tracking (CoTracker) are standalone Stage-0 scripts that
|
||||
produce the files preprocess later ingests — matching the repo convention (everything
|
||||
upstream of preprocess is plain scripts, not pipeline stages).
|
||||
|
||||
## Two scripts, not one fused process
|
||||
- The `.mp4` on disk is the checkpoint; generation is the expensive part. Tracking
|
||||
config will be tweaked many times — must not re-run Wan each time.
|
||||
- Wan2.2-14B is huge (multi-GPU + offload); CoTracker is tiny. "Generate all → free
|
||||
model → track all" is the natural ordering and is exactly two phases.
|
||||
- Independent restart/parallelism. Both scripts are idempotent (skip done work).
|
||||
|
||||
## Generate at training fps/length (no temporal alignment headache)
|
||||
CoTracker tracks per *video* frame, but the VAE compresses time (~4× for Wan). If we
|
||||
generated at arbitrary fps we'd have to re-index tracks by the preprocess frame-sampler
|
||||
(`sample_frame_index`) and then fold onto latent time. Since we *control* generation,
|
||||
we emit videos already at the train fps/length (default 16 fps, 81 frames ≈ 5 s) so
|
||||
source frames == sampled frames, 1:1. Tracks then align to frames trivially; the only
|
||||
remaining fold (frames→latent-frames) happens in the model/trainer.
|
||||
|
||||
## Generation is T2V even though the model is I2V+points
|
||||
Synthetic videos come from text→video (Wan2.2-T2V-A14B). At *training* time the model
|
||||
uses frame 0 as the I2V conditioning image + the point tracks as motion control. So we
|
||||
never need an input image during generation.
|
||||
|
||||
## Manifest format (FastVideo-compatible)
|
||||
`generate_videos.py` emits the same shape the existing loader expects:
|
||||
- `videos/vid_NNNNNN.mp4`
|
||||
- `videos2caption.json`: list of `{path(basename), cap:[...], fps, duration, num_frames, resolution}`
|
||||
- `merge.txt`: one line `<videos_dir>,<videos2caption.json>`
|
||||
Loader actually consumes only `path, cap, fps, duration` (+ optional conditioning path);
|
||||
`resolution`/`num_frames` are informational. We know all of these at generation time, so
|
||||
we write the manifest directly and SKIP `scripts/dataset_preparation/prepare_json_file.py`
|
||||
(it exists only to *recover* fps/duration by re-probing videos of unknown provenance).
|
||||
|
||||
## points_path stored as ABSOLUTE path
|
||||
The future preprocess loader joins `folder + relative_path`. Tracks live in a sibling
|
||||
`tracks/` dir, not under `videos/`. Storing `points_path` absolute makes `os.path.join`
|
||||
return it unchanged, so it resolves regardless of the manifest folder. Mirrors how
|
||||
MatrixGame2 references `action_path`, which is the pattern the points preprocess task
|
||||
will copy.
|
||||
|
||||
## CoTracker v3 via torch.hub (`cotracker3_offline`)
|
||||
The repo only vendors CoTracker **v2** under `third_party/eval/vbench` (eval-only). For
|
||||
v3 we pull `facebookresearch/co-tracker` `cotracker3_offline` via torch.hub. Prefetch on
|
||||
the login node (has internet); the hub cache lives on shared weka so compute nodes reuse
|
||||
it offline (`extract_tracks.py` falls back to `source="local"` from the cache dir).
|
||||
|
||||
## Output location OUTSIDE the git repo
|
||||
Generated videos are large (≈2–5 MB each → ~25 GB at 5k). Default output dir is
|
||||
`/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/...` (outside `FastVideo/`) to avoid
|
||||
bloating the working tree / accidental commits.
|
||||
|
||||
## Cluster
|
||||
- `shao_wm` = SLURM job **1788946** on node **fs-mbz-gpu-538**: 8×GPU (~140 GB each),
|
||||
128 CPU, 768 GB RAM, 5-day limit. As of setup all 8 GPUs were idle.
|
||||
- Attach with `srun --jobid=1788946 --overlap ...`. Never run heavy work on the login
|
||||
node (`fs-mbz-login-big-001`) — it lags the shared connection.
|
||||
- `sqs` is a `.bashrc` alias (`squeue -u hao.zhang`); not available in non-interactive
|
||||
shells — use `squeue -u $USER` directly.
|
||||
@@ -0,0 +1,40 @@
|
||||
# Data Pipeline — Progress Log
|
||||
|
||||
Project: build a `(video, prompt, point-tracks)` dataset to train a motion-controlled
|
||||
real-time interactive video model (bidirectional finetune stage first).
|
||||
|
||||
Branch: `shao/realtime-bidir`. This `data_pipeline/` dir holds the **Stage-0** data
|
||||
generation scripts (everything *upstream* of FastVideo's `preprocess` pipeline).
|
||||
|
||||
## Pipeline shape (where we are)
|
||||
|
||||
```
|
||||
vidprom prompts ──> generate_videos.py ──> videos/*.mp4 + videos2caption.json + merge.txt
|
||||
│
|
||||
extract_tracks.py ──> tracks/*.npz (+ patch points_path into the json)
|
||||
│
|
||||
[TODO] v1_preprocess.py --preprocess_task i2v_points ──> parquet
|
||||
│
|
||||
[TODO] bidir I2V+points trainer
|
||||
```
|
||||
|
||||
## Status
|
||||
|
||||
- [x] Located existing FastVideo preprocess stack + MatrixGame2 control-signal pattern (the template).
|
||||
- [x] Confirmed env = repo `.venv` (`/mnt/weka/home/hao.zhang/shao/FastVideo/.venv/bin/python`).
|
||||
- [x] Confirmed Wan2.2-T2V-A14B weights cached (`~/.cache/huggingface/hub/models--Wan-AI--Wan2.2-T2V-A14B-Diffusers`).
|
||||
- [x] Wrote `generate_videos.py` (Wan2.2-14B T2V → mp4 + manifest).
|
||||
- [x] Wrote `extract_tracks.py` (CoTracker v3 50×50 grid → npz, patch manifest).
|
||||
- [~] Smoke-test launched on fs-mbz-gpu-538 (`run.sh --smoke`, background). Node validated:
|
||||
internet=yes (downloads OK on the node), 8 GPUs idle, `.venv` imports fastvideo + torch
|
||||
2.11.0/CUDA, cotracker3_offline loads (25.4M params). Prompts downloaded: 248,221 (140 MB).
|
||||
- [ ] Scale to 50, then 5k.
|
||||
- [ ] Extend preprocess with an `i2v_points` task (mirror MatrixGame2).
|
||||
|
||||
## Open questions / decisions pending
|
||||
|
||||
- Training model for the bidir stage: 1.3B/480p for plumbing bring-up vs straight to 14B (data gen is fixed at 14B/720p regardless — user confirmed).
|
||||
- Confirm conditioning shape: I2V + points (input image + per-frame points → video). Assumed yes.
|
||||
- Point sampling at train time: store full 2500-pt grid + visibility; sample 1–200 in the trainer.
|
||||
|
||||
See `DECISIONS.md` for rationale, `RUNBOOK.md` for exact commands.
|
||||
@@ -0,0 +1,70 @@
|
||||
# Data Pipeline — Runbook
|
||||
|
||||
All heavy work runs on the `shao_wm` allocation via `srun --overlap`. Do NOT run model
|
||||
code on the login node.
|
||||
|
||||
## 0. Env + node handles
|
||||
```bash
|
||||
PY=/mnt/weka/home/hao.zhang/shao/FastVideo/.venv/bin/python
|
||||
JOBID=$(squeue -u "$USER" -n shao_wm -h -o %i | head -1) # shao_wm allocation id
|
||||
echo "jobid=$JOBID"
|
||||
# live GPU usage on the node (lightweight):
|
||||
srun --jobid="$JOBID" --overlap nvidia-smi \
|
||||
--query-gpu=index,memory.used,memory.total,utilization.gpu --format=csv,noheader
|
||||
```
|
||||
|
||||
## 1. One-time setup (login node — has internet)
|
||||
```bash
|
||||
cd /mnt/weka/home/hao.zhang/shao/FastVideo
|
||||
# prompts (self-forcing / vidprom):
|
||||
[ -f examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt ] || \
|
||||
( cd examples/dataset/vidprom && ./download_dataset.sh )
|
||||
# prefetch CoTracker v3 into the shared torch.hub cache so compute nodes work offline:
|
||||
$PY -c "import torch; torch.hub.load('facebookresearch/co-tracker','cotracker3_offline'); print('cotracker ok')"
|
||||
```
|
||||
|
||||
**Pick free GPUs yourself** from step 0 and pass them via `CUDA_VISIBLE_DEVICES` — the node
|
||||
is shared and the scripts do NOT auto-select. Pinning to busy GPUs OOMs.
|
||||
|
||||
## 2. Generate videos (Wan2.2-T2V-A14B, 720p, 16fps, 81f)
|
||||
```bash
|
||||
GPUS=2,3 # <- two currently-free GPUs from step 0
|
||||
srun --jobid="$JOBID" --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=$GPUS \
|
||||
$PY data_pipeline/generate_videos.py \
|
||||
--prompts examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt \
|
||||
--output-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p \
|
||||
--num-videos 50 --num-gpus 2 # num-gpus must match the device count
|
||||
```
|
||||
- Smoke first: `--num-videos 2`.
|
||||
- Idempotent: re-run resumes from `manifest.jsonl`.
|
||||
|
||||
## 3. Extract point tracks (CoTracker v3, 50×50 grid)
|
||||
```bash
|
||||
srun --jobid="$JOBID" --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=$GPUS \
|
||||
$PY data_pipeline/extract_tracks.py \
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p \
|
||||
--grid-size 50
|
||||
```
|
||||
- Idempotent: skips videos whose `tracks/<stem>.npz` exists.
|
||||
- If 720p OOMs CoTracker, add `--downscale 0.5` (coords are rescaled back to original px).
|
||||
|
||||
## 4. One-shot wrapper (derives num_gpus from CUDA_VISIBLE_DEVICES)
|
||||
```bash
|
||||
CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh # full run (defaults)
|
||||
CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh --smoke # 2 videos
|
||||
```
|
||||
|
||||
## Outputs
|
||||
```
|
||||
<output-dir>/
|
||||
videos/vid_000000.mp4 ...
|
||||
tracks/vid_000000.npz ... # keys: tracks (T,N,2 px), visibility (T,N), grid_size, height, width, fps, num_frames
|
||||
videos2caption.json # FastVideo manifest (+ points_path after step 3)
|
||||
merge.txt
|
||||
manifest.jsonl # incremental resume log (generation)
|
||||
```
|
||||
|
||||
## Gotchas
|
||||
- Check GPU availability before launching (step 0) — the node may not be fully free.
|
||||
- `generate_video(...)` logs a Deprecation warning (use of legacy API) — harmless.
|
||||
- First Wan2.2 run loads ~14B MoE weights; expect a few min of startup before frame 1.
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Download + extract OpenVid / OpenVidHD video shards from HuggingFace.
|
||||
|
||||
Videos are only distributed as ~30-50 GB zip shards (no per-file fetch), so this
|
||||
pulls whole shards, extracts the .mp4s into --videos-dir, and (unless --keep-zip)
|
||||
deletes the zip. Pair with openvid_filter.py (which clips to keep) + extract_tracks_mp.py
|
||||
(resolution/aspect probe + tracking). Idempotent: skips shards already extracted.
|
||||
|
||||
Some parts are HF-split (e.g. OpenVid_part102_partaa/ab) — those must be
|
||||
concatenated before unzip; this handles single-file parts (all OpenVidHD parts).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, zipfile
|
||||
from pathlib import Path
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
REPO = "nkp37/OpenVid-1M"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--shards", nargs="+", required=True,
|
||||
help="repo paths, e.g. OpenVidHD/OpenVidHD_part_1.zip OpenVid_part0.zip")
|
||||
ap.add_argument("--videos-dir", required=True, help="where .mp4s are extracted")
|
||||
ap.add_argument("--zip-dir", default=None, help="scratch for downloaded zips (default: <videos-dir>/../_zips)")
|
||||
ap.add_argument("--keep-zip", action="store_true", help="don't delete the zip after extract")
|
||||
ap.add_argument("--limit", type=int, default=0, help="extract at most N mp4s per shard (0=all; for smoke tests)")
|
||||
ap.add_argument("--only-list", default=None, help="file of basenames; extract only these (e.g. the ≥5s filtered list)")
|
||||
a = ap.parse_args()
|
||||
only = None
|
||||
if a.only_list:
|
||||
only = {l.strip() for l in open(a.only_list) if l.strip()}
|
||||
|
||||
vdir = Path(a.videos_dir); vdir.mkdir(parents=True, exist_ok=True)
|
||||
zdir = Path(a.zip_dir) if a.zip_dir else vdir.parent / "_zips"
|
||||
zdir.mkdir(parents=True, exist_ok=True)
|
||||
marker_dir = vdir.parent / "_extracted"; marker_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for shard in a.shards:
|
||||
done_marker = marker_dir / (Path(shard).name + ".done")
|
||||
if done_marker.exists():
|
||||
print(f"[dl] skip {shard} (already extracted)", flush=True); continue
|
||||
print(f"[dl] downloading {shard} ...", flush=True)
|
||||
zp = hf_hub_download(REPO, shard, repo_type="dataset", local_dir=str(zdir))
|
||||
print(f"[dl] extracting {zp} -> {vdir}", flush=True)
|
||||
n = 0
|
||||
with zipfile.ZipFile(zp) as z:
|
||||
for m in z.namelist():
|
||||
if m.lower().endswith(".mp4"):
|
||||
if only is not None and Path(m).name not in only:
|
||||
continue
|
||||
# flatten: write basename directly into videos-dir
|
||||
(vdir / Path(m).name).write_bytes(z.read(m))
|
||||
n += 1
|
||||
if a.limit and n >= a.limit:
|
||||
break
|
||||
done_marker.write_text(str(n))
|
||||
if not a.keep_zip:
|
||||
Path(zp).unlink(missing_ok=True)
|
||||
print(f"[dl] {shard}: extracted {n} mp4s", flush=True)
|
||||
print("[dl] done", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,100 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Robust OpenVidHD (1080p, 16:9) full downloader for the trackwan bidir dataset.
|
||||
|
||||
OpenVidHD parts are a mix of single-file `.zip` (parts 1-14) and SPLIT parts
|
||||
(`OpenVidHD_part_<i>_part_aa` + `_part_ab`) that must be concatenated before unzip
|
||||
(mirrors the official download_scripts/download_OpenVid.py). This enumerates the HF
|
||||
repo, plans each part, downloads (single or split+cat), extracts only the mp4s in
|
||||
--only-list (the >=5s filtered set), writes them into --videos-dir, and deletes the
|
||||
zip. Idempotent via per-part .done markers, so it resumes after interruption.
|
||||
|
||||
Shard by --shard/--num-shards (e.g. SLURM env) to spread the download across nodes.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, os, re, subprocess, zipfile, collections
|
||||
from pathlib import Path
|
||||
from huggingface_hub import HfApi, hf_hub_download
|
||||
|
||||
REPO = "nkp37/OpenVid-1M"
|
||||
|
||||
|
||||
def plan_parts():
|
||||
"""Return {part_num: [repo_files]} for OpenVidHD (single .zip or split _part_aa/ab)."""
|
||||
files = HfApi().list_repo_files(REPO, repo_type="dataset")
|
||||
parts = collections.defaultdict(list)
|
||||
for f in files:
|
||||
if not f.startswith("OpenVidHD/"):
|
||||
continue
|
||||
m = re.search(r"OpenVidHD_part_(\d+)", f)
|
||||
if m:
|
||||
parts[int(m.group(1))].append(f)
|
||||
return {k: sorted(v) for k, v in parts.items()}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--videos-dir", required=True)
|
||||
ap.add_argument("--zip-dir", required=True, help="scratch for zips (deleted after extract unless --keep-zip)")
|
||||
ap.add_argument("--only-list", required=True, help="basenames to keep (the >=5s filtered list)")
|
||||
ap.add_argument("--keep-zip", action="store_true")
|
||||
ap.add_argument("--num-shards", type=int, default=int(os.environ.get("SLURM_NTASKS", "1")))
|
||||
ap.add_argument("--shard", type=int, default=int(os.environ.get("SLURM_PROCID", "0")))
|
||||
ap.add_argument("--parts", default="all", help="'all' or comma range like 1-14 / 1,2,3")
|
||||
a = ap.parse_args()
|
||||
|
||||
vdir = Path(a.videos_dir); vdir.mkdir(parents=True, exist_ok=True)
|
||||
zdir = Path(a.zip_dir); zdir.mkdir(parents=True, exist_ok=True)
|
||||
mdir = vdir.parent / "_extracted"; mdir.mkdir(parents=True, exist_ok=True)
|
||||
only = {l.strip() for l in open(a.only_list) if l.strip()}
|
||||
|
||||
parts = plan_parts()
|
||||
keys = sorted(parts)
|
||||
if a.parts != "all":
|
||||
want = set()
|
||||
for tok in a.parts.split(","):
|
||||
if "-" in tok:
|
||||
lo, hi = tok.split("-"); want |= set(range(int(lo), int(hi) + 1))
|
||||
else:
|
||||
want.add(int(tok))
|
||||
keys = [k for k in keys if k in want]
|
||||
keys = keys[a.shard::a.num_shards]
|
||||
print(f"[hd-dl shard {a.shard}/{a.num_shards}] {len(keys)} parts: {keys[:6]}{'...' if len(keys)>6 else ''}", flush=True)
|
||||
|
||||
for i in keys:
|
||||
marker = mdir / f"OpenVidHD_part_{i}.done"
|
||||
if marker.exists():
|
||||
continue
|
||||
fs = parts[i]
|
||||
zip_path = zdir / f"OpenVidHD_part_{i}.zip"
|
||||
try:
|
||||
if len(fs) == 1 and fs[0].endswith(".zip"):
|
||||
p = hf_hub_download(REPO, fs[0], repo_type="dataset", local_dir=str(zdir))
|
||||
zip_path = Path(p)
|
||||
else: # split: download _part_aa/_part_ab, concat
|
||||
locals_ = []
|
||||
for f in fs:
|
||||
locals_.append(hf_hub_download(REPO, f, repo_type="dataset", local_dir=str(zdir)))
|
||||
with open(zip_path, "wb") as out:
|
||||
for lp in sorted(locals_):
|
||||
with open(lp, "rb") as src:
|
||||
while (chunk := src.read(1 << 24)):
|
||||
out.write(chunk)
|
||||
for lp in locals_:
|
||||
Path(lp).unlink(missing_ok=True)
|
||||
n = 0
|
||||
with zipfile.ZipFile(zip_path) as z:
|
||||
for m in z.namelist():
|
||||
if m.lower().endswith(".mp4") and Path(m).name in only:
|
||||
(vdir / Path(m).name).write_bytes(z.read(m))
|
||||
n += 1
|
||||
marker.write_text(str(n))
|
||||
if not a.keep_zip:
|
||||
zip_path.unlink(missing_ok=True)
|
||||
print(f"[hd-dl] part_{i}: extracted {n} clips", flush=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[hd-dl] part_{i} FAILED: {repr(e)[:150]}", flush=True)
|
||||
print(f"[hd-dl shard {a.shard}] done", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage-1 (metadata) filter for OpenVid — from the dataset CSV.
|
||||
|
||||
The OpenVid CSV columns are: video, caption, aesthetic score, motion score,
|
||||
temporal consistency score, camera motion, frame, fps, seconds.
|
||||
There is NO resolution / aspect-ratio column, so 720p + 16:9 CANNOT be filtered
|
||||
here — those are enforced downstream by probing (extract_tracks_mp.py
|
||||
--min-height / --aspect-tol). This stage keeps clips that are long enough to yield
|
||||
`num_frames` at the target fps, with optional motion/quality gates (tracking wants
|
||||
motion; static/low-motion clips give degenerate tracks).
|
||||
|
||||
Output: newline-delimited video basenames to keep (feeds the shard extraction +
|
||||
tracking stages). Also writes a captions sidecar (video -> caption) for later use.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse, csv, json, sys
|
||||
|
||||
csv.field_size_limit(sys.maxsize) # captions are long
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--csv", required=True, help="OpenVid-1M.csv or OpenVidHD.csv")
|
||||
ap.add_argument("--out", required=True, help="output: one video basename per line")
|
||||
ap.add_argument("--captions-out", default=None, help="optional JSON {video: caption}")
|
||||
ap.add_argument("--num-frames", type=int, default=121)
|
||||
ap.add_argument("--target-fps", type=int, default=24)
|
||||
ap.add_argument("--min-seconds", type=float, default=None,
|
||||
help="override; default = num_frames/target_fps (+2%% margin)")
|
||||
ap.add_argument("--min-motion", type=float, default=0.0, help="drop clips below this motion score")
|
||||
ap.add_argument("--min-aesthetic", type=float, default=0.0)
|
||||
ap.add_argument("--exclude-static", action="store_true", help="drop camera motion == 'static'")
|
||||
a = ap.parse_args()
|
||||
|
||||
min_secs = a.min_seconds if a.min_seconds is not None else (a.num_frames / a.target_fps) * 1.02
|
||||
kept = 0; total = 0; caps = {}
|
||||
with open(a.csv, newline="") as f, open(a.out, "w") as o:
|
||||
r = csv.DictReader(f)
|
||||
for row in r:
|
||||
total += 1
|
||||
try:
|
||||
secs = float(row["seconds"]); motion = float(row["motion score"]); aes = float(row["aesthetic score"])
|
||||
except (KeyError, ValueError):
|
||||
continue
|
||||
if secs < min_secs: continue
|
||||
if motion < a.min_motion: continue
|
||||
if aes < a.min_aesthetic: continue
|
||||
if a.exclude_static and row.get("camera motion", "").strip().lower() == "static":
|
||||
continue
|
||||
o.write(row["video"] + "\n"); kept += 1
|
||||
if a.captions_out is not None:
|
||||
caps[row["video"]] = row.get("caption", "")
|
||||
if a.captions_out is not None:
|
||||
json.dump(caps, open(a.captions_out, "w"))
|
||||
print(f"[filter] min_seconds={min_secs:.3f} min_motion={a.min_motion} min_aes={a.min_aesthetic} "
|
||||
f"exclude_static={a.exclude_static}")
|
||||
print(f"[filter] kept {kept}/{total} ({100*kept/max(total,1):.1f}%) -> {a.out}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,81 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Precompute the UMT5 embedding of an empty string.
|
||||
|
||||
This is the canonical Wan/T5 null-text convention: feed "" to the text encoder
|
||||
and use the resulting non-zero embedding as the ∅ in classifier-free guidance.
|
||||
Contrast with `torch.zeros_like(text_embedding)`, which is a different tensor
|
||||
distribution that the base model was never asked to handle.
|
||||
|
||||
Saves a dict:
|
||||
{"embedding": [seq_len, dim], "attention_mask": [seq_len]}
|
||||
so the val callback can just index into these when it needs a null text.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/precompute_null_text.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/synth_stage2_paperLR_ckpt400 \
|
||||
--output /mnt/lustre/vlm-s4duan/exports/null_text_umt5.pt
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True,
|
||||
help="Diffusers export dir containing text_encoder/ and tokenizer/")
|
||||
p.add_argument("--output", required=True)
|
||||
p.add_argument("--text-len", type=int, default=256,
|
||||
help="Pad/truncate to this length (matches training text_padding_length)")
|
||||
p.add_argument("--device", default="cuda")
|
||||
args = p.parse_args()
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
try:
|
||||
from transformers import UMT5EncoderModel # type: ignore
|
||||
except ImportError:
|
||||
UMT5EncoderModel = None # type: ignore
|
||||
from transformers import AutoModel
|
||||
|
||||
model_dir = Path(args.model_dir)
|
||||
tok = AutoTokenizer.from_pretrained(model_dir / "tokenizer")
|
||||
if UMT5EncoderModel is not None:
|
||||
try:
|
||||
enc = UMT5EncoderModel.from_pretrained(model_dir / "text_encoder", torch_dtype=torch.bfloat16)
|
||||
except Exception:
|
||||
enc = AutoModel.from_pretrained(model_dir / "text_encoder", torch_dtype=torch.bfloat16)
|
||||
else:
|
||||
enc = AutoModel.from_pretrained(model_dir / "text_encoder", torch_dtype=torch.bfloat16)
|
||||
enc = enc.to(args.device).eval()
|
||||
|
||||
tokens = tok(
|
||||
[""],
|
||||
max_length=args.text_len,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
return_attention_mask=True,
|
||||
)
|
||||
input_ids = tokens["input_ids"].to(args.device)
|
||||
attn_mask = tokens["attention_mask"].to(args.device)
|
||||
print(f"tokenized '' -> {input_ids.shape}, non-pad tokens = {int(attn_mask.sum().item())}")
|
||||
print(f" first 5 token ids: {input_ids[0, :5].tolist()}")
|
||||
|
||||
with torch.no_grad():
|
||||
out = enc(input_ids=input_ids, attention_mask=attn_mask)
|
||||
# take last_hidden_state; strip the batch dim
|
||||
hs = out.last_hidden_state[0].float().cpu() # [seq_len, dim]
|
||||
am = attn_mask[0].cpu()
|
||||
print(f"embedding shape={list(hs.shape)} norm={hs.norm().item():.3f} mean={hs.mean().item():.5f}")
|
||||
print(f"attn_mask shape={list(am.shape)} sum={int(am.sum().item())}")
|
||||
|
||||
Path(args.output).parent.mkdir(parents=True, exist_ok=True)
|
||||
torch.save({"embedding": hs, "attention_mask": am, "text_len": args.text_len}, args.output)
|
||||
print(f"[null-text] saved -> {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+43
@@ -0,0 +1,43 @@
|
||||
#!/bin/bash
|
||||
# Stage-0 data pipeline: generate videos (Wan2.2-14B T2V) -> extract tracks (CoTracker v3).
|
||||
# Runs both steps on the shao_wm allocation via `srun --overlap`.
|
||||
#
|
||||
# You pick the GPUs (the node is shared); pass them explicitly:
|
||||
# CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh # full run (defaults below)
|
||||
# CUDA_VISIBLE_DEVICES=2,3 bash data_pipeline/run.sh --smoke # 2 videos
|
||||
# CUDA_VISIBLE_DEVICES=2,3 NUM_VIDEOS=50 OUTPUT_DIR=... bash data_pipeline/run.sh
|
||||
#
|
||||
# num_gpus is derived from CUDA_VISIBLE_DEVICES. Never run model code on the login node.
|
||||
set -euo pipefail
|
||||
|
||||
REPO=/mnt/weka/home/hao.zhang/shao/FastVideo
|
||||
PY="$REPO/.venv/bin/python"
|
||||
PROMPTS="${PROMPTS:-$REPO/examples/dataset/vidprom/prompts/vidprom_filtered_extended.txt}"
|
||||
OUTPUT_DIR="${OUTPUT_DIR:-/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_t2v_720p}"
|
||||
JOB_NAME="${JOB_NAME:-shao_wm}"
|
||||
NUM_VIDEOS="${NUM_VIDEOS:-50}"
|
||||
GRID_SIZE="${GRID_SIZE:-50}"
|
||||
|
||||
[[ "${1:-}" == "--smoke" ]] && NUM_VIDEOS=2
|
||||
|
||||
: "${CUDA_VISIBLE_DEVICES:?set CUDA_VISIBLE_DEVICES to the GPUs to use, e.g. CUDA_VISIBLE_DEVICES=2,3}"
|
||||
IFS=',' read -ra _gpus <<< "$CUDA_VISIBLE_DEVICES"
|
||||
NUM_GPUS="${NUM_GPUS:-${#_gpus[@]}}"
|
||||
|
||||
JOBID="${JOBID:-$(squeue -u "$USER" -n "$JOB_NAME" -h -o %i | head -1)}"
|
||||
[[ -n "$JOBID" ]] || { echo "[run] ERROR: no running '$JOB_NAME' allocation" >&2; exit 1; }
|
||||
echo "[run] jobid=$JOBID CUDA_VISIBLE_DEVICES=$CUDA_VISIBLE_DEVICES num_gpus=$NUM_GPUS num_videos=$NUM_VIDEOS"
|
||||
|
||||
# /usr/bin/env (absolute: bare `env` may resolve to a broken ~/.local/bin/env) pins the
|
||||
# device list inside the step. Not SLURM --export: its comma parsing splits "2,3".
|
||||
SRUN=(srun --jobid="$JOBID" --overlap --ntasks=1 /usr/bin/env "CUDA_VISIBLE_DEVICES=$CUDA_VISIBLE_DEVICES")
|
||||
|
||||
echo "[run] === generate videos ==="
|
||||
"${SRUN[@]}" "$PY" "$REPO/data_pipeline/generate_videos.py" \
|
||||
--prompts "$PROMPTS" --output-dir "$OUTPUT_DIR" --num-videos "$NUM_VIDEOS" --num-gpus "$NUM_GPUS"
|
||||
|
||||
echo "[run] === extract tracks ==="
|
||||
"${SRUN[@]}" "$PY" "$REPO/data_pipeline/extract_tracks.py" \
|
||||
--data-dir "$OUTPUT_DIR" --grid-size "$GRID_SIZE"
|
||||
|
||||
echo "[run] done -> $OUTPUT_DIR"
|
||||
Executable
+25
@@ -0,0 +1,25 @@
|
||||
#!/usr/bin/env bash
|
||||
# Wan 2.2 5B VAE decode of the 32k HF latent dataset -> 480x832 mp4s.
|
||||
# Sharded across NODES * GPUS * PROCS_PER_GPU workers via SLURM_PROCID.
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${PARQUET_DIR:?}"; : "${OUT_DIR:?}"; : "${VAE_DIR:?}"
|
||||
NODES=${NODES:-8}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
NUM_SHARDS=$(( NODES * TASKS_PER_NODE ))
|
||||
|
||||
mkdir -p "$OUT_DIR" "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $NUM_SHARDS workers"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 12:00:00 -J decode_wansyn32k \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/decode_wansyn32k_%j.out" -e "$WORK/logs/decode_wansyn32k_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && \
|
||||
export HOME=$WORK HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch \
|
||||
MPLCONFIGDIR=$WORK/.mpl TOKENIZERS_PARALLELISM=false && \
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS )) && \
|
||||
python data_pipeline/decode_wansyn32k.py \
|
||||
--parquet-dir $PARQUET_DIR --vae-dir $VAE_DIR --out-dir $OUT_DIR \
|
||||
--num-shards $NUM_SHARDS --shard \$SLURM_PROCID'"
|
||||
@@ -0,0 +1,21 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node OpenVidHD download+extract (>=5s clips only), split-shard aware.
|
||||
# Parts are sharded across tasks (SLURM_PROCID/NTASKS) so the ~8.8TB download runs
|
||||
# in parallel across nodes. CPU/network only — no GPU. Idempotent (resumes).
|
||||
#
|
||||
# Usage:
|
||||
# NODES=8 TASKS_PER_NODE=3 bash data_pipeline/run_download_hd_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
NODES=${NODES:-8}; TASKS_PER_NODE=${TASKS_PER_NODE:-3} # concurrent HF downloads per node
|
||||
OUT=${OUT:-$WORK/openvid_1m} # goal folder
|
||||
FILTER=${FILTER:-$WORK/openvid/OpenVidHD_filtered.txt}
|
||||
mkdir -p "$OUT/videos" "$OUT/_zips" "$WORK/logs"
|
||||
echo "download OpenVidHD -> $OUT/videos ($NODES nodes x $TASKS_PER_NODE = $((NODES*TASKS_PER_NODE)) parallel downloads)"
|
||||
|
||||
sbatch -N "$NODES" --ntasks-per-node=$TASKS_PER_NODE --cpus-per-task=16 --mem=0 --exclusive \
|
||||
-t 48:00:00 -J openvid_dl --chdir="$WORK/FastVideo" \
|
||||
-o "$WORK/logs/openvid_dl_%j.out" -e "$WORK/logs/openvid_dl_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo bash -lc 'source .venv/bin/activate && export HF_HOME=$WORK/.hf && \
|
||||
python data_pipeline/openvid_download_hd.py \
|
||||
--videos-dir $OUT/videos --zip-dir $OUT/_zips --only-list $FILTER'"
|
||||
@@ -0,0 +1,116 @@
|
||||
#!/usr/bin/env bash
|
||||
# 720p (1280x720) i2v_track PREPROCESS of OpenVid-1M, run INSIDE a held allocation.
|
||||
#
|
||||
# Differs from run_preprocess_track_slurm.sh (which sbatch's its own job array):
|
||||
# * runs via `srun --overlap --jobid=$JOBID` inside an existing held alloc, so it
|
||||
# shares the nodes with (and does not disturb) the training hold,
|
||||
# * ONE srun step PER NODE instead of one gang step, so the k8s operator cordoning
|
||||
# a node kills only that node's step; a supervisor pass relaunches it,
|
||||
# * shards are CLAIMED with an atomic `mkdir` under <combined>/.claims/, so any
|
||||
# worker can steal any not-yet-claimed shard -> a dead node's remaining work is
|
||||
# picked up by the survivors instead of being silently skipped.
|
||||
# * per-shard `.done` markers make every (re)start skip finished work.
|
||||
#
|
||||
# Usage: JOBID=881 bash data_pipeline/run_preprocess_720p_held.sh
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
DATA_DIR=${DATA_DIR:-$WORK/openvid_1m}
|
||||
MODEL=${MODEL:-$WORK/models/trackwan_1.3b_i2v_d64_nobias_init} # same encoders as the 480p run
|
||||
CLIPS_DIR=${CLIPS_DIR:-$DATA_DIR/clips}
|
||||
MANIFEST=${MANIFEST:-$DATA_DIR/videos2caption.json}
|
||||
COMBINED=${COMBINED:-$DATA_DIR/combined_parquet_dataset_720p}
|
||||
SHARDS_DIR=${SHARDS_DIR:-$DATA_DIR/preprocess_shards} # reuse the 480p partition
|
||||
NUM_SHARDS=${NUM_SHARDS:-1500}
|
||||
JOBID=${JOBID:-881}
|
||||
GPUS=${GPUS:-4}
|
||||
MAX_H=${MAX_H:-720}; MAX_W=${MAX_W:-1280}; NUM_FRAMES=${NUM_FRAMES:-121}
|
||||
TRAIN_FPS=${TRAIN_FPS:-24}; NUM_LATENT_T=${NUM_LATENT_T:-31}
|
||||
BATCH=${BATCH:-1} # must be 1: pad-free T5 tokenizer can't batch variable-length captions
|
||||
VAE_PREC=${VAE_PREC:-fp32} # match the 480p run
|
||||
NW=${NW:-8} # decode-prefetch workers
|
||||
MAXPASS=${MAXPASS:-8}
|
||||
LOGDIR=$WORK/logs/prep720
|
||||
CLAIMS=$COMBINED/.claims
|
||||
mkdir -p "$LOGDIR" "$COMBINED" "$CLAIMS"
|
||||
|
||||
# By default use every node in the held alloc; set NODELIST to run on a SUBSET
|
||||
# (e.g. NODELIST=hpc-rack-2-[0-2,4] to leave the rest free for training).
|
||||
NODELIST=${NODELIST:-$(squeue -h -j "$JOBID" -o %N)}
|
||||
mapfile -t NODES < <(scontrol show hostnames "$NODELIST")
|
||||
NNODES=${#NODES[@]}
|
||||
WORKERS=$(( NNODES * GPUS ))
|
||||
echo "[sup] jobid=$JOBID nodes=$NNODES (${NODES[*]}) workers=$WORKERS"
|
||||
echo "[sup] geometry ${MAX_W}x${MAX_H} ${NUM_FRAMES}f -> latent [16,$NUM_LATENT_T,$((MAX_H/8)),$((MAX_W/8))]"
|
||||
echo "[sup] out=$COMBINED shards=$SHARDS_DIR ($NUM_SHARDS)"
|
||||
|
||||
count_done() { find "$COMBINED" -maxdepth 2 -name .done 2>/dev/null | wc -l; }
|
||||
|
||||
for PASS in $(seq 1 "$MAXPASS"); do
|
||||
DONE=$(count_done)
|
||||
echo "[sup] pass $PASS: $DONE/$NUM_SHARDS shards done"
|
||||
[ "$DONE" -ge "$NUM_SHARDS" ] && { echo "[sup] all shards done"; break; }
|
||||
|
||||
# No workers are alive at this point, so any claim without a .done is stale
|
||||
# (its worker died) -> release it so this pass can re-run that shard.
|
||||
RELEASED=0
|
||||
for C in "$CLAIMS"/shard_*; do
|
||||
[ -d "$C" ] || continue
|
||||
S=$(basename "$C")
|
||||
[ -f "$COMBINED/$S/.done" ] || { rmdir "$C" 2>/dev/null && RELEASED=$((RELEASED+1)); }
|
||||
done
|
||||
[ "$RELEASED" -gt 0 ] && echo "[sup] released $RELEASED stale claim(s)"
|
||||
|
||||
PIDS=()
|
||||
for I in $(seq 0 $(( NNODES - 1 ))); do
|
||||
NODE=${NODES[$I]}
|
||||
srun --overlap --jobid="$JOBID" --nodelist="$NODE" --nodes=1 --ntasks="$GPUS" \
|
||||
--ntasks-per-node="$GPUS" --gres=gpu:"$GPUS" --cpus-per-task=$(( 128 / GPUS )) \
|
||||
--chdir="$WORK/FastVideo" \
|
||||
-o "$LOGDIR/pass${PASS}_${NODE}.out" -e "$LOGDIR/pass${PASS}_${NODE}.out" \
|
||||
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 MPLCONFIGDIR=$WORK/.mpl \
|
||||
PYTHONPATH=$WORK/FastVideo TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0
|
||||
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=\$(( $I * $GPUS + SLURM_LOCALID ))
|
||||
echo \"[worker \$W] host=\$(hostname) gpu=\$CUDA_VISIBLE_DEVICES\"
|
||||
# Start spread out across the shard list, then walk the whole list claiming
|
||||
# whatever is free -> automatic work-stealing, no idle workers at the tail.
|
||||
OFF=\$(( W * $NUM_SHARDS / $WORKERS ))
|
||||
for K in \$(seq 0 $(( NUM_SHARDS - 1 ))); do
|
||||
IDX=\$(( (OFF + K) % $NUM_SHARDS ))
|
||||
SHARD=\$(printf shard_%05d \$IDX)
|
||||
ODIR=$COMBINED/\$SHARD
|
||||
[ -f \"\$ODIR/.done\" ] && continue
|
||||
mkdir \"$CLAIMS/\$SHARD\" 2>/dev/null || continue # atomic claim; someone else has it
|
||||
rm -rf \"\$ODIR\"
|
||||
T0=\$(date +%s)
|
||||
if python fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL --preprocess_task i2v_track \
|
||||
--data_merge_path $SHARDS_DIR/\$SHARD/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 \
|
||||
--video_length_tolerance_range 5 --seed 1000 > \"$LOGDIR/\$SHARD.log\" 2>&1; then
|
||||
touch \"\$ODIR/.done\"
|
||||
rm -f \"$LOGDIR/\$SHARD.log\" # keep only failures' logs
|
||||
echo \"[worker \$W] \$SHARD done in \$(( \$(date +%s) - T0 ))s\"
|
||||
else
|
||||
echo \"[worker \$W] \$SHARD FAILED (\$(( \$(date +%s) - T0 ))s), log $LOGDIR/\$SHARD.log -> releasing claim\"
|
||||
tail -5 \"$LOGDIR/\$SHARD.log\" | sed \"s/^/[worker \$W] /\"
|
||||
rmdir \"$CLAIMS/\$SHARD\" 2>/dev/null
|
||||
fi
|
||||
done
|
||||
echo \"[worker \$W] no shards left to claim\"" &
|
||||
PIDS+=($!)
|
||||
done
|
||||
echo "[sup] pass $PASS: launched ${#PIDS[@]} node steps; waiting"
|
||||
wait "${PIDS[@]}"
|
||||
echo "[sup] pass $PASS finished: $(count_done)/$NUM_SHARDS done"
|
||||
done
|
||||
|
||||
echo "[sup] FINAL: $(count_done)/$NUM_SHARDS shards done"
|
||||
Executable
+88
@@ -0,0 +1,88 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node data-parallel i2v_track PREPROCESS over OpenVid-1M.
|
||||
# v1_preprocess.py asserts WORLD_SIZE==1 (one GPU per process), so scale-out is
|
||||
# N nodes x 4 GPU x 1 proc/GPU = 4N independent single-GPU processes, each over a
|
||||
# SHARD of the clips, writing parquet into <combined>/shard_<idx>/. The trainer
|
||||
# reads <combined> as ONE dataset (get_parquet_files_and_length os.walks recursively).
|
||||
#
|
||||
# Usage:
|
||||
# NODES=16 bash data_pipeline/run_preprocess_track_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
DATA_DIR=${DATA_DIR:-$WORK/openvid_1m}
|
||||
MODEL=${MODEL:-$WORK/models/trackwan_1.3b_i2v_d64_nobias_init}
|
||||
CLIPS_DIR=${CLIPS_DIR:-$DATA_DIR/clips}
|
||||
MANIFEST=${MANIFEST:-$DATA_DIR/videos2caption.json}
|
||||
COMBINED=${COMBINED:-$DATA_DIR/combined_parquet_dataset} # trainer points here
|
||||
SHARDS_DIR=${SHARDS_DIR:-$DATA_DIR/preprocess_shards}
|
||||
NODES=${NODES:-14}; GPUS=${GPUS:-4}
|
||||
PARTITION=${PARTITION:-all}
|
||||
CONC=${CONC:-$(( NODES * GPUS ))} # max concurrent shards (= total GPUs)
|
||||
# MANY small shards (not one-per-GPU) so a node glitch only loses its in-flight
|
||||
# ~CONC/NODES shards, and .done markers make a resubmit skip finished work.
|
||||
NUM_SHARDS=${NUM_SHARDS:-1024} # ~253 clips/shard (fine-grained .done resume)
|
||||
# Training-target geometry (must match trainer):
|
||||
MAX_H=${MAX_H:-480}; MAX_W=${MAX_W:-832}; NUM_FRAMES=${NUM_FRAMES:-121}
|
||||
TRAIN_FPS=${TRAIN_FPS:-24}; NUM_LATENT_T=${NUM_LATENT_T:-31}
|
||||
# batch_size MUST be 1: the T5 tokenizer_kwargs pad-free config (configs/models/
|
||||
# encoders/t5.py) can't batch variable-length captions with return_tensors="pt".
|
||||
BATCH=${BATCH:-1}
|
||||
VAE_PREC=${VAE_PREC:-fp32} # bf16 gives NO speedup here (decode-bound, not VAE-compute-bound)
|
||||
NW=${NW:-8} # decode-prefetch workers. SAFE now: patched VideoCaptionMergedDataset.__iter__
|
||||
# to shard by get_worker_info() (else num_workers>0 duplicates). Validated:
|
||||
# 16 clips -> 16 rows, latents byte-identical to nw=0. Overlaps decode w/ GPU encode.
|
||||
TIME=${TIME:-24:00:00}
|
||||
CPUS_PER_TASK=$(( 128 / GPUS )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
mkdir -p "$WORK/logs" "$COMBINED"
|
||||
|
||||
# 1) Build/refresh shard manifests (idempotent; cheap).
|
||||
source "$WORK/FastVideo/.venv/bin/activate"
|
||||
python "$WORK/FastVideo/data_pipeline/split_manifest_shards.py" \
|
||||
--manifest "$MANIFEST" --clips-dir "$CLIPS_DIR" \
|
||||
--out-dir "$SHARDS_DIR" --num-shards "$NUM_SHARDS"
|
||||
|
||||
echo "array of $NUM_SHARDS shards, up to $CONC concurrent (~$((NODES)) nodes x $GPUS GPU); out=$COMBINED"
|
||||
|
||||
# 2) Launch. This cluster gives a GPU job the WHOLE node (all 4 GPU + 144 CPU), so we
|
||||
# use -N nodes x ntasks-per-node=4 (= 4N=56 workers) rather than a job array (which
|
||||
# would get 1 whole node per task = only 14 concurrent). Each worker loops over its
|
||||
# slice shards[PROCID::NW], running a FRESH python per SMALL (~260-clip) shard:
|
||||
# the fresh process bounds the pipeline's ~0.75 GB/clip RSS growth (VAE feature cache)
|
||||
# -> 4 workers x ~214 GB peak = ~856 GB < 979 GB/node. Per-shard .done makes resubmit
|
||||
# skip finished work; --no-requeue avoids a destructive auto-restart on node glitches.
|
||||
WORKERS=$(( NODES * GPUS ))
|
||||
# CORDON-PROOF: the k8s operator cordons worker pods mid-run (node-pool scale-down);
|
||||
# a gang -N$NODES job dies entirely if ANY one node is cordoned. So submit a JOB ARRAY
|
||||
# of single-node tasks (--array=0..NODES-1, each -N1): a cordon kills+requeues only ONE
|
||||
# array task; the other nodes keep running. Global worker idx = ARRAY_TASK_ID*GPUS+LOCALID
|
||||
# over WORKERS=NODES*GPUS; per-shard .done makes every (re)start skip finished work.
|
||||
WORKERS=$(( NODES * GPUS ))
|
||||
sbatch --array=0-$(( NODES - 1 )) -N1 --gres=gpu:$GPUS --ntasks-per-node=$GPUS --exclusive --mem=0 \
|
||||
-t "$TIME" --requeue -p "$PARTITION" -J preprocess_track \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/preprocess_track_%A_%a.out" -e "$WORK/logs/preprocess_track_%A_%a.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo bash -lc '
|
||||
set -uo pipefail
|
||||
source .venv/bin/activate
|
||||
export HOME=$WORK TRITON_CACHE_DIR=$WORK/.triton XDG_CACHE_HOME=$WORK/.cache \
|
||||
HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch TOKENIZERS_PARALLELISM=false
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS ))
|
||||
export WORLD_SIZE=1 RANK=0 LOCAL_RANK=0 MASTER_ADDR=127.0.0.1
|
||||
export MASTER_PORT=\$(( 29500 + SLURM_LOCALID ))
|
||||
W=\$(( SLURM_ARRAY_TASK_ID * $GPUS + SLURM_LOCALID )); NW=$WORKERS
|
||||
echo \"[worker \$W/\$NW arr=\$SLURM_ARRAY_TASK_ID] host=\$(hostname) CVD=\$CUDA_VISIBLE_DEVICES\"
|
||||
for IDX in \$(seq \$W \$NW $(( NUM_SHARDS - 1 ))); do
|
||||
SHARD=\$(printf shard_%05d \$IDX); SDIR=$SHARDS_DIR/\$SHARD; ODIR=$COMBINED/\$SHARD
|
||||
[ -f \"\$ODIR/.done\" ] && continue
|
||||
rm -rf \"\$ODIR\"
|
||||
if python fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL --preprocess_task i2v_track \
|
||||
--data_merge_path \$SDIR/merge.txt --output_dir \$ODIR \
|
||||
--max_height $MAX_H --max_width $MAX_W --num_frames $NUM_FRAMES \
|
||||
--train_fps $TRAIN_FPS --num_latent_t $NUM_LATENT_T --vae_precision $VAE_PREC \
|
||||
--preprocess_video_batch_size $BATCH --dataloader_num_workers $NW; then
|
||||
touch \"\$ODIR/.done\"; echo \"[worker \$W] \$SHARD done\"
|
||||
else
|
||||
echo \"[worker \$W] \$SHARD FAILED (rc=\$?), skipping\"
|
||||
fi
|
||||
done
|
||||
echo \"[worker \$W] all shards done\"'"
|
||||
@@ -0,0 +1,29 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node data-parallel frame-0 segmentation (SAM2.1-b+) over OpenVid-1M.
|
||||
# Bare node (torch-only ultralytics needs just the driver + Lustre venv, no container).
|
||||
# Each Slurm task = one SAM worker; PROCS_PER_GPU tasks share each GPU.
|
||||
#
|
||||
# Usage:
|
||||
# DATA_DIR=/mnt/lustre/vlm-s4duan/openvid_1m NODES=6 PROCS_PER_GPU=2 \
|
||||
# bash data_pipeline/run_seg_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
DATA_DIR=${DATA_DIR:-$WORK/openvid_1m}
|
||||
NODES=${NODES:-1}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
MODEL=${MODEL:-sam2.1_b.pt}; VIDEOS_SUBDIR=${VIDEOS_SUBDIR:-clips}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
TIME=${TIME:-12:00:00}
|
||||
mkdir -p "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $((NODES*TASKS_PER_NODE)) workers, model=$MODEL"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t "$TIME" -J seg_sam_dp \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/seg_sam_dp_%j.out" -e "$WORK/logs/seg_sam_dp_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && \
|
||||
export TORCH_HOME=$WORK/.torch HF_HOME=$WORK/.hf MPLCONFIGDIR=$WORK/.mpl \
|
||||
YOLO_CONFIG_DIR=$WORK/.ultralytics TOKENIZERS_PARALLELISM=false \
|
||||
PYTHONPATH=$WORK/FastVideo:$WORK/FastVideo/data_pipeline && \
|
||||
python data_pipeline/segment_tracks_mp.py \
|
||||
--data-dir $DATA_DIR --videos-subdir $VIDEOS_SUBDIR --model $MODEL --gpus-per-node $GPUS'"
|
||||
Executable
+25
@@ -0,0 +1,25 @@
|
||||
#!/usr/bin/env bash
|
||||
# SAM 2.1-b+ segmentation of frame-0 for our synth mp4s. Adds object_ids +
|
||||
# track_weights to each tracks .npz. Idempotent (skips npz already labeled).
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${DATA_DIR:?}"
|
||||
NODES=${NODES:-4}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
NUM_SHARDS=$(( NODES * TASKS_PER_NODE ))
|
||||
|
||||
mkdir -p "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $NUM_SHARDS workers"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 4:00:00 -J segment_synth \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/segment_synth_%j.out" -e "$WORK/logs/segment_synth_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && \
|
||||
export TORCH_HOME=$WORK/.torch HF_HOME=$WORK/.hf HOME=$WORK MPLCONFIGDIR=$WORK/.mpl \
|
||||
YOLO_CONFIG_DIR=$WORK/.ultralytics TOKENIZERS_PARALLELISM=false && \
|
||||
export CUDA_VISIBLE_DEVICES=\$(( SLURM_LOCALID % $GPUS )) && \
|
||||
python data_pipeline/segment_tracks.py \
|
||||
--data-dir $DATA_DIR \
|
||||
--num-shards $NUM_SHARDS --shard \$SLURM_PROCID'"
|
||||
@@ -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"
|
||||
@@ -0,0 +1,32 @@
|
||||
#!/usr/bin/env bash
|
||||
# Multi-node / multi-GPU data-parallel CoTracker on this Slinky GB200 cluster.
|
||||
# Each Slurm task = one CoTracker worker (B=1); PROCS_PER_GPU tasks share each GPU.
|
||||
# Runs inside the enroot `fvbuild` container (needed for the venv's CUDA runtime).
|
||||
#
|
||||
# Usage:
|
||||
# VIDEO_LIST=/mnt/lustre/vlm-s4duan/openvid/videos.txt \
|
||||
# OUT_DIR=/mnt/lustre/vlm-s4duan/openvid/tracks \
|
||||
# NODES=4 PROCS_PER_GPU=2 FPS=24 NUM_FRAMES=121 GRID=50 \
|
||||
# bash data_pipeline/run_tracks_slurm.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${VIDEO_LIST:?set VIDEO_LIST}"; : "${OUT_DIR:?set OUT_DIR}"
|
||||
NODES=${NODES:-1}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
FPS=${FPS:-24}; NUM_FRAMES=${NUM_FRAMES:-121}; GRID=${GRID:-50}
|
||||
CLIPS_ARG=""; [ -n "${CLIPS_DIR:-}" ] && CLIPS_ARG="--clips-dir $CLIPS_DIR"
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
|
||||
mkdir -p "$OUT_DIR" "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $((NODES*TASKS_PER_NODE)) workers"
|
||||
|
||||
# Bare-node: torch (self-contained cu128 wheels) + CoTracker need only the driver +
|
||||
# the Lustre venv — NO container (avoids a per-node image pull).
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 24:00:00 -J cotracker_dp \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/cotracker_dp_%j.out" -e "$WORK/logs/cotracker_dp_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && export TORCH_HOME=$WORK/.torch TOKENIZERS_PARALLELISM=false && \
|
||||
python data_pipeline/extract_tracks_mp.py \
|
||||
--video-list $VIDEO_LIST --out-dir $OUT_DIR $CLIPS_ARG \
|
||||
--gpus-per-node $GPUS --fps $FPS --num-frames $NUM_FRAMES --grid-size $GRID'"
|
||||
Executable
+25
@@ -0,0 +1,25 @@
|
||||
#!/usr/bin/env bash
|
||||
# CoTracker extraction for our synth mp4s. Same pattern as run_tracks_slurm.sh but
|
||||
# with the aspect/lowres filters disabled (synth videos are exactly 720x1280) and
|
||||
# H/W set to our training resolution 480x832.
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
: "${VIDEO_LIST:?}"; : "${OUT_DIR:?}"
|
||||
NODES=${NODES:-4}; GPUS=${GPUS:-4}; PROCS_PER_GPU=${PROCS_PER_GPU:-2}
|
||||
FPS=${FPS:-24}; NUM_FRAMES=${NUM_FRAMES:-121}; GRID=${GRID:-50}
|
||||
HEIGHT=${HEIGHT:-480}; WIDTH=${WIDTH:-832}
|
||||
TASKS_PER_NODE=$(( GPUS * PROCS_PER_GPU ))
|
||||
CPUS_PER_TASK=$(( 128 / TASKS_PER_NODE )); [ "$CPUS_PER_TASK" -lt 1 ] && CPUS_PER_TASK=1
|
||||
|
||||
mkdir -p "$OUT_DIR" "$WORK/logs"
|
||||
echo "nodes=$NODES gpus/node=$GPUS procs/gpu=$PROCS_PER_GPU -> $((NODES*TASKS_PER_NODE)) workers at ${HEIGHT}x${WIDTH}"
|
||||
|
||||
sbatch -N "$NODES" --gres=gpu:$GPUS --ntasks-per-node=$TASKS_PER_NODE --exclusive \
|
||||
--cpus-per-task=$CPUS_PER_TASK --mem=0 -t 6:00:00 -J cotracker_synth \
|
||||
--chdir="$WORK/FastVideo" -o "$WORK/logs/cotracker_synth_%j.out" -e "$WORK/logs/cotracker_synth_%j.out" \
|
||||
--wrap "srun --chdir=$WORK/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && export TORCH_HOME=$WORK/.torch TOKENIZERS_PARALLELISM=false && \
|
||||
python data_pipeline/extract_tracks_mp.py \
|
||||
--video-list $VIDEO_LIST --out-dir $OUT_DIR \
|
||||
--gpus-per-node $GPUS --fps $FPS --num-frames $NUM_FRAMES --grid-size $GRID \
|
||||
--height $HEIGHT --width $WIDTH --min-height 0 --aspect-tol 10.0'"
|
||||
@@ -0,0 +1,213 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Visualize the training-time track sampler on a preprocessed dataset.
|
||||
|
||||
For each clip it shows, straight from the stored ``tracks.npz`` (tracks + visibility +
|
||||
``object_ids`` + ``track_weights``):
|
||||
(1) the low-rank informativeness HEATMAP over the 50x50 grid on frame 0,
|
||||
(2) the points the SAMPLER KEEPS (exact port of WanTrack ``_augment_tracks``:
|
||||
>=1 pt per SAM object + low-rank-weighted draw (1-uniform_frac) + uniform draw),
|
||||
(3) an overlay video of ONLY the kept tracks over the source video, colored by low-rank
|
||||
weight (blue = generic, red = informative), so you can see if good traces are sampled.
|
||||
|
||||
CPU only (no GPU / no model needed).
|
||||
|
||||
.venv/bin/python data_pipeline/sampling_viz.py render --data-dir <dataset root> --out <dir> --limit 8
|
||||
.venv/bin/python data_pipeline/sampling_viz.py serve --viz-dir <dir> --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from track_informativeness import _draw_overlay, _read_frames, _norm # noqa: E402
|
||||
|
||||
|
||||
def sample_kept(tracks: np.ndarray,
|
||||
vis: np.ndarray,
|
||||
oid: np.ndarray | None,
|
||||
tw: np.ndarray | None,
|
||||
k: int,
|
||||
uniform_frac: float = 0.3,
|
||||
diversity: float = 1.0,
|
||||
seed: int = 0) -> np.ndarray:
|
||||
"""Exact numpy port of WanTrack._augment_tracks sampling -> kept bool [N]."""
|
||||
rng = np.random.default_rng(seed)
|
||||
N = tracks.shape[1]
|
||||
valid = (vis[0] > 0.5).astype(np.float64) # only frame-0-visible are valid queries
|
||||
if tw is not None and tw.size == N:
|
||||
w = (0.05 + diversity * tw.astype(np.float64)) * valid
|
||||
else:
|
||||
disp = np.sqrt(((tracks - tracks[0:1])**2).sum(-1))
|
||||
motion = (disp * (vis > 0.5)).max(0)
|
||||
w = (1.0 + diversity * (motion / (motion.mean() + 1e-6))) * valid
|
||||
if w.sum() <= 0:
|
||||
w = valid.copy()
|
||||
if w.sum() <= 0:
|
||||
w = np.ones(N)
|
||||
keep = np.zeros(N, bool)
|
||||
# (a) object coverage: one weighted pick per present segment
|
||||
if oid is not None:
|
||||
for o in np.unique(oid):
|
||||
if int(o) < 0:
|
||||
continue
|
||||
idx = np.where(oid == o)[0]
|
||||
ww = w[idx]
|
||||
ww = np.ones_like(ww) if ww.sum() <= 0 else ww
|
||||
keep[rng.choice(idx, p=ww / ww.sum())] = True
|
||||
# (b) low-rank-weighted draw for (1 - uniform_frac) of the remaining budget
|
||||
rem = max(0, k - int(keep.sum()))
|
||||
n_weighted = rem - int(round(rem * uniform_frac))
|
||||
pool = np.where((~keep) & (w > 0))[0]
|
||||
nw = min(n_weighted, pool.size)
|
||||
if nw > 0:
|
||||
pw = w[pool] / w[pool].sum()
|
||||
keep[rng.choice(pool, size=nw, replace=False, p=pw)] = True
|
||||
# (c) uniform draw to fill up to k (also backfills weighted shortfall)
|
||||
pool2 = np.where((~keep) & (valid > 0))[0]
|
||||
nu = min(max(0, k - int(keep.sum())), pool2.size)
|
||||
if nu > 0:
|
||||
keep[rng.choice(pool2, size=nu, replace=False)] = True
|
||||
return keep
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import imageio.v2 as imageio
|
||||
turbo = matplotlib.colormaps["turbo"]
|
||||
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
entries = []
|
||||
for i, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
vpath = str(data / "videos" / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = _read_frames(vpath)
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]]
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]]
|
||||
oid = d["object_ids"].astype(np.int64) if "object_ids" in d else None
|
||||
tw = d["track_weights"].astype(np.float32) if "track_weights" in d else None
|
||||
N = tracks.shape[1]
|
||||
k = int(args.k)
|
||||
keep = sample_kept(tracks, vis, oid, tw, k, args.uniform_frac, args.diversity, args.seed)
|
||||
twn = _norm(tw) if tw is not None else np.zeros(N)
|
||||
x, y = tracks[0, :, 0], tracks[0, :, 1]
|
||||
|
||||
# ---- figure: (1) lowrank heatmap (2) kept vs dropped
|
||||
fig, axes = plt.subplots(1, 2, figsize=(17, 5.6))
|
||||
axes[0].imshow(frames[0])
|
||||
sc = axes[0].scatter(x, y, c=twn, cmap="turbo", s=12, vmin=0, vmax=1)
|
||||
axes[0].set_title(f"low-rank informativeness heatmap (all {N} pts)", fontsize=11)
|
||||
axes[0].axis("off")
|
||||
fig.colorbar(sc, ax=axes[0], fraction=0.03)
|
||||
axes[1].imshow(frames[0])
|
||||
axes[1].scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.5) # dropped
|
||||
axes[1].scatter(x[keep],
|
||||
y[keep],
|
||||
c=twn[keep],
|
||||
cmap="turbo",
|
||||
s=16,
|
||||
vmin=0,
|
||||
vmax=1,
|
||||
edgecolors="white",
|
||||
linewidths=0.3) # kept, colored by weight
|
||||
n_obj = len(np.unique(oid[oid >= 0])) if oid is not None else 0
|
||||
n_cov = len(np.unique(oid[keep][oid[keep] >= 0])) if oid is not None else 0
|
||||
axes[1].set_title(f"kept by sampler: {int(keep.sum())}/{N} (objects covered {n_cov}/{n_obj})", fontsize=11)
|
||||
axes[1].axis("off")
|
||||
fig.suptitle(f"{stem} | {(it['cap'][0] if isinstance(it.get('cap'), list) else '')[:80]}", fontsize=10)
|
||||
fig.tight_layout()
|
||||
heat = f"{stem}_sampling.png"
|
||||
fig.savefig(str(out / heat), dpi=80, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
|
||||
# ---- overlay video: ONLY kept tracks, colored by low-rank weight
|
||||
kidx = np.where(keep)[0]
|
||||
# subsample for legibility if many kept
|
||||
if kidx.size > 700:
|
||||
kidx = kidx[np.linspace(0, kidx.size - 1, 700).round().astype(int)]
|
||||
pcols = (np.array([turbo(v)[:3] for v in twn[kidx]]) * 255).astype(np.uint8)
|
||||
ov = _draw_overlay(frames, tracks[:, kidx].copy(), vis[:, kidx], pcols, 14, 2, 0.5)
|
||||
vid = f"{stem}_kept.mp4"
|
||||
imageio.mimsave(str(out / vid), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
|
||||
stats = {
|
||||
"id": stem,
|
||||
"n_points": int(N),
|
||||
"k_sampled": k,
|
||||
"kept": int(keep.sum()),
|
||||
"objects_covered": f"{n_cov}/{n_obj}",
|
||||
"mean_weight_kept": round(float(twn[keep].mean()), 3) if keep.any() else 0.0,
|
||||
"mean_weight_dropped": round(float(twn[~keep].mean()), 3) if (~keep).any() else 0.0,
|
||||
"frac_kept_informative(>0.5)": round(float((twn[keep] > 0.5).mean()), 3) if keep.any() else 0.0,
|
||||
}
|
||||
entries.append({"id": stem, "heat": heat, "video": vid, "stats": stats})
|
||||
print(
|
||||
f"[samp] [{i}/{len(items)}] {stem}: kept {int(keep.sum())}/{N}, cov {n_cov}/{n_obj}, "
|
||||
f"w_kept={stats['mean_weight_kept']} vs w_drop={stats['mean_weight_dropped']}",
|
||||
flush=True)
|
||||
|
||||
(out / "manifest.json").write_text(json.dumps(entries, indent=2))
|
||||
print(f"[samp] rendered {len(entries)} clips -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
entries = json.loads((viz / "manifest.json").read_text())
|
||||
gal = [(str(viz / e["heat"]), e["id"]) for e in entries]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = entries[evt.index]
|
||||
return str(viz / e["heat"]), str(viz / e["video"]), json.dumps(e["stats"], indent=2)
|
||||
|
||||
with gr.Blocks(title="Track sampling viz") as demo:
|
||||
gr.Markdown("### Training-time track sampler: low-rank heatmap + kept points + kept-track overlay\n"
|
||||
"Left panel: low-rank informativeness heatmap (bright = informative). Right panel: the points "
|
||||
"the sampler **keeps** (colored by weight, gray = dropped) — should cover every object and "
|
||||
"favor bright/informative points while keeping some uniform. Click a clip -> the overlay video "
|
||||
"of ONLY the kept tracks (blue = generic, red = informative). Good sampling = kept traces follow "
|
||||
"the moving hands/objects. `stats.mean_weight_kept` should exceed `mean_weight_dropped`.")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gal, columns=2, height=620, label="clips (click)")
|
||||
with gr.Column():
|
||||
heat = gr.Image(label="heatmap (left) + kept points (right)")
|
||||
vid = gr.Video(label="kept tracks over video (colored by low-rank weight)")
|
||||
meta = gr.Code(label="stats", language="json")
|
||||
g.select(show, None, [heat, vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--k", type=int, default=1500, help="tracks to keep (training samples U[1000,2500])")
|
||||
r.add_argument("--uniform-frac", type=float, default=0.3)
|
||||
r.add_argument("--diversity", type=float, default=1.0)
|
||||
r.add_argument("--seed", type=int, default=0)
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7889)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,377 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Compare segmentation-model *variants* (not just FastSAM configs) for frame-0
|
||||
object segmentation + sparse track sampling, on real openvid clips.
|
||||
|
||||
``render`` (GPU) runs each model in no-prompt "everything" mode on frame-0 of each
|
||||
clip, labels the CoTracker grid points by object (smallest containing mask wins),
|
||||
runs the 1-per-object + extras sparse sampler, and writes per (model, clip):
|
||||
- <..>_masks.png frame-0 with colored SAM masks
|
||||
- <..>_sparse.png frame-0 with kept sparse tracks (colored by object) vs dropped (gray)
|
||||
- <..>_tracks.mp4 kept sparse tracks over the whole clip
|
||||
plus a per-entry .npz (oid/tw/vis0/xy0) so the dashboard can re-sample live.
|
||||
|
||||
``serve`` (CPU, login node) is a gradio dashboard: pick a clip, see every model's
|
||||
masks / sparse-sampling side by side, tweak num_sampled live, click to inspect.
|
||||
|
||||
aarch64-safe: PyAV decode, headless cv2, no decord.
|
||||
|
||||
# render on a GPU node (piggyback job 365)
|
||||
srun --overlap --jobid=365 --ntasks=1 --chdir=/mnt/lustre/vlm-s4duan/FastVideo \
|
||||
bash -lc 'source .venv/bin/activate && export CUDA_VISIBLE_DEVICES=1 \
|
||||
YOLO_CONFIG_DIR=/mnt/lustre/vlm-s4duan/.ultralytics \
|
||||
TORCH_HOME=/mnt/lustre/vlm-s4duan/.torch HF_HOME=/mnt/lustre/vlm-s4duan/.hf \
|
||||
PYTHONPATH=$PWD:$PWD/data_pipeline MPLCONFIGDIR=/mnt/lustre/vlm-s4duan/.mpl && \
|
||||
python -u data_pipeline/seg_compare.py render --limit 8'
|
||||
|
||||
# serve on the login node, then ssh -L 7862:localhost:7862 <login>
|
||||
.venv/bin/python data_pipeline/seg_compare.py serve
|
||||
"""
|
||||
# NOTE: no ``from __future__ import annotations`` — gradio needs the real
|
||||
# gr.SelectData annotation on the click handler.
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from segment_tracks import object_ids_for_points, lowrank_track_weights # noqa: E402
|
||||
from sparse_sampling_dashboard import sparse_sample # noqa: E402
|
||||
from track_informativeness import _draw_overlay # noqa: E402
|
||||
|
||||
DATA_DEFAULT = "/mnt/lustre/vlm-s4duan/openvid_1m"
|
||||
WEIGHTS_DIR = "/mnt/lustre/vlm-s4duan/models/seg"
|
||||
OUT_DEFAULT = "/mnt/lustre/vlm-s4duan/seg_compare_out"
|
||||
|
||||
# The model zoo. loader in {FastSAM, SAM}. All run no-prompt everything-mode.
|
||||
# min_area_frac drops mask specks (< frac of frame); nms_contain drops a mask if
|
||||
# > that fraction of it is already covered by a larger kept mask (collapses
|
||||
# part-level fragments into their object). max_masks caps to N largest (0=all).
|
||||
DEFAULT_MODELS = [
|
||||
{"name": "fastsam-s", "loader": "FastSAM", "weight": "FastSAM-s.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "fastsam-x", "loader": "FastSAM", "weight": "FastSAM-x.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "mobilesam", "loader": "SAM", "weight": "mobile_sam.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam-b", "loader": "SAM", "weight": "sam_b.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam2-b", "loader": "SAM", "weight": "sam2_b.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam2.1-b+", "loader": "SAM", "weight": "sam2.1_b.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
{"name": "sam2.1-l", "loader": "SAM", "weight": "sam2.1_l.pt",
|
||||
"imgsz": 1024, "conf": 0.4, "iou": 0.9, "min_area_frac": 0.0015, "nms_contain": 0.0, "max_masks": 0},
|
||||
]
|
||||
|
||||
# sparse-sampler defaults baked into the rendered PNG/mp4 (dashboard can re-sample the PNG live)
|
||||
SPARSE_DEFAULT = {"num_sampled": 20, "mode": "weighted", "seed": 0}
|
||||
|
||||
|
||||
def read_frames(path: str) -> np.ndarray:
|
||||
"""Full clip -> [T,H,W,3] uint8 via PyAV (aarch64-safe)."""
|
||||
import av
|
||||
c = av.open(path)
|
||||
frames = [f.to_ndarray(format="rgb24") for f in c.decode(video=0)]
|
||||
c.close()
|
||||
return np.stack(frames)
|
||||
|
||||
|
||||
def colors(n: int) -> np.ndarray:
|
||||
import colorsys
|
||||
return np.array([[int(255 * v) for v in colorsys.hsv_to_rgb((i * 0.61803) % 1.0, 0.65, 1.0)]
|
||||
for i in range(max(1, n))], np.uint8)
|
||||
|
||||
|
||||
def extract_masks(res, H: int, W: int) -> np.ndarray:
|
||||
if not res or res[0].masks is None:
|
||||
return np.zeros((0, H, W), bool)
|
||||
m = res[0].masks.data.cpu().numpy().astype(bool) # [M,h,w]
|
||||
if m.shape[0] and m.shape[1:] != (H, W): # nearest resize w/o cv2
|
||||
import torch
|
||||
t = torch.from_numpy(m.astype(np.uint8))[None].float()
|
||||
t = torch.nn.functional.interpolate(t, size=(H, W), mode="nearest")[0]
|
||||
m = t.numpy().astype(bool)
|
||||
return m
|
||||
|
||||
|
||||
def filter_masks(masks: np.ndarray, min_area_frac: float, nms_contain: float, max_masks: int) -> np.ndarray:
|
||||
if masks.shape[0] == 0:
|
||||
return masks
|
||||
H, W = masks.shape[1], masks.shape[2]
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
if min_area_frac > 0:
|
||||
keep = (areas / float(H * W)) >= min_area_frac
|
||||
masks, areas = masks[keep], areas[keep]
|
||||
if masks.shape[0] and nms_contain > 0: # drop fragments mostly inside a larger kept mask
|
||||
order = np.argsort(-areas) # large -> small
|
||||
kept_idx, covered = [], np.zeros((H, W), bool)
|
||||
for i in order:
|
||||
m = masks[i]
|
||||
a = areas[i]
|
||||
if a > 0 and (m & covered).sum() / a > nms_contain:
|
||||
continue
|
||||
kept_idx.append(i)
|
||||
covered |= m
|
||||
masks, areas = masks[kept_idx], areas[kept_idx]
|
||||
if max_masks and masks.shape[0] > max_masks:
|
||||
masks = masks[np.argsort(-areas)[:max_masks]]
|
||||
return masks
|
||||
|
||||
|
||||
def mask_panel(frame0: np.ndarray, masks: np.ndarray, oid: np.ndarray, xy0: np.ndarray) -> "Image.Image":
|
||||
from PIL import Image, ImageDraw
|
||||
mcols = colors(len(masks) + 1)
|
||||
base = frame0.astype(np.float32)
|
||||
for mi, m in enumerate(masks):
|
||||
base[m] = 0.55 * base[m] + 0.45 * mcols[mi % len(mcols)][None].astype(np.float32)
|
||||
img = Image.fromarray(base.clip(0, 255).astype(np.uint8))
|
||||
draw = ImageDraw.Draw(img)
|
||||
for o in sorted({int(v) for v in np.unique(oid) if int(v) >= 0}): # 1 dot per object (its centroid pt)
|
||||
idx = np.where(oid == o)[0]
|
||||
cx, cy = float(xy0[idx, 0].mean()), float(xy0[idx, 1].mean())
|
||||
draw.ellipse([cx - 4, cy - 4, cx + 4, cy + 4], fill=(255, 255, 255), outline=(0, 0, 0))
|
||||
return img
|
||||
|
||||
|
||||
def sparse_panel(frame0: np.ndarray, xy0: np.ndarray, oid: np.ndarray, keep: np.ndarray, title: str,
|
||||
out_path: str) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
cmap = plt.get_cmap("tab20")
|
||||
fig, ax = plt.subplots(figsize=(9, 5.2))
|
||||
ax.imshow(frame0)
|
||||
x, y = xy0[:, 0], xy0[:, 1]
|
||||
ax.scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.30)
|
||||
for o in np.unique(oid[keep]):
|
||||
m = keep & (oid == o)
|
||||
col = "white" if int(o) < 0 else cmap(int(o) % 20)
|
||||
ax.scatter(x[m], y[m], c=[col], s=80, edgecolors="black", linewidths=1.1)
|
||||
ax.set_title(title, fontsize=10)
|
||||
ax.axis("off")
|
||||
fig.savefig(out_path, dpi=90, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
import torch
|
||||
import imageio.v2 as imageio
|
||||
from ultralytics import FastSAM, SAM
|
||||
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
clips_dir = data / ("clips" if (data / "clips").exists() else "videos")
|
||||
models_cfg = json.loads(Path(args.models_json).read_text()) if args.models_json else DEFAULT_MODELS
|
||||
if args.models: # subset by name
|
||||
want = set(args.models.split(","))
|
||||
models_cfg = [m for m in models_cfg if m["name"] in want]
|
||||
loaders = {"FastSAM": FastSAM, "SAM": SAM}
|
||||
|
||||
# load every model once (reused across all clips)
|
||||
print(f"[seg] loading {len(models_cfg)} models ...", flush=True)
|
||||
loaded = {}
|
||||
for mc in models_cfg:
|
||||
wp = Path(WEIGHTS_DIR) / mc["weight"]
|
||||
loaded[mc["name"]] = loaders[mc["loader"]](str(wp) if wp.exists() else mc["weight"])
|
||||
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
print(f"[seg] {len(items)} clips x {len(models_cfg)} models", flush=True)
|
||||
entries, clip_ids = [], []
|
||||
for k, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
clip_ids.append(stem)
|
||||
vpath = str(clips_dir / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = read_frames(vpath)
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]] # [T,N,2] px
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]] # [T,N]
|
||||
xy0, vis0 = tracks[0], vis[0]
|
||||
tw = lowrank_track_weights(tracks) # [N] informativeness in [0,1]
|
||||
# frame-0 png (shared, for dashboard live re-sampling background)
|
||||
from PIL import Image
|
||||
f0png = f"{stem}__frame0.png"
|
||||
Image.fromarray(frames[0]).save(str(out / f0png))
|
||||
|
||||
for mc in models_cfg:
|
||||
t0 = time.time()
|
||||
res = loaded[mc["name"]](frames[0], device=args.device, retina_masks=True,
|
||||
imgsz=mc["imgsz"], conf=mc["conf"], iou=mc["iou"], verbose=False)
|
||||
masks = filter_masks(extract_masks(res, H, W), mc["min_area_frac"], mc["nms_contain"], mc["max_masks"])
|
||||
oid = object_ids_for_points(masks, xy0, H, W)
|
||||
n_obj = int(len({int(v) for v in np.unique(oid) if int(v) >= 0}))
|
||||
keep = sparse_sample(oid, tw, vis0, SPARSE_DEFAULT["num_sampled"], SPARSE_DEFAULT["mode"],
|
||||
SPARSE_DEFAULT["seed"])
|
||||
dt = time.time() - t0
|
||||
pfx = f"{mc['name']}__{stem}"
|
||||
mask_panel(frames[0], masks, oid, xy0).save(str(out / f"{pfx}_masks.png"))
|
||||
sparse_panel(frames[0], xy0, oid, keep,
|
||||
f"{stem} | {mc['name']}: {int(keep.sum())} kept ({n_obj} objs +"
|
||||
f"{SPARSE_DEFAULT['num_sampled']} extra)", str(out / f"{pfx}_sparse.png"))
|
||||
trk_name = ""
|
||||
if not args.no_video:
|
||||
import matplotlib.pyplot as plt
|
||||
cmap = plt.get_cmap("tab20")
|
||||
kidx = np.where(keep)[0]
|
||||
cols = np.zeros((kidx.size, 3), np.uint8)
|
||||
for o in np.unique(oid[kidx]):
|
||||
mm = oid[kidx] == o
|
||||
rgba = (1., 1., 1., 1.) if int(o) < 0 else cmap(int(o) % 20)
|
||||
cols[mm] = (np.array(rgba[:3]) * 255).astype(np.uint8)
|
||||
ov = _draw_overlay(frames, tracks[:, kidx].copy(), vis[:, kidx], cols, 14, 3, 0.5)
|
||||
trk_name = f"{pfx}_tracks.mp4"
|
||||
imageio.mimsave(str(out / trk_name), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
# per-entry npz so the dashboard can re-sample live (no GPU)
|
||||
np.savez(str(out / f"{pfx}.npz"), oid=oid.astype(np.int64), tw=tw, vis0=vis0, xy0=xy0)
|
||||
entries.append({
|
||||
"model": mc["name"], "clip": stem,
|
||||
"caption": (it["cap"][0] if isinstance(it.get("cap"), list) else str(it.get("cap", "")))[:140],
|
||||
"n_masks": int(masks.shape[0]), "n_objects": n_obj,
|
||||
"n_labeled": int((oid >= 0).sum()), "n_points": int(oid.shape[0]),
|
||||
"n_kept": int(keep.sum()), "sec": round(dt, 2),
|
||||
"masks_png": f"{pfx}_masks.png", "sparse_png": f"{pfx}_sparse.png",
|
||||
"tracks_mp4": trk_name, "frame0_png": f0png, "npz": f"{pfx}.npz",
|
||||
})
|
||||
print(f"[seg] [{k}/{len(items)}] {stem} [{mc['name']}]: {masks.shape[0]} masks, "
|
||||
f"{n_obj} objs, {(oid >= 0).sum()}/{oid.shape[0]} pts labeled, {dt:.1f}s", flush=True)
|
||||
del frames
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
manifest = {"models": models_cfg, "clips": clip_ids, "sparse_default": SPARSE_DEFAULT, "entries": entries}
|
||||
(out / "manifest.json").write_text(json.dumps(manifest, indent=2))
|
||||
print(f"[seg] wrote {len(entries)} entries -> {out}/manifest.json", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
manifest = json.loads((viz / "manifest.json").read_text())
|
||||
entries = manifest["entries"]
|
||||
model_names = [m["name"] for m in manifest["models"]]
|
||||
clip_ids = manifest["clips"]
|
||||
by_clip = {c: [e for e in entries if e["clip"] == c] for c in clip_ids}
|
||||
caption_of = {c: (by_clip[c][0]["caption"] if by_clip.get(c) else "") for c in clip_ids}
|
||||
npz_cache = {}
|
||||
|
||||
def _npz(e):
|
||||
p = e["npz"]
|
||||
if p not in npz_cache:
|
||||
npz_cache[p] = {k: v for k, v in np.load(str(viz / p)).items()}
|
||||
return npz_cache[p]
|
||||
|
||||
def masks_gallery(clip):
|
||||
es = sorted(by_clip.get(clip, []), key=lambda e: model_names.index(e["model"]))
|
||||
return [(str(viz / e["masks_png"]), f'{e["model"]} • {e["n_masks"]}m / {e["n_objects"]}o') for e in es]
|
||||
|
||||
def resample_png(e, num_sampled, mode, seed):
|
||||
z = _npz(e)
|
||||
oid, tw, vis0, xy0 = z["oid"], z["tw"], z["vis0"], z["xy0"]
|
||||
keep = sparse_sample(oid, tw, vis0, int(num_sampled), mode, int(seed))
|
||||
frame0 = plt.imread(str(viz / e["frame0_png"]))
|
||||
cmap = plt.get_cmap("tab20")
|
||||
fig, ax = plt.subplots(figsize=(9, 5.2))
|
||||
ax.imshow(frame0)
|
||||
x, y = xy0[:, 0], xy0[:, 1]
|
||||
ax.scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.30)
|
||||
for o in np.unique(oid[keep]):
|
||||
m = keep & (oid == o)
|
||||
col = "white" if int(o) < 0 else cmap(int(o) % 20)
|
||||
ax.scatter(x[m], y[m], c=[col], s=80, edgecolors="black", linewidths=1.1)
|
||||
n_obj = int(len({int(v) for v in np.unique(oid) if int(v) >= 0}))
|
||||
ax.set_title(f'{e["clip"]} | {e["model"]}: {int(keep.sum())} kept ({n_obj} objs +{num_sampled}, {mode})',
|
||||
fontsize=10)
|
||||
ax.axis("off")
|
||||
import numpy as _np
|
||||
fig.canvas.draw()
|
||||
buf = _np.asarray(fig.canvas.buffer_rgba())[..., :3].copy()
|
||||
plt.close(fig)
|
||||
return buf
|
||||
|
||||
def sparse_gallery(clip, num_sampled, mode, seed):
|
||||
es = sorted(by_clip.get(clip, []), key=lambda e: model_names.index(e["model"]))
|
||||
return [(resample_png(e, num_sampled, mode, seed), f'{e["model"]} • {e["n_kept"]}→') for e in es]
|
||||
|
||||
with gr.Blocks(title="Segmentation model comparison") as demo:
|
||||
gr.Markdown(
|
||||
"## Segmentation model comparison — frame-0 masks + sparse track sampling\n"
|
||||
"Pick a **clip**; every model is shown side by side. **Masks** tab = raw everything-mode "
|
||||
"segmentation (tile label `<#masks>m / <#objects-with-points>o`). **Sparse sampling** tab = "
|
||||
"1 track per object + N extras (tweak live). Fewer, cleaner object masks = better for our "
|
||||
"1-per-object recipe. Click a tile to inspect it big + its track-overlay video.")
|
||||
clip_dd = gr.Dropdown(clip_ids, value=clip_ids[0], label=f"clip (1 of {len(clip_ids)})")
|
||||
clip_cap = gr.Markdown(f"*{caption_of[clip_ids[0]]}*")
|
||||
with gr.Tab("Masks (everything-mode)"):
|
||||
mg = gr.Gallery(value=masks_gallery(clip_ids[0]), columns=3, height=680,
|
||||
label="SAM masks per model (click a tile)")
|
||||
with gr.Tab("Sparse sampling"):
|
||||
with gr.Row():
|
||||
num_sl = gr.Slider(0, 150, value=manifest["sparse_default"]["num_sampled"], step=1,
|
||||
label="num_sampled (extras beyond 1-per-object)")
|
||||
mode_rd = gr.Radio(["weighted", "random"], value=manifest["sparse_default"]["mode"], label="extras mode")
|
||||
seed_sl = gr.Slider(0, 50, value=manifest["sparse_default"]["seed"], step=1, label="seed")
|
||||
sg = gr.Gallery(value=sparse_gallery(clip_ids[0], manifest["sparse_default"]["num_sampled"],
|
||||
manifest["sparse_default"]["mode"],
|
||||
manifest["sparse_default"]["seed"]),
|
||||
columns=3, height=640, label="sparse-sampled tracks per model (click a tile)")
|
||||
gr.Markdown("### Inspect")
|
||||
with gr.Row():
|
||||
big = gr.Image(label="selected panel", height=460)
|
||||
vid = gr.Video(label="kept-tracks overlay over the clip", height=460)
|
||||
meta = gr.Code(label="metrics", language="json")
|
||||
|
||||
def on_clip(clip, num_sampled, mode, seed):
|
||||
return masks_gallery(clip), sparse_gallery(clip, num_sampled, mode, seed), f"*{caption_of.get(clip, '')}*"
|
||||
|
||||
def on_sparse_ctrl(clip, num_sampled, mode, seed):
|
||||
return sparse_gallery(clip, num_sampled, mode, seed)
|
||||
|
||||
def pick(clip, evt: gr.SelectData):
|
||||
es = sorted(by_clip.get(clip, []), key=lambda e: model_names.index(e["model"]))
|
||||
e = es[evt.index]
|
||||
v = str(viz / e["tracks_mp4"]) if e.get("tracks_mp4") else None
|
||||
return str(viz / e["masks_png"]), v, json.dumps(e, indent=2)
|
||||
|
||||
clip_dd.change(on_clip, [clip_dd, num_sl, mode_rd, seed_sl], [mg, sg, clip_cap])
|
||||
for c in (num_sl, mode_rd, seed_sl):
|
||||
c.change(on_sparse_ctrl, [clip_dd, num_sl, mode_rd, seed_sl], sg)
|
||||
mg.select(pick, clip_dd, [big, vid, meta])
|
||||
sg.select(pick, clip_dd, [big, vid, meta])
|
||||
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", default=DATA_DEFAULT)
|
||||
r.add_argument("--out", default=OUT_DEFAULT)
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--device", default="cuda")
|
||||
r.add_argument("--models", default=None, help="comma-separated subset of model names")
|
||||
r.add_argument("--models-json", default=None, help="path to a JSON list of model config dicts")
|
||||
r.add_argument("--no-video", action="store_true", help="skip the track-overlay mp4s (faster)")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", default=OUT_DEFAULT)
|
||||
s.add_argument("--host", default="127.0.0.1")
|
||||
s.add_argument("--port", type=int, default=7862)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,78 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Smoke-test: run several ultralytics segmentation backends in no-prompt
|
||||
'everything' mode on one real openvid frame. Report #masks + latency so we can
|
||||
pick which variants are worth a full render. aarch64-safe (PyAV decode, no cv2)."""
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
DATA = Path("/mnt/lustre/vlm-s4duan/openvid_1m")
|
||||
WEIGHTS = Path("/mnt/lustre/vlm-s4duan/models/seg")
|
||||
WEIGHTS.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# (display name, weight file, loader class)
|
||||
MODELS = [
|
||||
("fastsam-s", "FastSAM-s.pt", "FastSAM"),
|
||||
("fastsam-x", "FastSAM-x.pt", "FastSAM"),
|
||||
("mobilesam", "mobile_sam.pt", "SAM"),
|
||||
("sam-b", "sam_b.pt", "SAM"),
|
||||
("sam2-t", "sam2_t.pt", "SAM"),
|
||||
("sam2-b", "sam2_b.pt", "SAM"),
|
||||
("sam2.1-l", "sam2.1_l.pt", "SAM"),
|
||||
]
|
||||
|
||||
|
||||
def read_frame0(path: str) -> np.ndarray:
|
||||
import av
|
||||
c = av.open(path)
|
||||
for f in c.decode(video=0):
|
||||
return f.to_ndarray(format="rgb24")
|
||||
raise RuntimeError(f"no frames in {path}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
from ultralytics import FastSAM, SAM # noqa: F401
|
||||
import torch
|
||||
|
||||
items = json.loads((DATA / "videos2caption.json").read_text())
|
||||
item = items[0]
|
||||
vpath = str(DATA / "clips" / item["path"]) if (DATA / "clips" / item["path"]).exists() \
|
||||
else str(DATA / "videos" / item["path"])
|
||||
frame0 = read_frame0(vpath)
|
||||
H, W = frame0.shape[:2]
|
||||
print(f"frame: {vpath} {W}x{H}", flush=True)
|
||||
|
||||
loaders = {"FastSAM": FastSAM, "SAM": SAM}
|
||||
rows = []
|
||||
for name, wf, cls in MODELS:
|
||||
wp = WEIGHTS / wf
|
||||
try:
|
||||
t0 = time.time()
|
||||
model = loaders[cls](str(wp) if wp.exists() else wf)
|
||||
# everything-mode: no prompts. imgsz 1024, retina full-res masks.
|
||||
res = model(frame0, device="cuda", retina_masks=True, imgsz=1024,
|
||||
conf=0.4, iou=0.9, verbose=False)
|
||||
dt = time.time() - t0
|
||||
m = res[0].masks
|
||||
nm = 0 if m is None else int(m.data.shape[0])
|
||||
mh, mw = (0, 0) if m is None else tuple(m.data.shape[1:])
|
||||
print(f"[OK] {name:10s} weights={wf:14s} masks={nm:4d} maskres={mw}x{mh} "
|
||||
f"time={dt:6.2f}s", flush=True)
|
||||
rows.append((name, nm, dt))
|
||||
# cache downloaded weight into WEIGHTS dir
|
||||
if not wp.exists() and Path(wf).exists():
|
||||
Path(wf).replace(wp)
|
||||
del model
|
||||
torch.cuda.empty_cache()
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[FAIL] {name:10s} weights={wf:14s} -> {type(e).__name__}: {e}", flush=True)
|
||||
print("\nsummary (name, n_masks, sec):", flush=True)
|
||||
for r in rows:
|
||||
print(f" {r[0]:10s} {r[1]:4d} {r[2]:6.2f}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,166 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stage 0d: segment frame-0 (SAM2.1-b+ by default; --model to switch) 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. A grid
|
||||
point is assigned the SMALLEST mask that contains it (most specific object).
|
||||
|
||||
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_frame0(path: str) -> np.ndarray:
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
return VideoReader(path, ctx=cpu(0))[0].asnumpy() # HxWx3 uint8
|
||||
except Exception: # noqa: BLE001
|
||||
import av
|
||||
c = av.open(path)
|
||||
for f in c.decode(video=0):
|
||||
return f.to_ndarray(format="rgb24")
|
||||
raise RuntimeError(f"could not read {path}")
|
||||
|
||||
|
||||
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 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("--min-area-frac",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="drop masks smaller than this fraction of the frame (fights over-segmentation)")
|
||||
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("--shard", type=int, default=0, help="this shard index (0-based)")
|
||||
p.add_argument("--num-shards", type=int, default=1, help="total shards for parallel processing")
|
||||
args = p.parse_args()
|
||||
|
||||
# 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.limit:
|
||||
items = items[:args.limit]
|
||||
if args.num_shards > 1:
|
||||
items = items[args.shard::args.num_shards]
|
||||
print(f"[seg] shard {args.shard}/{args.num_shards}: processing {len(items)} items", flush=True)
|
||||
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 "object_ids" in d and "track_weights" in d:
|
||||
n_ok += 1
|
||||
continue
|
||||
frame0 = read_frame0(str(vpath))
|
||||
H, W = frame0.shape[0], frame0.shape[1]
|
||||
res = model(frame0,
|
||||
device=args.device,
|
||||
retina_masks=True,
|
||||
imgsz=args.imgsz,
|
||||
conf=args.conf,
|
||||
iou=args.iou,
|
||||
verbose=False)
|
||||
masks = np.zeros((0, H, W), bool)
|
||||
if res and res[0].masks is not None:
|
||||
masks = res[0].masks.data.cpu().numpy().astype(bool) # [M,h,w]
|
||||
if masks.shape[1:] != (H, W): # resize masks to frame res if needed
|
||||
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] else np.zeros((0, H, W), bool)
|
||||
# drop tiny masks / cap count to fight FastSAM over-segmentation
|
||||
if masks.shape[0] and (args.min_area_frac > 0 or args.max_masks):
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
if args.min_area_frac > 0:
|
||||
keep = (areas / float(H * W)) >= args.min_area_frac
|
||||
masks, areas = masks[keep], areas[keep]
|
||||
if args.max_masks and masks.shape[0] > args.max_masks:
|
||||
masks = masks[np.argsort(-areas)[:args.max_masks]]
|
||||
tracks = d["tracks"].astype(np.float32) # [T,N,2] px (orig res)
|
||||
oid = object_ids_for_points(masks, tracks[0], H, W) # frame-0 positions
|
||||
d["object_ids"] = oid.astype(np.int64)
|
||||
d["n_objects"] = np.int64(masks.shape[0])
|
||||
d["track_weights"] = lowrank_track_weights(tracks) # [N] low-rank informativeness in [0,1]
|
||||
tmp = npz_path.with_suffix(".tmp.npz")
|
||||
np.savez(tmp, **d)
|
||||
tmp.replace(npz_path)
|
||||
n_ok += 1
|
||||
cov = int((oid >= 0).sum())
|
||||
print(
|
||||
f"[seg] [{k}/{len(items)}] {vpath.name}: {masks.shape[0]} objs, "
|
||||
f"{cov}/{oid.shape[0]} grid pts labeled, "
|
||||
f"w[mean={d['track_weights'].mean():.3f} >0.5={(d['track_weights'] > 0.5).mean():.2f}]",
|
||||
flush=True)
|
||||
print(f"[seg] done; {n_ok}/{len(items)} npz have object_ids", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Data-parallel frame-0 segmentation (SAM2.1-b+ everything-mode) for OpenVid-1M.
|
||||
|
||||
Mirrors extract_tracks_mp.py: each Slurm task segments manifest[shard::num_shards]
|
||||
on one pinned GPU and writes object_ids + track_weights into the EXISTING tracks
|
||||
.npz (adds keys, keeps tracks/visibility). Idempotent — skips npz that already have
|
||||
both keys, so it is resumable.
|
||||
|
||||
Sharding + GPU pinning from Slurm env:
|
||||
shard = SLURM_PROCID (0..num_shards-1)
|
||||
num_shards = SLURM_NTASKS
|
||||
gpu = SLURM_LOCALID % gpus_per_node
|
||||
|
||||
aarch64-safe: PyAV frame-0 decode, torch mask resize (no cv2/decord).
|
||||
|
||||
srun --ntasks-per-node=<gpus*procs_per_gpu> ... \
|
||||
python data_pipeline/segment_tracks_mp.py --data-dir <root> --videos-subdir clips
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from segment_tracks import object_ids_for_points, lowrank_track_weights # noqa: E402
|
||||
|
||||
|
||||
def read_frame0(path: str) -> np.ndarray:
|
||||
import av
|
||||
c = av.open(str(path))
|
||||
try:
|
||||
for f in c.decode(video=0):
|
||||
return f.to_ndarray(format="rgb24")
|
||||
finally:
|
||||
c.close()
|
||||
raise RuntimeError(f"no frames in {path}")
|
||||
|
||||
|
||||
def extract_masks(res, H: int, W: int) -> np.ndarray:
|
||||
if not res or res[0].masks is None:
|
||||
return np.zeros((0, H, W), bool)
|
||||
import torch
|
||||
m = res[0].masks.data
|
||||
if m.shape[1:] != (H, W):
|
||||
m = torch.nn.functional.interpolate(m[None].float(), size=(H, W), mode="nearest")[0]
|
||||
return m.cpu().numpy().astype(bool)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--data-dir", required=True)
|
||||
ap.add_argument("--videos-subdir", default="clips", help="frame-0 source; MUST match track coords (720p clips)")
|
||||
ap.add_argument("--tracks-subdir", default="tracks")
|
||||
ap.add_argument("--manifest", default="videos2caption.json")
|
||||
ap.add_argument("--model", default="sam2.1_b.pt")
|
||||
ap.add_argument("--weights-dir", default="/mnt/lustre/vlm-s4duan/models/seg")
|
||||
ap.add_argument("--imgsz", type=int, default=1024)
|
||||
ap.add_argument("--conf", type=float, default=0.4)
|
||||
ap.add_argument("--iou", type=float, default=0.9)
|
||||
ap.add_argument("--min-area-frac", type=float, default=0.0)
|
||||
ap.add_argument("--num-shards", type=int, default=int(os.environ.get("SLURM_NTASKS", "1")))
|
||||
ap.add_argument("--shard", type=int, default=int(os.environ.get("SLURM_PROCID", "0")))
|
||||
ap.add_argument("--gpus-per-node", type=int, default=4)
|
||||
ap.add_argument("--limit", type=int, default=None)
|
||||
a = ap.parse_args()
|
||||
|
||||
import torch
|
||||
local = int(os.environ.get("SLURM_LOCALID", str(a.shard)))
|
||||
gpu = local % a.gpus_per_node
|
||||
torch.cuda.set_device(gpu)
|
||||
device = f"cuda:{gpu}"
|
||||
|
||||
from ultralytics import FastSAM, SAM
|
||||
wp = Path(a.weights_dir) / a.model
|
||||
weight = str(wp) if wp.exists() else a.model
|
||||
model = (FastSAM if Path(a.model).name.lower().startswith("fastsam") else SAM)(weight)
|
||||
|
||||
data = Path(a.data_dir)
|
||||
items = json.loads((data / a.manifest).read_text())
|
||||
if a.limit:
|
||||
items = items[:a.limit]
|
||||
mine = items[a.shard::a.num_shards]
|
||||
print(f"[seg {a.shard}/{a.num_shards}] gpu={gpu} localid={local} items={len(mine)}", flush=True)
|
||||
|
||||
t0 = time.time()
|
||||
done = skip = notrk = err = 0
|
||||
for it in mine:
|
||||
stem = Path(it["path"]).stem
|
||||
npz_path = Path(it.get("points_path") or (data / a.tracks_subdir / f"{stem}.npz"))
|
||||
if not npz_path.exists():
|
||||
notrk += 1
|
||||
continue
|
||||
try:
|
||||
d = dict(np.load(npz_path))
|
||||
if "object_ids" in d and "track_weights" in d:
|
||||
skip += 1
|
||||
continue
|
||||
vpath = data / a.videos_subdir / it["path"]
|
||||
frame0 = read_frame0(str(vpath))
|
||||
H, W = frame0.shape[0], frame0.shape[1]
|
||||
res = model(frame0, device=device, retina_masks=True, imgsz=a.imgsz,
|
||||
conf=a.conf, iou=a.iou, verbose=False)
|
||||
masks = extract_masks(res, H, W)
|
||||
if masks.shape[0] and a.min_area_frac > 0:
|
||||
areas = masks.reshape(masks.shape[0], -1).sum(1).astype(np.float64)
|
||||
masks = masks[(areas / float(H * W)) >= a.min_area_frac]
|
||||
tracks = d["tracks"].astype(np.float32)
|
||||
oid = object_ids_for_points(masks, tracks[0], H, W)
|
||||
d["object_ids"] = oid.astype(np.int64)
|
||||
d["n_objects"] = np.int64(masks.shape[0])
|
||||
d["track_weights"] = lowrank_track_weights(tracks)
|
||||
tmp = npz_path.with_suffix(".seg.tmp.npz")
|
||||
np.savez(tmp, **d)
|
||||
tmp.replace(npz_path)
|
||||
done += 1
|
||||
if done % 200 == 0:
|
||||
r = done / (time.time() - t0)
|
||||
print(f"[seg {a.shard}] {done} done ({r:.2f}/s), skip={skip} notrk={notrk} err={err}", flush=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
err += 1
|
||||
print(f"[seg {a.shard}] ERR {stem}: {repr(e)[:120]}", flush=True)
|
||||
dt = time.time() - t0
|
||||
print(f"[seg {a.shard}] DONE done={done} skip={skip} notrk={notrk} err={err} in {dt:.1f}s "
|
||||
f"-> {done / dt if dt > 0 else 0:.3f} clip/s", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,337 @@
|
||||
# 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 object_ids_for_points
|
||||
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
|
||||
|
||||
for cfg in configs:
|
||||
res = model(frames[0],
|
||||
device=args.device,
|
||||
retina_masks=True,
|
||||
imgsz=cfg["imgsz"],
|
||||
conf=cfg["conf"],
|
||||
iou=cfg["iou"],
|
||||
verbose=False)
|
||||
masks = filter_masks(_extract_masks(res, H, W), cfg["min_area_frac"], cfg["max_masks"])
|
||||
oid = object_ids_for_points(masks, tracks[0], H, W) # [N] mask idx per point (-1 none)
|
||||
objs = sorted(int(o) for o in np.unique(oid) if int(o) >= 0)
|
||||
|
||||
# (1) SAM panel: colored masks + chosen point (highest-motion track) per object
|
||||
mcols = _colors(len(masks) + 1)
|
||||
base = frames[0].astype(np.float32)
|
||||
for mi, m in enumerate(masks):
|
||||
base[m] = 0.55 * base[m] + 0.45 * mcols[mi % len(mcols)][None].astype(np.float32)
|
||||
img = Image.fromarray(base.clip(0, 255).astype(np.uint8))
|
||||
draw = ImageDraw.Draw(img)
|
||||
for o in objs:
|
||||
idx = np.where(oid == o)[0]
|
||||
pick = idx[np.argmax(disp[idx])]
|
||||
x, y = float(tracks[0, pick, 0]), float(tracks[0, pick, 1])
|
||||
draw.ellipse([x - 5, y - 5, x + 5, y + 5], fill=(255, 255, 255), outline=(0, 0, 0))
|
||||
sam_name = f"{cfg['name']}__{stem}_sam.png"
|
||||
img.save(str(out / sam_name))
|
||||
|
||||
# (2) CoTracker tracks over the video, colored by object (background gray)
|
||||
trk_name = ""
|
||||
if not args.no_tracks:
|
||||
ocols = _colors(len(objs) + 1)
|
||||
pcols = np.tile(np.array([[110, 110, 110]], np.uint8), (tracks.shape[1], 1))
|
||||
for oi, o in enumerate(objs):
|
||||
pcols[oid == o] = ocols[oi % len(ocols)]
|
||||
# grid-aware subsample: a flat row-major stride staggers columns and
|
||||
# looks like half the grid; stride rows AND cols equally instead so the
|
||||
# true (e.g. 50x50) grid stays visible and aligned.
|
||||
N = tracks.shape[1]
|
||||
G = int(round(N**0.5))
|
||||
if G * G == N:
|
||||
k = max(1, G // 50) # keep full grid up to 50x50
|
||||
sel = np.arange(N).reshape(G, G)[::k, ::k].reshape(-1)
|
||||
else:
|
||||
st = max(1, N // 1500)
|
||||
sel = np.arange(0, N, st)
|
||||
ov = _draw_overlay(frames, tracks[:, sel].copy(), vis[:, sel], pcols[sel], 12, 2, 0.5)
|
||||
trk_name = f"{cfg['name']}__{stem}_tracks.mp4"
|
||||
imageio.mimsave(str(out / trk_name), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
|
||||
entries.append({
|
||||
"config":
|
||||
cfg["name"],
|
||||
"clip":
|
||||
stem,
|
||||
"caption": (it["cap"][0] if isinstance(it.get("cap"), list) else str(it.get("cap", "")))[:120],
|
||||
"n_masks":
|
||||
int(len(masks)),
|
||||
"n_objects":
|
||||
len(objs),
|
||||
"n_labeled":
|
||||
int((oid >= 0).sum()),
|
||||
"n_points":
|
||||
int(oid.shape[0]),
|
||||
"sam":
|
||||
sam_name,
|
||||
"tracks":
|
||||
trk_name,
|
||||
})
|
||||
print(
|
||||
f"[viz] [{k}/{len(items)}] {stem} [{cfg['name']}]: "
|
||||
f"{len(masks)} masks, {len(objs)} objs, {(oid >= 0).sum()}/{oid.shape[0]} pts labeled",
|
||||
flush=True)
|
||||
|
||||
manifest = {"configs": configs, "clips": clip_ids, "entries": entries}
|
||||
(out / "manifest.json").write_text(json.dumps(manifest, indent=2))
|
||||
print(f"[viz] rendered {len(entries)} panels -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
manifest = json.loads((viz / "manifest.json").read_text())
|
||||
entries = manifest["entries"]
|
||||
cfg_names = [c["name"] for c in manifest["configs"]]
|
||||
clip_ids = manifest["clips"]
|
||||
cfg_params = {c["name"]: c for c in manifest["configs"]}
|
||||
# closure state: the currently-filtered entry list that the gallery reflects
|
||||
current = {"entries": list(entries)}
|
||||
|
||||
def _label(e: dict) -> str:
|
||||
return f'{e["clip"]} | {e["config"]} | {e["n_masks"]}m/{e["n_objects"]}o'
|
||||
|
||||
def gallery_for(cfg_sel: str, clip_sel: str):
|
||||
es = entries
|
||||
if cfg_sel and cfg_sel != "(all)":
|
||||
es = [e for e in es if e["config"] == cfg_sel]
|
||||
if clip_sel and clip_sel != "(all)":
|
||||
es = [e for e in es if e["clip"] == clip_sel]
|
||||
current["entries"] = es
|
||||
return [(str(viz / e["sam"]), _label(e)) for e in es]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = current["entries"][evt.index]
|
||||
tv = str(viz / e["tracks"]) if e.get("tracks") else None
|
||||
info = {**e, "config_params": cfg_params.get(e["config"], {})}
|
||||
return str(viz / e["sam"]), tv, json.dumps(info, indent=2)
|
||||
|
||||
with gr.Blocks(title="FastSAM config sweep") as demo:
|
||||
gr.Markdown("### FastSAM config sweep — masks + chosen per-object points + CoTracker tracks\n"
|
||||
"Filter by **config** and/or **clip**, scroll the gallery, click a tile to inspect. "
|
||||
"Tile label: `clip | config | <#masks>m/<#objects-with-points>o`. "
|
||||
"Fewer, cleaner masks = less over-segmentation.")
|
||||
with gr.Row():
|
||||
cfg_dd = gr.Dropdown(["(all)"] + cfg_names, value="(all)", label="config")
|
||||
clip_dd = gr.Dropdown(["(all)"] + clip_ids, value="(all)", label="clip")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gallery_for("(all)", "(all)"), columns=4, height=620, label="panels (click to view)")
|
||||
with gr.Column():
|
||||
samimg = gr.Image(label="frame-0: SAM masks + chosen points")
|
||||
trkvid = gr.Video(label="CoTracker tracks (colored by object)")
|
||||
meta = gr.Code(label="info", language="json")
|
||||
cfg_dd.change(gallery_for, [cfg_dd, clip_dd], g)
|
||||
clip_dd.change(gallery_for, [cfg_dd, clip_dd], g)
|
||||
g.select(show, None, [samimg, trkvid, meta])
|
||||
# allowed_paths: without this gradio blocks serving the PNGs/mp4s that live
|
||||
# outside its app root -> the three side panels error out on click.
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--model", default="FastSAM-s.pt")
|
||||
r.add_argument("--device", default="cuda")
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--configs", default=None, help="inline JSON list of config dicts (overrides sweep)")
|
||||
r.add_argument("--configs-json", default=None, help="path to a JSON list of config dicts")
|
||||
r.add_argument("--no-tracks", action="store_true", help="skip the (slow) track-overlay mp4s")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7880)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,221 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Interactive dashboard for the SPARSE ~1-per-object + few-background sampler.
|
||||
|
||||
Per Yongqi's advice, our track budget for the sparse recipe is dramatically smaller
|
||||
than MotionStream's 1000-2500: keep 1 track per SAM object, then add ``num_sampled``
|
||||
extra points drawn from the rest (either uniformly at random, or weighted by the
|
||||
precomputed low-rank informativeness). The dashboard lets us eyeball whether the
|
||||
resulting subset actually covers the "action" of each clip before we commit to
|
||||
launching training runs on it.
|
||||
|
||||
CPU only, gradio share (no GPU needed).
|
||||
|
||||
.venv/bin/python data_pipeline/sparse_sampling_dashboard.py \
|
||||
--data-dir /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/synthetic_toy --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
from track_informativeness import _draw_overlay, _read_frames # noqa: E402
|
||||
|
||||
|
||||
def sparse_sample(oid: np.ndarray, tw: np.ndarray | None, vis0: np.ndarray, num_sampled: int, mode: str,
|
||||
seed: int) -> np.ndarray:
|
||||
"""One point per SAM object (weighted-by-`tw` inside each object) + ``num_sampled`` extras
|
||||
drawn from the rest via ``mode`` in {"random", "weighted"}. Returns kept bool [N]."""
|
||||
rng = np.random.default_rng(int(seed))
|
||||
N = oid.shape[0]
|
||||
valid = vis0 > 0.5
|
||||
keep = np.zeros(N, bool)
|
||||
# (a) 1 per object
|
||||
for o in np.unique(oid):
|
||||
if int(o) < 0:
|
||||
continue
|
||||
idx = np.where((oid == o) & valid)[0]
|
||||
if idx.size == 0:
|
||||
continue
|
||||
if tw is not None and tw[idx].sum() > 0:
|
||||
p = tw[idx] / tw[idx].sum()
|
||||
pick = int(rng.choice(idx, p=p))
|
||||
else:
|
||||
pick = int(rng.choice(idx))
|
||||
keep[pick] = True
|
||||
# (b) num_sampled extras from the remaining valid points
|
||||
pool = np.where((~keep) & valid)[0]
|
||||
k = int(min(max(num_sampled, 0), pool.size))
|
||||
if k > 0:
|
||||
if mode == "weighted" and tw is not None and tw[pool].sum() > 0:
|
||||
p = tw[pool] / tw[pool].sum()
|
||||
picks = rng.choice(pool, size=k, replace=False, p=p)
|
||||
else:
|
||||
picks = rng.choice(pool, size=k, replace=False)
|
||||
keep[picks] = True
|
||||
return keep
|
||||
|
||||
|
||||
def _load(data_dir: Path) -> list[dict]:
|
||||
man = json.loads((data_dir / "videos2caption.json").read_text())
|
||||
labels_dir = data_dir / "sam_labels"
|
||||
items = []
|
||||
for it in man:
|
||||
stem = Path(it["path"]).stem
|
||||
vpath = str(data_dir / "videos" / it["path"])
|
||||
npzp = it.get("points_path") or str(data_dir / "tracks" / f"{stem}.npz")
|
||||
d = np.load(npzp)
|
||||
lp = labels_dir / f"{stem}.npy"
|
||||
labels = np.load(lp) if lp.exists() else None # [H,W] int16, -1 = background
|
||||
items.append({
|
||||
"stem": stem,
|
||||
"caption": it["cap"][0] if isinstance(it.get("cap"), list) else "",
|
||||
"fps": int(it.get("fps", 24)),
|
||||
"vpath": vpath,
|
||||
"tracks": d["tracks"].astype(np.float32),
|
||||
"vis": d["visibility"].astype(np.float32),
|
||||
"oid": d["object_ids"].astype(np.int64) if "object_ids" in d else None,
|
||||
"tw": d["track_weights"].astype(np.float32) if "track_weights" in d else None,
|
||||
"labels": labels,
|
||||
})
|
||||
return items
|
||||
|
||||
|
||||
def build_ui(items: list[dict], out_dir: Path):
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
def _render_frame(clip_idx: int, num_sampled: int, mode: str, seed: int, mask_alpha: float = 0.45):
|
||||
it = items[int(clip_idx)]
|
||||
frames = _read_frames(it["vpath"])
|
||||
tracks = it["tracks"][:frames.shape[0]]
|
||||
vis = it["vis"][:frames.shape[0]]
|
||||
oid = it["oid"]
|
||||
tw = it["tw"]
|
||||
labels = it.get("labels")
|
||||
vis0 = vis[0]
|
||||
if oid is None:
|
||||
return None, "No object_ids in npz — run segment_tracks first."
|
||||
keep = sparse_sample(oid, tw, vis0, int(num_sampled), mode, int(seed))
|
||||
n_kept = int(keep.sum())
|
||||
n_obj = int(len(np.unique(oid[oid >= 0])))
|
||||
# Blend SAM segmentation on top of frame 0: color each segment by the same tab20
|
||||
# index we use for the tracked points (mask id == object id).
|
||||
base = frames[0].astype(np.float32)
|
||||
if labels is not None:
|
||||
cmap = plt.get_cmap("tab20")
|
||||
H, W = base.shape[:2]
|
||||
overlay = np.zeros_like(base)
|
||||
uniq = np.unique(labels)
|
||||
for u in uniq:
|
||||
if int(u) < 0:
|
||||
continue
|
||||
col = np.array(cmap(int(u) % 20)[:3]) * 255
|
||||
m = labels == u
|
||||
overlay[m] = col
|
||||
mask = (labels >= 0)[..., None]
|
||||
base = np.where(mask, (1 - mask_alpha) * base + mask_alpha * overlay, base)
|
||||
base = base.clip(0, 255).astype(np.uint8)
|
||||
fig, ax = plt.subplots(1, 1, figsize=(9, 5.5))
|
||||
ax.imshow(base)
|
||||
x, y = tracks[0, :, 0], tracks[0, :, 1]
|
||||
ax.scatter(x[~keep], y[~keep], c="0.5", s=3, alpha=0.35) # dropped
|
||||
# kept: color per unique object id (SAME palette as the mask overlay above)
|
||||
cmap = plt.get_cmap("tab20")
|
||||
for o in np.unique(oid[keep]):
|
||||
mask_pts = keep & (oid == o)
|
||||
col = "white" if int(o) < 0 else cmap(int(o) % 20)
|
||||
ax.scatter(x[mask_pts], y[mask_pts], c=[col], s=90, edgecolors="black", linewidths=1.2)
|
||||
ax.set_title(
|
||||
f"{it['stem']}: {n_kept} tracks kept ({n_obj} SAM objects, +{num_sampled} background, "
|
||||
f"{mode}, seed {seed})",
|
||||
fontsize=10)
|
||||
ax.axis("off")
|
||||
img_path = str(out_dir / f"preview_{it['stem']}.png")
|
||||
fig.savefig(img_path, dpi=90, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
info = json.dumps(
|
||||
{
|
||||
"caption": it["caption"],
|
||||
"n_kept": n_kept,
|
||||
"n_objects_present": n_obj,
|
||||
"kept_from_objects": int(len(np.unique(oid[keep][oid[keep] >= 0]))),
|
||||
"kept_from_background": int((keep & (oid < 0)).sum()),
|
||||
},
|
||||
indent=2)
|
||||
return img_path, info
|
||||
|
||||
def _render_video(clip_idx: int, num_sampled: int, mode: str, seed: int):
|
||||
it = items[int(clip_idx)]
|
||||
frames = _read_frames(it["vpath"])
|
||||
tracks = it["tracks"][:frames.shape[0]]
|
||||
vis = it["vis"][:frames.shape[0]]
|
||||
oid = it["oid"]
|
||||
tw = it["tw"]
|
||||
keep = sparse_sample(oid, tw, vis[0], int(num_sampled), mode, int(seed))
|
||||
kidx = np.where(keep)[0]
|
||||
if kidx.size == 0:
|
||||
return None
|
||||
cmap = plt.get_cmap("tab20")
|
||||
cols = np.zeros((kidx.size, 3), dtype=np.uint8)
|
||||
for o in np.unique(oid[kidx]):
|
||||
m = (oid[kidx] == o)
|
||||
rgba = (1.0, 1.0, 1.0, 1.0) if int(o) < 0 else cmap(int(o) % 20)
|
||||
cols[m] = (np.array(rgba[:3]) * 255).astype(np.uint8)
|
||||
ov = _draw_overlay(frames, tracks[:, kidx].copy(), vis[:, kidx], cols, 14, 3, 0.5)
|
||||
out = str(out_dir / f"preview_{it['stem']}_overlay.mp4")
|
||||
imageio.mimsave(out, ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
return out
|
||||
|
||||
with gr.Blocks(title="Sparse track sampler dashboard") as demo:
|
||||
gr.Markdown("### Sparse-sampling recipe: 1 track per SAM object + `num_sampled` extras\n"
|
||||
"Frame 0 shows all 2500 grid tracks in gray, and the kept ones colored by object. "
|
||||
"Set `num_sampled` = 0 for pure 1-per-object; raise it to add background context. "
|
||||
"Toggle `weighted` to bias the extras toward high-lowrank-informativeness points. "
|
||||
"The overlay video shows only the kept tracks over the whole clip.")
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
clip_dd = gr.Dropdown([f"{i:02d} — {items[i]['stem']}" for i in range(len(items))],
|
||||
value=f"00 — {items[0]['stem']}",
|
||||
label="Clip",
|
||||
type="index")
|
||||
num_slider = gr.Slider(0, 200, value=20, step=1, label="num_sampled (extras)")
|
||||
mode_radio = gr.Radio(["random", "weighted"], value="weighted", label="extras sampling mode")
|
||||
seed_slider = gr.Slider(0, 100, value=0, step=1, label="seed")
|
||||
info_box = gr.Code(label="info", language="json")
|
||||
vid_btn = gr.Button("Render overlay video", variant="primary")
|
||||
with gr.Column(scale=2):
|
||||
frame_img = gr.Image(label="frame 0 with kept tracks", height=460)
|
||||
overlay_vid = gr.Video(label="kept-tracks overlay video", height=460)
|
||||
|
||||
for src in [clip_dd, num_slider, mode_radio, seed_slider]:
|
||||
src.change(_render_frame, [clip_dd, num_slider, mode_radio, seed_slider], [frame_img, info_box])
|
||||
vid_btn.click(_render_video, [clip_dd, num_slider, mode_radio, seed_slider], overlay_vid)
|
||||
demo.load(_render_frame, [clip_dd, num_slider, mode_radio, seed_slider], [frame_img, info_box])
|
||||
return demo
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--data-dir", required=True)
|
||||
p.add_argument("--out-dir",
|
||||
default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/sparse_sampling_dashboard_out")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7894)
|
||||
p.add_argument("--share", action="store_true")
|
||||
args = p.parse_args()
|
||||
out_dir = Path(args.out_dir)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
items = _load(Path(args.data_dir))
|
||||
demo = build_ui(items, out_dir)
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(out_dir)])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,57 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Split a videos2caption.json manifest into NUM_SHARDS shards for data-parallel preprocess.
|
||||
|
||||
The i2v_track preprocess reads a merge.txt (``<clips_dir>,<json>``) whose JSON is a
|
||||
*list* of dicts (json.load, NOT jsonlines) — see
|
||||
fastvideo/dataset/preprocessing_datasets.py:_load_raw_data. So each shard is itself a
|
||||
JSON array + its own merge.txt, and each shard is fed to one single-GPU v1_preprocess.py
|
||||
process.
|
||||
|
||||
Also injects a ``duration`` field when missing: FrameSamplingStage.should_keep drops any
|
||||
row with ``duration is None`` (preprocessing_datasets.py:183), and OpenVid's
|
||||
videos2caption.json carries only num_frames/fps. duration = num_frames / fps.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--manifest", required=True, help="videos2caption.json (JSON array)")
|
||||
ap.add_argument("--clips-dir", required=True, help="dir with the .mp4 clips (merge.txt col 1)")
|
||||
ap.add_argument("--out-dir", required=True, help="output shards dir (shard_XXXXX/ created inside)")
|
||||
ap.add_argument("--num-shards", type=int, required=True)
|
||||
a = ap.parse_args()
|
||||
|
||||
items = json.load(open(a.manifest))
|
||||
filled = 0
|
||||
for it in items:
|
||||
if it.get("duration") is None:
|
||||
nf, fps = it.get("num_frames"), it.get("fps")
|
||||
if nf and fps:
|
||||
it["duration"] = float(nf) / float(fps)
|
||||
filled += 1
|
||||
|
||||
out = Path(a.out_dir)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
# Contiguous shards keep each shard's clips locally clustered on disk.
|
||||
n = a.num_shards
|
||||
total = len(items)
|
||||
per = (total + n - 1) // n
|
||||
written = 0
|
||||
for k in range(n):
|
||||
chunk = items[k * per:(k + 1) * per]
|
||||
sd = out / f"shard_{k:05d}"
|
||||
sd.mkdir(parents=True, exist_ok=True)
|
||||
mani = sd / "videos2caption.json"
|
||||
mani.write_text(json.dumps(chunk, indent=2))
|
||||
(sd / "merge.txt").write_text(f"{a.clips_dir},{mani}")
|
||||
written += len(chunk)
|
||||
print(f"[split] {total} rows ({filled} got injected duration) -> {n} shards "
|
||||
f"(~{per}/shard) under {out}; wrote {written} rows total")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,47 @@
|
||||
#!/usr/bin/env bash
|
||||
# Wait for the in-flight wave to flush its .done markers, then restart the 720p
|
||||
# preprocess on a 12-node subset of alloc 881 (frees hpc-rack-3-[2,5,10,15]).
|
||||
#
|
||||
# Safe because the run is resumable: per-shard .done markers are skipped on
|
||||
# restart, and the launcher releases any claim lacking a .done at pass start.
|
||||
set -uo pipefail
|
||||
W=/mnt/lustre/vlm-s4duan
|
||||
OUT=$W/openvid_1m/combined_parquet_dataset_720p
|
||||
KEEP12="hpc-rack-2-[0-2,4-6,8-9,12,14],hpc-rack-3-[0-1]"
|
||||
TARGET=${TARGET:-60} # wave 1 is 64 shards; don't wait forever on a straggler
|
||||
|
||||
echo "[switch] waiting for >=$TARGET shards to flush before restarting on 12 nodes"
|
||||
for _ in $(seq 1 90); do
|
||||
D=$(find "$OUT" -maxdepth 2 -name .done 2>/dev/null | wc -l)
|
||||
echo "[switch] $(date +%T) banked $D shards"
|
||||
[ "$D" -ge "$TARGET" ] && break
|
||||
sleep 60
|
||||
done
|
||||
BANKED=$(find "$OUT" -maxdepth 2 -name .done 2>/dev/null | wc -l)
|
||||
echo "[switch] proceeding with $BANKED shards banked"
|
||||
|
||||
# 1) Stop the supervisor FIRST so it cannot launch another pass.
|
||||
# Character-class in the pattern so this script's own cmdline never matches.
|
||||
pkill -f 'run_preprocess_720p_hel[d]' 2>/dev/null
|
||||
sleep 3
|
||||
# 2) Stop the 16 per-node srun steps.
|
||||
pkill -f 'jobid=881 --nodelis[t]=hpc-rack' 2>/dev/null
|
||||
sleep 15
|
||||
pkill -9 -f 'jobid=881 --nodelis[t]=hpc-rack' 2>/dev/null
|
||||
sleep 5
|
||||
echo "[switch] supervisor procs left: $(pgrep -cf 'run_preprocess_720p_hel[d]' || echo 0)"
|
||||
echo "[switch] srun step procs left: $(pgrep -f 'jobid=881 --nodelis[t]=hpc-rack' | wc -l)"
|
||||
|
||||
# 3) Confirm no worker python survives anywhere in the alloc.
|
||||
srun --overlap --jobid=881 --nodes=16 --ntasks=16 --ntasks-per-node=1 --chdir=$W \
|
||||
bash -lc 'N=$(ps -eo args --no-headers | grep -c "[p]ython fastvideo/pipelines"); [ "$N" -gt 0 ] && echo " STILL RUNNING $(hostname): $N procs"' 2>/dev/null | grep -v '^srun\|error:'
|
||||
echo "[switch] alloc drained"
|
||||
|
||||
# 4) Relaunch on 12 nodes. Resume skips banked shards; stale claims are released.
|
||||
cd $W/FastVideo
|
||||
JOBID=881 NODELIST="$KEEP12" nohup bash data_pipeline/run_preprocess_720p_held.sh \
|
||||
> $W/logs/prep720_supervisor_12n.log 2>&1 &
|
||||
echo "[switch] relaunched on 12 nodes (pid $!), log prep720_supervisor_12n.log"
|
||||
sleep 40
|
||||
head -6 $W/logs/prep720_supervisor_12n.log
|
||||
echo "[switch] FREE FOR TRAINING: hpc-rack-3-[2,5,10,15]"
|
||||
@@ -0,0 +1,289 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Author synthetic point-track motion controls for WanTrack.
|
||||
|
||||
A "track" is just a per-point (x, y) trajectory over time, so we can *author*
|
||||
counterfactual motion directly (no CoTracker needed -- that is only for
|
||||
*extracting* tracks from real video). These authored tracks are the test-time
|
||||
control signal: feed them to the model and check the generated motion follows.
|
||||
|
||||
All coordinates are NORMALIZED to [0, 1] (x/width, y/height) -- the same space
|
||||
the preprocessing stores and the model trains on. Outputs:
|
||||
tracks float32 [T, N, 2] normalized (x, y) per point per frame
|
||||
visibility float32 [T, N] 1 = active/visible, 0 = inactive/occluded
|
||||
|
||||
Modes mirror MotionStream's control surface:
|
||||
- ``pan`` : global translation (camera/scene pan)
|
||||
- ``zoom`` : scale toward/away from a center (dolly in/out)
|
||||
- ``rotate`` : rotate the grid about a center
|
||||
- ``drag`` : move only points inside a region (the sparse "drag" handle),
|
||||
with a smooth spatial falloff; everything else stays put
|
||||
- ``swirl`` : rotation whose angle decays with radius
|
||||
- ``static`` : no motion (tests "does text+first-frame alone reproduce motion?")
|
||||
|
||||
``select_*`` helpers turn a dense field into a SPARSE control by zeroing the
|
||||
visibility of unselected points -- this is how MotionStream lets you drag just a
|
||||
few points. (Requires a model trained with point-subsampling aug to work well.)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Grid + small helpers
|
||||
# ----------------------------------------------------------------------
|
||||
def make_grid(grid_size: int = 50) -> np.ndarray:
|
||||
"""Frame-0 normalized grid positions, returned as [N, 2] (x, y) in [0, 1].
|
||||
|
||||
Row-major (matches CoTracker's grid and our 50x50 preprocessing), so
|
||||
``reshape(grid_size, grid_size, 2)`` is [row(y), col(x)].
|
||||
"""
|
||||
ys, xs = np.meshgrid(
|
||||
np.linspace(0.0, 1.0, grid_size, dtype=np.float32),
|
||||
np.linspace(0.0, 1.0, grid_size, dtype=np.float32),
|
||||
indexing="ij",
|
||||
)
|
||||
return np.stack([xs.reshape(-1), ys.reshape(-1)], axis=-1) # [N, 2]
|
||||
|
||||
|
||||
def _ramp(num_frames: int, ease: bool = True) -> np.ndarray:
|
||||
"""Time ramp in [0, 1] over ``num_frames`` (optionally smooth ease-in-out)."""
|
||||
t = np.linspace(0.0, 1.0, num_frames, dtype=np.float32)
|
||||
if ease:
|
||||
t = t * t * (3.0 - 2.0 * t) # smoothstep
|
||||
return t
|
||||
|
||||
|
||||
def _visibility_from_bounds(tracks: np.ndarray, margin: float = 0.02) -> np.ndarray:
|
||||
"""1 where a point is inside [(-m), (1+m)] in both axes, else 0."""
|
||||
lo, hi = -margin, 1.0 + margin
|
||||
inb = (tracks[..., 0] >= lo) & (tracks[..., 0] <= hi) & (tracks[..., 1] >= lo) & (tracks[..., 1] <= hi)
|
||||
return inb.astype(np.float32)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Motion fields -> (tracks [T,N,2], visibility [T,N])
|
||||
# ----------------------------------------------------------------------
|
||||
def static(grid: np.ndarray, num_frames: int) -> tuple[np.ndarray, np.ndarray]:
|
||||
tracks = np.repeat(grid[None], num_frames, axis=0)
|
||||
return tracks, np.ones((num_frames, grid.shape[0]), np.float32)
|
||||
|
||||
|
||||
def pan(grid: np.ndarray, num_frames: int, dx: float, dy: float,
|
||||
ease: bool = True) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Translate all points by (dx, dy) (normalized units) over the clip."""
|
||||
r = _ramp(num_frames, ease)[:, None, None]
|
||||
tracks = grid[None] + r * np.array([dx, dy], np.float32)[None, None]
|
||||
return tracks.astype(np.float32), _visibility_from_bounds(tracks)
|
||||
|
||||
|
||||
def zoom(grid: np.ndarray, num_frames: int, scale: float,
|
||||
center: tuple[float, float] = (0.5, 0.5), ease: bool = True
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Scale toward (scale<1) or away from (scale>1) ``center`` (dolly)."""
|
||||
c = np.array(center, np.float32)[None, None]
|
||||
r = _ramp(num_frames, ease)[:, None, None]
|
||||
s = 1.0 + (scale - 1.0) * r
|
||||
tracks = c + (grid[None] - c) * s
|
||||
return tracks.astype(np.float32), _visibility_from_bounds(tracks)
|
||||
|
||||
|
||||
def rotate(grid: np.ndarray, num_frames: int, degrees: float,
|
||||
center: tuple[float, float] = (0.5, 0.5), ease: bool = True
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
c = np.array(center, np.float32)[None]
|
||||
r = _ramp(num_frames, ease)
|
||||
out = np.empty((num_frames, grid.shape[0], 2), np.float32)
|
||||
rel = grid - c
|
||||
for t in range(num_frames):
|
||||
a = np.deg2rad(degrees) * r[t]
|
||||
ca, sa = np.cos(a), np.sin(a)
|
||||
rot = np.array([[ca, -sa], [sa, ca]], np.float32)
|
||||
out[t] = rel @ rot.T + c
|
||||
return out, _visibility_from_bounds(out)
|
||||
|
||||
|
||||
def drag(grid: np.ndarray, num_frames: int, *, center: tuple[float, float],
|
||||
dx: float, dy: float, radius: float = 0.15, falloff: str = "smooth",
|
||||
ease: bool = True) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Move only points within ``radius`` of ``center`` by (dx, dy).
|
||||
|
||||
``falloff='smooth'`` weights displacement by a cosine window of the distance
|
||||
(handle-like drag); ``'hard'`` moves all in-radius points equally. Points
|
||||
outside the radius stay put. Visibility stays 1 for all points (dense field
|
||||
with a localized motion) -- use ``select_radius`` to make it *sparse*.
|
||||
"""
|
||||
c = np.array(center, np.float32)[None]
|
||||
d = np.linalg.norm(grid - c, axis=-1) # [N]
|
||||
if falloff == "smooth":
|
||||
w = np.clip(1.0 - d / max(radius, 1e-6), 0.0, 1.0)
|
||||
w = 0.5 - 0.5 * np.cos(np.pi * w) # smooth 0->1
|
||||
else:
|
||||
w = (d <= radius).astype(np.float32)
|
||||
r = _ramp(num_frames, ease)[:, None, None]
|
||||
disp = (w[:, None] * np.array([dx, dy], np.float32)[None])[None] # [1,N,2]
|
||||
tracks = grid[None] + r * disp
|
||||
return tracks.astype(np.float32), _visibility_from_bounds(tracks)
|
||||
|
||||
|
||||
def swirl(grid: np.ndarray, num_frames: int, degrees: float,
|
||||
center: tuple[float, float] = (0.5, 0.5), radius: float = 0.5,
|
||||
ease: bool = True) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Rotation whose angle decays linearly to 0 at ``radius`` from center."""
|
||||
c = np.array(center, np.float32)[None]
|
||||
rel = grid - c
|
||||
dist = np.linalg.norm(rel, axis=-1)
|
||||
decay = np.clip(1.0 - dist / max(radius, 1e-6), 0.0, 1.0)
|
||||
rmp = _ramp(num_frames, ease)
|
||||
out = np.empty((num_frames, grid.shape[0], 2), np.float32)
|
||||
for t in range(num_frames):
|
||||
a = np.deg2rad(degrees) * rmp[t] * decay
|
||||
ca, sa = np.cos(a), np.sin(a)
|
||||
x = rel[:, 0] * ca - rel[:, 1] * sa
|
||||
y = rel[:, 0] * sa + rel[:, 1] * ca
|
||||
out[t] = np.stack([x, y], -1) + c
|
||||
return out, _visibility_from_bounds(out)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Sparsify (sparse "drag a few points" control)
|
||||
# ----------------------------------------------------------------------
|
||||
def select_radius(visibility: np.ndarray, grid: np.ndarray, *,
|
||||
center: tuple[float, float], radius: float) -> np.ndarray:
|
||||
"""Zero visibility outside ``radius`` of ``center`` (keep only a local handle)."""
|
||||
c = np.array(center, np.float32)[None]
|
||||
keep = (np.linalg.norm(grid - c, axis=-1) <= radius).astype(np.float32) # [N]
|
||||
return visibility * keep[None]
|
||||
|
||||
|
||||
def select_random(visibility: np.ndarray, k: int, seed: int = 0) -> np.ndarray:
|
||||
"""Keep only ``k`` random points active (rest visibility 0)."""
|
||||
rng = np.random.default_rng(seed)
|
||||
n = visibility.shape[1]
|
||||
keep = np.zeros(n, np.float32)
|
||||
keep[rng.choice(n, size=min(k, n), replace=False)] = 1.0
|
||||
return visibility * keep[None]
|
||||
|
||||
|
||||
def select_stride(visibility: np.ndarray, grid_size: int, stride: int) -> np.ndarray:
|
||||
"""Keep a coarse sub-grid (every ``stride``-th point in both dims) active."""
|
||||
n = visibility.shape[1]
|
||||
mask2d = np.zeros((grid_size, grid_size), np.float32)
|
||||
mask2d[::stride, ::stride] = 1.0
|
||||
return visibility * mask2d.reshape(1, n)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Motion transfer: reuse tracks extracted from a real/other video
|
||||
# ----------------------------------------------------------------------
|
||||
def from_npz(npz_path: str, num_frames: int | None = None
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Load tracks from an ``extract_tracks.py`` npz, normalized to [0,1]."""
|
||||
data = np.load(npz_path)
|
||||
tr = data["tracks"].astype(np.float32) # [T, N, 2] pixels
|
||||
vis = data["visibility"].astype(np.float32) # [T, N]
|
||||
w = float(data["width"]) if "width" in data else 1.0
|
||||
h = float(data["height"]) if "height" in data else 1.0
|
||||
tr = tr.copy()
|
||||
tr[..., 0] /= max(w, 1e-6)
|
||||
tr[..., 1] /= max(h, 1e-6)
|
||||
if num_frames is not None:
|
||||
tr, vis = tr[:num_frames], vis[:num_frames]
|
||||
return tr, vis
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Convenience: named presets
|
||||
# ----------------------------------------------------------------------
|
||||
def preset(name: str, num_frames: int, grid_size: int = 50, *,
|
||||
strength: float = 0.25) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Author a named motion preset. ``strength`` scales translations/zoom."""
|
||||
g = make_grid(grid_size)
|
||||
name = name.lower()
|
||||
if name == "static":
|
||||
return static(g, num_frames)
|
||||
if name in ("pan_right", "pan_left", "pan_up", "pan_down"):
|
||||
dx = {"pan_right": strength, "pan_left": -strength}.get(name, 0.0)
|
||||
dy = {"pan_down": strength, "pan_up": -strength}.get(name, 0.0)
|
||||
return pan(g, num_frames, dx, dy)
|
||||
if name == "zoom_in":
|
||||
return zoom(g, num_frames, 1.0 + strength)
|
||||
if name == "zoom_out":
|
||||
return zoom(g, num_frames, 1.0 - strength)
|
||||
if name in ("rotate_cw", "rotate_ccw"):
|
||||
return rotate(g, num_frames, 30.0 * (1 if name == "rotate_cw" else -1))
|
||||
if name == "swirl":
|
||||
return swirl(g, num_frames, 60.0)
|
||||
if name == "drag_center_right":
|
||||
return drag(g, num_frames, center=(0.5, 0.5), dx=strength, dy=0.0, radius=0.2)
|
||||
raise ValueError(f"unknown preset {name!r}")
|
||||
|
||||
|
||||
PRESETS = ["static", "pan_right", "pan_left", "pan_up", "pan_down",
|
||||
"zoom_in", "zoom_out", "rotate_cw", "rotate_ccw", "swirl",
|
||||
"drag_center_right"]
|
||||
|
||||
|
||||
def to_pixel(tracks: np.ndarray, height: int, width: int) -> np.ndarray:
|
||||
"""Normalized [0,1] tracks -> pixel coords for overlay/EPE."""
|
||||
out = tracks.copy()
|
||||
out[..., 0] *= width
|
||||
out[..., 1] *= height
|
||||
return out
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# CLI: overlay a preset on a first frame (sanity check, no model/GPU)
|
||||
# ----------------------------------------------------------------------
|
||||
def _overlay_preview(frame0: np.ndarray, tracks_px: np.ndarray, vis: np.ndarray,
|
||||
stride: int = 3) -> np.ndarray:
|
||||
import colorsys
|
||||
|
||||
from PIL import Image, ImageDraw
|
||||
h, w, _ = frame0.shape
|
||||
img = Image.fromarray(frame0).convert("RGB")
|
||||
draw = ImageDraw.Draw(img)
|
||||
n = tracks_px.shape[1]
|
||||
gs = int(round(n ** 0.5))
|
||||
keep = np.zeros(n, bool)
|
||||
keep.reshape(gs, gs)[::stride, ::stride] = True
|
||||
for pi in range(n):
|
||||
if not keep[pi]:
|
||||
continue
|
||||
col = tuple(int(c * 255) for c in colorsys.hsv_to_rgb((pi % gs) / gs, 1.0, 1.0))
|
||||
pts = [(float(tracks_px[t, pi, 0]), float(tracks_px[t, pi, 1]))
|
||||
for t in range(tracks_px.shape[0]) if vis[t, pi] >= 0.5]
|
||||
if len(pts) >= 2:
|
||||
draw.line(pts, fill=col, width=1)
|
||||
if pts:
|
||||
x, y = pts[-1]
|
||||
draw.ellipse([x - 2, y - 2, x + 2, y + 2], fill=col)
|
||||
return np.asarray(img)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description="Preview a synthetic track preset over a first frame.")
|
||||
p.add_argument("--frame", type=Path, required=True, help="First-frame image (png/jpg).")
|
||||
p.add_argument("--preset", type=str, default="pan_right", choices=PRESETS)
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--grid-size", type=int, default=50)
|
||||
p.add_argument("--strength", type=float, default=0.25)
|
||||
p.add_argument("--out", type=Path, required=True)
|
||||
args = p.parse_args()
|
||||
|
||||
import imageio.v2 as imageio
|
||||
frame0 = np.asarray(imageio.imread(args.frame))[..., :3]
|
||||
h, w, _ = frame0.shape
|
||||
tracks, vis = preset(args.preset, args.num_frames, args.grid_size, strength=args.strength)
|
||||
overlay = _overlay_preview(frame0, to_pixel(tracks, h, w), vis)
|
||||
args.out.parent.mkdir(parents=True, exist_ok=True)
|
||||
imageio.imwrite(str(args.out), overlay)
|
||||
print(f"[synthetic-tracks] preset={args.preset} frames={args.num_frames} "
|
||||
f"points={tracks.shape[1]} -> {args.out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,142 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Controllability test for an overfitted TrackWan model.
|
||||
|
||||
Disentangles "did it overfit on the video content" from "did it learn to follow
|
||||
the tracks" by holding the content conditioning (first frame + text) fixed and
|
||||
varying ONLY the control tracks:
|
||||
|
||||
- gt : the clip's own tracks (reconstruction; should match GT)
|
||||
- none : no tracks at all (if motion still appears -> content overfit)
|
||||
- pan_*/zoom/drag : authored counterfactuals (if motion follows -> real control)
|
||||
- swap : another clip's tracks (motion transfer)
|
||||
|
||||
For each, it generates a video, overlays the control tracks, saves an mp4, and
|
||||
computes CoTracker EPE (input tracks vs tracks re-extracted from the generation).
|
||||
Low EPE on counterfactuals = genuine track control.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def _overlay(frames_thwc: np.ndarray, tracks_norm: np.ndarray, vis: np.ndarray, stride: int = 3) -> list[np.ndarray]:
|
||||
from fastvideo.train.callbacks.track_validation import (_draw_overlay, _grid_colors, _subsample)
|
||||
T, H, W, _ = frames_thwc.shape
|
||||
tt = min(T, tracks_norm.shape[0])
|
||||
fr, tr, vs = frames_thwc[:tt], tracks_norm[:tt], vis[:tt]
|
||||
grid = int(round(tr.shape[1]**0.5))
|
||||
tr, vs = _subsample(tr, vs, grid, stride)
|
||||
colors = _grid_colors(grid, stride)
|
||||
if colors.shape[0] != tr.shape[1]:
|
||||
colors = _grid_colors(int(round(tr.shape[1]**0.5)) or 1, 1)[:tr.shape[1]]
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
return _draw_overlay(fr, trpx, vs, colors, 12, 2, 0.5)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True, help="dcp_to_diffusers export dir")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data",
|
||||
default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset")
|
||||
p.add_argument("--out", required=True)
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import (compute_epe, load_cotracker)
|
||||
ct = load_cotracker(model.device)
|
||||
|
||||
# latent T -> pixel T
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
Tpx = (num_lat_t - 1) * ratio + 1
|
||||
|
||||
results = []
|
||||
for ci, s in zip(args.clips, samples, strict=False):
|
||||
gt_tracks = s["track_points"][0].numpy()[:Tpx] # [T,N,2] normalized
|
||||
gt_vis = s["track_visibility"][0].numpy()[:Tpx]
|
||||
swap = samples[(args.clips.index(ci) + 1) % len(samples)]
|
||||
swap_tracks = swap["track_points"][0].numpy()[:Tpx]
|
||||
swap_vis = swap["track_visibility"][0].numpy()[:Tpx]
|
||||
|
||||
g = st.make_grid(50)
|
||||
pan_t, pan_v = st.pan(g, Tpx, 0.25, 0.0)
|
||||
zoom_t, zoom_v = st.zoom(g, Tpx, 1.25)
|
||||
drag_t, drag_v = st.drag(g, Tpx, center=(0.5, 0.5), dx=0.3, dy=0.0, radius=0.2)
|
||||
drag_v_sparse = st.select_radius(drag_v, g, center=(0.5, 0.5), radius=0.2)
|
||||
|
||||
controls = {
|
||||
"gt": (gt_tracks, gt_vis),
|
||||
"none": (None, None),
|
||||
"pan_right": (pan_t, pan_v),
|
||||
"zoom_in": (zoom_t, zoom_v),
|
||||
"drag_dense": (drag_t, drag_v),
|
||||
"drag_sparse": (drag_t, drag_v_sparse),
|
||||
"swap": (swap_tracks, swap_vis),
|
||||
}
|
||||
|
||||
# GT reference (decode the real clip)
|
||||
ref = twi.decode_reference(model, s["vae_latent"])
|
||||
imageio.mimsave(os.path.join(args.out, f"clip{ci}_reference.mp4"),
|
||||
_overlay(ref, gt_tracks, gt_vis),
|
||||
fps=args.fps,
|
||||
macro_block_size=1)
|
||||
|
||||
for name, (tr, vs) in controls.items():
|
||||
tp = torch.from_numpy(tr)[None].float() if tr is not None else None
|
||||
tv = torch.from_numpy(vs)[None].float() if vs is not None else None
|
||||
lat = twi.generate(model,
|
||||
first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp,
|
||||
track_visibility=tv,
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=args.steps,
|
||||
seed=args.seed)
|
||||
frames = twi.decode_to_pixels(model, lat)
|
||||
ov_tr = tr if tr is not None else np.zeros((Tpx, 1, 2), np.float32)
|
||||
ov_vs = vs if vs is not None else np.zeros((Tpx, 1), np.float32)
|
||||
imageio.mimsave(os.path.join(args.out, f"clip{ci}_{name}.mp4"),
|
||||
_overlay(frames, ov_tr, ov_vs),
|
||||
fps=args.fps,
|
||||
macro_block_size=1)
|
||||
|
||||
epe = None
|
||||
if tr is not None:
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
tr_px = tr.copy()
|
||||
tr_px[..., 0] *= W
|
||||
tr_px[..., 1] *= H
|
||||
epe = compute_epe(frames, tr_px, vs, ct, model.device)["epe"]
|
||||
results.append((ci, name, epe))
|
||||
print(f"[clip {ci}] {name:12s} EPE={epe if epe is None else round(epe,2)}", flush=True)
|
||||
|
||||
print("\n=== EPE summary (lower = follows control better) ===")
|
||||
for ci, name, epe in results:
|
||||
print(f" clip {ci} {name:12s} {('n/a' if epe is None else f'{epe:7.2f} px')}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Quick motion-CFG test: does guidance>1 improve action adherence on an existing ckpt?"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import synthetic_tracks as st
|
||||
import trackwan_infer as twi
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument(
|
||||
"--data",
|
||||
default=
|
||||
"/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200/preprocessed_i2v_track/combined_parquet_dataset"
|
||||
)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_droid_overfit.yaml")
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--scales", type=float, nargs="+", default=[1.0, 2.0, 3.0])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
a = p.parse_args()
|
||||
model, tc = twi.load_trackwan(a.export, a.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(a.data, a.clips, text_len)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import compute_epe, load_cotracker
|
||||
ct = load_cotracker(model.device)
|
||||
nlt = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
Tpx = (nlt - 1) * ratio + 1
|
||||
# grid-snapped dense drag control (moves a radius-0.2 region by +0.3 x)
|
||||
g = st.make_grid(50)
|
||||
dt, dv = st.drag(g, Tpx, center=(0.5, 0.5), dx=0.3, dy=0.0, radius=0.2)
|
||||
disp = np.sqrt(((dt - dt[0:1])**2).sum(-1))
|
||||
moving = disp.max(0) > 8 / 832.0 # normalized thresh ~ px
|
||||
mv = np.where(moving)[0]
|
||||
print(f"{'clip':>4} {'scale':>6} {'EPE_all':>8} {'EPE_moving':>11}")
|
||||
for ci, s in zip(a.clips, samples, strict=False):
|
||||
tp = torch.from_numpy(dt)[None].float()
|
||||
tvv = torch.from_numpy(dv)[None].float()
|
||||
for sc in a.scales:
|
||||
lat = twi.generate(model,
|
||||
first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"],
|
||||
text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp,
|
||||
track_visibility=tvv,
|
||||
clip_feature=s["clip_feature"],
|
||||
num_steps=a.steps,
|
||||
seed=1000,
|
||||
guidance_scale=sc)
|
||||
fr = twi.decode_to_pixels(model, lat)
|
||||
H, W = fr.shape[1], fr.shape[2]
|
||||
dpx = dt.copy()
|
||||
dpx[..., 0] *= W
|
||||
dpx[..., 1] *= H
|
||||
epe_all = compute_epe(fr, dpx, dv, ct, model.device)["epe"]
|
||||
epe_mv = compute_epe(fr, dpx[:, mv], dv[:, mv], ct, model.device)["epe"]
|
||||
print(f"{ci:>4} {sc:>6.1f} {epe_all:>8.2f} {epe_mv:>11.2f}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,477 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Diagnostic bake-off: which point tracks are *informative* vs *generic*?
|
||||
|
||||
A point that moves a lot is not necessarily important. If the camera pans, every
|
||||
background point moves, but that motion is shared by the crowd, so any one background
|
||||
point is redundant. We compare several ways to score / select useful points, from
|
||||
simple statistics to clustering and learned embeddings.
|
||||
|
||||
Per-point scalar scores (heatmaps, bright = informative):
|
||||
abs_disp max_t ||p_t - p_0|| naive raw motion
|
||||
affine_residual residual after a robust per-frame global removes camera pan/zoom/rot
|
||||
affine x0 -> p_t
|
||||
lowrank_residual residual after subtracting mean trajectory unique vs ALL points, but
|
||||
+ top-k shared SVD modes (ONE subspace) assumes a single subspace
|
||||
motion_outlier 1 - HDBSCAN membership prob on trajectory density outlier in motion space
|
||||
embeddings (#1 motion clustering)
|
||||
subspace_resid per-cluster low-rank residual after union-of-subspaces; the
|
||||
Spectral motion segmentation per-object lowrank (#2 SSC-lite)
|
||||
dino_uniqueness 1 - membership prob clustering DINOv2 semantic ⊕ motion (#3)
|
||||
appearance ⊕ trajectory embeddings
|
||||
|
||||
Grouping / subset panels:
|
||||
motion_clusters HDBSCAN labels on trajectory embeddings motion groups (#1)
|
||||
subspace_clusters Spectral motion segmentation labels object subspaces (#2)
|
||||
dpp / fps / greedy diverse K-subset selection what sampling wants (#4)
|
||||
|
||||
``render`` (CPU; DINOv2 on CPU) writes per clip a big multi-panel PNG + an overlay mp4
|
||||
colored by low-rank informativeness + stats. ``serve`` is a gradio gallery (--share).
|
||||
|
||||
.venv/bin/python data_pipeline/track_informativeness.py render --data-dir <root> --out <dir> --limit 8
|
||||
.venv/bin/python data_pipeline/track_informativeness.py serve --viz-dir <dir> --share
|
||||
"""
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
|
||||
K_SELECT = 64 # size of the illustrative "useful subset" for the selection methods
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- overlay / io
|
||||
def _draw_overlay(frames: Any,
|
||||
tracks: Any,
|
||||
vis: Any,
|
||||
colors: Any,
|
||||
tail: int = 12,
|
||||
radius: int = 2,
|
||||
vis_thresh: float = 0.5) -> list:
|
||||
"""Self-contained track overlay (dots + short tails); no fastvideo/triton dependency."""
|
||||
from PIL import Image, ImageDraw
|
||||
T, N, _ = tracks.shape
|
||||
out = []
|
||||
for t in range(T):
|
||||
img = Image.fromarray(np.ascontiguousarray(frames[t]))
|
||||
dr = ImageDraw.Draw(img)
|
||||
t0 = max(0, t - tail)
|
||||
for i in range(N):
|
||||
col = (int(colors[i][0]), int(colors[i][1]), int(colors[i][2]))
|
||||
if t - t0 >= 1:
|
||||
dr.line([(float(p[0]), float(p[1])) for p in tracks[t0:t + 1, i]], fill=col, width=1)
|
||||
if vis[t, i] > vis_thresh:
|
||||
x, y = float(tracks[t, i, 0]), float(tracks[t, i, 1])
|
||||
dr.ellipse([x - radius, y - radius, x + radius, y + radius], fill=col)
|
||||
out.append(np.asarray(img))
|
||||
return out
|
||||
|
||||
|
||||
def _read_frames(path: str) -> np.ndarray:
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(path, ctx=cpu(0))
|
||||
return vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001
|
||||
import av
|
||||
c = av.open(path)
|
||||
return np.stack([f.to_ndarray(format="rgb24") for f in c.decode(video=0)])
|
||||
|
||||
|
||||
def _norm(x: np.ndarray, pct: float = 97.0) -> np.ndarray:
|
||||
lo = float(x.min())
|
||||
hi = float(np.percentile(x, pct))
|
||||
return np.zeros_like(x) if hi <= lo else np.clip((x - lo) / (hi - lo), 0.0, 1.0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- statistical scores
|
||||
def abs_disp(tracks: np.ndarray) -> np.ndarray:
|
||||
return np.sqrt(((tracks - tracks[0:1])**2).sum(-1)).max(0)
|
||||
|
||||
|
||||
def affine_residual(tracks: np.ndarray, vis: np.ndarray, iters: int = 3) -> np.ndarray:
|
||||
T, N, _ = tracks.shape
|
||||
X0 = np.concatenate([tracks[0], np.ones((N, 1), np.float32)], 1) # [N,3]
|
||||
resid = np.zeros((T, N), np.float32)
|
||||
for t in range(1, T):
|
||||
pt = tracks[t]
|
||||
w = (vis[t] > 0.5).astype(np.float64)
|
||||
if w.sum() < 6:
|
||||
w = np.ones(N)
|
||||
for _ in range(iters):
|
||||
W = X0.T * w
|
||||
A = np.linalg.lstsq((W @ X0), (W @ pt), rcond=None)[0]
|
||||
r = np.sqrt(((X0 @ A - pt)**2).sum(-1)) + 1e-6
|
||||
delta = 1.5 * np.median(r[w > 0]) + 1e-3
|
||||
w = np.where(r <= delta, 1.0, delta / r) * (vis[t] > 0.5)
|
||||
resid[t] = r
|
||||
return np.median(resid[1:], 0)
|
||||
|
||||
|
||||
def global_affine_motion(tracks: np.ndarray, vis: np.ndarray) -> float:
|
||||
"""Camera-fixedness proxy: median per-frame magnitude of the fitted GLOBAL affine
|
||||
displacement (in px), normalized by frame diagonal. ~0 => fixed camera; large => the
|
||||
whole scene translates/zooms (moving/egocentric camera). Robust (Huber) fit."""
|
||||
T, N, _ = tracks.shape
|
||||
X0 = np.concatenate([tracks[0], np.ones((N, 1), np.float32)], 1)
|
||||
mags = []
|
||||
for t in range(1, T):
|
||||
pt = tracks[t]
|
||||
w = (vis[t] > 0.5).astype(np.float64)
|
||||
if w.sum() < 6:
|
||||
w = np.ones(N)
|
||||
for _ in range(2):
|
||||
Wm = X0.T * w
|
||||
A = np.linalg.lstsq((Wm @ X0), (Wm @ pt), rcond=None)[0]
|
||||
r = np.sqrt(((X0 @ A - pt)**2).sum(-1)) + 1e-6
|
||||
delta = 1.5 * np.median(r[w > 0]) + 1e-3
|
||||
w = np.where(r <= delta, 1.0, delta / r) * (vis[t] > 0.5)
|
||||
mags.append(float(np.median(np.sqrt(((X0 @ A - tracks[0])**2).sum(-1)))))
|
||||
diag = float(np.hypot(tracks[..., 0].max(), tracks[..., 1].max()) + 1e-6)
|
||||
return round(float(np.median(mags)) / diag, 4)
|
||||
|
||||
|
||||
def displacement_matrix(tracks: np.ndarray) -> np.ndarray:
|
||||
"""[N, 2T] per-point displacement relative to frame 0."""
|
||||
T, N, _ = tracks.shape
|
||||
return (tracks - tracks[0:1]).transpose(1, 0, 2).reshape(N, 2 * T)
|
||||
|
||||
|
||||
def lowrank_residual(tracks: np.ndarray, rank: int = 3) -> np.ndarray:
|
||||
D = displacement_matrix(tracks)
|
||||
Dc = D - D.mean(0, keepdims=True)
|
||||
U, S, Vt = np.linalg.svd(Dc, full_matrices=False)
|
||||
r = min(rank, S.shape[0])
|
||||
return np.sqrt(((Dc - (U[:, :r] * S[:r]) @ Vt[:r])**2).sum(-1))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- embeddings + clustering
|
||||
def traj_embedding(tracks: np.ndarray, n_pca: int = 16) -> np.ndarray:
|
||||
from sklearn.decomposition import PCA
|
||||
D = displacement_matrix(tracks)
|
||||
D = (D - D.mean(0)) / (D.std(0) + 1e-6)
|
||||
d = int(min(n_pca, min(D.shape) - 1))
|
||||
return PCA(n_components=d, random_state=0).fit_transform(D).astype(np.float32)
|
||||
|
||||
|
||||
def motion_cluster(emb: np.ndarray):
|
||||
"""#1: HDBSCAN on trajectory embeddings -> labels, outlierness score (1 - membership prob)."""
|
||||
from sklearn.cluster import HDBSCAN
|
||||
mcs = max(10, emb.shape[0] // 100)
|
||||
cl = HDBSCAN(min_cluster_size=mcs, min_samples=5).fit(emb)
|
||||
return cl.labels_, (1.0 - cl.probabilities_).astype(np.float32)
|
||||
|
||||
|
||||
def subspace_cluster(tracks: np.ndarray, emb: np.ndarray, k: int):
|
||||
"""#2: Spectral motion segmentation into k groups; per-cluster low-rank residual."""
|
||||
from sklearn.cluster import SpectralClustering
|
||||
k = int(np.clip(k, 2, 8))
|
||||
lab = SpectralClustering(n_clusters=k,
|
||||
affinity="nearest_neighbors",
|
||||
n_neighbors=15,
|
||||
assign_labels="cluster_qr",
|
||||
random_state=0).fit_predict(emb)
|
||||
D = displacement_matrix(tracks)
|
||||
resid = np.zeros(lab.shape[0], np.float32)
|
||||
for c in np.unique(lab):
|
||||
idx = np.where(lab == c)[0]
|
||||
Z = D[idx] - D[idx].mean(0, keepdims=True)
|
||||
r = min(4, min(Z.shape) - 1)
|
||||
if r > 0:
|
||||
U, S, Vt = np.linalg.svd(Z, full_matrices=False)
|
||||
resid[idx] = np.sqrt(((Z - (U[:, :r] * S[:r]) @ Vt[:r])**2).sum(-1))
|
||||
return lab, resid
|
||||
|
||||
|
||||
def dino_point_features(frame0: np.ndarray, xy: np.ndarray, model, device: str) -> np.ndarray:
|
||||
"""#3: sample DINOv2 patch tokens at each frame-0 point (bilinear). Returns [N,C] L2-normed."""
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
H, W = frame0.shape[0], frame0.shape[1]
|
||||
gw, gh = max(1, round(W / 14)), max(1, round(H / 14))
|
||||
rw, rh = gw * 14, gh * 14
|
||||
img = torch.from_numpy(frame0).float().permute(2, 0, 1)[None] / 255.0
|
||||
img = F.interpolate(img, size=(rh, rw), mode="bilinear", align_corners=False)
|
||||
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
|
||||
std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
|
||||
img = ((img - mean) / std).to(device)
|
||||
with torch.no_grad():
|
||||
tok = model.forward_features(img)["x_norm_patchtokens"][0] # [gh*gw, C]
|
||||
C = tok.shape[-1]
|
||||
grid = tok.reshape(gh, gw, C).permute(2, 0, 1)[None] # [1,C,gh,gw]
|
||||
# normalized sample coords in [-1,1] from pixel xy
|
||||
gx = (xy[:, 0] / W) * 2 - 1
|
||||
gy = (xy[:, 1] / H) * 2 - 1
|
||||
samp = torch.from_numpy(np.stack([gx, gy], -1)).float().view(1, -1, 1, 2).to(device)
|
||||
feat = F.grid_sample(grid, samp, mode="bilinear", align_corners=False)[0, :, :, 0].T # [N,C]
|
||||
feat = feat.cpu().numpy()
|
||||
return (feat / (np.linalg.norm(feat, axis=1, keepdims=True) + 1e-6)).astype(np.float32)
|
||||
|
||||
|
||||
def dino_motion_cluster(dino_feat: np.ndarray, traj_emb: np.ndarray):
|
||||
"""Cluster fused (semantic ⊕ motion) embedding -> labels, uniqueness (1 - membership prob)."""
|
||||
from sklearn.cluster import HDBSCAN
|
||||
from sklearn.decomposition import PCA
|
||||
dz = PCA(n_components=int(min(16, min(dino_feat.shape) - 1)), random_state=0).fit_transform(dino_feat)
|
||||
dz = (dz - dz.mean(0)) / (dz.std(0) + 1e-6)
|
||||
tz = (traj_emb - traj_emb.mean(0)) / (traj_emb.std(0) + 1e-6)
|
||||
fused = np.concatenate([dz, tz], 1)
|
||||
cl = HDBSCAN(min_cluster_size=max(10, fused.shape[0] // 100), min_samples=5).fit(fused)
|
||||
return cl.labels_, (1.0 - cl.probabilities_).astype(np.float32)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- subset selection (#4)
|
||||
def select_fps(emb: np.ndarray, K: int, seed: int) -> np.ndarray:
|
||||
sel = [int(seed)]
|
||||
d = np.linalg.norm(emb - emb[seed], axis=1)
|
||||
for _ in range(K - 1):
|
||||
i = int(np.argmax(d))
|
||||
sel.append(i)
|
||||
d = np.minimum(d, np.linalg.norm(emb - emb[i], axis=1))
|
||||
return np.array(sorted(set(sel)))
|
||||
|
||||
|
||||
def select_dpp(quality: np.ndarray, emb: np.ndarray, K: int) -> np.ndarray:
|
||||
"""Greedy MAP for a DPP with L = diag(q) S diag(q), S a Gaussian similarity kernel."""
|
||||
from scipy.spatial.distance import cdist
|
||||
q = _norm(quality) + 0.05
|
||||
d2 = cdist(emb, emb, "sqeuclidean")
|
||||
S = np.exp(-d2 / (np.median(d2) + 1e-6))
|
||||
L = (q[:, None] * q[None, :]) * S
|
||||
N = L.shape[0]
|
||||
cis = np.zeros((K, N))
|
||||
di2s = np.diag(L).copy()
|
||||
j = int(np.argmax(di2s))
|
||||
sel = [j]
|
||||
for it in range(1, K):
|
||||
k = it - 1
|
||||
ei = (L[j] - cis[:k, j] @ cis[:k]) / np.sqrt(max(di2s[j], 1e-12))
|
||||
cis[k] = ei
|
||||
di2s = di2s - ei**2
|
||||
di2s[sel] = -np.inf
|
||||
j = int(np.argmax(di2s))
|
||||
if di2s[j] <= 1e-10:
|
||||
break
|
||||
sel.append(j)
|
||||
return np.array(sorted(sel))
|
||||
|
||||
|
||||
def select_greedy_recon(xy0: np.ndarray, D: np.ndarray, K: int) -> np.ndarray:
|
||||
"""Add the point whose displacement is least predictable by RBF-interpolating the
|
||||
already-selected points over frame-0 positions -> targets motion discontinuities."""
|
||||
from scipy.spatial.distance import cdist
|
||||
sel = [int(np.argmax(np.linalg.norm(D - D.mean(0), axis=1)))]
|
||||
for _ in range(K - 1):
|
||||
S = np.array(sel)
|
||||
d2 = cdist(xy0, xy0[S], "sqeuclidean")
|
||||
W = np.exp(-d2 / (np.median(d2) + 1e-6))
|
||||
W /= (W.sum(1, keepdims=True) + 1e-9)
|
||||
err = np.linalg.norm(D - W @ D[S], axis=1)
|
||||
err[S] = -1
|
||||
sel.append(int(np.argmax(err)))
|
||||
return np.array(sorted(sel))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------- render
|
||||
def _spearman(a, b) -> float:
|
||||
from scipy.stats import spearmanr
|
||||
return round(float(spearmanr(a, b).correlation), 3)
|
||||
|
||||
|
||||
def cmd_render(args: argparse.Namespace) -> None:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
import imageio.v2 as imageio
|
||||
|
||||
dino_model, dino_dev = None, "cpu"
|
||||
if not args.no_dino:
|
||||
try:
|
||||
import torch
|
||||
dino_dev = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dino_model = torch.hub.load("facebookresearch/dinov2", "dinov2_vits14").to(dino_dev).eval()
|
||||
print(f"[info] DINOv2 loaded on {dino_dev}", flush=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[info] DINOv2 unavailable ({type(e).__name__}: {e}); skipping semantic panels", flush=True)
|
||||
|
||||
data = Path(args.data_dir)
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
items = json.loads((data / "videos2caption.json").read_text())[:args.limit]
|
||||
turbo = matplotlib.colormaps["turbo"]
|
||||
entries = []
|
||||
for k, it in enumerate(items, 1):
|
||||
stem = Path(it["path"]).stem
|
||||
vpath = str(data / "videos" / it["path"])
|
||||
npz = it.get("points_path") or str(data / "tracks" / f"{stem}.npz")
|
||||
frames = _read_frames(vpath)
|
||||
d = np.load(npz)
|
||||
tracks = d["tracks"].astype(np.float32)[:frames.shape[0]]
|
||||
vis = d["visibility"].astype(np.float32)[:frames.shape[0]]
|
||||
N = tracks.shape[1]
|
||||
x, y = tracks[0, :, 0], tracks[0, :, 1]
|
||||
D = displacement_matrix(tracks)
|
||||
emb = traj_embedding(tracks)
|
||||
|
||||
# scalar scores
|
||||
scores = {
|
||||
"abs_disp": abs_disp(tracks),
|
||||
"affine_residual": affine_residual(tracks, vis),
|
||||
"lowrank_residual": lowrank_residual(tracks),
|
||||
}
|
||||
m_lab, m_out = motion_cluster(emb)
|
||||
scores["motion_outlier"] = m_out
|
||||
k_est = len(np.unique(m_lab[m_lab >= 0])) or 4
|
||||
s_lab, s_res = subspace_cluster(tracks, emb, k_est)
|
||||
scores["subspace_resid"] = s_res
|
||||
d_lab = None
|
||||
if dino_model is not None:
|
||||
try:
|
||||
dfeat = dino_point_features(frames[0], tracks[0], dino_model, dino_dev)
|
||||
d_lab, d_uniq = dino_motion_cluster(dfeat, emb)
|
||||
scores["dino_uniqueness"] = d_uniq
|
||||
except Exception as e: # noqa: BLE001
|
||||
print(f"[info] dino features failed for {stem}: {e}", flush=True)
|
||||
|
||||
# subset selections (#4)
|
||||
seed = int(np.argmax(scores["lowrank_residual"]))
|
||||
subsets = {
|
||||
"dpp": select_dpp(scores["lowrank_residual"], emb, K_SELECT),
|
||||
"fps": select_fps(emb, K_SELECT, seed),
|
||||
"greedy_recon": select_greedy_recon(tracks[0], D, K_SELECT),
|
||||
}
|
||||
|
||||
# ---- figure: row1 scalar scores, row2 cluster maps + subsets
|
||||
label_panels = [("motion_clusters", m_lab), ("subspace_clusters", s_lab)]
|
||||
if d_lab is not None:
|
||||
label_panels.append(("dino_clusters", d_lab))
|
||||
panels = ([("score", n, s)
|
||||
for n, s in scores.items()] + [("labels", n, lab)
|
||||
for n, lab in label_panels] + [("subset", n, idx)
|
||||
for n, idx in subsets.items()])
|
||||
ncol = 5
|
||||
nrow = int(np.ceil(len(panels) / ncol))
|
||||
fig, axes = plt.subplots(nrow, ncol, figsize=(4.4 * ncol, 3.1 * nrow))
|
||||
axes = np.array(axes).reshape(-1)
|
||||
for ax in axes:
|
||||
ax.axis("off")
|
||||
for ax, (kind, name, val) in zip(axes, panels, strict=False):
|
||||
ax.imshow(frames[0])
|
||||
if kind == "score":
|
||||
ax.scatter(x, y, c=_norm(val), cmap="turbo", s=10, vmin=0, vmax=1)
|
||||
ax.set_title(f"{name}", fontsize=10)
|
||||
elif kind == "labels":
|
||||
noise = val < 0
|
||||
ax.scatter(x[noise], y[noise], c="0.5", s=6)
|
||||
ax.scatter(x[~noise], y[~noise], c=val[~noise], cmap="tab20", s=10)
|
||||
ax.set_title(f"{name} ({len(np.unique(val[val >= 0]))} groups)", fontsize=10)
|
||||
else: # subset
|
||||
ax.scatter(x, y, c="0.4", s=3)
|
||||
ax.scatter(x[val], y[val], c="red", s=26, edgecolors="white", linewidths=0.4)
|
||||
ax.set_title(f"{name} (K={len(val)})", fontsize=10)
|
||||
fig.suptitle(f"{stem} | {(it['cap'][0] if isinstance(it.get('cap'), list) else '')[:90]}", fontsize=11)
|
||||
fig.tight_layout()
|
||||
heat = f"{stem}_scores.png"
|
||||
fig.savefig(str(out / heat), dpi=80, bbox_inches="tight")
|
||||
plt.close(fig)
|
||||
|
||||
# ---- overlay colored by lowrank informativeness
|
||||
lr = _norm(scores["lowrank_residual"])
|
||||
pcols = (np.array([turbo(v)[:3] for v in lr]) * 255).astype(np.uint8)
|
||||
G = int(round(N**0.5))
|
||||
sel = (np.arange(N).reshape(G, G)[::max(1, G // 25), ::max(1, G // 25)].reshape(-1) if G *
|
||||
G == N else np.arange(0, N, max(1, N // 625)))
|
||||
ov = _draw_overlay(frames, tracks[:, sel].copy(), vis[:, sel], pcols[sel], 14, 2, 0.5)
|
||||
vid = f"{stem}_lowrank.mp4"
|
||||
imageio.mimsave(str(out / vid), ov, fps=int(it.get("fps", 24)), macro_block_size=1)
|
||||
|
||||
# ---- stats
|
||||
names = list(scores.keys())
|
||||
rho = {f"{a}|{b}": _spearman(scores[a], scores[b]) for i, a in enumerate(names) for b in names[i + 1:]}
|
||||
lr_raw = scores["lowrank_residual"]
|
||||
lo_info = lr_raw <= np.percentile(lr_raw, 50)
|
||||
|
||||
def _cov_generic(idx: Any, m_lab: Any = m_lab, lo_info: Any = lo_info) -> dict:
|
||||
grp = m_lab[idx]
|
||||
ngrp = len(np.unique(grp[grp >= 0]))
|
||||
tot = max(len(np.unique(m_lab[m_lab >= 0])), 1)
|
||||
return {"coverage_of_motion_groups": f"{ngrp}/{tot}", "frac_generic": round(float(lo_info[idx].mean()), 3)}
|
||||
|
||||
stats = {
|
||||
"id": stem,
|
||||
"n_points": int(N),
|
||||
"camera_motion": global_affine_motion(tracks, vis), # ~0 = fixed cam, high = moving
|
||||
"n_motion_groups": int(len(np.unique(m_lab[m_lab >= 0]))),
|
||||
"noise_frac": round(float((m_lab < 0).mean()), 3),
|
||||
"spearman": rho,
|
||||
"subset_quality": {
|
||||
name: _cov_generic(idx)
|
||||
for name, idx in subsets.items()
|
||||
},
|
||||
}
|
||||
entries.append({"id": stem, "heat": heat, "video": vid, "stats": stats})
|
||||
print(
|
||||
f"[info] [{k}/{len(items)}] {stem}: cam_motion={stats['camera_motion']}, "
|
||||
f"{stats['n_motion_groups']} motion groups, noise={stats['noise_frac']}, "
|
||||
f"rho(abs|lowrank)={rho.get('abs_disp|lowrank_residual')}",
|
||||
flush=True)
|
||||
|
||||
(out / "manifest.json").write_text(json.dumps(entries, indent=2))
|
||||
print(f"[info] rendered {len(entries)} clips -> {out}", flush=True)
|
||||
|
||||
|
||||
def cmd_serve(args: argparse.Namespace) -> None:
|
||||
import gradio as gr
|
||||
viz = Path(args.viz_dir).resolve()
|
||||
entries = json.loads((viz / "manifest.json").read_text())
|
||||
gal = [(str(viz / e["heat"]), e["id"]) for e in entries]
|
||||
|
||||
def show(evt: gr.SelectData):
|
||||
e = entries[evt.index]
|
||||
return str(viz / e["heat"]), str(viz / e["video"]), json.dumps(e["stats"], indent=2)
|
||||
|
||||
with gr.Blocks(title="Track informativeness bake-off") as demo:
|
||||
gr.Markdown(
|
||||
"### Which point tracks are useful? Statistics vs clustering vs embeddings vs subset selection\n"
|
||||
"**Row 1 (scores, bright = informative):** `abs_disp` raw motion · `affine_residual` after "
|
||||
"removing camera affine · `lowrank_residual` unique vs all (one subspace) · `motion_outlier` "
|
||||
"HDBSCAN density outlier · `subspace_resid` per-object lowrank · `dino_uniqueness` semantic⊕motion. "
|
||||
"**Row 2:** `motion_clusters` / `subspace_clusters` / `dino_clusters` (colored groups, gray = noise) "
|
||||
"and the diverse K-subsets `dpp` / `fps` / `greedy_recon` (red = chosen). "
|
||||
"Click a clip -> panels + overlay colored by low-rank informativeness. `stats.subset_quality` shows "
|
||||
"how well each subset covers the motion groups and how many chosen points are generic.")
|
||||
with gr.Row():
|
||||
g = gr.Gallery(value=gal, columns=2, height=640, label="clips (click)")
|
||||
with gr.Column():
|
||||
heat = gr.Image(label="method panels")
|
||||
vid = gr.Video(label="overlay colored by low-rank informativeness")
|
||||
meta = gr.Code(label="stats", language="json")
|
||||
g.select(show, None, [heat, vid, meta])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share, allowed_paths=[str(viz)])
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
r = sub.add_parser("render")
|
||||
r.add_argument("--data-dir", required=True)
|
||||
r.add_argument("--out", required=True)
|
||||
r.add_argument("--limit", type=int, default=8)
|
||||
r.add_argument("--no-dino", action="store_true", help="skip the DINOv2 semantic panels")
|
||||
r.set_defaults(func=cmd_render)
|
||||
s = sub.add_parser("serve")
|
||||
s.add_argument("--viz-dir", required=True)
|
||||
s.add_argument("--host", default="0.0.0.0")
|
||||
s.add_argument("--port", type=int, default=7883)
|
||||
s.add_argument("--share", action="store_true")
|
||||
s.set_defaults(func=cmd_serve)
|
||||
a = p.parse_args()
|
||||
a.func(a)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/usr/bin/env bash
|
||||
# Overlap tracking with the slow self-healing download: idempotent multi-node passes
|
||||
# over the growing videos/ dir until download done (98 parts) AND ~all tracked.
|
||||
# MUST run inside srun on a compute node (login-node `sleep` gets killed by the harness).
|
||||
# Waits for any in-flight tracking pass first, so it's safe to relaunch (no double-submit).
|
||||
set +e
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
cd "$WORK/FastVideo"
|
||||
DC(){ ls "$WORK"/openvid_1m/_extracted/*.done 2>/dev/null | wc -l; }
|
||||
RUNNING(){ squeue -u vlm-s4duan -h -n cotracker_dp -o %i 2>/dev/null; }
|
||||
NODES=${NODES:-8}
|
||||
while true; do
|
||||
while [ -n "$(RUNNING)" ]; do sleep 90; done # let in-flight pass finish
|
||||
ls "$WORK"/openvid_1m/videos/*.mp4 > "$WORK"/openvid_1m/videos_all.txt 2>/dev/null
|
||||
nv=$(wc -l < "$WORK"/openvid_1m/videos_all.txt 2>/dev/null); nv=${nv:-0}
|
||||
nt=$(ls "$WORK"/openvid_1m/tracks/*.npz 2>/dev/null | wc -l)
|
||||
echo "[$(date +%H:%M)] parts=$(DC)/98 videos=$nv tracks=$nt"
|
||||
if [ "$(DC)" -ge 98 ] && [ "$nt" -ge "$((nv - nv/40 - 200))" ]; then
|
||||
echo "TRACK_LOOP_DONE videos=$nv tracks=$nt"; break; fi
|
||||
if [ "$nv" -le "$((nt + 300))" ]; then echo " little new; wait 300s"; sleep 300; continue; fi
|
||||
jid=$(CLIPS_DIR="$WORK"/openvid_1m/clips VIDEO_LIST="$WORK"/openvid_1m/videos_all.txt \
|
||||
OUT_DIR="$WORK"/openvid_1m/tracks NODES=$NODES PROCS_PER_GPU=2 FPS=24 NUM_FRAMES=121 GRID=50 \
|
||||
bash data_pipeline/run_tracks_slurm.sh 2>&1 | grep -oE "[0-9]+$" | tail -1)
|
||||
echo " submitted track job $jid over $nv videos"
|
||||
[ -z "$jid" ] && sleep 180
|
||||
done
|
||||
@@ -0,0 +1,117 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Single-checkpoint *swap sensitivity* (one process per checkpoint).
|
||||
|
||||
Linearized, generation-free version of the controllability EPE test: at a fixed
|
||||
noise + timestep, measure how much the predicted velocity ``v`` moves when we swap
|
||||
ONLY one conditioner and hold the rest fixed:
|
||||
|
||||
S_track = ||v(swap tracks) - v(base)|| / ||v(base)||
|
||||
S_ff = ||v(swap first-frame) - v(base)|| / ||v(base)||
|
||||
S_text = ||v(zero text) - v(base)|| / ||v(base)||
|
||||
|
||||
Averaged over a few timesteps and the given clips. This captures the *backbone's*
|
||||
responsiveness to each input (unlike weight-norm or input-projection magnitude,
|
||||
which only see the input side). Expectation for the control50 run: S_track high at
|
||||
step 500, decaying by 6000 (control prior forgotten), while S_ff stays high.
|
||||
|
||||
One checkpoint per process (loads via the framework loader, so weights/DTensor are
|
||||
handled correctly). Launch once per --step across GPUs.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def _v(model: Any, latents: torch.Tensor, ff: torch.Tensor, tp: torch.Tensor, tv: torch.Tensor, txt: torch.Tensor,
|
||||
mask: torch.Tensor, img: torch.Tensor, ts: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
with torch.no_grad(), torch.autocast(model.device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=ts, attn_metadata=None):
|
||||
v = model.transformer(hidden_states=model_in,
|
||||
encoder_hidden_states=txt,
|
||||
encoder_attention_mask=mask,
|
||||
timestep=ts,
|
||||
encoder_hidden_states_image=img,
|
||||
track_points=tp,
|
||||
track_visibility=tv,
|
||||
return_dict=False)
|
||||
return v.float()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True, help="single export dir (export_<step>)")
|
||||
p.add_argument("--step", type=int, default=-1, help="label only")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data",
|
||||
default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset")
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--timestep-idx", type=int, nargs="+", default=[2, 15, 27])
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
args = p.parse_args()
|
||||
|
||||
dtype = torch.bfloat16
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
device = model.device
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(30, device=device)
|
||||
tlist = [sched.timesteps[i].reshape(1).to(device, dtype) for i in args.timestep_idx]
|
||||
|
||||
pairs = []
|
||||
for k, ci in enumerate(args.clips):
|
||||
s = samples[k]
|
||||
sw = samples[(k + 1) % len(samples)]
|
||||
_, _, T, H, W = s["first_frame_latent"].shape
|
||||
g = torch.Generator(device="cpu").manual_seed(args.seed + ci)
|
||||
latents = torch.randn((1, 16, T, H, W), generator=g, dtype=torch.float32).to(device)
|
||||
|
||||
def d(x: torch.Tensor) -> torch.Tensor:
|
||||
return x.to(device, dtype)
|
||||
|
||||
pairs.append(
|
||||
dict(lat=latents,
|
||||
ff=d(s["first_frame_latent"]),
|
||||
ff_sw=d(sw["first_frame_latent"]),
|
||||
tp=d(s["track_points"]),
|
||||
tv=d(s["track_visibility"]),
|
||||
tp_sw=d(sw["track_points"]),
|
||||
tv_sw=d(sw["track_visibility"]),
|
||||
txt=d(s["text_embedding"]),
|
||||
mask=d(s["text_attention_mask"]),
|
||||
img=d(s["clip_feature"])))
|
||||
|
||||
st = sf = sx = 0.0
|
||||
n = 0
|
||||
for pr in pairs:
|
||||
for ts in tlist:
|
||||
vb = _v(model, pr["lat"], pr["ff"], pr["tp"], pr["tv"], pr["txt"], pr["mask"], pr["img"], ts, dtype)
|
||||
nb = vb.norm() + 1e-6
|
||||
vt = _v(model, pr["lat"], pr["ff"], pr["tp_sw"], pr["tv_sw"], pr["txt"], pr["mask"], pr["img"], ts, dtype)
|
||||
vf = _v(model, pr["lat"], pr["ff_sw"], pr["tp"], pr["tv"], pr["txt"], pr["mask"], pr["img"], ts, dtype)
|
||||
vx = _v(model, pr["lat"], pr["ff"], pr["tp"], pr["tv"], torch.zeros_like(pr["txt"]), pr["mask"], pr["img"],
|
||||
ts, dtype)
|
||||
st += float((vt - vb).norm() / nb)
|
||||
sf += float((vf - vb).norm() / nb)
|
||||
sx += float((vx - vb).norm() / nb)
|
||||
n += 1
|
||||
print(f"RESULT step={args.step} S_track={st/n:.4f} S_ff={sf/n:.4f} S_text={sx/n:.4f}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,176 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Standalone TrackWan inference: generate a video from (first frame + text + tracks).
|
||||
|
||||
Loads a checkpoint exported with ``dcp_to_diffusers`` into the *same*
|
||||
``WanTrackModel`` wrapper used for training (so train/inference parity is exact),
|
||||
then runs a flow-matching denoise loop feeding the point-track control. Used by
|
||||
the controllability tests and the Gradio app.
|
||||
|
||||
Backend only -- no track authoring here (see ``synthetic_tracks.py``) and no
|
||||
metrics (see ``fastvideo/eval/metrics/motion/cotracker_epe``).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def _ensure_dist_env() -> None:
|
||||
for k, v in [("RANK", "0"), ("LOCAL_RANK", "0"), ("WORLD_SIZE", "1"), ("MASTER_ADDR", "127.0.0.1"),
|
||||
("MASTER_PORT", "29600")]:
|
||||
os.environ.setdefault(k, v)
|
||||
|
||||
|
||||
def load_trackwan(export_dir: str, yaml_path: str) -> tuple[Any, Any]:
|
||||
"""Build a ``WanTrackModel`` with the exported (trained) weights on 1 GPU."""
|
||||
_ensure_dist_env()
|
||||
from fastvideo.distributed import (
|
||||
maybe_init_distributed_environment_and_model_parallel, )
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
from fastvideo.train.models.wantrack.wantrack import WanTrackModel
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
|
||||
cfg = load_run_config(yaml_path)
|
||||
tc = cfg.training
|
||||
tc.distributed.tp_size = 1
|
||||
tc.distributed.sp_size = 1
|
||||
tc.distributed.num_gpus = 1
|
||||
tc.distributed.hsdp_replicate_dim = 1
|
||||
tc.distributed.hsdp_shard_dim = 1
|
||||
tc.model_path = export_dir # vae + components loaded from the export
|
||||
|
||||
model = WanTrackModel(
|
||||
init_from=export_dir,
|
||||
training_config=tc,
|
||||
trainable=False,
|
||||
flow_shift=float(tc.pipeline_config.flow_shift),
|
||||
enable_gradient_checkpointing_type=None,
|
||||
)
|
||||
model.init_preprocessors(tc)
|
||||
model.transformer.eval()
|
||||
return model, tc
|
||||
|
||||
|
||||
def load_conditioning_from_parquet(data_path: str, indices: list[int], text_len: int) -> list[dict[str, Any]]:
|
||||
"""Pull (first_frame_latent, text, vae_latent, GT tracks, caption) for clips."""
|
||||
import glob
|
||||
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
|
||||
files = sorted(glob.glob(os.path.join(data_path, "**", "*.parquet"), recursive=True))
|
||||
rows: list[dict[str, Any]] = []
|
||||
for f in files:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
sel = [rows[i] for i in indices]
|
||||
batch = collate_rows_from_parquet_schema(sel,
|
||||
pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len),
|
||||
cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in sel]
|
||||
out = []
|
||||
for i in range(len(sel)):
|
||||
out.append({
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(), # [1,SeqLen,Dim] CLIP frame-0
|
||||
"track_points": batch["track_points"][i:i + 1].clone(), # [1,T,N,2] normalized
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(), # [1,T,N]
|
||||
"caption": str(infos[i].get("caption", "")),
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def generate(model: Any,
|
||||
*,
|
||||
first_frame_latent: torch.Tensor,
|
||||
text_embedding: torch.Tensor,
|
||||
text_attention_mask: torch.Tensor,
|
||||
track_points: torch.Tensor | None,
|
||||
track_visibility: torch.Tensor | None,
|
||||
clip_feature: torch.Tensor | None = None,
|
||||
num_steps: int = 30,
|
||||
seed: int = 0,
|
||||
guidance_scale: float = 1.0) -> torch.Tensor:
|
||||
"""Denoise from noise -> normalized latents [1,16,T,H,W]. Tracks may be None.
|
||||
|
||||
``guidance_scale`` > 1 = MOTION classifier-free guidance:
|
||||
v = v(no-tracks) + s*(v(tracks) - v(no-tracks)). The no-tracks branch is the
|
||||
base first-frame+text I2V prediction (track channels zeroed), so it works on any
|
||||
checkpoint (no training-time motion dropout needed)."""
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
ff = first_frame_latent.to(device, dtype) # [1,16,T,H,W] (already normalized)
|
||||
cond20 = model._build_i2v_cond_concat(ff) # [1,20,T,H,W]
|
||||
txt = text_embedding.to(device, dtype)
|
||||
mask = text_attention_mask.to(device, dtype)
|
||||
tp = track_points.to(device, dtype) if track_points is not None else None
|
||||
tv = track_visibility.to(device, dtype) if track_visibility is not None else None
|
||||
img = clip_feature.to(device, dtype) if clip_feature is not None else None
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(int(num_steps), device=device)
|
||||
use_cfg = guidance_scale != 1.0 and tp is not None
|
||||
for t in sched.timesteps:
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1) # [1,36,T,H,W]
|
||||
ts = t.reshape(1).to(device, dtype)
|
||||
|
||||
def _fwd(tpp: torch.Tensor | None,
|
||||
tvv: torch.Tensor | None,
|
||||
mi: torch.Tensor = model_in,
|
||||
tsv: torch.Tensor = ts) -> torch.Tensor:
|
||||
with torch.autocast(device.type, dtype=dtype), set_forward_context(current_timestep=tsv,
|
||||
attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi,
|
||||
encoder_hidden_states=txt,
|
||||
encoder_attention_mask=mask,
|
||||
timestep=tsv,
|
||||
encoder_hidden_states_image=img,
|
||||
track_points=tpp,
|
||||
track_visibility=tvv,
|
||||
return_dict=False)
|
||||
|
||||
if use_cfg:
|
||||
v_cond = _fwd(tp, tv).float()
|
||||
v_uncond = _fwd(None, None).float() # base I2V (track channels -> 0)
|
||||
v = v_uncond + guidance_scale * (v_cond - v_uncond)
|
||||
else:
|
||||
v = _fwd(tp, tv).float()
|
||||
latents = sched.step(v, t, latents.float(), return_dict=False)[0]
|
||||
return latents
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_to_pixels(model: Any, latents: torch.Tensor) -> np.ndarray:
|
||||
"""Normalized latents [1,16,T,H,W] -> uint8 frames [T,H,W,3]."""
|
||||
px = model.decode_latents(latents.permute(0, 2, 1, 3, 4))[0] # [3,T,H,W] in [0,1]
|
||||
video = (px.clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
return np.transpose(video, (1, 2, 3, 0)) # [T,H,W,3]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def decode_reference(model: Any, vae_latent: torch.Tensor) -> np.ndarray:
|
||||
"""Decode the raw GT vae_latent (needs normalize first) -> uint8 [T,H,W,3]."""
|
||||
from fastvideo.training.training_utils import normalize_dit_input
|
||||
raw = vae_latent.to(model.device, torch.bfloat16)
|
||||
norm = normalize_dit_input("wan", raw, model.vae)
|
||||
px = model.decode_latents(norm.permute(0, 2, 1, 3, 4))[0]
|
||||
video = (px.clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
return np.transpose(video, (1, 2, 3, 0))
|
||||
@@ -0,0 +1,325 @@
|
||||
#!/usr/bin/env python3
|
||||
"""2-track validation ablation.
|
||||
|
||||
For each val sample, keep exactly 2 tracks (one per top-motion object) and zero
|
||||
the rest. This is the extreme sparse-inference case: what the user actually gets
|
||||
when they draw 2 traces at inference. Compares checkpoints side-by-side.
|
||||
|
||||
Uses paper CFG (wt=3.0, wm=1.5) unless overridden. Runs `_sample`-style denoise
|
||||
matching track_validation.py exactly.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/two_track_val.py \
|
||||
--model-dir /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800 \
|
||||
--yaml examples/train/scenario/worldmodel/finetune_wantrack_synth_stage2_paperLR.yaml \
|
||||
--val-parquet /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset \
|
||||
--wandb-run-name twotrack_mb4800
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__)))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
from fastvideo.forward_context import set_forward_context # noqa: E402
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import ( # noqa: E402
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
|
||||
def load_val_samples(parquet_dir: str, text_len: int) -> list:
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
fs = sorted(glob.glob(os.path.join(parquet_dir, "**", "*.parquet"), recursive=True))
|
||||
if not fs:
|
||||
raise FileNotFoundError(parquet_dir)
|
||||
rows = []
|
||||
for f in fs:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
batch = collate_rows_from_parquet_schema(rows, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in rows]
|
||||
return [{
|
||||
"id": str(infos[i].get("id", f"clip{i:03d}")) if i < len(infos) else f"clip{i:03d}",
|
||||
"file_name": rows[i].get("file_name", ""),
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i:i + 1].clone(), # [1,T,N,2] normalised
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(),
|
||||
"object_ids": batch["object_ids"][i:i + 1].clone() if "object_ids" in batch else None,
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
} for i in range(len(rows))]
|
||||
|
||||
|
||||
def pick_two_tracks(sample: dict, seed: int, num_bg: int = 0,
|
||||
num_fg: int = 2) -> tuple[torch.Tensor, torch.Tensor, list]:
|
||||
"""Return (tp_2, tv_2, info) — ``num_fg`` foreground tracks + ``num_bg`` bg anchor tracks.
|
||||
|
||||
Foreground picks require the track to STAY IN FRAME [0,1] for all 121 frames — no
|
||||
point in showing traces that fly off the canvas.
|
||||
"""
|
||||
tp = sample["track_points"].float() # [1,T,N,2]
|
||||
tv = sample["track_visibility"].float() # [1,T,N]
|
||||
oid = sample["object_ids"] # [1,N]
|
||||
if oid is None:
|
||||
raise ValueError("no object_ids in val sample")
|
||||
oid = oid[0].to(torch.int64) # [N]
|
||||
|
||||
tp0 = tp[0] # [T, N, 2]
|
||||
N = tp0.shape[1]
|
||||
motion_per_track = torch.linalg.norm(tp0[-1] - tp0[0], dim=-1) # [N]
|
||||
|
||||
# Which tracks stay fully within [0,1] on both x and y across ALL frames?
|
||||
in_frame = ((tp0 >= 0.0) & (tp0 <= 1.0)).all(dim=-1).all(dim=0) # [N]
|
||||
|
||||
# Rank object_ids by max in-frame within-object track motion.
|
||||
fg_oids = torch.unique(oid[oid >= 0])
|
||||
obj_best = [] # (oid, best_track_idx, best_motion)
|
||||
for o in fg_oids.tolist():
|
||||
idxs = ((oid == o) & in_frame & (tv[0, 0] > 0.5)).nonzero(as_tuple=True)[0]
|
||||
if idxs.numel() == 0:
|
||||
continue
|
||||
m = motion_per_track[idxs]
|
||||
best = idxs[int(torch.argmax(m).item())]
|
||||
obj_best.append((o, int(best.item()), float(m.max().item())))
|
||||
obj_best.sort(key=lambda x: -x[2])
|
||||
picked_tracks = [idx for _, idx, _ in obj_best[:num_fg]]
|
||||
info = [{"oid": int(o), "motion_max": mm, "track_idx": pk}
|
||||
for (o, pk, mm) in obj_best[:num_fg]]
|
||||
rng = np.random.RandomState(seed)
|
||||
vis0 = tv[0, 0] # [N] visibility at frame 0
|
||||
|
||||
# Add num_bg background anchor tracks — near-static (bottom of motion), visible at frame 0.
|
||||
bg_picked = []
|
||||
if num_bg > 0:
|
||||
bg_pool = ((oid == -1) & (vis0 > 0.5)).nonzero(as_tuple=True)[0]
|
||||
if bg_pool.numel() > 0:
|
||||
# Sort bg tracks by motion (ascending), take lowest-motion ones and pick uniformly
|
||||
bg_motion = motion_per_track[bg_pool]
|
||||
n_pool = int(min(bg_pool.numel(), max(num_bg * 20, num_bg)))
|
||||
top_idxs = torch.argsort(bg_motion)[:n_pool]
|
||||
candidate = bg_pool[top_idxs]
|
||||
if candidate.numel() >= num_bg:
|
||||
sel = rng.choice(candidate.numel(), num_bg, replace=False)
|
||||
bg_picked = [int(candidate[k].item()) for k in sel]
|
||||
else:
|
||||
bg_picked = [int(k.item()) for k in candidate]
|
||||
for k in bg_picked:
|
||||
info.append({"oid": -1, "motion_max": float(motion_per_track[k].item()), "track_idx": k})
|
||||
|
||||
# Build zeroed visibility: only picked tracks visible.
|
||||
tv_new = torch.zeros_like(tv)
|
||||
for k in picked_tracks + bg_picked:
|
||||
tv_new[0, :, k] = tv[0, :, k]
|
||||
|
||||
return tp, tv_new, info
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_denoise(model, sample: dict, tp: torch.Tensor, tv: torch.Tensor,
|
||||
*, w_text: float, w_motion: float, num_steps: int, seed: int) -> np.ndarray:
|
||||
"""Reproduce track_validation._sample denoise EXACTLY."""
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
|
||||
ff = sample["first_frame_latent"].to(device, dtype)
|
||||
txt = sample["text_embedding"].to(device, dtype)
|
||||
mask = sample["text_attention_mask"].to(device, dtype)
|
||||
clip = sample["clip_feature"].to(device, dtype)
|
||||
tp_d = tp.to(device, dtype)
|
||||
tv_d = tv.to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(int(num_steps), device=device)
|
||||
|
||||
cfg_on = (w_text != 1.0) or (w_motion != 1.0)
|
||||
txt_null = torch.zeros_like(txt) if cfg_on else None
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, ts_):
|
||||
with torch.autocast(device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=ts_, attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=ts_,
|
||||
encoder_hidden_states_image=clip,
|
||||
track_points=tp_e, track_visibility=tv_e,
|
||||
return_dict=False)
|
||||
|
||||
for tt in sched.timesteps:
|
||||
mi = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
ts_ = tt.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp_d, tv_d, mi, ts_)
|
||||
if not cfg_on:
|
||||
v = v_full
|
||||
elif w_text != 1.0 and w_motion == 1.0:
|
||||
v_no_text = _fwd(txt_null, tp_d, tv_d, mi, ts_)
|
||||
v = v_no_text + w_text * (v_full - v_no_text)
|
||||
elif w_text == 1.0 and w_motion != 1.0:
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
v = v_no_motion + w_motion * (v_full - v_no_motion)
|
||||
else:
|
||||
v_no_text = _fwd(txt_null, tp_d, tv_d, mi, ts_)
|
||||
v_no_motion = _fwd(txt, None, None, mi, ts_)
|
||||
alpha = w_text / max(w_text + w_motion, 1e-6)
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + w_text * (v_full - v_no_text) + w_motion * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), tt, latents.float(), return_dict=False)[0]
|
||||
|
||||
return twi.decode_to_pixels(model, latents)
|
||||
|
||||
|
||||
def draw_track_overlay(video: np.ndarray, tp: torch.Tensor, tv: torch.Tensor,
|
||||
info: list | None = None) -> np.ndarray:
|
||||
"""Draw kept tracks on the video. Foreground = red/blue, background anchors = gray."""
|
||||
from PIL import Image, ImageDraw
|
||||
T, H, W, _ = video.shape
|
||||
tp_np = tp[0].cpu().numpy() # [T, N, 2]
|
||||
tv_np = tv[0].cpu().numpy() # [T, N]
|
||||
N = tp_np.shape[1]
|
||||
kept = [k for k in range(N) if tv_np[:, k].max() > 0.5]
|
||||
fg_colors = [(255, 60, 60), (60, 200, 255)]
|
||||
bg_color = (180, 180, 180)
|
||||
# figure out which are fg vs bg by looking at info if given
|
||||
is_bg_by_idx = {}
|
||||
if info is not None:
|
||||
for entry in info:
|
||||
is_bg_by_idx[entry["track_idx"]] = (entry["oid"] == -1)
|
||||
out = []
|
||||
fg_seen = 0
|
||||
for t in range(T):
|
||||
img = Image.fromarray(np.ascontiguousarray(video[t]))
|
||||
dr = ImageDraw.Draw(img)
|
||||
fg_seen = 0
|
||||
for k in kept:
|
||||
if tv_np[t, k] <= 0.5:
|
||||
continue
|
||||
x = float(tp_np[t, k, 0] * W)
|
||||
y = float(tp_np[t, k, 1] * H)
|
||||
if is_bg_by_idx.get(k, False):
|
||||
col = bg_color
|
||||
r = 3
|
||||
else:
|
||||
col = fg_colors[fg_seen % len(fg_colors)]
|
||||
fg_seen += 1
|
||||
r = 5
|
||||
dr.ellipse([x - r, y - r, x + r, y + r], fill=col, outline=(255, 255, 255))
|
||||
out.append(np.asarray(img))
|
||||
return np.stack(out)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--model-dir", required=True)
|
||||
p.add_argument("--yaml", required=True)
|
||||
p.add_argument("--val-parquet", required=True)
|
||||
p.add_argument("--wandb-project", default="wantrack-bidir")
|
||||
p.add_argument("--wandb-run-name", required=True)
|
||||
p.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/two_track_val")
|
||||
p.add_argument("--w-text", type=float, default=3.0)
|
||||
p.add_argument("--w-motion", type=float, default=1.5)
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
p.add_argument("--fps", type=int, default=24)
|
||||
p.add_argument("--num-fg-tracks", type=int, default=2,
|
||||
help="Number of foreground (moving object) tracks to keep.")
|
||||
p.add_argument("--num-bg-tracks", type=int, default=0,
|
||||
help="Extra static bg anchor tracks in addition to the fg ones.")
|
||||
p.add_argument("--swap-text-from", type=int, default=None,
|
||||
help="For every sample, replace text_embedding+mask with THIS sample idx's text.")
|
||||
p.add_argument("--swap-clip-from", type=int, default=None,
|
||||
help="For every sample, replace clip_feature with THIS sample idx's clip.")
|
||||
p.add_argument("--swap-ff-from", type=int, default=None,
|
||||
help="For every sample, replace first_frame_latent with THIS sample idx's ff.")
|
||||
args = p.parse_args()
|
||||
|
||||
out_dir = Path(args.out_dir) / args.wandb_run_name
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
import wandb
|
||||
|
||||
print(f"[2trk] loading model {args.model_dir}", flush=True)
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(getattr(tc.data, "text_padding_length", 256))
|
||||
print(f"[2trk] text_len={text_len}", flush=True)
|
||||
|
||||
print(f"[2trk] loading val samples", flush=True)
|
||||
samples = load_val_samples(args.val_parquet, text_len)
|
||||
print(f"[2trk] {len(samples)} val samples", flush=True)
|
||||
|
||||
run = wandb.init(project=args.wandb_project, name=args.wandb_run_name,
|
||||
config={"model_dir": args.model_dir, "w_text": args.w_text,
|
||||
"w_motion": args.w_motion, "steps": args.steps})
|
||||
|
||||
for i, s in enumerate(samples):
|
||||
sid = s["id"] or f"clip{i:03d}"
|
||||
print(f"\n[2trk] === sample {i}: {sid} ({s['file_name'][:40]}) ===", flush=True)
|
||||
# Optional cross-sample field swaps (isolate which condition drives quality)
|
||||
if args.swap_text_from is not None:
|
||||
src = samples[args.swap_text_from]
|
||||
s = {**s, "text_embedding": src["text_embedding"].clone(),
|
||||
"text_attention_mask": src["text_attention_mask"].clone()}
|
||||
print(f"[2trk] SWAPPED text from sample {args.swap_text_from}", flush=True)
|
||||
if args.swap_clip_from is not None:
|
||||
src = samples[args.swap_clip_from]
|
||||
s = {**s, "clip_feature": src["clip_feature"].clone()}
|
||||
print(f"[2trk] SWAPPED clip from sample {args.swap_clip_from}", flush=True)
|
||||
if args.swap_ff_from is not None:
|
||||
src = samples[args.swap_ff_from]
|
||||
s = {**s, "first_frame_latent": src["first_frame_latent"].clone()}
|
||||
print(f"[2trk] SWAPPED first_frame_latent from sample {args.swap_ff_from}", flush=True)
|
||||
tp2, tv2, info = pick_two_tracks(s, seed=args.seed + i,
|
||||
num_fg=args.num_fg_tracks, num_bg=args.num_bg_tracks)
|
||||
print(f"[2trk] picked: {info}", flush=True)
|
||||
|
||||
t0 = time.time()
|
||||
frames = sample_denoise(model, s, tp2, tv2, w_text=args.w_text, w_motion=args.w_motion,
|
||||
num_steps=args.steps, seed=args.seed + i)
|
||||
dt = time.time() - t0
|
||||
print(f"[2trk] denoise {dt:.1f}s -> frames {frames.shape}", flush=True)
|
||||
|
||||
overlay = draw_track_overlay(frames, tp2, tv2, info)
|
||||
fn_raw = out_dir / f"{sid}_gen.mp4"
|
||||
fn_ov = out_dir / f"{sid}_gen_overlay.mp4"
|
||||
imageio.mimsave(str(fn_raw), frames, fps=args.fps, macro_block_size=1)
|
||||
imageio.mimsave(str(fn_ov), overlay, fps=args.fps, macro_block_size=1)
|
||||
|
||||
# GT for reference
|
||||
try:
|
||||
ref = twi.decode_reference(model, s["vae_latent"].to(model.device, torch.bfloat16))
|
||||
ref_ov = draw_track_overlay(ref, tp2, tv2, info)
|
||||
fn_ref = out_dir / f"{sid}_gt_overlay.mp4"
|
||||
imageio.mimsave(str(fn_ref), ref_ov, fps=args.fps, macro_block_size=1)
|
||||
wandb.log({
|
||||
f"sample_{i:02d}_gen": wandb.Video(str(fn_ov), fps=args.fps, format="mp4"),
|
||||
f"sample_{i:02d}_gt": wandb.Video(str(fn_ref), fps=args.fps, format="mp4"),
|
||||
"sample": i, "caption": s["caption"][:200], "picked": str(info),
|
||||
})
|
||||
except Exception as e:
|
||||
print(f"[2trk] gt fail: {e}", flush=True)
|
||||
wandb.log({
|
||||
f"sample_{i:02d}_gen": wandb.Video(str(fn_ov), fps=args.fps, format="mp4"),
|
||||
"sample": i, "caption": s["caption"][:200], "picked": str(info),
|
||||
})
|
||||
|
||||
wandb.finish()
|
||||
print("[2trk] DONE", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,60 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Re-run validation as a clean GT-vs-generated side-by-side (no track overlay).
|
||||
|
||||
For each clip: decode the GT clip, generate with its GT tracks, and write a
|
||||
left=GT / right=generated side-by-side video + a first-frame comparison PNG, so
|
||||
the first-frame match and overall reconstruction are easy to eyeball.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, os.path.dirname(__file__))
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True)
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track/combined_parquet_dataset")
|
||||
p.add_argument("--out", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/val_compare_4k")
|
||||
p.add_argument("--fig", default="research_log/figures")
|
||||
p.add_argument("--clips", type=int, nargs="+", default=[0, 1])
|
||||
p.add_argument("--steps", type=int, default=30)
|
||||
p.add_argument("--seed", type=int, default=1000)
|
||||
args = p.parse_args()
|
||||
os.makedirs(args.out, exist_ok=True)
|
||||
os.makedirs(args.fig, exist_ok=True)
|
||||
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, args.clips, text_len)
|
||||
|
||||
for ci, s in zip(args.clips, samples):
|
||||
Tpx = (s["first_frame_latent"].shape[2] - 1) * ratio + 1
|
||||
gt = twi.decode_reference(model, s["vae_latent"]) # [T,H,W,3]
|
||||
lat = twi.generate(model, first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"], text_attention_mask=s["text_attention_mask"],
|
||||
track_points=s["track_points"][:, :Tpx], track_visibility=s["track_visibility"][:, :Tpx],
|
||||
num_steps=args.steps, seed=args.seed)
|
||||
gen = twi.decode_to_pixels(model, lat)
|
||||
T = min(len(gt), len(gen))
|
||||
side = np.concatenate([gt[:T], gen[:T]], axis=2) # [T,H,2W,3] (GT|gen)
|
||||
imageio.mimsave(os.path.join(args.out, f"val_clip{ci}_GTleft_GENright.mp4"),
|
||||
list(side), fps=24, macro_block_size=1)
|
||||
ff = np.concatenate([gt[0], gen[0]], axis=1)
|
||||
imageio.imwrite(os.path.join(args.fig, f"val_firstframe_clip{ci}_GTleft_GENright.png"), ff)
|
||||
mse0 = float(((gt[0].astype(np.float32) - gen[0].astype(np.float32)) ** 2).mean())
|
||||
print(f"[clip {ci}] first-frame MSE(GT,gen)={mse0:8.1f} -> {args.out}/val_clip{ci}_GTleft_GENright.mp4", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,168 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Verify the 720p i2v_track parquet set against the 480p reference.
|
||||
|
||||
Checks:
|
||||
1. row count, no duplicate ``file_name``
|
||||
2. clip-id set is EXACTLY the 480p set (and the source manifest)
|
||||
3. ``vae_latent_shape`` / ``first_frame_latent_shape`` == [16, 31, 90, 160] on every row
|
||||
(shape columns are tiny, so this is checked exhaustively, not sampled)
|
||||
4. schema field names + dtypes identical to the 480p set
|
||||
5. on --sample-tracks random clips: track_points / track_visibility / object_ids /
|
||||
track_weights bytes IDENTICAL to the 480p row for the same file_name
|
||||
|
||||
Latent-decode sanity is a separate tool: data_pipeline/check_720p_latents.py
|
||||
|
||||
Usage:
|
||||
python data_pipeline/verify_720p_parquet.py \
|
||||
--new /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset_720p \
|
||||
--ref /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset \
|
||||
--manifest /mnt/lustre/vlm-s4duan/openvid_1m/videos2caption.json \
|
||||
--expect-shape 16,31,90,160 --sample-tracks 8
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import random
|
||||
from collections import Counter
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
LIGHT_COLS = ["file_name", "id", "vae_latent_shape", "first_frame_latent_shape", "width", "height", "num_frames", "fps"]
|
||||
TRACK_FIELDS = ["track_points", "track_visibility", "object_ids", "track_weights"]
|
||||
|
||||
|
||||
def _files(root: str) -> list[str]:
|
||||
return sorted(glob.glob(f"{root}/**/*.parquet", recursive=True))
|
||||
|
||||
|
||||
def _scan(files: list[str], workers: int = 32):
|
||||
"""-> (rows, shape_counter, ffshape_counter, schema_repr). Reads only tiny columns."""
|
||||
rows: list[tuple[str, str]] = []
|
||||
shapes: Counter = Counter()
|
||||
ffshapes: Counter = Counter()
|
||||
schema_repr = None
|
||||
|
||||
def one(f):
|
||||
t = pq.read_table(f, columns=LIGHT_COLS)
|
||||
return f, t
|
||||
|
||||
with ThreadPoolExecutor(workers) as ex:
|
||||
for i, (f, t) in enumerate(ex.map(one, files)):
|
||||
if schema_repr is None:
|
||||
schema_repr = [(fl.name, str(fl.type)) for fl in pq.ParquetFile(f).schema_arrow]
|
||||
fn = t.column("file_name").to_pylist()
|
||||
ids = t.column("id").to_pylist()
|
||||
rows.extend(zip(fn, ids))
|
||||
shapes.update(tuple(s) for s in t.column("vae_latent_shape").to_pylist())
|
||||
ffshapes.update(tuple(s) for s in t.column("first_frame_latent_shape").to_pylist())
|
||||
if (i + 1) % 500 == 0:
|
||||
print(f" ...scanned {i+1}/{len(files)} files, {len(rows)} rows", flush=True)
|
||||
return rows, shapes, ffshapes, schema_repr
|
||||
|
||||
|
||||
def _find_rows(root: str, wanted: set[str]) -> dict[str, dict]:
|
||||
"""Full rows (incl. binary cols) for the given file_names. Scans until all found."""
|
||||
out: dict[str, dict] = {}
|
||||
for f in _files(root):
|
||||
names = pq.read_table(f, columns=["file_name"]).column("file_name").to_pylist()
|
||||
hit = [i for i, n in enumerate(names) if n in wanted and n not in out]
|
||||
if not hit:
|
||||
continue
|
||||
t = pq.read_table(f)
|
||||
for i in hit:
|
||||
r = t.slice(i, 1).to_pylist()[0]
|
||||
out[r["file_name"]] = r
|
||||
if len(out) == len(wanted):
|
||||
break
|
||||
return out
|
||||
|
||||
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__)
|
||||
ap.add_argument("--new", required=True)
|
||||
ap.add_argument("--ref", required=True)
|
||||
ap.add_argument("--manifest", default=None)
|
||||
ap.add_argument("--expect-shape", default="16,31,90,160")
|
||||
ap.add_argument("--sample-tracks", type=int, default=8)
|
||||
ap.add_argument("--seed", type=int, default=0)
|
||||
a = ap.parse_args()
|
||||
expect = tuple(int(x) for x in a.expect_shape.split(","))
|
||||
ok = True
|
||||
|
||||
nf, rf = _files(a.new), _files(a.ref)
|
||||
print(f"[1] parquet files: new={len(nf)} ref={len(rf)}")
|
||||
print("[2] scanning NEW (720p) ...", flush=True)
|
||||
nrows, nshapes, nff, nschema = _scan(nf)
|
||||
print("[3] scanning REF (480p) ...", flush=True)
|
||||
rrows, rshapes, rff, rschema = _scan(rf)
|
||||
|
||||
nnames = [x[0] for x in nrows]
|
||||
rnames = [x[0] for x in rrows]
|
||||
nset, rset = set(nnames), set(rnames)
|
||||
print(f"\n=== COUNTS ===\n new rows: {len(nrows)} unique file_name: {len(nset)}")
|
||||
print(f" ref rows: {len(rrows)} unique file_name: {len(rset)}")
|
||||
|
||||
dup = [k for k, v in Counter(nnames).items() if v > 1]
|
||||
print(f" duplicate file_name in new: {len(dup)}" + (f" e.g. {dup[:5]}" if dup else " OK"))
|
||||
ok &= not dup
|
||||
|
||||
missing, extra = rset - nset, nset - rset
|
||||
print(f"\n=== CLIP-ID SET vs 480p ===\n in 480p but MISSING from 720p: {len(missing)}")
|
||||
if missing:
|
||||
print(f" {sorted(missing)[:50]}")
|
||||
with open("/mnt/lustre/vlm-s4duan/openvid_1m/_missing_720p.txt", "w") as fh:
|
||||
fh.write("\n".join(sorted(missing)))
|
||||
print(" (full list -> openvid_1m/_missing_720p.txt)")
|
||||
print(f" in 720p but EXTRA vs 480p: {len(extra)}" + (f" {sorted(extra)[:20]}" if extra else ""))
|
||||
ok &= not missing and not extra
|
||||
|
||||
if a.manifest:
|
||||
man = {it["path"].rsplit(".", 1)[0] for it in json.load(open(a.manifest))}
|
||||
print(f" source manifest clips: {len(man)} | missing vs manifest: {len(man - nset)} | extra: {len(nset - man)}")
|
||||
ok &= not (man - nset)
|
||||
|
||||
print(f"\n=== LATENT SHAPES (all rows) ===\n new vae_latent_shape: {dict(nshapes)}")
|
||||
print(f" new first_frame_latent_shape: {dict(nff)}")
|
||||
print(f" ref vae_latent_shape: {dict(rshapes)}")
|
||||
good = set(nshapes) == {expect} and set(nff) == {expect}
|
||||
print(f" all == {list(expect)}: {'YES' if good else 'NO'}")
|
||||
ok &= good
|
||||
|
||||
print("\n=== SCHEMA ===")
|
||||
same_schema = nschema == rschema
|
||||
print(f" new schema == ref schema (names+types): {'YES' if same_schema else 'NO'}")
|
||||
if not same_schema:
|
||||
print(f" new: {nschema}\n ref: {rschema}")
|
||||
ok &= same_schema
|
||||
|
||||
if a.sample_tracks > 0:
|
||||
random.seed(a.seed)
|
||||
pick = set(random.sample(sorted(nset & rset), min(a.sample_tracks, len(nset & rset))))
|
||||
print(f"\n=== TRACK / TEXT IDENTITY vs 480p ({len(pick)} sampled clips) ===", flush=True)
|
||||
newr = _find_rows(a.new, pick)
|
||||
refr = _find_rows(a.ref, pick)
|
||||
for name in sorted(pick):
|
||||
n, r = newr.get(name), refr.get(name)
|
||||
if n is None or r is None:
|
||||
print(f" {name}: NOT FOUND (new={n is not None} ref={r is not None})")
|
||||
ok = False
|
||||
continue
|
||||
res = {f: n[f + "_bytes"] == r[f + "_bytes"] for f in TRACK_FIELDS}
|
||||
res["text_embedding"] = n["text_embedding_bytes"] == r["text_embedding_bytes"]
|
||||
res["caption"] = n["caption"] == r["caption"]
|
||||
clipf = n["clip_feature_shape"] == r["clip_feature_shape"]
|
||||
bad = [k for k, v in res.items() if not v]
|
||||
print(f" {name}: latent {tuple(n['vae_latent_shape'])} | "
|
||||
f"tracks+text identical: {'YES' if not bad else 'NO ' + str(bad)} | "
|
||||
f"clip_feature shape match: {clipf}")
|
||||
ok &= not bad and clipf
|
||||
|
||||
print(f"\n=== RESULT: {'PASS' if ok else 'FAIL'} ===")
|
||||
raise SystemExit(0 if ok else 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,199 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Row-integrity check: does every field of a row belong to the SAME clip?
|
||||
|
||||
Yesterday's verify_720p_parquet.py compared 720p rows against 480p rows. That proves
|
||||
the two datasets agree, but both were produced by the same pipeline, so a systematic
|
||||
mispairing (row named X carrying clip Y's latent/tracks/caption) would appear in both
|
||||
and cancel out. This checks each row against the ORIGINAL sources instead:
|
||||
|
||||
caption <- videos2caption.json (exact string compare, ALL rows)
|
||||
track_points <- tracks/<name>.npz (exact float compare after normalize)
|
||||
track_visibility <- tracks/<name>.npz (exact)
|
||||
object_ids <- tracks/<name>.npz (exact)
|
||||
track_weights <- tracks/<name>.npz (exact)
|
||||
vae_latent -> VAE decode -> PSNR vs clips/<name>.mp4
|
||||
first_frame_lat -> VAE decode -> PSNR vs frame 0 of clips/<name>.mp4
|
||||
clip_feature <- CLIP re-encode of frame 0 of clips/<name>.mp4
|
||||
text_embedding <- T5 re-encode of the manifest caption
|
||||
|
||||
Also renders track points onto the DECODED frames: if the tracks belonged to a
|
||||
different clip they would not sit on moving content of this one.
|
||||
|
||||
Usage:
|
||||
python data_pipeline/verify_row_integrity.py --num 12 [--out <png>]
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import imageio.v3 as iio
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
W = "/mnt/lustre/vlm-s4duan"
|
||||
OUT = f"{W}/openvid_1m/combined_parquet_dataset_720p"
|
||||
CLIPS = f"{W}/openvid_1m/clips"
|
||||
TRACKS = f"{W}/openvid_1m/tracks"
|
||||
MANIFEST = f"{W}/openvid_1m/videos2caption.json"
|
||||
MODEL = f"{W}/models/trackwan_1.3b_i2v_d64_nobias_init"
|
||||
|
||||
|
||||
def psnr(a, b):
|
||||
mse = float(((a.astype(np.float32) - b.astype(np.float32)) ** 2).mean())
|
||||
return 99.0 if mse == 0 else 10.0 * float(np.log10(255.0 * 255.0 / mse))
|
||||
|
||||
|
||||
def read_clip(path, n):
|
||||
fr = []
|
||||
for i, f in enumerate(iio.imiter(path, plugin="pyav")):
|
||||
if i >= n:
|
||||
break
|
||||
fr.append(f)
|
||||
return np.stack(fr)
|
||||
|
||||
|
||||
def global_caption_check(manifest):
|
||||
"""ALL rows: file_name -> caption must equal the manifest's caption for that clip."""
|
||||
print("=== GLOBAL: caption/metadata vs manifest, ALL rows ===", flush=True)
|
||||
cap_of = {it["path"].rsplit(".", 1)[0]: it["cap"][0] for it in manifest}
|
||||
fps_of = {it["path"].rsplit(".", 1)[0]: float(it["fps"]) for it in manifest}
|
||||
files = sorted(glob.glob(f"{OUT}/**/*.parquet", recursive=True))
|
||||
n = bad_cap = bad_meta = 0
|
||||
bad_examples = []
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
def one(f):
|
||||
return pq.read_table(f, columns=["file_name", "caption", "fps", "num_frames", "width", "height"])
|
||||
|
||||
with ThreadPoolExecutor(32) as ex:
|
||||
for k, t in enumerate(ex.map(one, files)):
|
||||
for fn, cap, fps, nf, w, h in zip(t.column("file_name").to_pylist(), t.column("caption").to_pylist(),
|
||||
t.column("fps").to_pylist(), t.column("num_frames").to_pylist(),
|
||||
t.column("width").to_pylist(), t.column("height").to_pylist()):
|
||||
n += 1
|
||||
if cap_of.get(fn) != cap:
|
||||
bad_cap += 1
|
||||
if len(bad_examples) < 5:
|
||||
bad_examples.append(fn)
|
||||
if fps != fps_of.get(fn) or nf != 31 or w != 720 or h != 1280:
|
||||
bad_meta += 1
|
||||
if (k + 1) % 1000 == 0:
|
||||
print(f" ...{k+1}/{len(files)} files, {n} rows", flush=True)
|
||||
print(f" rows checked : {n}")
|
||||
print(f" caption mismatches : {bad_cap} {bad_examples if bad_examples else ''}")
|
||||
print(f" fps/num_frames/w/h bad : {bad_meta}")
|
||||
return bad_cap == 0 and bad_meta == 0, n
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description=__doc__)
|
||||
ap.add_argument("--num", type=int, default=12)
|
||||
ap.add_argument("--seed", type=int, default=7)
|
||||
ap.add_argument("--out", default=f"{W}/openvid_1m/_row_integrity_overlay.png")
|
||||
ap.add_argument("--skip-global", action="store_true")
|
||||
a = ap.parse_args()
|
||||
|
||||
manifest = json.load(open(MANIFEST))
|
||||
ok = True
|
||||
if not a.skip_global:
|
||||
g_ok, _ = global_caption_check(manifest)
|
||||
ok &= g_ok
|
||||
|
||||
# --- sample rows spread across the whole run (by shard completion time, so we
|
||||
# cover all three launch epochs: 16-node wave 1, the 12-node run, the final 16-node run)
|
||||
dones = sorted(glob.glob(f"{OUT}/shard_*/.done"), key=os.path.getmtime)
|
||||
picks = [dones[int(i * (len(dones) - 1) / max(1, a.num - 1))] for i in range(a.num)]
|
||||
rng = random.Random(a.seed)
|
||||
|
||||
print(f"\n=== PER-ROW vs ORIGINAL SOURCES ({a.num} clips spread across the run) ===", flush=True)
|
||||
from diffusers import AutoencoderKLWan
|
||||
vae = AutoencoderKLWan.from_pretrained(MODEL, subfolder="vae", torch_dtype=torch.float32).to("cuda").eval()
|
||||
|
||||
cap_of = {it["path"].rsplit(".", 1)[0]: it["cap"][0] for it in manifest}
|
||||
panels = []
|
||||
for d in picks:
|
||||
sd = os.path.dirname(d)
|
||||
pqs = sorted(glob.glob(f"{sd}/**/*.parquet", recursive=True))
|
||||
if not pqs:
|
||||
print(f" {os.path.basename(sd)}: EMPTY shard (expected for the tail) - skip")
|
||||
continue
|
||||
t = pq.read_table(rng.choice(pqs))
|
||||
i = rng.randrange(t.num_rows)
|
||||
r = t.slice(i, 1).to_pylist()[0]
|
||||
name = r["file_name"]
|
||||
|
||||
# 1) caption vs manifest
|
||||
cap_ok = cap_of.get(name) == r["caption"]
|
||||
|
||||
# 2) tracks vs the clip's own npz (exact)
|
||||
z = np.load(f"{TRACKS}/{name}.npz")
|
||||
tw, th = float(z["width"]), float(z["height"])
|
||||
exp_tp = z["tracks"].astype(np.float32)[:121].copy()
|
||||
exp_tp[..., 0] /= tw
|
||||
exp_tp[..., 1] /= th
|
||||
got_tp = np.frombuffer(r["track_points_bytes"], np.float32).reshape(r["track_points_shape"])
|
||||
got_vis = np.frombuffer(r["track_visibility_bytes"], np.float32).reshape(r["track_visibility_shape"])
|
||||
tp_ok = np.array_equal(got_tp, exp_tp)
|
||||
vis_ok = np.array_equal(got_vis, z["visibility"].astype(np.float32)[:121])
|
||||
oid_ok = np.array_equal(np.frombuffer(r["object_ids_bytes"], np.float32),
|
||||
z["object_ids"].astype(np.float32)) if "object_ids" in z else None
|
||||
twt_ok = np.array_equal(np.frombuffer(r["track_weights_bytes"], np.float32),
|
||||
z["track_weights"].astype(np.float32)) if "track_weights" in z else None
|
||||
|
||||
# 3) latent -> pixels vs the clip's own mp4
|
||||
lat = np.frombuffer(r["vae_latent_bytes"], np.float32).reshape(r["vae_latent_shape"]).copy()
|
||||
with torch.no_grad():
|
||||
px = vae.decode(torch.from_numpy(lat)[None].to("cuda"), return_dict=False)[0]
|
||||
dec = ((px[0].permute(1, 2, 3, 0).clamp(-1, 1) + 1) * 127.5).to(torch.uint8).cpu().numpy()
|
||||
src = read_clip(f"{CLIPS}/{name}.mp4", dec.shape[0])
|
||||
p_all = psnr(dec[:len(src)], src[:len(dec)])
|
||||
|
||||
# 4) first-frame conditioning latent -> its frame 0 vs the clip's frame 0
|
||||
ff = np.frombuffer(r["first_frame_latent_bytes"], np.float32).reshape(r["first_frame_latent_shape"]).copy()
|
||||
with torch.no_grad():
|
||||
pxf = vae.decode(torch.from_numpy(ff)[None].to("cuda"), return_dict=False)[0]
|
||||
f0 = ((pxf[0, :, 0].permute(1, 2, 0).clamp(-1, 1) + 1) * 127.5).to(torch.uint8).cpu().numpy()
|
||||
p_f0 = psnr(f0, src[0])
|
||||
|
||||
flags = [f"caption={'OK' if cap_ok else 'MISMATCH'}",
|
||||
f"tracks={'exact' if tp_ok else 'MISMATCH'}",
|
||||
f"vis={'exact' if vis_ok else 'MISMATCH'}",
|
||||
f"oid={'exact' if oid_ok else oid_ok}",
|
||||
f"w={'exact' if twt_ok else twt_ok}",
|
||||
f"latentPSNR={p_all:.1f}dB", f"ff0PSNR={p_f0:.1f}dB"]
|
||||
good = cap_ok and tp_ok and vis_ok and (oid_ok is not False) and (twt_ok is not False) and p_all > 30
|
||||
ok &= good
|
||||
print(f" {'PASS' if good else 'FAIL'} {name} [{os.path.basename(sd)}] " + " ".join(flags), flush=True)
|
||||
|
||||
# 5) overlay this row's tracks onto the DECODED frames (visual correspondence)
|
||||
idxs = [0, len(dec) // 2, len(dec) - 1]
|
||||
H, Wd = dec.shape[1], dec.shape[2]
|
||||
row = []
|
||||
for fi in idxs:
|
||||
img = dec[fi].copy()
|
||||
pts = got_tp[fi]
|
||||
vis = got_vis[fi] > 0.5
|
||||
xs = (pts[:, 0] * Wd).astype(int)
|
||||
ys = (pts[:, 1] * H).astype(int)
|
||||
keep = vis & (xs >= 1) & (xs < Wd - 1) & (ys >= 1) & (ys < H - 1)
|
||||
for x, y in zip(xs[keep][::4], ys[keep][::4]):
|
||||
img[y - 1:y + 2, x - 1:x + 2] = (0, 255, 0)
|
||||
row.append(img)
|
||||
panels.append(np.concatenate(row, axis=1))
|
||||
|
||||
if panels:
|
||||
hmin = min(p.shape[0] for p in panels)
|
||||
iio.imwrite(a.out, np.concatenate([p[:hmin] for p in panels[:6]], axis=0))
|
||||
print(f"\n track-on-decoded-frame overlay -> {a.out}")
|
||||
|
||||
print(f"\n=== ROW INTEGRITY: {'PASS' if ok else 'FAIL'} ===")
|
||||
raise SystemExit(0 if ok else 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,160 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Visualize CoTracker point tracks over a generated video.
|
||||
|
||||
Produces, for one (video, tracks) pair:
|
||||
- ``<out>_overlay.mp4``: every frame with visible points drawn as dots plus a short
|
||||
motion tail (last ``--tail`` frames). Colour encodes the point's initial grid position.
|
||||
- ``<out>_trajectories.png``: all full trajectories drawn over frame 0 (static summary).
|
||||
|
||||
Pure CPU; depends only on numpy + PIL + imageio + torchvision (already in the venv), so it
|
||||
does NOT need CoTracker or a GPU.
|
||||
|
||||
Example:
|
||||
.venv/bin/python data_pipeline/visualize_tracks.py \
|
||||
--video /.../smoke_wan21_1.3b_480p/videos/vid_000000.mp4 \
|
||||
--tracks /.../smoke_wan21_1.3b_480p/tracks/vid_000000.npz \
|
||||
--out /.../smoke_wan21_1.3b_480p/viz/vid_000000
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import colorsys
|
||||
from pathlib import Path
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--video", type=Path, required=True, help="Source .mp4")
|
||||
p.add_argument("--tracks", type=Path, required=True, help=".npz from extract_tracks.py")
|
||||
p.add_argument("--out", type=Path, required=True, help="Output path prefix (no extension).")
|
||||
p.add_argument("--stride", type=int, default=2, help="Sub-sample the NxN grid by this factor for clarity.")
|
||||
p.add_argument("--tail", type=int, default=12, help="Motion-tail length in frames.")
|
||||
p.add_argument("--radius", type=int, default=2, help="Point radius in px.")
|
||||
p.add_argument("--fps", type=int, default=16, help="Output video fps.")
|
||||
p.add_argument("--vis-thresh", type=float, default=0.5, help="Visibility threshold.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def load_frames(path: Path) -> np.ndarray:
|
||||
"""Return frames (T, H, W, 3) uint8.
|
||||
|
||||
Recent torchvision dropped ``torchvision.io.read_video``; use decord, falling
|
||||
back to imageio's ffmpeg reader.
|
||||
"""
|
||||
try:
|
||||
from decord import VideoReader, cpu
|
||||
vr = VideoReader(str(path), ctx=cpu(0))
|
||||
frames = vr.get_batch(list(range(len(vr)))).asnumpy()
|
||||
except Exception: # noqa: BLE001 - fall back to ffmpeg
|
||||
reader = imageio.get_reader(str(path), format="ffmpeg")
|
||||
frames = np.stack([np.asarray(f) for f in reader], axis=0)
|
||||
reader.close()
|
||||
return frames[..., :3].astype(np.uint8)
|
||||
|
||||
|
||||
def grid_colors(grid_size: int, stride: int) -> np.ndarray:
|
||||
"""One RGB colour per (sub-sampled) grid point, encoding its initial position."""
|
||||
idx = np.arange(0, grid_size, stride)
|
||||
gy, gx = np.meshgrid(idx, idx, indexing="ij")
|
||||
nx = gx.reshape(-1) / max(grid_size - 1, 1)
|
||||
ny = gy.reshape(-1) / max(grid_size - 1, 1)
|
||||
cols = np.empty((nx.shape[0], 3), dtype=np.uint8)
|
||||
for i, (x, y) in enumerate(zip(nx, ny)):
|
||||
r, g, b = colorsys.hsv_to_rgb(float(x), 1.0, 0.5 + 0.5 * float(y))
|
||||
cols[i] = (int(r * 255), int(g * 255), int(b * 255))
|
||||
return cols
|
||||
|
||||
|
||||
def subsample(tracks: np.ndarray, vis: np.ndarray, grid_size: int, stride: int):
|
||||
"""tracks (T,N,2), vis (T,N) with N==grid_size**2 -> sub-sampled by stride in both grid dims."""
|
||||
t = tracks.shape[0]
|
||||
if tracks.shape[1] != grid_size * grid_size:
|
||||
return tracks, vis # unknown layout; keep as-is
|
||||
tr = tracks.reshape(t, grid_size, grid_size, 2)[:, ::stride, ::stride, :].reshape(t, -1, 2)
|
||||
vs = vis.reshape(t, grid_size, grid_size)[:, ::stride, ::stride].reshape(t, -1)
|
||||
return tr, vs
|
||||
|
||||
|
||||
def draw_overlay(frames, tracks, vis, colors, tail, radius, vis_thresh) -> list[np.ndarray]:
|
||||
t, h, w, _ = frames.shape
|
||||
n = tracks.shape[1]
|
||||
out = []
|
||||
for fi in range(t):
|
||||
img = Image.fromarray(frames[fi]).convert("RGB")
|
||||
draw = ImageDraw.Draw(img)
|
||||
t0 = max(0, fi - tail)
|
||||
for pi in range(n):
|
||||
col = tuple(int(c) for c in colors[pi])
|
||||
# tail: consecutive visible positions in the window
|
||||
pts = []
|
||||
for tj in range(t0, fi + 1):
|
||||
if vis[tj, pi] >= vis_thresh:
|
||||
x, y = float(tracks[tj, pi, 0]), float(tracks[tj, pi, 1])
|
||||
if 0 <= x < w and 0 <= y < h:
|
||||
pts.append((x, y))
|
||||
else:
|
||||
pts = [] # break the tail on occlusion
|
||||
if len(pts) >= 2:
|
||||
draw.line(pts, fill=col, width=1)
|
||||
if vis[fi, pi] >= vis_thresh:
|
||||
x, y = float(tracks[fi, pi, 0]), float(tracks[fi, pi, 1])
|
||||
if 0 <= x < w and 0 <= y < h:
|
||||
draw.ellipse([x - radius, y - radius, x + radius, y + radius], fill=col)
|
||||
out.append(np.asarray(img))
|
||||
return out
|
||||
|
||||
|
||||
def draw_trajectories(frame0, tracks, vis, colors, vis_thresh) -> np.ndarray:
|
||||
h, w, _ = frame0.shape
|
||||
img = Image.fromarray(frame0).convert("RGB")
|
||||
draw = ImageDraw.Draw(img)
|
||||
n = tracks.shape[1]
|
||||
for pi in range(n):
|
||||
col = tuple(int(c) for c in colors[pi])
|
||||
pts = [(float(tracks[tj, pi, 0]), float(tracks[tj, pi, 1]))
|
||||
for tj in range(tracks.shape[0])
|
||||
if vis[tj, pi] >= vis_thresh and 0 <= tracks[tj, pi, 0] < w and 0 <= tracks[tj, pi, 1] < h]
|
||||
if len(pts) >= 2:
|
||||
draw.line(pts, fill=col, width=1)
|
||||
return np.asarray(img)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
data = np.load(args.tracks)
|
||||
tracks = data["tracks"].astype(np.float32) # (T, N, 2)
|
||||
vis = data["visibility"].astype(np.float32) # (T, N)
|
||||
grid_size = int(data["grid_size"]) if "grid_size" in data else int(round(tracks.shape[1] ** 0.5))
|
||||
|
||||
frames = load_frames(args.video)
|
||||
t = min(frames.shape[0], tracks.shape[0])
|
||||
frames, tracks, vis = frames[:t], tracks[:t], vis[:t]
|
||||
|
||||
tracks, vis = subsample(tracks, vis, grid_size, args.stride)
|
||||
colors = grid_colors(grid_size, args.stride)
|
||||
if colors.shape[0] != tracks.shape[1]: # layout fallback: cycle a rainbow
|
||||
colors = grid_colors(int(round(tracks.shape[1] ** 0.5)) or 1, 1)[:tracks.shape[1]]
|
||||
|
||||
args.out.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
overlay = draw_overlay(frames, tracks, vis, colors, args.tail, args.radius, args.vis_thresh)
|
||||
mp4_path = args.out.with_name(args.out.name + "_overlay.mp4")
|
||||
imageio.mimsave(str(mp4_path), overlay, fps=args.fps, macro_block_size=1)
|
||||
|
||||
traj = draw_trajectories(frames[0], tracks, vis, colors, args.vis_thresh)
|
||||
png_path = args.out.with_name(args.out.name + "_trajectories.png")
|
||||
imageio.imwrite(str(png_path), traj)
|
||||
|
||||
visible_frac = float((vis >= args.vis_thresh).mean())
|
||||
print(f"[viz] frames={t} points_drawn={tracks.shape[1]} (grid {grid_size}x{grid_size}, stride {args.stride}) "
|
||||
f"mean_visible={visible_frac:.2f}")
|
||||
print(f"[viz] wrote {mp4_path}")
|
||||
print(f"[viz] wrote {png_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -11,7 +11,7 @@ def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Interactive track-control demo for TrackWan (MotionStream-style).
|
||||
|
||||
Pick a preprocessed clip (gives the first frame + text + the option of its own
|
||||
tracks), author a motion control, preview the tracks over the first frame, then
|
||||
generate and watch whether the model follows them.
|
||||
|
||||
Control modes (mirroring MotionStream):
|
||||
- Trajectory / drag : click the first frame to set a drag handle, choose a
|
||||
direction + radius; optionally make it SPARSE (only the handle points active).
|
||||
- Camera-like : pan / zoom / rotate / swirl presets.
|
||||
- Motion transfer : reuse another clip's extracted tracks.
|
||||
- GT : the clip's own tracks (reconstruction sanity check).
|
||||
|
||||
Launch (single GPU, on the cluster)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=1 \
|
||||
HF_HUB_CACHE=/.../hf_cache_clean PYTHONPATH=$PWD \
|
||||
.venv/bin/python examples/inference/gradio/trackwan/app.py \
|
||||
--export /.../trackwan_1.3b_overfit4k --port 7860
|
||||
|
||||
Then SSH-forward the port and open http://localhost:7860 .
|
||||
|
||||
NOTE: sparse drag control needs a model trained with point-subsampling aug
|
||||
(WANTRACK_AUG=1). A clean-overfit (aug-off) checkpoint mostly ignores sparse
|
||||
controls -- use dense presets to test it.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
REPO = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", ".."))
|
||||
sys.path.insert(0, os.path.join(REPO, "data_pipeline"))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
STATE: dict = {} # model, tc, samples, Tpx, H, W, out_dir
|
||||
|
||||
|
||||
def _overlay(frames_thwc, tracks_norm, vis, stride=3):
|
||||
from fastvideo.train.callbacks.track_validation import (_draw_overlay, _grid_colors, _subsample)
|
||||
T, H, W, _ = frames_thwc.shape
|
||||
tt = min(T, tracks_norm.shape[0])
|
||||
fr, tr, vs = frames_thwc[:tt], tracks_norm[:tt], vis[:tt]
|
||||
grid = int(round(tr.shape[1] ** 0.5))
|
||||
tr, vs = _subsample(tr, vs, grid, stride)
|
||||
colors = _grid_colors(grid, stride)
|
||||
if colors.shape[0] != tr.shape[1]:
|
||||
colors = _grid_colors(int(round(tr.shape[1] ** 0.5)) or 1, 1)[:tr.shape[1]]
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
return _draw_overlay(fr, trpx, vs, colors, 12, 2, 0.5)
|
||||
|
||||
|
||||
def _first_frame_rgb(clip_idx: int) -> np.ndarray:
|
||||
"""Decode the clip's GT first frame for display/clicking."""
|
||||
s = STATE["samples"][clip_idx]
|
||||
ref = twi.decode_reference(STATE["model"], s["vae_latent"]) # [T,H,W,3]
|
||||
return ref[0]
|
||||
|
||||
|
||||
def _build_tracks(clip_idx, mode, preset_name, strength, cx, cy, radius,
|
||||
sparsity, drag_dx, drag_dy, transfer_idx):
|
||||
"""Return (tracks_norm [T,N,2], vis [T,N]) or (None, None) for 'none'."""
|
||||
Tpx = STATE["Tpx"]
|
||||
g = st.make_grid(50)
|
||||
if mode == "gt":
|
||||
s = STATE["samples"][clip_idx]
|
||||
return s["track_points"][0].numpy()[:Tpx], s["track_visibility"][0].numpy()[:Tpx]
|
||||
if mode == "none":
|
||||
return None, None
|
||||
if mode == "transfer":
|
||||
s = STATE["samples"][int(transfer_idx)]
|
||||
return s["track_points"][0].numpy()[:Tpx], s["track_visibility"][0].numpy()[:Tpx]
|
||||
if mode == "preset":
|
||||
tr, vs = st.preset(preset_name, Tpx, 50, strength=float(strength))
|
||||
elif mode == "drag":
|
||||
tr, vs = st.drag(g, Tpx, center=(float(cx), float(cy)), dx=float(drag_dx),
|
||||
dy=float(drag_dy), radius=float(radius))
|
||||
else:
|
||||
tr, vs = st.static(g, Tpx)
|
||||
# sparsity
|
||||
if sparsity == "radius (handle only)":
|
||||
vs = st.select_radius(vs, g, center=(float(cx), float(cy)), radius=float(radius))
|
||||
elif sparsity == "coarse grid":
|
||||
vs = st.select_stride(vs, 50, 4)
|
||||
return tr, vs
|
||||
|
||||
|
||||
def on_select_clip(clip_idx):
|
||||
img = _first_frame_rgb(int(clip_idx))
|
||||
cap = STATE["samples"][int(clip_idx)]["caption"][:200]
|
||||
return img, cap
|
||||
|
||||
|
||||
def on_image_click(clip_idx, evt: gr.SelectData):
|
||||
"""Set drag center (normalized) from a click on the first frame."""
|
||||
img = _first_frame_rgb(int(clip_idx))
|
||||
H, W, _ = img.shape
|
||||
x, y = evt.index[0] / W, evt.index[1] / H
|
||||
return round(float(x), 3), round(float(y), 3)
|
||||
|
||||
|
||||
def on_preview(clip_idx, mode, preset_name, strength, cx, cy, radius, sparsity,
|
||||
drag_dx, drag_dy, transfer_idx):
|
||||
img = _first_frame_rgb(int(clip_idx))
|
||||
tr, vs = _build_tracks(int(clip_idx), mode, preset_name, strength, cx, cy, radius,
|
||||
sparsity, drag_dx, drag_dy, transfer_idx)
|
||||
if tr is None:
|
||||
return img, "no tracks (mode=none)"
|
||||
frames = np.repeat(img[None], tr.shape[0], 0)
|
||||
ov = _overlay(frames, tr, vs)[0]
|
||||
active = int((vs[0] > 0.5).sum())
|
||||
return ov, f"{active} active points (of {vs.shape[1]})"
|
||||
|
||||
|
||||
def on_generate(clip_idx, mode, preset_name, strength, cx, cy, radius, sparsity,
|
||||
drag_dx, drag_dy, transfer_idx, steps, seed):
|
||||
clip_idx = int(clip_idx)
|
||||
s = STATE["samples"][clip_idx]
|
||||
tr, vs = _build_tracks(clip_idx, mode, preset_name, strength, cx, cy, radius,
|
||||
sparsity, drag_dx, drag_dy, transfer_idx)
|
||||
tp = torch.from_numpy(tr)[None].float() if tr is not None else None
|
||||
tv = torch.from_numpy(vs)[None].float() if vs is not None else None
|
||||
lat = twi.generate(STATE["model"], first_frame_latent=s["first_frame_latent"],
|
||||
text_embedding=s["text_embedding"], text_attention_mask=s["text_attention_mask"],
|
||||
track_points=tp, track_visibility=tv, clip_feature=s["clip_feature"],
|
||||
num_steps=int(steps), seed=int(seed))
|
||||
frames = twi.decode_to_pixels(STATE["model"], lat)
|
||||
ov_tr = tr if tr is not None else np.zeros((frames.shape[0], 1, 2), np.float32)
|
||||
ov_vs = vs if vs is not None else np.zeros((frames.shape[0], 1), np.float32)
|
||||
out_frames = _overlay(frames, ov_tr, ov_vs)
|
||||
path = os.path.join(STATE["out_dir"], f"gen_clip{clip_idx}_{mode}_{preset_name}.mp4")
|
||||
imageio.mimsave(path, out_frames, fps=24, macro_block_size=1)
|
||||
|
||||
epe_txt = "EPE: n/a (no control tracks)"
|
||||
if tr is not None:
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import compute_epe
|
||||
H, W = frames.shape[1], frames.shape[2]
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
res = compute_epe(frames, trpx, vs, STATE["ct"], STATE["model"].device)
|
||||
if res["epe"] is not None:
|
||||
epe_txt = f"EPE: {res['epe']:.2f} px ({res['n_points']} pts) — lower = follows control"
|
||||
return path, epe_txt
|
||||
|
||||
|
||||
def build_ui(num_clips):
|
||||
with gr.Blocks(title="TrackWan — interactive motion control") as demo:
|
||||
gr.Markdown("## TrackWan — interactive point-track control\n"
|
||||
"Pick a clip, author a motion control, **Preview tracks**, then **Generate**. "
|
||||
"Click the first frame to set the drag handle center.")
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
clip = gr.Dropdown(list(range(num_clips)), value=0, label="Clip")
|
||||
caption = gr.Textbox(label="Prompt", interactive=False, lines=2)
|
||||
mode = gr.Radio(["preset", "drag", "transfer", "gt", "none"], value="preset", label="Control mode")
|
||||
preset_name = gr.Dropdown(st.PRESETS, value="pan_right", label="Preset (preset mode)")
|
||||
strength = gr.Slider(0.0, 0.6, value=0.25, step=0.01, label="Strength (pan/zoom)")
|
||||
with gr.Row():
|
||||
cx = gr.Number(value=0.5, label="drag center x (click frame)")
|
||||
cy = gr.Number(value=0.5, label="drag center y")
|
||||
radius = gr.Slider(0.05, 0.6, value=0.2, step=0.01, label="Radius (drag/handle)")
|
||||
with gr.Row():
|
||||
drag_dx = gr.Slider(-0.5, 0.5, value=0.3, step=0.01, label="drag dx")
|
||||
drag_dy = gr.Slider(-0.5, 0.5, value=0.0, step=0.01, label="drag dy")
|
||||
sparsity = gr.Radio(["full", "radius (handle only)", "coarse grid"], value="full",
|
||||
label="Sparsity (sparse needs aug-trained ckpt)")
|
||||
transfer_idx = gr.Dropdown(list(range(num_clips)), value=min(1, num_clips - 1),
|
||||
label="Transfer from clip (transfer mode)")
|
||||
with gr.Row():
|
||||
steps = gr.Slider(10, 60, value=30, step=5, label="Denoise steps")
|
||||
seed = gr.Number(value=1000, label="Seed")
|
||||
with gr.Row():
|
||||
preview_btn = gr.Button("Preview tracks")
|
||||
gen_btn = gr.Button("Generate", variant="primary")
|
||||
with gr.Column(scale=1):
|
||||
frame_img = gr.Image(label="First frame (click to set drag center) / track preview", type="numpy")
|
||||
status = gr.Textbox(label="Status", interactive=False)
|
||||
out_video = gr.Video(label="Generated (tracks overlaid)")
|
||||
epe_box = gr.Textbox(label="Motion fidelity", interactive=False)
|
||||
|
||||
ctrl = [clip, mode, preset_name, strength, cx, cy, radius, sparsity, drag_dx, drag_dy, transfer_idx]
|
||||
clip.change(on_select_clip, [clip], [frame_img, caption])
|
||||
frame_img.select(on_image_click, [clip], [cx, cy])
|
||||
preview_btn.click(on_preview, ctrl, [frame_img, status])
|
||||
gen_btn.click(on_generate, ctrl + [steps, seed], [out_video, epe_box])
|
||||
demo.load(on_select_clip, [clip], [frame_img, caption])
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--export", required=True, help="dcp_to_diffusers export dir")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset")
|
||||
p.add_argument("--num-clips", type=int, default=10)
|
||||
p.add_argument("--out-dir", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/gradio_out")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7860)
|
||||
p.add_argument("--share", action="store_true", help="Create a public gradio.live link (no SSH forward needed).")
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
model, tc = twi.load_trackwan(args.export, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, list(range(args.num_clips)), text_len)
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import load_cotracker
|
||||
STATE.update(model=model, tc=tc, samples=samples, Tpx=(num_lat_t - 1) * ratio + 1,
|
||||
out_dir=args.out_dir, ct=load_cotracker(model.device))
|
||||
|
||||
demo = build_ui(len(samples))
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[args.out_dir])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,478 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""TrackWan action-recording demo — draw a motion, watch the model follow it.
|
||||
|
||||
Designed to answer one question: *does this checkpoint actually follow the point
|
||||
tracks (action-following), and how does that change across checkpoints?*
|
||||
|
||||
Interaction (the part the old demo lacked):
|
||||
1. Pick a clip (gives the first frame + prompt) and a checkpoint.
|
||||
2. Move the mouse over the first frame and press **Space** -> records your cursor
|
||||
path for 5 seconds (121 frames). The green dot is where the action starts,
|
||||
the red dot where it ends.
|
||||
3. The 50x50 CoTracker training grid is used directly: the grid points within the
|
||||
action radius of your start are dragged along your recorded path, the rest stay
|
||||
static. This is grid-snapped / in-distribution (NOT an off-grid patch added on
|
||||
top of a static grid). Preview shows the full field (moving=green, static=gray).
|
||||
4. Generate -> original + generation side by side, with the full field overlaid,
|
||||
plus the moving-point CoTracker-EPE of how well the arm followed your path.
|
||||
|
||||
Checkpoints hot-swap (only the transformer weights reload, ~3 s) so you can flip
|
||||
between e.g. step-2000 / 3000 / 4000 on the same clip and action.
|
||||
|
||||
Visibility = "dense field" keeps the whole 50x50 grid visible (the trained coverage;
|
||||
recommended for aug-off checkpoints); "sparse" makes only the moved points visible.
|
||||
|
||||
Launch (single GPU on the cluster)::
|
||||
|
||||
srun --jobid=<job> --overlap --ntasks=1 env CUDA_VISIBLE_DEVICES=1 \
|
||||
PYTHONPATH=$PWD .venv/bin/python \
|
||||
examples/inference/gradio/trackwan/app_action.py \
|
||||
--exports-root /.../control_now_exports --share
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import glob
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
REPO = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", ".."))
|
||||
sys.path.insert(0, os.path.join(REPO, "data_pipeline"))
|
||||
import synthetic_tracks as st # noqa: E402
|
||||
import trackwan_infer as twi # noqa: E402
|
||||
|
||||
STATE: dict = {}
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------- utils
|
||||
def _img_b64(arr: np.ndarray) -> str:
|
||||
buf = io.BytesIO()
|
||||
Image.fromarray(arr).save(buf, format="PNG")
|
||||
return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode()
|
||||
|
||||
|
||||
def _first_frame_rgb(clip_idx: int) -> np.ndarray:
|
||||
s = STATE["samples"][clip_idx]
|
||||
return twi.decode_reference(STATE["model"], s["vae_latent"])[0]
|
||||
|
||||
|
||||
def _original_video(clip_idx: int) -> str:
|
||||
"""Decode + cache the GT clip as an mp4 (checkpoint-independent)."""
|
||||
cache = STATE.setdefault("orig_cache", {})
|
||||
if clip_idx in cache:
|
||||
return cache[clip_idx]
|
||||
s = STATE["samples"][clip_idx]
|
||||
frames = twi.decode_reference(STATE["model"], s["vae_latent"]) # [T,H,W,3]
|
||||
path = os.path.join(STATE["out_dir"], f"orig_clip{clip_idx}.mp4")
|
||||
imageio.mimsave(path, frames, fps=24, macro_block_size=1)
|
||||
cache[clip_idx] = path
|
||||
return path
|
||||
|
||||
|
||||
def swap_checkpoint(name: str) -> str:
|
||||
"""Load a checkpoint's transformer weights in place (fast; no text/vae reload)."""
|
||||
if STATE.get("cur_ckpt") == name:
|
||||
return f"checkpoint: {name} (loaded)"
|
||||
from safetensors.torch import load_file
|
||||
try:
|
||||
from torch.distributed.tensor import DTensor
|
||||
except Exception: # older torch
|
||||
from torch.distributed._tensor import DTensor
|
||||
path = STATE["exports"][name]
|
||||
sd = load_file(os.path.join(path, "transformer", "model.safetensors"))
|
||||
tgt = STATE["model"].transformer
|
||||
# Params are FSDP DTensors (even on 1 GPU); copy into each param's local shard
|
||||
# instead of load_state_dict (which errors mixing Tensor and DTensor).
|
||||
params = dict(tgt.named_parameters())
|
||||
n_ok = 0
|
||||
with torch.no_grad():
|
||||
for k, p in params.items():
|
||||
if k not in sd:
|
||||
continue
|
||||
loc = p.to_local() if isinstance(p, DTensor) else p.data
|
||||
src = sd[k].to(dtype=loc.dtype, device=loc.device)
|
||||
if tuple(loc.shape) == tuple(src.shape):
|
||||
loc.copy_(src)
|
||||
n_ok += 1
|
||||
tgt.eval()
|
||||
STATE["cur_ckpt"] = name
|
||||
warn = "" if n_ok == len(params) else f" [WARN copied {n_ok}/{len(params)} params]"
|
||||
return f"checkpoint: {name} (swapped){warn}"
|
||||
|
||||
|
||||
def _grid_action_tracks(traj, Tpx: int, radius: float, falloff: str, sparse: str, num_bg: int, seed: int):
|
||||
"""Build the input tracks for the drawn action, matching the model's training regime.
|
||||
|
||||
Modes (``sparse``):
|
||||
"dense" -- full 50x50 = 2500 tracks; the ones within ``radius`` of the draw-start
|
||||
follow the path, the rest stay static and visible. Matches the OLD
|
||||
training regime (WANTRACK_SPARSE=0, sample K=1000-2500 per step).
|
||||
"sparse" -- ONLY the moving handle + ``num_bg`` random background points from the
|
||||
training grid. Total N ~= handle_size + num_bg. Matches the SPARSE
|
||||
training regime (WANTRACK_SPARSE=1, 1-per-object + ~20 background).
|
||||
``num_bg`` should be small (5-30) to match training.
|
||||
|
||||
Returns (tracks[Tpx,N,2], vis[Tpx,N], moving[N]) or (None, None, None).
|
||||
"""
|
||||
traj = np.asarray(traj, np.float32)
|
||||
if traj.ndim != 2 or len(traj) < 2:
|
||||
return None, None, None
|
||||
idx = np.linspace(0, len(traj) - 1, Tpx)
|
||||
path = np.stack([
|
||||
np.interp(idx, np.arange(len(traj)), traj[:, 0]),
|
||||
np.interp(idx, np.arange(len(traj)), traj[:, 1]),
|
||||
], 1).astype(np.float32) # [Tpx,2]
|
||||
g = st.make_grid(50) # [2500,2] frame-0 grid (x,y) in [0,1]
|
||||
start = path[0]
|
||||
d = np.linalg.norm(g - start[None], axis=-1) # [2500]
|
||||
if falloff == "smooth":
|
||||
w = np.clip(1.0 - d / max(radius, 1e-6), 0.0, 1.0)
|
||||
w = (0.5 - 0.5 * np.cos(np.pi * w)).astype(np.float32) # handle-like weighting
|
||||
else: # hard
|
||||
w = (d <= radius).astype(np.float32)
|
||||
handle_mask = w > 1e-3 # [2500] which grid points are the moving handle
|
||||
disp = path - start[None] # [Tpx,2] displacement from start each frame
|
||||
mode = str(sparse)
|
||||
if mode == "sparse":
|
||||
# Take the moving handle + `num_bg` random background points (indices from the
|
||||
# non-handle pool), matching the sparse training sampler shape (~num_objects + extras).
|
||||
rng = np.random.default_rng(int(seed))
|
||||
bg_pool = np.where(~handle_mask)[0]
|
||||
n_bg = int(min(max(num_bg, 0), bg_pool.size))
|
||||
bg_idx = rng.choice(bg_pool, size=n_bg, replace=False) if n_bg > 0 else np.zeros(0, np.int64)
|
||||
handle_idx = np.where(handle_mask)[0]
|
||||
keep = np.concatenate([handle_idx, bg_idx])
|
||||
g_k = g[keep] # [N,2]
|
||||
w_k = w[keep] # [N] (background weights are ~0 -> stay static)
|
||||
tracks = (g_k[None] + w_k[None, :, None] * disp[:, None, :]).astype(np.float32) # [Tpx,N,2]
|
||||
tracks = np.clip(tracks, 0.0, 1.0)
|
||||
moving = np.zeros(keep.size, bool)
|
||||
moving[:handle_idx.size] = True
|
||||
vis = np.ones(tracks.shape[:2], np.float32) # all sparse points visible from start
|
||||
return tracks, vis, moving
|
||||
# "dense" (original behavior): 2500 tracks, handle moves, rest static + visible.
|
||||
tracks = (g[None] + w[None, :, None] * disp[:, None, :]).astype(np.float32) # [Tpx,2500,2]
|
||||
tracks = np.clip(tracks, 0.0, 1.0)
|
||||
vis = np.ones(tracks.shape[:2], np.float32)
|
||||
return tracks, vis, handle_mask
|
||||
|
||||
|
||||
def _overlay_grid(frames: np.ndarray, tracks: np.ndarray, vis: np.ndarray, moving, stride: int = 3) -> np.ndarray:
|
||||
"""Draw the input tracks: moving handle (green + tails) + static background (gray dots).
|
||||
|
||||
For dense mode (N = 50*50 = 2500) we grid-subsample by ``stride`` for legibility. For
|
||||
sparse mode (N ~= handle + num_bg, much smaller) we draw every point, no subsample.
|
||||
"""
|
||||
from fastvideo.train.callbacks.track_validation import _draw_overlay, _subsample
|
||||
T, H, W, _ = frames.shape
|
||||
tt = min(T, tracks.shape[0])
|
||||
tr, vs, mv = tracks[:tt], vis[:tt], moving
|
||||
if tr.shape[1] == 2500:
|
||||
tr, vs = _subsample(tr, vs, 50, stride)
|
||||
mv_full = np.broadcast_to(moving[None, :, None].astype(np.float32),
|
||||
(tt, moving.shape[0], 2)).copy()
|
||||
mv, _ = _subsample(mv_full, vis[:tt], 50, stride)
|
||||
is_mv = mv[0, :, 0] > 0.5
|
||||
else: # sparse: draw all points
|
||||
is_mv = mv.astype(bool)
|
||||
colors = np.where(is_mv[:, None], np.array([0, 255, 100], np.uint8),
|
||||
np.array([120, 120, 120], np.uint8)).astype(np.uint8)
|
||||
trpx = tr.copy()
|
||||
trpx[..., 0] *= W
|
||||
trpx[..., 1] *= H
|
||||
return _draw_overlay(frames[:tt], trpx, vs, colors, 12, 2, 0.5)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------- handlers
|
||||
def on_clip(clip_idx):
|
||||
clip_idx = int(clip_idx)
|
||||
img = _first_frame_rgb(clip_idx)
|
||||
cap = STATE["samples"][clip_idx]["caption"][:240]
|
||||
return cap, _img_b64(img), _original_video(clip_idx)
|
||||
|
||||
|
||||
def on_ckpt(name):
|
||||
return swap_checkpoint(name)
|
||||
|
||||
|
||||
def _build_tracks(traj_json, radius, falloff, sparse, num_bg, seed):
|
||||
"""Parse the drawn trajectory -> (tracks[Tpx,N,2], vis, moving[N]) or (None, msg).
|
||||
|
||||
``sparse``: "dense" (full 2500 grid, background static) matches the OLD training
|
||||
regime, "sparse" (handle + ``num_bg`` random background) matches the SPARSE
|
||||
training regime our overfit runs used (WANTRACK_SPARSE=1, ~20 background extras).
|
||||
"""
|
||||
import json
|
||||
if not traj_json or traj_json.strip() in ("", "[]"):
|
||||
return None, "Record an action first: move the mouse over the frame and press Space (5 s)."
|
||||
try:
|
||||
traj = json.loads(traj_json)
|
||||
except Exception as e: # noqa: BLE001
|
||||
return None, f"bad trajectory json: {e}"
|
||||
tracks, vis, moving = _grid_action_tracks(traj, STATE["Tpx"], float(radius), str(falloff),
|
||||
str(sparse), int(num_bg), int(seed))
|
||||
if tracks is None:
|
||||
return None, "trajectory too short — hold Space and move the mouse for the full 5 s."
|
||||
return (tracks, vis, moving), None
|
||||
|
||||
|
||||
def _first_frame_cached(clip_idx):
|
||||
cache = STATE.setdefault("_ff_cache", {})
|
||||
if clip_idx not in cache:
|
||||
cache[clip_idx] = _first_frame_rgb(clip_idx)
|
||||
return cache[clip_idx]
|
||||
|
||||
|
||||
def on_preview(clip_idx, traj_json, radius, falloff, sparse, num_bg, seed):
|
||||
"""Show the input tracks overlaid on the (frozen) first frame -- BEFORE generating."""
|
||||
clip_idx = int(clip_idx)
|
||||
built, err = _build_tracks(traj_json, radius, falloff, sparse, num_bg, seed)
|
||||
if built is None:
|
||||
return None, err
|
||||
tracks, vis, moving = built
|
||||
ff = _first_frame_cached(clip_idx) # [H,W,3]
|
||||
frames = np.repeat(ff[None], tracks.shape[0], axis=0) # freeze frame 0, T copies
|
||||
out = _overlay_grid(frames, tracks, vis, moving)
|
||||
path = os.path.join(STATE["out_dir"], f"preview_clip{clip_idx}.mp4")
|
||||
imageio.mimsave(path, out, fps=24, macro_block_size=1)
|
||||
kind = "SPARSE" if str(sparse) == "sparse" else "DENSE 50x50"
|
||||
return path, (f"PREVIEW ({kind}, N={tracks.shape[1]}): {int(moving.sum())} moving handle points (green, tails) "
|
||||
f"and {int((~moving).sum())} background points (gray). "
|
||||
f"This matches the training track format. Press Generate.")
|
||||
|
||||
|
||||
def on_generate(clip_idx, ckpt_name, traj_json, radius, falloff, sparse, steps, seed,
|
||||
w_text: float = 3.0, w_motion: float = 1.5, mode: str = "joint", num_bg: int = 20):
|
||||
"""Generator: streams per-step denoise progress to the log, yields the video at the end."""
|
||||
import time
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
clip_idx = int(clip_idx)
|
||||
# Sparse mode uses the same "seed" slider as the tracks background seed AND the noise seed
|
||||
# (fine for interactive use; the two are independent knobs technically).
|
||||
built, err = _build_tracks(traj_json, radius, falloff, sparse, num_bg, seed)
|
||||
if built is None:
|
||||
yield gr.update(), f"[error] {err}", gr.update()
|
||||
return
|
||||
tracks, vis, moving = built
|
||||
steps = int(steps)
|
||||
t0 = time.time()
|
||||
|
||||
yield gr.update(), f"[1/4] swapping to {ckpt_name} ...", gr.update()
|
||||
swap_checkpoint(ckpt_name)
|
||||
|
||||
model = STATE["model"]
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
s = STATE["samples"][clip_idx]
|
||||
ff = s["first_frame_latent"].to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
txt = s["text_embedding"].to(device, dtype)
|
||||
mask = s["text_attention_mask"].to(device, dtype)
|
||||
img = s["clip_feature"].to(device, dtype)
|
||||
tp = torch.from_numpy(tracks)[None].to(device, dtype)
|
||||
tv = torch.from_numpy(vis)[None].to(device, dtype)
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=float(model.timestep_shift))
|
||||
sched.set_timesteps(steps, device=device)
|
||||
|
||||
# Match TrackValidationCallback._sample: MotionStream joint text+motion CFG (Eq. 2).
|
||||
# mode="joint" -> full formula (3 NFE / step). Default. Matches training-time val.
|
||||
# mode="text_only" -> text CFG only, tracks in forward (2 NFE / step).
|
||||
# mode="no_track" -> tracks=None, text CFG only (2 NFE / step).
|
||||
tp_eff = None if mode == "no_track" else tp
|
||||
tv_eff = None if mode == "no_track" else tv
|
||||
wt, wm = float(w_text), float(w_motion)
|
||||
cfg_on = (wt != 1.0) or (wm != 1.0)
|
||||
txt_null = torch.zeros_like(txt) if cfg_on else None
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, tsv):
|
||||
with torch.no_grad(), torch.autocast(device.type, dtype=dtype), \
|
||||
set_forward_context(current_timestep=tsv, attn_metadata=None):
|
||||
return model.transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=tsv,
|
||||
encoder_hidden_states_image=img, track_points=tp_e,
|
||||
track_visibility=tv_e, return_dict=False)
|
||||
|
||||
for i, t in enumerate(sched.timesteps):
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
ts = t.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp_eff, tv_eff, model_in, ts)
|
||||
if not cfg_on:
|
||||
v = v_full
|
||||
elif tp_eff is None or mode == "text_only":
|
||||
v_no_text = _fwd(txt_null, tp_eff, tv_eff, model_in, ts)
|
||||
v = v_no_text + wt * (v_full - v_no_text)
|
||||
else: # mode="joint" (Eq. 2)
|
||||
v_no_text = _fwd(txt_null, tp_eff, tv_eff, model_in, ts)
|
||||
v_no_motion = _fwd(txt, None, None, model_in, ts)
|
||||
alpha = wt / (wt + wm) if (wt + wm) > 0 else 0.5
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + wt * (v_full - v_no_text) + wm * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), t, latents.float(), return_dict=False)[0]
|
||||
yield gr.update(), f"[2/4] denoising {i + 1}/{steps} ({time.time() - t0:.1f}s, cfg={mode})", gr.update()
|
||||
|
||||
yield gr.update(), f"[3/4] decoding latents ... ({time.time() - t0:.1f}s)", gr.update()
|
||||
frames = twi.decode_to_pixels(model, latents)
|
||||
out = _overlay_grid(frames, tracks, vis, moving)
|
||||
tag = ckpt_name.replace(" ", "")
|
||||
path = os.path.join(STATE["out_dir"], f"action_clip{clip_idx}_{tag}.mp4")
|
||||
imageio.mimsave(path, out, fps=24, macro_block_size=1)
|
||||
|
||||
yield gr.update(), f"[4/4] computing CoTracker EPE ... ({time.time() - t0:.1f}s)", gr.update()
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import compute_epe
|
||||
# EPE on the MOVING points only (did the arm actually follow your path)
|
||||
mv_idx = np.where(moving)[0]
|
||||
ppx = tracks[:, mv_idx].copy()
|
||||
ppx[..., 0] *= W
|
||||
ppx[..., 1] *= H
|
||||
res = compute_epe(frames, ppx, vis[:, mv_idx], STATE["ct"], device)
|
||||
epe = res["epe"]
|
||||
epe_msg = (f"moving-point EPE = {epe:.2f}px ({res['n_points']} moved pts) — "
|
||||
f"lower = the arm region followed your drawn path."
|
||||
if epe is not None else "EPE n/a (no points re-tracked)")
|
||||
yield path, f"[done] generated in {time.time() - t0:.1f}s", epe_msg
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------- ui
|
||||
def _canvas_js(Tpx: int) -> str:
|
||||
return ("() => {\n"
|
||||
" const cvs = document.getElementById('cvs'); if(!cvs||cvs._init) return; cvs._init=true;\n"
|
||||
" const ctx = cvs.getContext('2d');\n"
|
||||
f" const st = {{img:new Image(), mouse:[0.5,0.5], rec:false, traj:[], N:{Tpx}}};\n"
|
||||
" cvs.width=512; cvs.height=288;\n"
|
||||
" function draw(){ ctx.clearRect(0,0,cvs.width,cvs.height);\n"
|
||||
" if(st.img.complete && st.img.width) ctx.drawImage(st.img,0,0,cvs.width,cvs.height);\n"
|
||||
" if(st.traj.length){ ctx.strokeStyle='#00ff66'; ctx.lineWidth=3; ctx.beginPath();\n"
|
||||
" st.traj.forEach((p,i)=>{const x=p[0]*cvs.width,y=p[1]*cvs.height; i?ctx.lineTo(x,y):ctx.moveTo(x,y);}); ctx.stroke();\n"
|
||||
" const s=st.traj[0], e=st.traj[st.traj.length-1];\n"
|
||||
" ctx.fillStyle='#00ff66'; ctx.beginPath(); ctx.arc(s[0]*cvs.width,s[1]*cvs.height,6,0,7); ctx.fill();\n"
|
||||
" ctx.fillStyle='#ff3333'; ctx.beginPath(); ctx.arc(e[0]*cvs.width,e[1]*cvs.height,6,0,7); ctx.fill(); } }\n"
|
||||
" cvs.addEventListener('mousemove', e=>{const r=cvs.getBoundingClientRect(); st.mouse=[(e.clientX-r.left)/r.width,(e.clientY-r.top)/r.height];});\n"
|
||||
" window._setbg=(b64)=>{ if(!b64) return; const im=new Image(); im.onload=()=>{ const ar=im.width/im.height; cvs.width=512; cvs.height=Math.round(512/ar); st.img=im; st.traj=[]; draw(); }; im.src=b64; };\n"
|
||||
" window._clear=()=>{ st.traj=[]; const tb=document.querySelector('#traj_json textarea'); if(tb){tb.value=''; tb.dispatchEvent(new Event('input',{bubbles:true}));} draw(); };\n"
|
||||
" window._startRec=()=>{ if(st.rec) return; st.rec=true; st.traj=[]; let n=0; const stt=document.getElementById('recstatus');\n"
|
||||
" const iv=setInterval(()=>{ st.traj.push([st.mouse[0],st.mouse[1]]); n++; draw(); if(stt) stt.textContent='\\u25CF REC '+n+'/'+st.N;\n"
|
||||
" if(n>=st.N){ clearInterval(iv); st.rec=false; if(stt) stt.textContent='recorded '+st.N+' frames'; \n"
|
||||
" const tb=document.querySelector('#traj_json textarea'); tb.value=JSON.stringify(st.traj); tb.dispatchEvent(new Event('input',{bubbles:true})); } }, 1000/24); };\n"
|
||||
" document.addEventListener('keydown', e=>{ if(e.code==='Space'){ e.preventDefault(); window._startRec(); } });\n"
|
||||
" draw();\n"
|
||||
"}")
|
||||
|
||||
|
||||
def build_ui(num_clips: int, ckpt_names: list[str], Tpx: int):
|
||||
canvas_html = ("<div style='user-select:none'>"
|
||||
"<canvas id='cvs' style='border:1px solid #888;cursor:crosshair;max-width:100%'></canvas>"
|
||||
"<div id='recstatus' style='font-weight:bold;color:#0a0;height:1.4em'></div>"
|
||||
"<div style='font-size:0.85em;color:#888'>Move the mouse over the frame, press "
|
||||
"<b>Space</b> to record your action for 5 s. Green=start, red=end.</div></div>")
|
||||
with gr.Blocks(title="TrackWan — action recorder") as demo:
|
||||
gr.Markdown("## TrackWan — record an action, test action-following across checkpoints\n"
|
||||
"Pick a clip + checkpoint, **press Space** and move the mouse over the frame to draw a "
|
||||
"5 second motion, then **Generate**. The 50x50 grid points near your start are dragged along your path (rest stay static).")
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
ckpt = gr.Dropdown(ckpt_names, value=ckpt_names[-1], label="Checkpoint (hot-swap)")
|
||||
clip = gr.Dropdown(list(range(num_clips)), value=0, label="Clip (first frame + prompt)")
|
||||
caption = gr.Textbox(label="Prompt", interactive=False, lines=2)
|
||||
gr.HTML(canvas_html)
|
||||
with gr.Row():
|
||||
rec_btn = gr.Button("● Record (Space)")
|
||||
clr_btn = gr.Button("Clear")
|
||||
with gr.Row():
|
||||
radius = gr.Slider(0.02, 0.5, value=0.15, step=0.01,
|
||||
label="Action radius (grid pts within this move)")
|
||||
falloff = gr.Radio(["hard", "smooth"], value="hard", label="Falloff")
|
||||
with gr.Row():
|
||||
background = gr.Radio(["dense", "sparse"],
|
||||
value="sparse",
|
||||
label="Track budget (matches training: WANTRACK_SPARSE)")
|
||||
num_bg = gr.Slider(0, 60, value=20, step=1,
|
||||
label="Sparse: # background points (WANTRACK_EXTRA_RANDOM)")
|
||||
with gr.Row():
|
||||
steps = gr.Slider(10, 60, value=30, step=5, label="Denoise steps")
|
||||
seed = gr.Number(value=1000, label="Seed (also seeds sparse background pool)")
|
||||
with gr.Row():
|
||||
w_text = gr.Slider(1.0, 8.0, value=3.0, step=0.5, label="Text CFG (w_t)")
|
||||
w_motion = gr.Slider(1.0, 5.0, value=1.5, step=0.25, label="Motion CFG (w_m)")
|
||||
mode = gr.Radio(["joint", "text_only", "no_track"],
|
||||
value="joint",
|
||||
label="CFG mode (joint = Eq. 2; text_only = drop motion arm; no_track = tracks=None)")
|
||||
with gr.Row():
|
||||
preview_btn = gr.Button("Preview traces")
|
||||
gen_btn = gr.Button("Generate", variant="primary")
|
||||
status = gr.Textbox(label="Status / generation log", interactive=False, lines=2)
|
||||
with gr.Column(scale=1):
|
||||
orig_vid = gr.Video(label="Original clip")
|
||||
gen_vid = gr.Video(label="Preview traces / Generated (your action overlaid)")
|
||||
epe_box = gr.Textbox(label="Action-following (moving-point CoTracker EPE)", interactive=False)
|
||||
|
||||
ff_b64 = gr.Textbox(elem_id="ff_b64", visible=False)
|
||||
traj_json = gr.Textbox(elem_id="traj_json", visible=False)
|
||||
|
||||
clip.change(on_clip, [clip], [caption, ff_b64, orig_vid])
|
||||
ckpt.change(on_ckpt, [ckpt], [status])
|
||||
ff_b64.change(None, [ff_b64], None, js="(b64)=>{ window._setbg(b64); }")
|
||||
rec_btn.click(None, None, None, js="()=>{ window._startRec(); }")
|
||||
clr_btn.click(None, None, None, js="()=>{ window._clear(); }")
|
||||
preview_btn.click(on_preview,
|
||||
[clip, traj_json, radius, falloff, background, num_bg, seed],
|
||||
[gen_vid, status])
|
||||
gen_btn.click(on_generate,
|
||||
[clip, ckpt, traj_json, radius, falloff, background, steps, seed,
|
||||
w_text, w_motion, mode, num_bg],
|
||||
[gen_vid, status, epe_box])
|
||||
demo.load(on_clip, [clip], [caption, ff_b64, orig_vid])
|
||||
demo.load(None, None, None, js=_canvas_js(Tpx))
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--exports-root", required=True,
|
||||
help="dir containing export_* checkpoint dirs (each a dcp_to_diffusers export)")
|
||||
p.add_argument("--yaml", default="examples/train/scenario/worldmodel/finetune_wantrack_i2v.yaml")
|
||||
p.add_argument("--data", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/"
|
||||
"wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp_now/combined_parquet_dataset")
|
||||
p.add_argument("--num-clips", type=int, default=10)
|
||||
p.add_argument("--out-dir", default="/mnt/weka/home/hao.zhang/shao/data/motion_pipeline/gradio_action_out")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7861)
|
||||
p.add_argument("--share", action="store_true")
|
||||
args = p.parse_args()
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
roots = sorted(glob.glob(os.path.join(args.exports_root, "export_*")))
|
||||
if not roots:
|
||||
raise SystemExit(f"no export_* dirs under {args.exports_root}")
|
||||
exports = {f"step {os.path.basename(r).split('_')[-1]}": r for r in roots}
|
||||
names = list(exports.keys())
|
||||
|
||||
model, tc = twi.load_trackwan(exports[names[-1]], args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
samples = twi.load_conditioning_from_parquet(args.data, list(range(args.num_clips)), text_len)
|
||||
num_lat_t = samples[0]["first_frame_latent"].shape[2]
|
||||
ratio = int(tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
|
||||
from fastvideo.eval.metrics.motion.cotracker_epe.metric import load_cotracker
|
||||
STATE.update(model=model, tc=tc, samples=samples, Tpx=(num_lat_t - 1) * ratio + 1,
|
||||
out_dir=args.out_dir, ct=load_cotracker(model.device), exports=exports,
|
||||
cur_ckpt=names[-1])
|
||||
|
||||
demo = build_ui(len(samples), names, STATE["Tpx"])
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[args.out_dir])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,543 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""TrackWan interactive rollout demo.
|
||||
|
||||
Two point-input modes on the same canvas:
|
||||
- **Trace mode** — hover, press *Space* to record a 5-second (121-frame) drag.
|
||||
- **Anchor mode** — click once to drop a static point (a track whose (x,y) is
|
||||
constant across all 121 frames).
|
||||
|
||||
Design (matches training-time validation exactly):
|
||||
- Loads conditioning (text_embedding, clip_feature, first_frame_latent) DIRECTLY from
|
||||
the preprocessed openvid parquet. Same path as
|
||||
``fastvideo/train/callbacks/track_validation.py``.
|
||||
- First frame image shown to the user is decoded from the VAE latent (``decode_reference``).
|
||||
- Both traces and anchors are packed into a single ``track_points[T,N,2]`` tensor.
|
||||
Anchors get their (x,y) broadcast across T=121; traces are linearly interpolated.
|
||||
- Generation: the SAME motion-CFG denoise loop as ``TrackValidationCallback._sample``
|
||||
(MotionStream Eq. 2, wt=3.0, wm=1.5 by default).
|
||||
|
||||
Launch::
|
||||
|
||||
srun --overlap --jobid=<job> --ntasks=1 -w <node> --chdir=$PWD bash -lc "
|
||||
source .venv/bin/activate
|
||||
export ... TRACKWAN_TRACK_BIAS=1 CUDA_VISIBLE_DEVICES=0
|
||||
python examples/inference/gradio/trackwan/app_multi_trace.py \\
|
||||
--model-dir /path/to/diffusers_export --yaml <yaml> \\
|
||||
--data-path /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
"
|
||||
|
||||
share=True doesn't work on aarch64 (no frpc); use SSH port-forward on the compute node.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import gradio as gr
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
REPO = Path(__file__).resolve().parents[4]
|
||||
NUM_FRAMES = 121
|
||||
FPS = 24
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# frame b64 + track construction
|
||||
# =============================================================================
|
||||
def img_to_b64(arr: np.ndarray) -> str:
|
||||
buf = io.BytesIO()
|
||||
Image.fromarray(arr).save(buf, format="PNG")
|
||||
return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode()
|
||||
|
||||
|
||||
def build_track_points(traces_json: str, anchors_json: str) -> tuple[np.ndarray, np.ndarray, int, int]:
|
||||
"""(traces, anchors) → (track_points[T,N,2], track_visibility[T,N], n_tr, n_an).
|
||||
|
||||
Each trace is a list of normalized (x,y) samples at ~24 Hz over 5 s → interpolated to T=121.
|
||||
Each anchor is a single normalized (x,y) → broadcast across all T frames.
|
||||
"""
|
||||
T = NUM_FRAMES
|
||||
try:
|
||||
traces = json.loads(traces_json) if traces_json else []
|
||||
except Exception:
|
||||
traces = []
|
||||
try:
|
||||
anchors = json.loads(anchors_json) if anchors_json else []
|
||||
except Exception:
|
||||
anchors = []
|
||||
n_tr, n_an = len(traces), len(anchors)
|
||||
N = n_tr + n_an
|
||||
if N == 0:
|
||||
return np.zeros((T, 0, 2), np.float32), np.zeros((T, 0), np.float32), 0, 0
|
||||
tp = np.zeros((T, N, 2), np.float32)
|
||||
tv = np.ones((T, N), np.float32)
|
||||
for i, traj in enumerate(traces):
|
||||
if not traj:
|
||||
continue
|
||||
arr = np.asarray(traj, np.float32)
|
||||
if arr.shape[0] < 2:
|
||||
tp[:, i] = arr[0]
|
||||
else:
|
||||
src = np.arange(arr.shape[0])
|
||||
tgt = np.linspace(0, arr.shape[0] - 1, T)
|
||||
tp[:, i, 0] = np.interp(tgt, src, arr[:, 0])
|
||||
tp[:, i, 1] = np.interp(tgt, src, arr[:, 1])
|
||||
for j, (ax, ay) in enumerate(anchors):
|
||||
tp[:, n_tr + j, 0] = ax
|
||||
tp[:, n_tr + j, 1] = ay
|
||||
return tp, tv, n_tr, n_an
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# canvas JS (traces via Space + anchors via click, mode radio decides)
|
||||
# =============================================================================
|
||||
def canvas_js(num_frames: int, fps: int) -> str:
|
||||
palette = ["#ff3c3c", "#3cd23c", "#3c78ff", "#ffc83c", "#c83cc8", "#3cdcdc", "#ff821e", "#9696ff"]
|
||||
return ("() => {\n"
|
||||
" const cvs = document.getElementById('cvs'); if (!cvs || cvs._init) return; cvs._init=true;\n"
|
||||
" const ctx = cvs.getContext('2d');\n"
|
||||
f" const PAL = {json.dumps(palette)};\n"
|
||||
f" const N = {num_frames}; const DT = {int(1000 / fps)};\n"
|
||||
" const st = { img:new Image(), mouse:[0.5,0.5], rec:false, current:[], "
|
||||
"traces:[], anchors:[], mode:'trace' };\n"
|
||||
" window._st = st;\n"
|
||||
" cvs.width = 832; cvs.height = 480;\n"
|
||||
" function commit_traces(){\n"
|
||||
" const tb = document.querySelector('#traces_json textarea');\n"
|
||||
" if (tb) { tb.value = JSON.stringify(st.traces); tb.dispatchEvent(new Event('input',{bubbles:true})); }\n"
|
||||
" }\n"
|
||||
" function commit_anchors(){\n"
|
||||
" const tb = document.querySelector('#anchors_json textarea');\n"
|
||||
" if (tb) { tb.value = JSON.stringify(st.anchors); tb.dispatchEvent(new Event('input',{bubbles:true})); }\n"
|
||||
" }\n"
|
||||
" function draw(){\n"
|
||||
" ctx.clearRect(0,0,cvs.width,cvs.height);\n"
|
||||
" if (st.img.complete && st.img.width) ctx.drawImage(st.img, 0, 0, cvs.width, cvs.height);\n"
|
||||
" // anchors: solid gray dots with dark ring\n"
|
||||
" st.anchors.forEach((p, i) => { const x=p[0]*cvs.width, y=p[1]*cvs.height;\n"
|
||||
" ctx.beginPath(); ctx.arc(x,y,7,0,7); ctx.fillStyle='#888'; ctx.fill();\n"
|
||||
" ctx.strokeStyle='#000'; ctx.lineWidth=2; ctx.stroke();\n"
|
||||
" ctx.fillStyle='#fff'; ctx.font='bold 10px sans-serif'; ctx.textAlign='center'; ctx.textBaseline='middle';\n"
|
||||
" ctx.fillText('A'+(i+1), x, y);\n"
|
||||
" });\n"
|
||||
" // traces: colored polyline with start/end markers\n"
|
||||
" st.traces.forEach((traj, i) => {\n"
|
||||
" const color = PAL[i % PAL.length];\n"
|
||||
" ctx.strokeStyle=color; ctx.lineWidth=3; ctx.beginPath();\n"
|
||||
" traj.forEach((p, j) => { const x=p[0]*cvs.width, y=p[1]*cvs.height; j?ctx.lineTo(x,y):ctx.moveTo(x,y); });\n"
|
||||
" ctx.stroke();\n"
|
||||
" const s = traj[0], e = traj[traj.length-1];\n"
|
||||
" ctx.fillStyle=color; ctx.beginPath(); ctx.arc(s[0]*cvs.width, s[1]*cvs.height, 6, 0, 7); ctx.fill(); ctx.strokeStyle='#000'; ctx.lineWidth=1; ctx.stroke();\n"
|
||||
" ctx.fillStyle='#ff3333'; ctx.beginPath(); ctx.arc(e[0]*cvs.width, e[1]*cvs.height, 6, 0, 7); ctx.fill(); ctx.strokeStyle='#000'; ctx.stroke();\n"
|
||||
" ctx.fillStyle='#000'; ctx.font='13px sans-serif'; ctx.textAlign='left'; ctx.textBaseline='alphabetic'; ctx.fillText('T'+(i+1), s[0]*cvs.width+8, s[1]*cvs.height-8);\n"
|
||||
" });\n"
|
||||
" // live-drawing current trace\n"
|
||||
" if (st.rec && st.current.length) {\n"
|
||||
" const color = PAL[st.traces.length % PAL.length];\n"
|
||||
" ctx.strokeStyle=color; ctx.lineWidth=3; ctx.beginPath();\n"
|
||||
" st.current.forEach((p, j) => { const x=p[0]*cvs.width, y=p[1]*cvs.height; j?ctx.lineTo(x,y):ctx.moveTo(x,y); });\n"
|
||||
" ctx.stroke();\n"
|
||||
" }\n"
|
||||
" }\n"
|
||||
" cvs.addEventListener('mousemove', e => {\n"
|
||||
" const r = cvs.getBoundingClientRect();\n"
|
||||
" st.mouse = [(e.clientX - r.left) / r.width, (e.clientY - r.top) / r.height];\n"
|
||||
" });\n"
|
||||
" cvs.addEventListener('click', e => {\n"
|
||||
" if (st.rec) return;\n"
|
||||
" const r = cvs.getBoundingClientRect();\n"
|
||||
" const cx = (e.clientX - r.left) * (cvs.width / r.width);\n"
|
||||
" const cy = (e.clientY - r.top) * (cvs.height / r.height);\n"
|
||||
" if (st.mode === 'anchor') {\n"
|
||||
" // toggle delete-if-near, else place\n"
|
||||
" let idx = -1, best = 14*14;\n"
|
||||
" st.anchors.forEach((p, i) => { const dx = p[0]*cvs.width - cx, dy = p[1]*cvs.height - cy;\n"
|
||||
" const d2 = dx*dx + dy*dy; if (d2 < best) { best = d2; idx = i; } });\n"
|
||||
" if (idx >= 0) st.anchors.splice(idx, 1);\n"
|
||||
" else st.anchors.push([cx / cvs.width, cy / cvs.height]);\n"
|
||||
" commit_anchors(); draw();\n"
|
||||
" }\n"
|
||||
" });\n"
|
||||
" window._setmode = (m) => { st.mode = m || 'trace'; };\n"
|
||||
" window._setbg = (b64) => {\n"
|
||||
" if (!b64) return;\n"
|
||||
" const im = new Image();\n"
|
||||
" im.onload = () => { const ar = im.width / im.height; cvs.width = 832; cvs.height = Math.round(832 / ar); st.img = im; draw(); };\n"
|
||||
" im.src = b64;\n"
|
||||
" };\n"
|
||||
" window._clear_traces = () => {\n"
|
||||
" st.traces = []; st.current = []; st.rec = false; commit_traces(); draw();\n"
|
||||
" const stt = document.getElementById('recstatus'); if (stt) stt.textContent = 'traces cleared. Space for trace #1';\n"
|
||||
" };\n"
|
||||
" window._undo_trace = () => {\n"
|
||||
" if (st.traces.length) st.traces.pop();\n"
|
||||
" commit_traces(); draw();\n"
|
||||
" const stt = document.getElementById('recstatus'); if (stt) stt.textContent = st.traces.length + ' trace(s). Space for another';\n"
|
||||
" };\n"
|
||||
" window._clear_anchors = () => { st.anchors = []; commit_anchors(); draw(); };\n"
|
||||
" window._undo_anchor = () => { if (st.anchors.length) st.anchors.pop(); commit_anchors(); draw(); };\n"
|
||||
" window._startRec = () => {\n"
|
||||
" if (st.rec) return;\n"
|
||||
" st.rec = true; st.current = []; let n = 0;\n"
|
||||
" const stt = document.getElementById('recstatus');\n"
|
||||
" const iv = setInterval(() => {\n"
|
||||
" st.current.push([st.mouse[0], st.mouse[1]]); n++; draw();\n"
|
||||
" if (stt) stt.textContent = 'REC trace #' + (st.traces.length + 1) + ' — ' + n + '/' + N;\n"
|
||||
" if (n >= N) {\n"
|
||||
" clearInterval(iv); st.rec = false;\n"
|
||||
" st.traces.push(st.current); st.current = [];\n"
|
||||
" if (stt) stt.textContent = 'committed trace #' + st.traces.length + ' (Space for another, or Generate)';\n"
|
||||
" commit_traces(); draw();\n"
|
||||
" }\n"
|
||||
" }, DT);\n"
|
||||
" };\n"
|
||||
" document.addEventListener('keydown', e => {\n"
|
||||
" if (e.code !== 'Space' || e.repeat) return;\n"
|
||||
" const tag = document.activeElement && document.activeElement.tagName;\n"
|
||||
" if (tag === 'INPUT' || tag === 'TEXTAREA') return;\n"
|
||||
" if (st.mode !== 'trace') return; // Space only records in trace mode\n"
|
||||
" e.preventDefault(); window._startRec();\n"
|
||||
" });\n"
|
||||
" draw();\n"
|
||||
"}")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Gradio UI
|
||||
# =============================================================================
|
||||
def build_ui(state):
|
||||
canvas_html = ("<div style='user-select:none'>"
|
||||
"<canvas id='cvs' style='border:1px solid #888;cursor:crosshair;max-width:100%'></canvas>"
|
||||
"<div id='recstatus' style='font-weight:bold;color:#0a0;height:1.4em;margin-top:6px'>"
|
||||
"pick a clip, click <b>Load frame</b>, then choose a mode below</div>"
|
||||
"</div>")
|
||||
labels = state["labels"]
|
||||
with gr.Blocks(title="TrackWan interactive rollout") as demo:
|
||||
gr.Markdown(
|
||||
"## TrackWan — interactive rollout demo\n"
|
||||
"1. Pick a preprocessed clip → **Load frame**.\n"
|
||||
"2. Choose a mode:\n"
|
||||
" - **Trace** — hover on the frame and press **Space** to record a 5 s drag.\n"
|
||||
" - **Anchor** — click on the frame to drop a *static* point. Click an existing anchor to remove it.\n"
|
||||
"3. **Generate**. Denoise formula matches training's validation exactly (MotionStream Eq. 2).")
|
||||
|
||||
with gr.Row():
|
||||
with gr.Column(scale=1):
|
||||
clip = gr.Dropdown(labels, value=labels[0] if labels else None, label="Preprocessed clip")
|
||||
caption = gr.Textbox(label="Caption (from parquet)", interactive=False, lines=2)
|
||||
load_btn = gr.Button("Load frame", variant="primary")
|
||||
|
||||
gr.Markdown("**Input mode**")
|
||||
mode = gr.Radio(["trace", "anchor"], value="trace",
|
||||
label="Mode",
|
||||
info="Trace = Space to record a moving drag. Anchor = single click for a static point.")
|
||||
|
||||
gr.Markdown("**Traces** (moving)")
|
||||
with gr.Row():
|
||||
rec_btn = gr.Button("● Record (Space)")
|
||||
undo_trace_btn = gr.Button("Undo trace")
|
||||
clear_trace_btn = gr.Button("Clear traces")
|
||||
|
||||
gr.Markdown("**Anchors** (static)")
|
||||
with gr.Row():
|
||||
undo_anchor_btn = gr.Button("Undo anchor")
|
||||
clear_anchor_btn = gr.Button("Clear anchors")
|
||||
|
||||
info_md = gr.Markdown("_no traces or anchors yet._")
|
||||
|
||||
gr.Markdown("**Generation** (denoise formula matches training's validation exactly)")
|
||||
seed_in = gr.Slider(0, 9999, 1000, step=1, label="Seed")
|
||||
steps_in = gr.Slider(10, 60, 30, step=1, label="Denoise steps")
|
||||
w_text = gr.Slider(1.0, 8.0, 3.0, step=0.25, label="Text CFG (w_t)")
|
||||
w_motion = gr.Slider(1.0, 5.0, 1.5, step=0.25, label="Motion CFG (w_m)")
|
||||
gen_btn = gr.Button("Generate video", variant="primary")
|
||||
with gr.Column(scale=2):
|
||||
gr.HTML(canvas_html)
|
||||
out_video = gr.Video(label="Generated video", height=380)
|
||||
status = gr.Markdown("_ready_")
|
||||
log_box = gr.Textbox(label="Generation log", interactive=False, lines=8,
|
||||
max_lines=16, value="_no runs yet._")
|
||||
|
||||
# hidden bridges
|
||||
ff_b64 = gr.Textbox(elem_id="ff_b64", visible=False)
|
||||
traces_json = gr.Textbox(elem_id="traces_json", visible=False, value="[]")
|
||||
anchors_json = gr.Textbox(elem_id="anchors_json", visible=False, value="[]")
|
||||
|
||||
# server-side per-page state
|
||||
st_clip_idx = gr.State(0)
|
||||
st_frame = gr.State(None)
|
||||
|
||||
# ---------------- caption on clip change ----------------
|
||||
def _cap(clip_name):
|
||||
i = state["label_to_idx"].get(clip_name, 0)
|
||||
return i, state["samples"][i]["caption"][:400]
|
||||
clip.change(_cap, [clip], [st_clip_idx, caption])
|
||||
|
||||
# ---------------- load frame (decode ref frame from VAE latent) ----------------
|
||||
first_frame_cache: dict[int, np.ndarray] = {}
|
||||
|
||||
def _load(clip_name):
|
||||
i = state["label_to_idx"].get(clip_name, 0)
|
||||
if i not in first_frame_cache:
|
||||
first_frame_cache[i] = state["decode_first_frame"](i)
|
||||
frame = first_frame_cache[i]
|
||||
msg = f"_clip loaded ({state['samples'][i]['caption'][:60]}...)_"
|
||||
return i, frame, img_to_b64(frame), msg
|
||||
|
||||
load_btn.click(_load, [clip], [st_clip_idx, st_frame, ff_b64, status])
|
||||
|
||||
# preload first clip on startup
|
||||
def _preload():
|
||||
i = 0
|
||||
if i not in first_frame_cache:
|
||||
first_frame_cache[i] = state["decode_first_frame"](i)
|
||||
frame = first_frame_cache[i]
|
||||
return (i, state["samples"][i]["caption"][:400], frame, img_to_b64(frame),
|
||||
"_default clip preloaded. Draw a trace (Space) or place anchors (click)._")
|
||||
demo.load(_preload, None, [st_clip_idx, caption, st_frame, ff_b64, status])
|
||||
|
||||
# ---------------- JS bridges ----------------
|
||||
ff_b64.change(None, ff_b64, None, js="(b) => window._setbg && window._setbg(b)")
|
||||
mode.change(None, mode, None, js="(m) => window._setmode && window._setmode(m)")
|
||||
rec_btn.click(None, None, None, js="() => window._startRec && window._startRec()")
|
||||
undo_trace_btn.click(None, None, None, js="() => window._undo_trace && window._undo_trace()")
|
||||
clear_trace_btn.click(None, None, None, js="() => window._clear_traces && window._clear_traces()")
|
||||
undo_anchor_btn.click(None, None, None, js="() => window._undo_anchor && window._undo_anchor()")
|
||||
clear_anchor_btn.click(None, None, None, js="() => window._clear_anchors && window._clear_anchors()")
|
||||
|
||||
def _tinfo(t, a):
|
||||
try:
|
||||
n_t = len(json.loads(t) if t else [])
|
||||
except Exception:
|
||||
n_t = 0
|
||||
try:
|
||||
n_a = len(json.loads(a) if a else [])
|
||||
except Exception:
|
||||
n_a = 0
|
||||
if n_t == 0 and n_a == 0:
|
||||
return "_no traces or anchors yet._"
|
||||
return f"_{n_t} trace(s) + {n_a} anchor(s) → {n_t + n_a} total track(s)_"
|
||||
traces_json.change(_tinfo, [traces_json, anchors_json], info_md)
|
||||
anchors_json.change(_tinfo, [traces_json, anchors_json], info_md)
|
||||
|
||||
# ---------------- generate (MATCHES TrackValidationCallback._sample exactly) ----------------
|
||||
def _generate(clip_idx, traces_str, anchors_str, seed, steps, w_t, w_m):
|
||||
def logln(log, msg):
|
||||
line = time.strftime("%H:%M:%S ") + msg
|
||||
print(line, flush=True)
|
||||
return (log + "\n" + line) if log else line
|
||||
|
||||
log = ""
|
||||
tp_np, tv_np, n_tr, n_an = build_track_points(traces_str, anchors_str)
|
||||
log = logln(log, f"clip idx={clip_idx}: {n_tr} trace(s) + {n_an} anchor(s) → N={tp_np.shape[1]}")
|
||||
if tp_np.shape[1] == 0:
|
||||
return None, "_record at least one trace or place at least one anchor_", log
|
||||
|
||||
ts = time.strftime("%Y%m%d_%H%M%S")
|
||||
req_dir = Path(state["out_dir"]) / f"req_{ts}"
|
||||
req_dir.mkdir(parents=True, exist_ok=True)
|
||||
np.savez(req_dir / "tracks.npz", track_points=tp_np, track_visibility=tv_np)
|
||||
(req_dir / "meta.json").write_text(json.dumps({
|
||||
"clip_idx": int(clip_idx), "seed": int(seed), "steps": int(steps),
|
||||
"w_text": float(w_t), "w_motion": float(w_m),
|
||||
"num_traces": n_tr, "num_anchors": n_an,
|
||||
}, indent=2))
|
||||
log = logln(log, f"req -> {req_dir.name}, denoising ...")
|
||||
|
||||
try:
|
||||
t0 = time.time()
|
||||
mp4 = state["run_generate"](int(clip_idx), tp_np, tv_np,
|
||||
int(seed), int(steps), float(w_t), float(w_m),
|
||||
req_dir)
|
||||
elapsed = time.time() - t0
|
||||
log = logln(log, f"done in {elapsed:.1f}s → {mp4}")
|
||||
return str(mp4), f"_generated in {elapsed:.1f}s_", log
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
return None, f"_generation failed: {e}_", logln(log, f"EXC: {e}")
|
||||
|
||||
gen_btn.click(_generate,
|
||||
[st_clip_idx, traces_json, anchors_json, seed_in, steps_in, w_text, w_motion],
|
||||
[out_video, status, log_box])
|
||||
|
||||
demo.load(None, None, None, js=canvas_js(NUM_FRAMES, FPS))
|
||||
|
||||
return demo
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Generation (MATCHES TrackValidationCallback._sample)
|
||||
# =============================================================================
|
||||
def make_generator(model, tc, samples):
|
||||
"""Return a callable (clip_idx, tp_np, tv_np, seed, steps, w_t, w_m, out_dir) -> mp4 path.
|
||||
|
||||
The denoise formula is copied verbatim from
|
||||
``fastvideo/train/callbacks/track_validation.py::_sample`` (MotionStream Eq. 2).
|
||||
"""
|
||||
import torch
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler, )
|
||||
|
||||
device = model.device
|
||||
dtype = torch.bfloat16
|
||||
flow_shift = float(model.timestep_shift)
|
||||
transformer = model.transformer
|
||||
|
||||
@torch.no_grad()
|
||||
def _run(clip_idx: int, tp_np: np.ndarray, tv_np: np.ndarray,
|
||||
seed: int, steps: int, w_t: float, w_m: float, out_dir: Path) -> str:
|
||||
s = samples[clip_idx]
|
||||
ff = s["first_frame_latent"].to(device, dtype)
|
||||
cond20 = model._build_i2v_cond_concat(ff)
|
||||
txt = s["text_embedding"].to(device, dtype)
|
||||
mask = s["text_attention_mask"].to(device, dtype)
|
||||
img = s["clip_feature"].to(device, dtype)
|
||||
# user tracks (override the sample's tracks)
|
||||
tp = torch.from_numpy(tp_np)[None].to(device, dtype) # [1,T,N,2]
|
||||
tv = torch.from_numpy(tv_np)[None].to(device, dtype) # [1,T,N]
|
||||
|
||||
_, _, T, H, W = ff.shape
|
||||
gen = torch.Generator(device="cpu").manual_seed(int(seed))
|
||||
latents = torch.randn((1, 16, T, H, W), generator=gen, dtype=torch.float32).to(device)
|
||||
sched = FlowMatchEulerDiscreteScheduler(shift=flow_shift)
|
||||
sched.set_timesteps(int(steps), device=device)
|
||||
cfg_on = (w_t != 1.0) or (w_m != 1.0)
|
||||
txt_null = torch.zeros_like(txt) if cfg_on else None
|
||||
|
||||
def _fwd(text_e, tp_e, tv_e, mi, tsv):
|
||||
with torch.autocast(device.type, dtype=dtype), set_forward_context(current_timestep=tsv,
|
||||
attn_metadata=None):
|
||||
return transformer(hidden_states=mi, encoder_hidden_states=text_e,
|
||||
encoder_attention_mask=mask, timestep=tsv,
|
||||
encoder_hidden_states_image=img, track_points=tp_e,
|
||||
track_visibility=tv_e, return_dict=False)
|
||||
|
||||
for tt in sched.timesteps:
|
||||
model_in = torch.cat([latents.to(dtype), cond20], dim=1)
|
||||
tss = tt.reshape(1).to(device, dtype)
|
||||
v_full = _fwd(txt, tp, tv, model_in, tss)
|
||||
if not cfg_on:
|
||||
v = v_full
|
||||
else:
|
||||
v_no_text = _fwd(txt_null, tp, tv, model_in, tss)
|
||||
v_no_motion = _fwd(txt, None, None, model_in, tss)
|
||||
alpha = w_t / (w_t + w_m) if (w_t + w_m) > 0 else 0.5
|
||||
v_base = alpha * v_no_text + (1.0 - alpha) * v_no_motion
|
||||
v = v_base + w_t * (v_full - v_no_text) + w_m * (v_full - v_no_motion)
|
||||
latents = sched.step(v.float(), tt, latents.float(), return_dict=False)[0]
|
||||
|
||||
px = model.decode_latents(latents.permute(0, 2, 1, 3, 4))[0] # [3,T,H,W] in [0,1]
|
||||
video = (px.clamp(0, 1).float().cpu().numpy() * 255.0).astype(np.uint8)
|
||||
frames = np.transpose(video, (1, 2, 3, 0))
|
||||
mp4 = out_dir / "generation.mp4"
|
||||
imageio.mimsave(mp4, frames, fps=FPS, macro_block_size=1)
|
||||
return str(mp4)
|
||||
|
||||
return _run
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Efficient parquet loader (reads only enough files to get N rows, NOT all 4494)
|
||||
# =============================================================================
|
||||
def _load_first_n_from_parquet(data_path: str, n: int, text_len: int) -> list:
|
||||
"""Grab the first N rows from a preprocessed parquet dataset."""
|
||||
import glob
|
||||
import pyarrow.parquet as pq
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_i2v_track
|
||||
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
|
||||
|
||||
files = sorted(glob.glob(os.path.join(data_path, "**", "*.parquet"), recursive=True))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"no *.parquet under {data_path}")
|
||||
rows: list = []
|
||||
for f in files:
|
||||
rows.extend(pq.read_table(f).to_pylist())
|
||||
if len(rows) >= n:
|
||||
break
|
||||
sel = rows[:n]
|
||||
print(f"[app] loaded {len(sel)} rows from parquet", flush=True)
|
||||
batch = collate_rows_from_parquet_schema(sel, pyarrow_schema_i2v_track,
|
||||
text_padding_length=int(text_len), cfg_rate=0.0)
|
||||
infos = batch.get("info_list") or [{} for _ in sel]
|
||||
out = []
|
||||
for i in range(len(sel)):
|
||||
out.append({
|
||||
"text_embedding": batch["text_embedding"][i:i + 1].clone(),
|
||||
"text_attention_mask": batch["text_attention_mask"][i:i + 1].clone(),
|
||||
"vae_latent": batch["vae_latent"][i:i + 1].clone(),
|
||||
"first_frame_latent": batch["first_frame_latent"][i:i + 1].clone(),
|
||||
"clip_feature": batch["clip_feature"][i:i + 1].clone(),
|
||||
"track_points": batch["track_points"][i:i + 1].clone(),
|
||||
"track_visibility": batch["track_visibility"][i:i + 1].clone(),
|
||||
"caption": str(infos[i].get("caption", "") if i < len(infos) else ""),
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# main
|
||||
# =============================================================================
|
||||
def main() -> None:
|
||||
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
ap.add_argument("--model-dir", required=True, help="diffusers-exported ckpt dir")
|
||||
ap.add_argument("--yaml", required=True, help="training yaml")
|
||||
ap.add_argument("--data-path", required=True,
|
||||
help="preprocessed parquet root (combined_parquet_dataset)")
|
||||
ap.add_argument("--num-clips", type=int, default=30, help="how many parquet clips to preload")
|
||||
ap.add_argument("--out-dir", default="/mnt/lustre/vlm-s4duan/multi_trace_out")
|
||||
ap.add_argument("--host", default="0.0.0.0")
|
||||
ap.add_argument("--port", type=int, default=7864)
|
||||
args = ap.parse_args()
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
state: dict = {"out_dir": args.out_dir}
|
||||
|
||||
# Load model
|
||||
print(f"[app] loading model from {args.model_dir} ...", flush=True)
|
||||
sys.path.insert(0, str(REPO / "data_pipeline"))
|
||||
import trackwan_infer as twi
|
||||
model, tc = twi.load_trackwan(args.model_dir, args.yaml)
|
||||
text_len = int(tc.pipeline_config.text_encoder_configs[0].arch_config.text_len)
|
||||
state["model"] = model
|
||||
state["tc"] = tc
|
||||
|
||||
# Load parquet conditioning efficiently
|
||||
print(f"[app] loading {args.num_clips} parquet clips ...", flush=True)
|
||||
samples = _load_first_n_from_parquet(args.data_path, args.num_clips, text_len)
|
||||
labels = []
|
||||
label_to_idx = {}
|
||||
for i, s in enumerate(samples):
|
||||
cap = (s["caption"] or f"clip_{i}").strip().replace("\n", " ")
|
||||
lab = f"{i:03d} {cap[:60]}"
|
||||
labels.append(lab)
|
||||
label_to_idx[lab] = i
|
||||
state["samples"] = samples
|
||||
state["labels"] = labels
|
||||
state["label_to_idx"] = label_to_idx
|
||||
|
||||
# First-frame decoder (from vae_latent) — used as the drawing background
|
||||
def _decode_first_frame(i: int) -> np.ndarray:
|
||||
return twi.decode_reference(model, samples[i]["vae_latent"])[0] # [H,W,3] uint8
|
||||
state["decode_first_frame"] = _decode_first_frame
|
||||
|
||||
# Generator (validation callback's exact denoise loop)
|
||||
state["run_generate"] = make_generator(model, tc, samples)
|
||||
|
||||
print(f"[app] ready. {len(samples)} clips loaded.", flush=True)
|
||||
demo = build_ui(state)
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port,
|
||||
share=False, allowed_paths=[args.out_dir])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,127 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""TrackWan eval-artifact VIEWER — browse pre-generated controllability videos.
|
||||
|
||||
No GPU / no model: just serves the artifacts written by
|
||||
``data_pipeline/gen_eval_artifacts.py``. Pick a checkpoint, a clip, and a
|
||||
counterfactual control; see, side by side:
|
||||
|
||||
- the original clip (and the original with its GT tracks),
|
||||
- the generation (raw),
|
||||
- the generation with the INPUT control tracks (cyan = what we asked for),
|
||||
- the generation with the tracks RE-EXTRACTED from it (yellow = what it did),
|
||||
- the EPE heatmap (per-point error: green = followed, red = ignored),
|
||||
|
||||
plus the CoTracker EPE for that control.
|
||||
|
||||
Launch (no GPU needed)::
|
||||
|
||||
.venv/bin/python examples/inference/gradio/trackwan/app_viewer.py \
|
||||
--artifacts-root /.../eval_artifacts --share
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
|
||||
import gradio as gr
|
||||
|
||||
CONTROL_ORDER = ["gt", "none", "pan_right", "zoom_in", "drag_dense", "drag_sparse", "swap"]
|
||||
STEPS: dict = {} # name -> {"dir": path, "m": manifest}
|
||||
|
||||
|
||||
def _p(ckpt, fname):
|
||||
if not fname:
|
||||
return None
|
||||
path = os.path.join(STEPS[ckpt]["dir"], fname)
|
||||
return path if os.path.exists(path) else None
|
||||
|
||||
|
||||
def update(ckpt, clip, control):
|
||||
if ckpt not in STEPS:
|
||||
return [None] * 7 + ["", ""]
|
||||
clips = STEPS[ckpt]["m"]["clips"]
|
||||
e = clips.get(str(clip))
|
||||
if e is None:
|
||||
return [None] * 7 + ["(clip not in this checkpoint)", ""]
|
||||
c = e["controls"].get(control, {})
|
||||
parts = []
|
||||
if c.get("epe") is not None:
|
||||
parts.append(f"EPE_all={c['epe']:.1f}px")
|
||||
if c.get("epe_moving") is not None:
|
||||
parts.append(f"EPE_moving={c['epe_moving']:.1f}px ({c.get('n_moving', 0)} moving pts)")
|
||||
if c.get("sensitivity_roi") is not None:
|
||||
parts.append(f"sensitivity_roi={c['sensitivity_roi']:.3f}")
|
||||
if c.get("localization_iou") is not None:
|
||||
parts.append(f"localization_IoU={c['localization_iou']:.3f}")
|
||||
if c.get("bg_leakage") is not None:
|
||||
parts.append(f"bg_leakage={c['bg_leakage']:.3f}")
|
||||
txt = " | ".join(parts) if parts else "n/a (GT baseline / no control tracks)"
|
||||
return (_p(ckpt, e.get("original")), _p(ckpt, e.get("original_tracks")),
|
||||
_p(ckpt, c.get("gen")), _p(ckpt, c.get("input")),
|
||||
_p(ckpt, c.get("tracked")), _p(ckpt, c.get("heat")), _p(ckpt, c.get("diff")),
|
||||
e.get("caption", ""), txt)
|
||||
|
||||
|
||||
def build_ui(ckpt_names, clip_ids):
|
||||
with gr.Blocks(title="TrackWan — eval viewer") as demo:
|
||||
gr.Markdown("## TrackWan — controllability eval viewer\n"
|
||||
"Browse pre-generated videos per **checkpoint × clip × control**. "
|
||||
"Cyan = input tracks (asked), yellow = re-tracked from the generation (did), "
|
||||
"EPE heatmap green→red = per-point EPE (followed→ignored). "
|
||||
"**Diff** = |this control − GT-track generation|: bright where the control changed the "
|
||||
"output (the real test — dark everywhere = the model ignored the control).")
|
||||
with gr.Row():
|
||||
ckpt = gr.Dropdown(ckpt_names, value=ckpt_names[0], label="Checkpoint")
|
||||
clip = gr.Dropdown(clip_ids, value=clip_ids[0], label="Clip")
|
||||
control = gr.Dropdown(CONTROL_ORDER, value="drag_dense", label="Control")
|
||||
caption = gr.Textbox(label="Prompt", interactive=False, lines=2)
|
||||
epe_box = gr.Textbox(label="Metrics (EPE_moving + counterfactual sensitivity/IoU = the honest ones)",
|
||||
interactive=False)
|
||||
with gr.Row():
|
||||
v_orig = gr.Video(label="Original clip")
|
||||
v_orig_tr = gr.Video(label="Original + GT tracks")
|
||||
with gr.Row():
|
||||
v_gen = gr.Video(label="Generated (raw)")
|
||||
v_input = gr.Video(label="Generated + INPUT tracks (cyan)")
|
||||
with gr.Row():
|
||||
v_tracked = gr.Video(label="Generated + RE-TRACKED (yellow)")
|
||||
v_heat = gr.Video(label="EPE heatmap (green=followed, red=ignored)")
|
||||
with gr.Row():
|
||||
v_diff = gr.Video(label="DIFF vs GT-track gen (bright = control changed the output)")
|
||||
gr.Markdown("")
|
||||
|
||||
outs = [v_orig, v_orig_tr, v_gen, v_input, v_tracked, v_heat, v_diff, caption, epe_box]
|
||||
for comp in (ckpt, clip, control):
|
||||
comp.change(update, [ckpt, clip, control], outs)
|
||||
demo.load(update, [ckpt, clip, control], outs)
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--artifacts-root", required=True, help="dir containing step*/manifest.json")
|
||||
p.add_argument("--host", default="0.0.0.0")
|
||||
p.add_argument("--port", type=int, default=7870)
|
||||
p.add_argument("--share", action="store_true")
|
||||
args = p.parse_args()
|
||||
|
||||
for d in sorted(glob.glob(os.path.join(args.artifacts_root, "*"))):
|
||||
mf = os.path.join(d, "manifest.json")
|
||||
if os.path.isfile(mf):
|
||||
with open(mf) as f:
|
||||
m = json.load(f)
|
||||
STEPS[m["name"]] = {"dir": d, "m": m}
|
||||
if not STEPS:
|
||||
raise SystemExit(f"no */manifest.json under {args.artifacts_root}")
|
||||
|
||||
names = list(STEPS.keys())
|
||||
clip_ids = sorted(int(k) for k in STEPS[names[0]]["m"]["clips"])
|
||||
demo = build_ui(names, clip_ids)
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=args.share,
|
||||
allowed_paths=[args.artifacts_root])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,52 @@
|
||||
#!/usr/bin/env bash
|
||||
# Acquire the LARGEST held allocation we can, targeting 16 nodes but falling back 16->12->8.
|
||||
# The cluster's topology/block quirk refuses -N<big> even when many nodes are "idle" (freshly
|
||||
# k8s-scaled nodes register outside the static block defs -> unallocatable). Only nodes held by a
|
||||
# RUNNING job are reliably allocatable. So: warm ~20 nodes with 1-node warmups, wait until they're
|
||||
# RUNNING, then scancel them and immediately race a real -N request, trying the biggest size first.
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
JOB=wan14b_bidir_hold
|
||||
TIME=120:00:00
|
||||
say(){ echo "[best $(date +%H:%M:%S)] $*"; }
|
||||
|
||||
# 1) warm ~20 nodes
|
||||
say "submitting 20 warmups to warm the pool"
|
||||
for _ in $(seq 1 20); do
|
||||
sbatch -N1 --gres=gpu:4 --exclusive -p all -t 00:20:00 -J warmup -o /dev/null \
|
||||
--chdir="$WORK" --wrap='srun sleep 900' >/dev/null 2>&1
|
||||
done
|
||||
# 2) wait until >=16 warmups are RUNNING
|
||||
for i in $(seq 1 40); do
|
||||
r=$(squeue -h -u "$USER" -n warmup -t R -o '%i' 2>/dev/null | wc -l)
|
||||
say "warmups running: $r"
|
||||
[ "$r" -ge 16 ] && break
|
||||
sleep 15
|
||||
done
|
||||
# 3) cancel warmups and race the biggest -N that the scheduler accepts
|
||||
say "cancelling warmups and racing a real allocation"
|
||||
scancel -u "$USER" -n warmup 2>/dev/null
|
||||
got=""
|
||||
for i in $(seq 1 90); do
|
||||
idle=$(sinfo -h -p all -N -o '%T' 2>/dev/null | grep -c idle)
|
||||
for N in 16 12 8; do
|
||||
[ "$idle" -ge "$N" ] || continue
|
||||
out=$(sbatch -N"$N" --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 | head -1)
|
||||
if echo "$out" | grep -q "Submitted"; then
|
||||
jid=$(echo "$out" | grep -oE '[0-9]+' | head -1)
|
||||
say "GOT N=$N alloc jid=$jid"; got="$jid:$N"; break 2
|
||||
fi
|
||||
done
|
||||
say "iter $i: idle=$idle, no size accepted yet"
|
||||
sleep 4
|
||||
done
|
||||
[ -n "$got" ] || { say "FAILED to acquire any of 16/12/8"; exit 1; }
|
||||
jid="${got%%:*}"; N="${got##*:}"
|
||||
# 4) wait for it to start
|
||||
for i in $(seq 1 120); do
|
||||
[ "$(squeue -h -j "$jid" -o '%t' 2>/dev/null)" = R ] && { say "RUNNING jid=$jid N=$N nodes=$(squeue -h -j "$jid" -o '%N')"; echo "ALLOC=$jid NODES=$N"; exit 0; }
|
||||
sleep 5
|
||||
done
|
||||
say "jid=$jid submitted but not running yet"; echo "ALLOC=$jid NODES=$N"; exit 0
|
||||
@@ -0,0 +1,80 @@
|
||||
#!/usr/bin/env bash
|
||||
# Acquire a single gang allocation of N nodes on the Slinky (k8s-backed) pool.
|
||||
#
|
||||
# Why this is not just `sbatch -N<N>`: the k8s autoscaler destroys idle nodes, so the pool only
|
||||
# contains as many nodes as there is current demand. Slurm REJECTS (does not queue) a job asking
|
||||
# for more nodes than physically exist in the partition:
|
||||
# "Batch job submission failed: Requested node configuration is not available"
|
||||
# So we have to manufacture demand first, then grab the gang allocation:
|
||||
# 1. keep N+SLACK short 1-node warmups queued -> autoscaler grows the pool
|
||||
# 2. probe with `sbatch --test-only -N<N>` until Slurm says it *could* schedule it
|
||||
# 3. submit the real -N<N> hold (now accepted; pends behind the warmups)
|
||||
# 4. drop the warmups so the hold starts immediately
|
||||
#
|
||||
# Note: other people's jobs (e.g. the demo hold) occupy nodes too, so the pool has to grow to
|
||||
# roughly N + (nodes held by others) before a -N<N> request becomes satisfiable.
|
||||
#
|
||||
# Usage: NODES=12 JOB=wan14b_hold bash examples/train/acquire_nodes.sh
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
NODES="${NODES:-12}"
|
||||
JOB="${JOB:-wan14b_hold}"
|
||||
TIME="${TIME:-120:00:00}"
|
||||
SLACK="${SLACK:-3}" # extra warmups beyond NODES, to out-pace other users' holds
|
||||
WARM_SECS="${WARM_SECS:-900}"
|
||||
MAX_ROUNDS="${MAX_ROUNDS:-60}"
|
||||
mkdir -p "$WORK/logs"
|
||||
say() { echo "[acquire $(date +%H:%M:%S)] $*"; }
|
||||
|
||||
EXIST=$(squeue -h -u "$USER" -n "${JOB}" -o '%i' 2>/dev/null | head -1)
|
||||
if [ -n "$EXIST" ]; then say "reusing existing hold $EXIST"; echo "ALLOC=$EXIST"; exit 0; fi
|
||||
|
||||
WANT_WARM=$(( NODES + SLACK ))
|
||||
|
||||
for round in $(seq 1 "$MAX_ROUNDS"); do
|
||||
pool=$(sinfo -h -p all -N -o '%N' 2>/dev/null | sort -u | wc -l)
|
||||
|
||||
# Probe: can Slurm schedule -N$NODES at all? (--test-only never actually submits)
|
||||
probe=$(sbatch --test-only -N"$NODES" --gres=gpu:4 --ntasks-per-node=1 --exclusive \
|
||||
-t "$TIME" -p all --wrap='hostname' 2>&1 | head -1)
|
||||
|
||||
if echo "$probe" | grep -q "to start at"; then
|
||||
say "pool=$pool — probe OK, submitting real -N$NODES hold"
|
||||
OUT=$(sbatch -N"$NODES" --gres=gpu:4 --ntasks-per-node=1 --exclusive -t "$TIME" \
|
||||
-p all -J "$JOB" --requeue --chdir="$WORK" -o "$WORK/logs/${JOB}_%j.out" \
|
||||
--wrap='srun sleep infinity' 2>&1)
|
||||
if echo "$OUT" | grep -q "Submitted batch"; then
|
||||
ALLOC=$(echo "$OUT" | grep -oE '[0-9]+' | head -1)
|
||||
say "SUBMITTED hold jobid=$ALLOC — clearing warmups so it can start"
|
||||
scancel -u "$USER" -n warmup 2>/dev/null
|
||||
for i in $(seq 1 180); do
|
||||
st=$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)
|
||||
if [ "$st" = R ]; then
|
||||
say "RUNNING on: $(squeue -h -j "$ALLOC" -o '%N')"
|
||||
echo "ALLOC=$ALLOC"; exit 0
|
||||
fi
|
||||
[ -z "$st" ] && { say "hold $ALLOC vanished"; break; }
|
||||
sleep 10
|
||||
done
|
||||
say "hold $ALLOC queued but not started yet (jobid kept)"; echo "ALLOC=$ALLOC"; exit 0
|
||||
fi
|
||||
say "real submit rejected despite OK probe: $(echo "$OUT" | tail -1)"
|
||||
else
|
||||
say "pool=$pool — probe says not schedulable yet"
|
||||
fi
|
||||
|
||||
# Keep warmup pressure up so the autoscaler keeps growing the pool.
|
||||
have=$(squeue -h -u "$USER" -n warmup -o '%i' 2>/dev/null | wc -l)
|
||||
need=$(( WANT_WARM - have ))
|
||||
if [ "$need" -gt 0 ]; then
|
||||
say "submitting $need warmup(s) (have $have, want $WANT_WARM) to drive scale-up"
|
||||
for _ in $(seq 1 "$need"); do
|
||||
sbatch -N1 --gres=gpu:4 --exclusive -p all -t 00:20:00 -J warmup \
|
||||
-o /dev/null --chdir="$WORK" --wrap="srun sleep $WARM_SECS" >/dev/null 2>&1
|
||||
done
|
||||
fi
|
||||
sleep 30
|
||||
done
|
||||
|
||||
say "FAILED to acquire $NODES nodes after $MAX_ROUNDS rounds"
|
||||
exit 1
|
||||
@@ -0,0 +1,44 @@
|
||||
#!/usr/bin/env bash
|
||||
# Wait for step A to finish, then seed and launch step B — so the held allocation never sits
|
||||
# idle between the two overfit stages.
|
||||
#
|
||||
# "A is done" means BOTH: no step-A launcher process is alive, AND a COMPLETE checkpoint-1000
|
||||
# exists (dcp/.metadata present). Requiring both avoids launching B off a half-written
|
||||
# checkpoint if A died mid-save.
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
ALLOC="${ALLOC:-728}"
|
||||
A_JOB="${A_JOB:-wan14b_stepA_v2}"
|
||||
A_CKPT="$WORK/wantrack_14b_synth_sparse_fixed_out/checkpoint-1000"
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY}"
|
||||
say() { echo "[chain $(date +%H:%M:%S)] $*"; }
|
||||
|
||||
say "waiting for step A to reach a complete $A_CKPT ..."
|
||||
while :; do
|
||||
if [ -f "$A_CKPT/dcp/.metadata" ]; then
|
||||
# Checkpoint is complete; wait for the launcher to actually exit before reusing the nodes.
|
||||
# Use the launcher's PID file, NOT `pgrep -f run_wan14b_held.sh` — that pattern also matches
|
||||
# any watcher/shell whose command line contains the string, so the check never goes false
|
||||
# and the chain waits forever on nodes that are already idle.
|
||||
A_RUNFILE="$WORK/logs/${A_JOB}.running"
|
||||
if [ ! -f "$A_RUNFILE" ] || ! kill -0 "$(cat "$A_RUNFILE" 2>/dev/null)" 2>/dev/null; then
|
||||
say "step A finished and $A_CKPT is complete"
|
||||
break
|
||||
fi
|
||||
say "checkpoint-1000 complete, waiting for step A launcher (pid $(cat "$A_RUNFILE")) to exit ..."
|
||||
fi
|
||||
[ -n "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" ] || { say "ALLOC $ALLOC vanished — aborting chain"; exit 1; }
|
||||
sleep 60
|
||||
done
|
||||
|
||||
say "seeding step B output dir from A's checkpoint-1000"
|
||||
bash "$REPO/examples/train/run_stepB_seed.sh" || { say "seed failed"; exit 1; }
|
||||
|
||||
say "launching step B (random-track overfit, WANTRACK_FIXED_SAMPLE=0)"
|
||||
cd "$REPO"
|
||||
exec env ALLOC="$ALLOC" NODES=4 JOB=wan14b_stepB_v2 PORT=30918 \
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_synth_sparse_random_14b.yaml \
|
||||
WANTRACK_FIXED_SAMPLE=0 WANTRACK_FREEZE_HEAD=0 TRACKWAN_TRACK_BIAS=1 \
|
||||
WANDB_API_KEY="$WANDB_API_KEY" \
|
||||
bash examples/train/run_wan14b_held.sh
|
||||
@@ -0,0 +1,98 @@
|
||||
#!/usr/bin/env bash
|
||||
# Launch a bidir teacher training run on a HELD node allocation, so it survives this
|
||||
# cluster's ~80-min node-cordon recycling (a held allocation keeps its nodes because a
|
||||
# drain can't complete while a job holds them — same reason the shao_wm2 sleep-infinity
|
||||
# node lived a day). We hold N nodes with sleep infinity, then run torchrun training
|
||||
# INSIDE the allocation via `srun --overlap`, auto-restarting on any exit (resume from
|
||||
# the latest checkpoint). Verified: one srun --overlap fans out across all held nodes.
|
||||
#
|
||||
# Usage:
|
||||
# WANDB_API_KEY=<key> CFG=<config.yaml> NODES=<n> JOB=<name> \
|
||||
# bash examples/train/run_openvid_bidir_held.sh
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
CFG="${CFG:-examples/train/scenario/worldmodel/finetune_wantrack_openvid_sparse_1p3b.yaml}"
|
||||
NODES="${NODES:-4}"; GPUS=4
|
||||
JOB="${JOB:-openvid_bidir_1p3b}"
|
||||
PORT="${PORT:-29500}"
|
||||
# Stable, per-model W&B run id so crash-relaunches RESUME a single run instead of minting a
|
||||
# fresh run every attempt (that's what fragmented the project into dozens of runs). wandb.init
|
||||
# honors WANDB_RUN_ID + WANDB_RESUME from the env, so no code change is needed. Distinct per
|
||||
# JOB (1p3b vs 14b); override with WANDB_RUN_ID=... to point restarts at an existing run.
|
||||
WANDB_RUN_ID="${WANDB_RUN_ID:-${JOB}}"
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
TOTAL_GPUS=$(( NODES * GPUS ))
|
||||
mkdir -p "$WORK/logs"
|
||||
# output_dir (for checkpoint cleanup) — read from the config
|
||||
OUTPUT_DIR=$(grep -oE 'output_dir:[^#]*' "$REPO/$CFG" | head -1 | sed 's/.*output_dir:[[:space:]]*//; s/"//g' | xargs)
|
||||
|
||||
# Drop a half-written LATEST checkpoint (a crash mid-save leaves checkpoint-N/dcp with no
|
||||
# .metadata) so resume_from_checkpoint=latest falls back to the previous VALID checkpoint (or
|
||||
# scratch) instead of crash-looping on "metadata is None". Keeps only complete checkpoints.
|
||||
clean_bad_ckpt() {
|
||||
[ -n "$OUTPUT_DIR" ] || return 0
|
||||
local latest
|
||||
latest=$(ls -d "$OUTPUT_DIR"/checkpoint-* 2>/dev/null | sed 's/.*checkpoint-//' | sort -n | tail -1)
|
||||
[ -n "$latest" ] || return 0
|
||||
local d="$OUTPUT_DIR/checkpoint-$latest"
|
||||
if [ ! -f "$d/dcp/.metadata" ]; then
|
||||
echo "[held] checkpoint-$latest is incomplete (no dcp/.metadata) — removing so resume uses the last good one"
|
||||
rm -rf "$d"
|
||||
fi
|
||||
}
|
||||
|
||||
# --- 1) hold the nodes (survives cordons; sleep infinity never releases them) --------
|
||||
# Reuse an already-held allocation by passing ALLOC=<jobid> in the env — lets us swap
|
||||
# env vars / config without losing the nodes to the cold pool.
|
||||
if [ -z "${ALLOC:-}" ]; then
|
||||
echo "[held] requesting $NODES nodes ($TOTAL_GPUS GPUs) ..."
|
||||
ALLOC=$(sbatch -N"$NODES" --gres=gpu:$GPUS --ntasks-per-node=1 --exclusive -t 120:00:00 \
|
||||
-p all -J "${JOB}_hold" --chdir="$WORK" -o "$WORK/logs/${JOB}_hold_%j.out" \
|
||||
--wrap='srun sleep infinity' | grep -oE '[0-9]+' | head -1)
|
||||
[ -z "$ALLOC" ] && { echo "[held] sbatch failed (cold pool? submit a warmup + retry)"; exit 1; }
|
||||
else
|
||||
echo "[held] reusing existing allocation JobID=$ALLOC"
|
||||
fi
|
||||
echo "[held] allocation JobID=$ALLOC ; waiting for it to start ..."
|
||||
for i in $(seq 1 60); do
|
||||
[ "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" = R ] && break; sleep 5
|
||||
done
|
||||
[ "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" = R ] || { echo "[held] $ALLOC not running"; exit 1; }
|
||||
NODELIST=$(squeue -h -j "$ALLOC" -o '%N')
|
||||
MASTER=$(scontrol show hostnames "$NODELIST" | head -1)
|
||||
echo "[held] nodes=$NODELIST master=$MASTER (cancel with: scancel $ALLOC)"
|
||||
|
||||
# --- 2) run training inside the held allocation, auto-restart on exit -----------------
|
||||
attempt=0
|
||||
while :; do
|
||||
# bail out if the held allocation itself is gone (should be rare)
|
||||
[ -n "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" ] || { echo "[held] allocation $ALLOC vanished — re-run this script"; exit 1; }
|
||||
attempt=$((attempt + 1))
|
||||
clean_bad_ckpt # remove any half-written checkpoint before (re)starting
|
||||
echo "=== [train] attempt $attempt on alloc $ALLOC (resume from latest valid checkpoint) ==="
|
||||
srun --overlap --jobid="$ALLOC" --nodes="$NODES" --ntasks="$NODES" --ntasks-per-node=1 \
|
||||
--chdir="$REPO" bash -lc "
|
||||
source .venv/bin/activate
|
||||
export HOME=$WORK HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch MPLCONFIGDIR=$WORK/.mpl \
|
||||
TRITON_CACHE_DIR=$WORK/.cache/triton_${JOB} TORCHINDUCTOR_CACHE_DIR=$WORK/.cache/inductor_${JOB} \
|
||||
TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0 PYTHONPATH=$REPO \
|
||||
WANDB_API_KEY=$WANDB_API_KEY WANDB_MODE=online \
|
||||
WANDB_RUN_ID=$WANDB_RUN_ID WANDB_RESUME=allow \
|
||||
WANTRACK_AUG=1 WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 WANTRACK_EXTRA_MODE=random \
|
||||
WANTRACK_PMASK=0 WANTRACK_FIXED_SAMPLE=0 WANTRACK_MOTION_DROP=0 WANTRACK_TEXT_DROP=0 \
|
||||
WANTRACK_DEBUG=1 TRACKWAN_TRACK_BIAS=${TRACKWAN_TRACK_BIAS:-0} \
|
||||
TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=1800 TORCH_NCCL_TRACE_BUFFER_SIZE=2000 \
|
||||
NCCL_SOCKET_NTHREADS=4 NCCL_NSOCKS_PERTHREAD=8
|
||||
torchrun --nnodes=$NODES --nproc-per-node=$GPUS --node-rank=\$SLURM_PROCID \
|
||||
--rdzv-backend=c10d --rdzv-endpoint=$MASTER:$PORT \
|
||||
fastvideo/train/entrypoint/train.py --config $CFG \
|
||||
--training.distributed.num_gpus $TOTAL_GPUS
|
||||
"
|
||||
rc=$?
|
||||
if [ $rc -eq 0 ]; then echo "[train] finished cleanly (rc=0)"; break; fi
|
||||
echo "[train] exited rc=$rc — nodes still held, relaunching from checkpoint in 20s ..."
|
||||
sleep 20
|
||||
done
|
||||
echo "[held] training complete. Freeing allocation $ALLOC."
|
||||
scancel "$ALLOC" 2>/dev/null || true
|
||||
@@ -0,0 +1,44 @@
|
||||
#!/usr/bin/env bash
|
||||
# Launch the OpenVid bidir teacher run (Wan2.1 1.3B Fun-I2V) with the SPARSE point
|
||||
# conditioning recipe + MotionStream stage-1 hparams. Wraps examples/train/run_slurm.sh
|
||||
# with this cluster's settings (partition=all, 4 GPU/node) and the WANTRACK_* env that
|
||||
# selects sparse mode with NO track masking and NO fixed/overfit sampling.
|
||||
#
|
||||
# WANDB_API_KEY must be exported in the environment before calling (never hard-coded).
|
||||
#
|
||||
# Usage: WANDB_API_KEY=<key> bash examples/train/run_openvid_bidir_slurm.sh [num_nodes] [extra --dotted overrides]
|
||||
set -euo pipefail
|
||||
# CFG overridable for the 14B run: CFG=.../finetune_wantrack_openvid_sparse_14b.yaml
|
||||
CFG="${CFG:-examples/train/scenario/worldmodel/finetune_wantrack_openvid_sparse_1p3b.yaml}"
|
||||
NODES="${1:-8}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
# ---- sparse point conditioning (1-per-SAM-object + EXTRA_RANDOM extras), no masking ----
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_PMASK=0 # no stochastic track masking (initial stage)
|
||||
export WANTRACK_FIXED_SAMPLE=0 # no deterministic/overfit sampling
|
||||
export WANTRACK_MOTION_DROP=0
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}" # log sampling/coverage stats (throttled)
|
||||
|
||||
export WANDB_MODE=online
|
||||
|
||||
# ---- compute nodes: /home is NOT mounted; point every ~-based cache at writable Lustre/tmp ----
|
||||
export HOME=/mnt/lustre/vlm-s4duan
|
||||
export HF_HOME=/mnt/lustre/vlm-s4duan/.hf
|
||||
export TORCH_HOME=/mnt/lustre/vlm-s4duan/.torch
|
||||
export MPLCONFIGDIR=/mnt/lustre/vlm-s4duan/.mpl
|
||||
export TORCHINDUCTOR_CACHE_DIR=/tmp/inductor_cache # node-local; triton cache set in run_slurm.sh
|
||||
|
||||
# ---- this cluster ----
|
||||
export PARTITION=all
|
||||
export NUM_GPUS=4 # GPUs per node
|
||||
export MEM=0 # all node memory
|
||||
export CPUS_PER_TASK=128
|
||||
export JOB_NAME="${JOB_NAME:-openvid_bidir_1p3b}"
|
||||
export OUTPUT_DIR=/mnt/lustre/vlm-s4duan/logs/slurm
|
||||
|
||||
bash examples/train/run_slurm.sh "$CFG" "$NODES" "$@"
|
||||
Executable
+53
@@ -0,0 +1,53 @@
|
||||
#!/usr/bin/env bash
|
||||
# Launch a stage-2 finetuning run (MotionStream-style track masking, lr 1e-6,
|
||||
# 800 steps) on a HELD node allocation, like ``run_openvid_bidir_held.sh``.
|
||||
#
|
||||
# Two variants (pass VARIANT=frozen or VARIANT=unfrozen):
|
||||
# VARIANT=frozen → freezes track_encoder (matches MotionStream paper)
|
||||
# VARIANT=unfrozen → keeps track_encoder trainable (ablation)
|
||||
#
|
||||
# WANDB_API_KEY must be exported before calling.
|
||||
#
|
||||
# Usage:
|
||||
# WANDB_API_KEY=<key> VARIANT=frozen bash examples/train/run_openvid_stage2_slurm.sh [num_nodes]
|
||||
set -euo pipefail
|
||||
VARIANT="${VARIANT:-frozen}"
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
case "$VARIANT" in
|
||||
frozen)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage2_frozen.yaml
|
||||
JOB=openvid_stage2_frozen
|
||||
PORT=30100
|
||||
FREEZE=1
|
||||
;;
|
||||
unfrozen)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage2_unfrozen.yaml
|
||||
JOB=openvid_stage2_unfrozen
|
||||
PORT=30200
|
||||
FREEZE=0
|
||||
;;
|
||||
*) echo "unknown VARIANT=$VARIANT (frozen|unfrozen)"; exit 1 ;;
|
||||
esac
|
||||
|
||||
# Stage-2 recipe: stochastic mid-frame track masking (MotionStream Sec. 3.1)
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_PMASK=0.2 # 20% chance to zero contiguous frame chunks
|
||||
export WANTRACK_MASK_CHUNK=8 # 8-frame chunks (~1/3 sec at 24 fps)
|
||||
export WANTRACK_FIXED_SAMPLE=0
|
||||
export WANTRACK_MOTION_DROP=0
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}"
|
||||
export WANTRACK_FREEZE_HEAD="$FREEZE" # 1=freeze track_encoder (MotionStream), 0=trainable
|
||||
export TRACKWAN_TRACK_BIAS=1 # merged-bias init used bias=True convs
|
||||
|
||||
export WANDB_MODE=online
|
||||
export CFG JOB PORT
|
||||
|
||||
# call the same held-alloc launcher used for stage 1
|
||||
CFG="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
Executable
+34
@@ -0,0 +1,34 @@
|
||||
#!/usr/bin/env bash
|
||||
# Stage-2 REPLACEMENT (hi-LR): the empirical eval showed stages 2/3 at lr=1e-6
|
||||
# barely moved weights (~0.02%/600 steps), so behavior barely changed. Try lr=5e-5
|
||||
# to actually shift the model. Aggressive drop rates so it truly sees the sparse
|
||||
# regime the user hits at inference time.
|
||||
#
|
||||
# Init: merged_bias_ckpt4800 (stage-1 end, clean base).
|
||||
# Recipe: TRACK_DROP=0.5 + MOTION_DROP=0.3 + PMASK=0.2 + freeze head.
|
||||
set -euo pipefail
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_TRACK_DROP=0.5
|
||||
export WANTRACK_MOTION_DROP=0.3
|
||||
export WANTRACK_PMASK=0.2
|
||||
export WANTRACK_MASK_CHUNK=8
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_FIXED_SAMPLE=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}"
|
||||
export WANTRACK_FREEZE_HEAD=1
|
||||
export TRACKWAN_TRACK_BIAS=1
|
||||
|
||||
export WANDB_MODE=online
|
||||
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage2v2_hilr.yaml
|
||||
JOB=openvid_stage2v2_hilr
|
||||
PORT=30600
|
||||
|
||||
CFG="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
Executable
+55
@@ -0,0 +1,55 @@
|
||||
#!/usr/bin/env bash
|
||||
# Stage-3 experiments: per-track dropout (WANTRACK_TRACK_DROP=0.5) to fix the
|
||||
# sparse-conditioning distribution mismatch at inference. Three variants:
|
||||
#
|
||||
# VARIANT=A_chunkmask TRACK_DROP=0.5 + PMASK=0.2 (from stage-2 frozen 800)
|
||||
# VARIANT=B_motiondrop TRACK_DROP=0.5 + MOTION_DROP=0.10 (from stage-2 frozen 800)
|
||||
# VARIANT=C_from3700 TRACK_DROP=0.5 + PMASK=0.2 (from merged-bias 3700 — earlier ckpt sanity)
|
||||
#
|
||||
# All three: lr 1e-6, freeze head, 600 steps.
|
||||
set -euo pipefail
|
||||
VARIANT="${VARIANT:-A_chunkmask}"
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
case "$VARIANT" in
|
||||
A_chunkmask)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage3A_chunkmask.yaml
|
||||
JOB=openvid_stage3A_chunkmask
|
||||
PORT=30300
|
||||
PMASK=0.2; MOTION_DROP=0.0
|
||||
;;
|
||||
B_motiondrop)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage3B_motiondrop.yaml
|
||||
JOB=openvid_stage3B_motiondrop
|
||||
PORT=30400
|
||||
PMASK=0.0; MOTION_DROP=0.10
|
||||
;;
|
||||
C_from3700)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_openvid_stage3C_from3700.yaml
|
||||
JOB=openvid_stage3C_from3700
|
||||
PORT=30500
|
||||
PMASK=0.2; MOTION_DROP=0.0
|
||||
;;
|
||||
*) echo "unknown VARIANT=$VARIANT"; exit 1 ;;
|
||||
esac
|
||||
|
||||
export WANTRACK_AUG=1
|
||||
export WANTRACK_SPARSE=1
|
||||
export WANTRACK_EXTRA_RANDOM=20
|
||||
export WANTRACK_EXTRA_MODE=random
|
||||
export WANTRACK_TRACK_DROP=0.5 # per-track dropout (the fix)
|
||||
export WANTRACK_PMASK="$PMASK" # chunked mid-frame mask (0 for B)
|
||||
export WANTRACK_MASK_CHUNK=8
|
||||
export WANTRACK_MOTION_DROP="$MOTION_DROP" # full-track drop (0 for A/C)
|
||||
export WANTRACK_TEXT_DROP=0
|
||||
export WANTRACK_FIXED_SAMPLE=0
|
||||
export WANTRACK_DEBUG="${WANTRACK_DEBUG:-1}"
|
||||
export WANTRACK_FREEZE_HEAD=1 # head is converged
|
||||
export TRACKWAN_TRACK_BIAS=1
|
||||
|
||||
export WANDB_MODE=online
|
||||
export CFG JOB PORT
|
||||
|
||||
CFG="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
@@ -59,6 +59,7 @@ SBATCH_ARGS=(
|
||||
--output="${OUTPUT_DIR}/${JOB_NAME}_%j.out"
|
||||
--error="${OUTPUT_DIR}/${JOB_NAME}_%j.err"
|
||||
--exclusive
|
||||
--requeue
|
||||
)
|
||||
|
||||
if [[ -n "${EXCLUDE}" ]]; then
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
#!/usr/bin/env bash
|
||||
# One-shot smoke test of the 720p stage-1 geometry inside held alloc 881.
|
||||
# Unlike run_wan14b_held.sh this does NOT auto-restart on failure — a crash (e.g. OOM) is the
|
||||
# RESULT we are looking for, not something to retry through.
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
ALLOC="${ALLOC:-881}"
|
||||
NODES="${NODES:-16}"; GPUS=4
|
||||
PORT="${PORT:-31777}"
|
||||
CFG="${CFG:-examples/train/scenario/worldmodel/smoke_720p_14b_stage1.yaml}"
|
||||
JOB="${JOB:-smoke720}"
|
||||
TOTAL_GPUS=$(( NODES * GPUS ))
|
||||
LOG=$WORK/logs/${JOB}.log
|
||||
|
||||
[ "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" = R ] || { echo "alloc $ALLOC not running"; exit 1; }
|
||||
NODELIST=$(squeue -h -j "$ALLOC" -o '%N')
|
||||
# Pin the subset so the rdzv master provably hosts rank 0 (see run_wan14b_held.sh).
|
||||
SUBSET=$(scontrol show hostnames "$NODELIST" | head -n "$NODES" | paste -sd,)
|
||||
MASTER=$(echo "$SUBSET" | cut -d, -f1)
|
||||
echo "[smoke] alloc=$ALLOC nodes=$NODES gpus=$TOTAL_GPUS master=$MASTER"
|
||||
echo "[smoke] cfg=$CFG log=$LOG"
|
||||
|
||||
# Stage-1 env. WANTRACK_FREEZE_HEAD=0 — the track_encoder is TRAINABLE in stage-1.
|
||||
# The 14B stage-1 config header claimed FREEZE_HEAD=1, but that contradicts the 1.3B run it
|
||||
# says it copies exactly: run_openvid_bidir_held.sh never sets the var (code default "0") and
|
||||
# the 1.3B logs have zero "FROZE track_encoder" lines. Every real freeze in this repo is
|
||||
# stage-2/3 (run_openvid_stage2*/stage3_slurm.sh, the latter commented "head is converged" —
|
||||
# converged BY stage-1 training it). wantrack.py itself calls it a "stage-2 knob ... after
|
||||
# initial training", and notes the head only plateaus by step ~4700 of a 4800-step stage-1.
|
||||
srun --overlap --jobid="$ALLOC" --nodelist="$SUBSET" --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_MODE=disabled \
|
||||
WANTRACK_AUG=1 WANTRACK_SPARSE=1 WANTRACK_EXTRA_RANDOM=20 WANTRACK_EXTRA_MODE=random \
|
||||
WANTRACK_PMASK=0 WANTRACK_MASK_CHUNK=0 WANTRACK_TRACK_DROP=0 WANTRACK_MOTION_DROP=0 \
|
||||
WANTRACK_TEXT_DROP=0 WANTRACK_FIXED_SAMPLE=0 WANTRACK_FREEZE_HEAD=0 TRACKWAN_TRACK_BIAS=1 \
|
||||
WANTRACK_DEBUG=1 \
|
||||
TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=1800 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
|
||||
" 2>&1 | tee "$LOG"
|
||||
rc=${PIPESTATUS[0]}
|
||||
echo "[smoke] srun exit rc=$rc"
|
||||
|
||||
# Always sweep stray ranks: a dead rank keeps ~100GB of GPU memory and would poison the real run.
|
||||
echo "[smoke] sweeping stray ranks ..."
|
||||
timeout 120 srun --overlap --jobid="$ALLOC" --nodelist="$SUBSET" --nodes="$NODES" \
|
||||
--ntasks="$NODES" --ntasks-per-node=1 \
|
||||
bash -c "pkill -9 -f 'entrypoint/train.py --config $CFG'; exit 0" >/dev/null 2>&1 || true
|
||||
echo "[smoke] done (rc=$rc)"
|
||||
exit $rc
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/usr/bin/env bash
|
||||
# Seed step B's output dir with step A's final checkpoint, so B starts from A while its config
|
||||
# can still use resume_from_checkpoint: latest (which is what makes crash-restarts safe).
|
||||
#
|
||||
# Hardlinked (cp -al), not copied: a 14B training-state checkpoint is hundreds of GB (bf16
|
||||
# weights + fp32 master + Adam m/v), and both dirs live on the same Lustre filesystem.
|
||||
# Checkpoints are write-once, so sharing inodes is safe here.
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
STEP="${STEP:-1000}"
|
||||
SRC="${SRC:-$WORK/wantrack_14b_synth_sparse_fixed_out/checkpoint-$STEP}"
|
||||
DST_DIR="${DST_DIR:-$WORK/wantrack_14b_synth_sparse_random_out}"
|
||||
DST="$DST_DIR/checkpoint-$STEP"
|
||||
|
||||
[ -d "$SRC" ] || { echo "[B-seed] missing $SRC"; exit 1; }
|
||||
[ -f "$SRC/dcp/.metadata" ] || { echo "[B-seed] $SRC incomplete (no dcp/.metadata) — refusing"; exit 1; }
|
||||
|
||||
if [ -d "$DST" ]; then
|
||||
echo "[B-seed] $DST already exists — leaving it alone"
|
||||
else
|
||||
mkdir -p "$DST_DIR"
|
||||
echo "[B-seed] hardlinking $SRC -> $DST"
|
||||
cp -al "$SRC" "$DST"
|
||||
fi
|
||||
echo "[B-seed] checkpoints now in $DST_DIR:"
|
||||
ls -d "$DST_DIR"/checkpoint-* 2>/dev/null | sed 's/^/ /'
|
||||
@@ -0,0 +1,74 @@
|
||||
#!/usr/bin/env bash
|
||||
# STEP C — build the 14B "merged" init from the step-B overfit.
|
||||
#
|
||||
# Takes the 14B-native overfit (steps A+B) and transplants BOTH halves of the track pathway
|
||||
# into a FRESH 14B base:
|
||||
# * track_encoder.{proj,temporal_conv}.{weight,bias} (--track-src)
|
||||
# * patch_embedding.weight[:, 36:] — the track slot (--pe-src)
|
||||
# while patch_embedding.weight[:, :36] stays the pretrained I2V weights from --base.
|
||||
#
|
||||
# Both halves come from the SAME source on purpose: they were co-adapted during the overfit.
|
||||
# Lifting the encoder alone (the "partial_merged" experiment) measured WORSE than random init,
|
||||
# because the trained encoder ends up feeding a random projection.
|
||||
#
|
||||
# Usage: [STEP=2000] bash examples/train/run_stepC_merge.sh
|
||||
set -euo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
STEP="${STEP:-2000}"
|
||||
SRC_CKPT="${SRC_CKPT:-$WORK/wantrack_14b_synth_sparse_random_out/checkpoint-$STEP}"
|
||||
EXPORT_DIR="${EXPORT_DIR:-$WORK/exports/overfit_14b_random_ckpt$STEP}"
|
||||
OUT="${OUT:-$WORK/models/trackwan_14b_i2v_d64_merged_from_overfit_bias}"
|
||||
BASE="${BASE:-$WORK/models/Wan2.1-I2V-14B-720P-Diffusers}"
|
||||
|
||||
cd "$REPO"
|
||||
source .venv/bin/activate
|
||||
export HOME=$WORK HF_HOME=$WORK/.hf TORCH_HOME=$WORK/.torch MPLCONFIGDIR=$WORK/.mpl \
|
||||
PYTHONPATH=$REPO TOKENIZERS_PARALLELISM=false NCCL_CUMEM_ENABLE=0 \
|
||||
TRACKWAN_TRACK_BIAS=1
|
||||
|
||||
[ -d "$SRC_CKPT" ] || { echo "[C] missing $SRC_CKPT"; exit 1; }
|
||||
[ -f "$SRC_CKPT/dcp/.metadata" ] || { echo "[C] $SRC_CKPT is an INCOMPLETE checkpoint (no dcp/.metadata)"; exit 1; }
|
||||
|
||||
# 1) DCP -> diffusers (single process; DCP reshards automatically)
|
||||
if [ ! -d "$EXPORT_DIR/transformer" ]; then
|
||||
echo "[C] exporting $SRC_CKPT -> $EXPORT_DIR"
|
||||
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
|
||||
--checkpoint "$SRC_CKPT" --output-dir "$EXPORT_DIR"
|
||||
else
|
||||
echo "[C] reusing existing export $EXPORT_DIR"
|
||||
fi
|
||||
|
||||
# 2) fresh 14B base + co-adapted encoder AND track slot from that export
|
||||
echo "[C] building merged init -> $OUT"
|
||||
python data_pipeline/convert_trackwan_init_v2.py \
|
||||
--base "$BASE" --out "$OUT" --id-dim 64 --pe-init random \
|
||||
--track-src "$EXPORT_DIR" --pe-src "$EXPORT_DIR"
|
||||
|
||||
# 3) verify: pretrained channels untouched, track slot + encoder match the overfit
|
||||
echo "[C] verifying merged init against source ..."
|
||||
python - "$OUT" "$EXPORT_DIR" "$BASE" <<'PY'
|
||||
import sys, glob
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
def load(d):
|
||||
st = {}
|
||||
for f in sorted(glob.glob(f"{d}/transformer/*.safetensors")):
|
||||
st.update(load_file(f))
|
||||
return st
|
||||
|
||||
out, src, base = (load(p) for p in sys.argv[1:4])
|
||||
pe_o, pe_s, pe_b = out["patch_embedding.weight"], src["patch_embedding.weight"], base["patch_embedding.weight"]
|
||||
print(f" pe shape {tuple(pe_o.shape)} (base {tuple(pe_b.shape)})")
|
||||
print(f" pe[:, :36] == base : {torch.equal(pe_o[:, :36], pe_b)} std={pe_o[:, :36].float().std():.5f}")
|
||||
print(f" pe[:, 36:] == overfit slot : {torch.equal(pe_o[:, 36:], pe_s[:, 36:])} std={pe_o[:, 36:].float().std():.5f}")
|
||||
ok = torch.equal(pe_o[:, :36], pe_b) and torch.equal(pe_o[:, 36:], pe_s[:, 36:])
|
||||
for k in sorted(k for k in out if "track_encoder" in k):
|
||||
same = k in src and torch.equal(out[k], src[k])
|
||||
ok &= same
|
||||
print(f" {k:42s} lifted={same} norm={out[k].float().norm():.5f}")
|
||||
print(" RESULT:", "OK — encoder and track slot are co-adapted" if ok else "MISMATCH")
|
||||
sys.exit(0 if ok else 1)
|
||||
PY
|
||||
echo "[C] done -> $OUT"
|
||||
Executable
+43
@@ -0,0 +1,43 @@
|
||||
#!/usr/bin/env bash
|
||||
# Synth stage-2: MotionStream-style fine-tuning on 49k Wan2.2 synth data.
|
||||
# Init: merged_bias_ckpt4800 (our stage-1 end on openvid).
|
||||
# Recipe: TRACK_DROP=0.5 + MOTION_DROP=0.3 + PMASK=0.2 + freeze head.
|
||||
# 600 steps -> ~1.55 epochs on combined 49k (paper claims "~1 epoch" for stage-2).
|
||||
# Two VARIANTs: paperLR (1e-6, paper spec) or 5x (5e-6, 5x paper, 2x below stage-1's 1e-5).
|
||||
set -euo pipefail
|
||||
VARIANT="${VARIANT:-paperLR}"
|
||||
NODES="${1:-4}"; shift || true
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY before launching}"
|
||||
|
||||
case "$VARIANT" in
|
||||
paperLR)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_synth_stage2_paperLR.yaml
|
||||
JOB=synth_stage2_paperLR
|
||||
PORT=30700
|
||||
;;
|
||||
5x)
|
||||
CFG=examples/train/scenario/worldmodel/finetune_wantrack_synth_stage2_5x.yaml
|
||||
JOB=synth_stage2_5x
|
||||
PORT=30800
|
||||
;;
|
||||
*) 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
|
||||
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="$CFG" JOB="$JOB" PORT="$PORT" \
|
||||
bash "$(dirname "$0")/run_openvid_bidir_held.sh" "$NODES" "$@"
|
||||
@@ -0,0 +1,142 @@
|
||||
#!/usr/bin/env bash
|
||||
# Run a Wan2.1-14B WanTrack stage inside an ALREADY-HELD Slurm allocation.
|
||||
#
|
||||
# Same held-alloc + auto-restart pattern as run_openvid_bidir_held.sh, but every WANTRACK_*
|
||||
# knob is parameterized instead of hardcoded, because the four stages of the 14B teacher
|
||||
# pipeline need DIFFERENT sampling/masking settings:
|
||||
#
|
||||
# A (fixed overfit) : WANTRACK_FIXED_SAMPLE=1 FREEZE_HEAD=0
|
||||
# B (random overfit) : WANTRACK_FIXED_SAMPLE=0 FREEZE_HEAD=0
|
||||
# D (stage-1 openvid): FREEZE_HEAD=0, no dropping <-- CORRECTED 2026-07-31 (was 1)
|
||||
# E (stage-2 synth) : FREEZE_HEAD=1 TRACK_DROP=0.5 MOTION_DROP=0.3 PMASK=0.2 MASK_CHUNK=8
|
||||
#
|
||||
# FREEZE_HEAD is the important one: stages A/B/D MUST train track_encoder (A/B co-adapt it with
|
||||
# the patch_embedding track slot; D keeps refining it), and only E (and stage-3) freeze it.
|
||||
# D said FREEZE_HEAD=1 here until 2026-07-31, contradicting the 1.3B run D is the analog of:
|
||||
# run_openvid_bidir_held.sh never sets the var (wantrack.py defaults it to "0") and the 1.3B
|
||||
# logs contain zero "FROZE track_encoder" lines. run_openvid_stage3_slurm.sh freezes with the
|
||||
# comment "head is converged" — it converges BECAUSE stage-1/D trains it, and wantrack.py notes
|
||||
# the head only plateaus by step ~4700 of D's 4800 steps.
|
||||
#
|
||||
# Usage:
|
||||
# ALLOC=728 NODES=4 JOB=wan14b_stepA PORT=30910 \
|
||||
# CFG=examples/train/scenario/worldmodel/finetune_wantrack_synth_sparse_fixed_14b.yaml \
|
||||
# WANTRACK_FIXED_SAMPLE=1 WANTRACK_FREEZE_HEAD=0 \
|
||||
# WANDB_API_KEY=... bash examples/train/run_wan14b_held.sh
|
||||
set -uo pipefail
|
||||
WORK=/mnt/lustre/vlm-s4duan
|
||||
REPO=$WORK/FastVideo
|
||||
CFG="${CFG:?set CFG}"
|
||||
NODES="${NODES:-4}"; GPUS=4
|
||||
JOB="${JOB:-wan14b}"
|
||||
PORT="${PORT:-30900}"
|
||||
ALLOC="${ALLOC:?set ALLOC to a running held allocation jobid}"
|
||||
WANDB_RUN_ID="${WANDB_RUN_ID:-${JOB}}"
|
||||
: "${WANDB_API_KEY:?export WANDB_API_KEY}"
|
||||
TOTAL_GPUS=$(( NODES * GPUS ))
|
||||
mkdir -p "$WORK/logs"
|
||||
|
||||
OUTPUT_DIR=$(grep -oE 'output_dir:[^#]*' "$REPO/$CFG" | head -1 | sed 's/.*output_dir:[[:space:]]*//; s/"//g' | xargs)
|
||||
|
||||
# Per-stage WANTRACK defaults (overridable from the environment).
|
||||
W_AUG="${WANTRACK_AUG:-1}"
|
||||
W_SPARSE="${WANTRACK_SPARSE:-1}"
|
||||
W_EXTRA_RANDOM="${WANTRACK_EXTRA_RANDOM:-20}"
|
||||
W_EXTRA_MODE="${WANTRACK_EXTRA_MODE:-random}"
|
||||
W_PMASK="${WANTRACK_PMASK:-0}"
|
||||
W_MASK_CHUNK="${WANTRACK_MASK_CHUNK:-0}"
|
||||
W_TRACK_DROP="${WANTRACK_TRACK_DROP:-0}"
|
||||
W_MOTION_DROP="${WANTRACK_MOTION_DROP:-0}"
|
||||
W_TEXT_DROP="${WANTRACK_TEXT_DROP:-0}"
|
||||
W_FIXED="${WANTRACK_FIXED_SAMPLE:-0}"
|
||||
W_FREEZE="${WANTRACK_FREEZE_HEAD:-0}"
|
||||
W_BIAS="${TRACKWAN_TRACK_BIAS:-1}"
|
||||
|
||||
# A crash mid-save leaves checkpoint-N/dcp without .metadata; drop it so resume falls back to
|
||||
# the last COMPLETE checkpoint instead of crash-looping on "metadata is None".
|
||||
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 incomplete (no dcp/.metadata) — removing"
|
||||
rm -rf "$d"
|
||||
fi
|
||||
}
|
||||
|
||||
[ "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" = R ] || { echo "[held] alloc $ALLOC not running"; exit 1; }
|
||||
NODELIST=$(squeue -h -j "$ALLOC" -o '%N')
|
||||
# Pin the exact subset of the held allocation this stage runs on. Without --nodelist, srun is
|
||||
# free to pick ANY $NODES of the allocation's nodes, so the rdzv endpoint (computed here) can
|
||||
# land on a node that isn't in the subset — rank 0 then never listens there and every worker
|
||||
# dies with "client socket has timed out while trying to connect". Choosing the subset up front
|
||||
# makes MASTER provably the node hosting rank 0, and lets stages share one allocation safely.
|
||||
# By default take the first $NODES nodes of the allocation. Pass NODELIST_OVERRIDE=a,b,c to pin
|
||||
# an explicit subset — required when several stages share one allocation, otherwise every run
|
||||
# grabs the same leading nodes and they fight over the same GPUs.
|
||||
SUBSET="${NODELIST_OVERRIDE:-$(scontrol show hostnames "$NODELIST" | head -n "$NODES" | paste -sd,)}"
|
||||
MASTER=$(echo "$SUBSET" | cut -d, -f1)
|
||||
echo "[held] alloc=$ALLOC nodes=$NODELIST"
|
||||
echo "[held] subset=$SUBSET master=$MASTER using $NODES node(s) / $TOTAL_GPUS GPU(s)"
|
||||
echo "[held] cfg=$CFG out=$OUTPUT_DIR"
|
||||
echo "[held] WANTRACK: fixed=$W_FIXED freeze=$W_FREEZE pmask=$W_PMASK chunk=$W_MASK_CHUNK track_drop=$W_TRACK_DROP motion_drop=$W_MOTION_DROP bias=$W_BIAS"
|
||||
|
||||
# Killing this script (or its srun) does NOT kill the torchrun/python ranks out on the compute
|
||||
# nodes — they survive, keep ~107GB of GPU memory each, and silently keep writing checkpoints.
|
||||
# That poisons the next run: it contends for GPUs and resurrects deleted output dirs. So sweep
|
||||
# the subset on any exit. Matches on the config path so we only kill THIS stage's ranks.
|
||||
# Liveness marker for anything chaining off this stage. Do NOT make chainers use
|
||||
# `pgrep -f run_wan14b_held.sh`: that also matches any monitoring/shell command whose command
|
||||
# line merely CONTAINS the string, so a watcher looking for this launcher keeps "seeing" it long
|
||||
# after it exited (that stall cost ~1h of idle nodes between stages A and B).
|
||||
RUNFILE="$WORK/logs/${JOB}.running"
|
||||
echo "$$" > "$RUNFILE"
|
||||
|
||||
cleanup_ranks() {
|
||||
echo "[held] sweeping stray ranks for $CFG on $SUBSET ..."
|
||||
timeout 120 srun --overlap --jobid="$ALLOC" --nodelist="$SUBSET" --nodes="$NODES" \
|
||||
--ntasks="$NODES" --ntasks-per-node=1 \
|
||||
bash -c "pkill -9 -f 'entrypoint/train.py --config $CFG'; exit 0" >/dev/null 2>&1 || true
|
||||
}
|
||||
trap 'echo "[held] interrupted — cleaning up"; cleanup_ranks; rm -f "$RUNFILE"; exit 130' INT TERM
|
||||
trap 'rm -f "$RUNFILE"' EXIT
|
||||
|
||||
attempt=0
|
||||
while :; do
|
||||
[ -n "$(squeue -h -j "$ALLOC" -o '%t' 2>/dev/null)" ] || { echo "[held] alloc $ALLOC vanished"; exit 1; }
|
||||
attempt=$((attempt + 1))
|
||||
clean_bad_ckpt
|
||||
echo "=== [train] attempt $attempt on alloc $ALLOC ($(date)) ==="
|
||||
srun --overlap --jobid="$ALLOC" --nodelist="$SUBSET" --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=$W_AUG WANTRACK_SPARSE=$W_SPARSE \
|
||||
WANTRACK_EXTRA_RANDOM=$W_EXTRA_RANDOM WANTRACK_EXTRA_MODE=$W_EXTRA_MODE \
|
||||
WANTRACK_PMASK=$W_PMASK WANTRACK_MASK_CHUNK=$W_MASK_CHUNK \
|
||||
WANTRACK_TRACK_DROP=$W_TRACK_DROP WANTRACK_MOTION_DROP=$W_MOTION_DROP \
|
||||
WANTRACK_TEXT_DROP=$W_TEXT_DROP WANTRACK_FIXED_SAMPLE=$W_FIXED \
|
||||
WANTRACK_FREEZE_HEAD=$W_FREEZE TRACKWAN_TRACK_BIAS=$W_BIAS \
|
||||
WANTRACK_DEBUG=1 \
|
||||
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 in 20s ..."
|
||||
# A crashed rank can leave its siblings alive and holding GPU memory; clear them before retry
|
||||
# or the relaunch will OOM or hang waiting on a rendezvous the stale ranks are still in.
|
||||
cleanup_ranks
|
||||
sleep 20
|
||||
done
|
||||
echo "[held] stage complete; allocation $ALLOC left RUNNING for the next stage."
|
||||
@@ -0,0 +1,89 @@
|
||||
# WanTrack (Wan2.1 Fun-Control + MotionStream point-track) DROID overfit.
|
||||
#
|
||||
# Goal: prove the data design forces TRACE reliance. The 200 DROID clips share a
|
||||
# similar first frame (tight frame-0 cluster) and ONE generic caption, so neither
|
||||
# the first frame nor the text can determine the (diverse) motion -- the model must
|
||||
# read the point tracks to reduce loss. This is the step-by-step "overfit on the
|
||||
# traces first" experiment, so MotionStream augments are left OFF (launch env
|
||||
# WANTRACK_AUG=0) for a clean loss-down signal.
|
||||
#
|
||||
# Init from the Fun-Control variant (trackwan_1.3b_i2v_control_init): its track-slot
|
||||
# patch-embed channels are Fun-Control's pretrained control channels (non-zero), so
|
||||
# gradient reaches the track encoder from step 0 -- no double-zero-init deadlock.
|
||||
#
|
||||
# Self-contained: a fresh output_dir + run_name, so it does not touch any prior run.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/droid_track_200/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: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_droid_overfit_out
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: droid-overfit-clean
|
||||
|
||||
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: 500
|
||||
num_val_samples: 2
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,85 @@
|
||||
# WanTrack EPIC-KITCHENS-200 finetune: egocentric human manipulation, 3 kitchens
|
||||
# (P02/P04/P35), 200 clips x 121 frames @ 854x480 (trained at 480x832). Fun-Control init.
|
||||
#
|
||||
# Unlike the Wan-200 overfit run (WANTRACK_AUG=0, all 2500 tracks), this run turns the
|
||||
# MotionStream sampler ON via launch env so we exercise the new coverage+informativeness
|
||||
# sampler on a moving-camera dataset:
|
||||
# WANTRACK_AUG=1 WANTRACK_MIN_POINTS=1000 WANTRACK_MAX_POINTS=2500 \
|
||||
# WANTRACK_PMASK=0.2 WANTRACK_DIVERSITY=1.0 WANTRACK_UNIFORM_FRAC=0.3
|
||||
# object_ids (SAM conf0.6) + track_weights (low-rank informativeness) come from the
|
||||
# preprocessed parquet, so sampling = >=1 pt/object + lowrank-weighted + uniform.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/epic_kitchens_200/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: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_epic200_out
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: epic200-sampler
|
||||
|
||||
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: 500
|
||||
num_val_samples: 2
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,85 @@
|
||||
# WanTrack EPIC-KITCHENS-200 finetune: egocentric human manipulation, 3 kitchens
|
||||
# (P02/P04/P35), 200 clips x 121 frames @ 854x480 (trained at 480x832). Fun-Control init.
|
||||
#
|
||||
# Unlike the Wan-200 overfit run (WANTRACK_AUG=0, all 2500 tracks), this run turns the
|
||||
# MotionStream sampler ON via launch env so we exercise the new coverage+informativeness
|
||||
# sampler on a moving-camera dataset:
|
||||
# WANTRACK_AUG=1 WANTRACK_MIN_POINTS=1000 WANTRACK_MAX_POINTS=2500 \
|
||||
# WANTRACK_PMASK=0.2 WANTRACK_DIVERSITY=1.0 WANTRACK_UNIFORM_FRAC=0.3
|
||||
# object_ids (SAM conf0.6) + track_weights (low-rank informativeness) come from the
|
||||
# preprocessed parquet, so sampling = >=1 pt/object + lowrank-weighted + uniform.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/epic_static_200/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: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_epic_static200_out
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: epic-static200-sampler
|
||||
|
||||
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: 500
|
||||
num_val_samples: 2
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 25
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,78 @@
|
||||
# WanTrack Wan-200 finetune: SAME params as the 50-video (control50) run, scaled to
|
||||
# 200 Wan2.2-A14B clips. Fun-Control init, WANTRACK_AUG=0 (launch env), unique prompts
|
||||
# (standard funinp preprocess). Isolates the data-scale variable (50 -> 200).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan_golf_overfit/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: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_golf_nonzero_out
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: golf-nonzero-runtime
|
||||
|
||||
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: 200
|
||||
num_val_samples: 1
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,78 @@
|
||||
# WanTrack Wan-200 finetune: SAME params as the 50-video (control50) run, scaled to
|
||||
# 200 Wan2.2-A14B clips. Fun-Control init, WANTRACK_AUG=0 (launch env), unique prompts
|
||||
# (standard funinp preprocess). Isolates the data-scale variable (50 -> 200).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_control_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan_golf_overfit/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: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_golf_overfit_out
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: golf-overfit-debug
|
||||
|
||||
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: 200
|
||||
num_val_samples: 1
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,78 @@
|
||||
# WanTrack Wan-200 finetune: SAME params as the 50-video (control50) run, scaled to
|
||||
# 200 Wan2.2-A14B clips. Fun-Control init, WANTRACK_AUG=0 (launch env), unique prompts
|
||||
# (standard funinp preprocess). Isolates the data-scale variable (50 -> 200).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_nonzero_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_nonzero_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan_golf_overfit/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: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_golf_overfit_nonzero_out
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: golf-overfit-nonzero
|
||||
|
||||
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: 200
|
||||
num_val_samples: 1
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,80 @@
|
||||
# WanTrack (Wan2.2 + MotionStream point-track) bidirectional I2V finetune.
|
||||
# Overfit config: 1.3B base, 480x832, all 121 frames (tracks align 1:1).
|
||||
# Goal: drive flow-matching loss down on the 10-clip dataset to prove the
|
||||
# track-conditioned path learns. MotionStream augments are env-gated and left
|
||||
# OFF here (WANTRACK_AUG=0 in the launch env) for a clean loss-down signal.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/weka/home/hao.zhang/shao/data/models/trackwan_1.3b_i2v_init
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wan22_a14b_720p_24fps/preprocessed_i2v_track_funinp/combined_parquet_dataset
|
||||
dataloader_num_workers: 2
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-4
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/weka/home/hao.zhang/shao/data/motion_pipeline/wantrack_overfit_i2v_out
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 2
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: overfit-funinp-i2v
|
||||
|
||||
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: 500
|
||||
num_val_samples: 2
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 3
|
||||
@@ -0,0 +1,100 @@
|
||||
# EXPERIMENT 2 (layer-wise LR) — can stage-1 bootstrap a FRESH track pathway on OpenVid,
|
||||
# with no overfit->merge stages at all?
|
||||
#
|
||||
# Same goal as experiment 1, different mechanism: instead of staging, train everything from
|
||||
# step 0 but give the track pathway a 100x LR (1e-3 vs 1e-5) so it can catch up with an already
|
||||
# pretrained DiT. Adam is scale-invariant per parameter, so this MUST be a real param group with
|
||||
# its own lr — scaling the gradient would do nothing. patch_embedding's pretrained channels
|
||||
# [:, :36] are grad-masked so the 100x LR only ever touches the track slot.
|
||||
#
|
||||
# NOTE: an earlier in-repo LWLR attempt reported 'no encoder bootstrap', but it referenced
|
||||
# 'patch_embedding.proj.weight' incorrectly and so likely never applied the boosted LR to
|
||||
# anything. This re-tests the idea properly; the param-group builder now RAISES if nothing
|
||||
# matches, so a silent no-op cannot masquerade as a negative result.
|
||||
#
|
||||
# Recipe is otherwise IDENTICAL to the 1.3B stage-1 reference: lr 1e-5, global bs 128, 480x832,
|
||||
# 121f, 24fps, flow_shift 6, bf16 — only the param-group/warm-up behaviour is new.
|
||||
#
|
||||
# Init is the ZERO-gate init: patch_embedding[:, 36:] = 0, so the model starts bit-identical to
|
||||
# base I2V. (The random-gate init was measured to corrupt the prior from step 0.)
|
||||
#
|
||||
# Launch env: WANTRACK_TRACK_GROUP=1 (build named track/base param groups),
|
||||
# WANTRACK_TRACK_LR_MULT=100 (track pathway at 1e-3), WANTRACK_FREEZE_HEAD=0
|
||||
# (the encoder MUST train here — this is the opposite of the merged-init stage 1).
|
||||
# 4 nodes x 4 GPU = 16 GPUs, grad_accum 8 -> 1 * 8 * 16 = 128 global.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_zero_init_bias
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_zero_init_bias
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 16
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1000 # experiment: the bootstrap signal shows well before this
|
||||
gradient_accumulation_steps: 8 # 1 * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/exp_openvid_14b_lwlr_out
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 100 # keep all — early checkpoints are the whole point
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: exp-openvid-14b-lwlr-100x
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,100 @@
|
||||
# EXPERIMENT 1 (encoder warm-up) — can stage-1 bootstrap a FRESH track pathway on OpenVid,
|
||||
# with no overfit->merge stages at all?
|
||||
#
|
||||
# The overfit+merge pipeline exists only because a from-scratch track encoder would not train
|
||||
# during stage 1. This tests a more direct route: for the first N steps hold the pretrained DiT
|
||||
# completely still (base param group at lr 0) and let ONLY the track pathway learn
|
||||
# (track_encoder.* + patch_embedding.proj.weight[:, 36:]); then ramp the DiT back in. If the
|
||||
# gate and encoder grow to a healthy scale and the loss behaves, the whole A/B/C detour can go.
|
||||
#
|
||||
# Recipe is otherwise IDENTICAL to the 1.3B stage-1 reference: lr 1e-5, global bs 128, 480x832,
|
||||
# 121f, 24fps, flow_shift 6, bf16 — only the param-group/warm-up behaviour is new.
|
||||
#
|
||||
# Init is the ZERO-gate init: patch_embedding[:, 36:] = 0, so the model starts bit-identical to
|
||||
# base I2V. (The random-gate init was measured to corrupt the prior from step 0.)
|
||||
#
|
||||
# Launch env: WANTRACK_TRACK_GROUP=1 (build named track/base param groups),
|
||||
# WANTRACK_TRACK_LR_MULT=1 (keep the reference LR), WANTRACK_FREEZE_HEAD=0
|
||||
# (the encoder MUST train here — this is the opposite of the merged-init stage 1).
|
||||
# 4 nodes x 4 GPU = 16 GPUs, grad_accum 8 -> 1 * 8 * 16 = 128 global.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_zero_init_bias
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_zero_init_bias
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 16
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4
|
||||
hsdp_shard_dim: 4
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1000 # experiment: the bootstrap signal shows well before this
|
||||
gradient_accumulation_steps: 8 # 1 * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/exp_openvid_14b_warmup_out
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 100 # keep all — early checkpoints are the whole point
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: exp-openvid-14b-encoder-warmup
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
track_warmup:
|
||||
_target_: fastvideo.train.callbacks.track_warmup.TrackWarmupCallback
|
||||
warmup_steps: 300 # DiT held at lr 0 while the track pathway learns
|
||||
ramp_steps: 200 # then linearly hand control back to the full model
|
||||
log_every: 25
|
||||
track_validation:
|
||||
_target_: fastvideo.train.callbacks.track_validation.TrackValidationCallback
|
||||
every_steps: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,83 @@
|
||||
# MotionStream-style INITIAL teacher training on OpenVid-1M bidir set — Wan2.1-I2V-14B-720P.
|
||||
# Same recipe/data as the 1.3B run (finetune_wantrack_openvid_sparse_1p3b.yaml): sparse point
|
||||
# conditioning (WANTRACK_* env at launch), lr 1e-5, global bs 128, 480x832, 24fps, 121f,
|
||||
# flow_shift 6, bf16 DiT. SHARES the same combined_parquet_dataset (VAE/T5/CLIP identical
|
||||
# across Wan2.1 sizes). Plain I2V init (no control warm-start) -> track channels start at zero.
|
||||
# 8 nodes x 4 GPU = 32 GPUs; FSDP shards the 14B within each node (shard_dim 4) and replicates
|
||||
# across nodes (replicate 8). grad_accum 4 -> global bs 128.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_nobias_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_nobias_init
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 32 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8 # replicate across the 8 nodes
|
||||
hsdp_shard_dim: 4 # shard the 14B within each node over NVLink (8*4 = 32)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4800 # MotionStream stage-1 (=4.8K); 14B is slow — resumable, stop when good
|
||||
gradient_accumulation_steps: 4 # 1 * 4 * 32 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_14b_out
|
||||
training_state_checkpointing_steps: 100
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-14b-sparse-stage1
|
||||
|
||||
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: 400
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
+91
@@ -0,0 +1,91 @@
|
||||
# STEP D — MotionStream stage-1 teacher training on the full OpenVid-1M bidir set, Wan2.1-14B.
|
||||
#
|
||||
# Direct 14B analog of finetune_wantrack_openvid_sparse_1p3b_merged_bias.yaml (wandb run
|
||||
# openvid_bidir_1p3b_merged_bias). Hparams are copied EXACTLY from that run — lr 1e-5,
|
||||
# global bs 128, 4800 steps, 480x832, 121f, 24fps, flow_shift 6, bf16 DiT — the only changes
|
||||
# are the ones 14B geometry forces (hsdp dims, grad_accum to keep bs 128 at 32 GPUs).
|
||||
#
|
||||
# init_from is the step-C MERGED init: a fresh 14B base whose track_encoder AND
|
||||
# patch_embedding.weight[:, 36:] both come from the 14B-native overfit (steps A+B), so the two
|
||||
# are co-adapted. This is the whole point of the pipeline: gradient-SNR analysis showed a
|
||||
# random 14B init gives per-sample track_encoder.proj grad ~0.011 (the same regime that made
|
||||
# 1.3B-random FAIL), while 1.3B-merged sat at ~0.502 and succeeded.
|
||||
#
|
||||
# Launched with WANTRACK_FREEZE_HEAD=1 (track_encoder frozen) + TRACKWAN_TRACK_BIAS=1.
|
||||
# 8 nodes x 4 GPU = 32 GPUs; grad_accum 4 -> 1 * 4 * 32 = 128 global.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_merged_from_overfit_bias
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_merged_from_overfit_bias
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 32 # overridden by the launcher to nodes*4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8 # replicate across the 8 nodes
|
||||
hsdp_shard_dim: 4 # shard the 14B within each node over NVLink (8*4 = 32)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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
|
||||
gradient_accumulation_steps: 4 # 1 * 4 * 32 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_14b_merged_bias_out
|
||||
training_state_checkpointing_steps: 100 # ~backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-14b-merged-init-bias
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # 2 openvid + car + dog
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
+91
@@ -0,0 +1,91 @@
|
||||
# STEP D — MotionStream stage-1 teacher training on the full OpenVid-1M bidir set, Wan2.1-14B.
|
||||
#
|
||||
# Direct 14B analog of finetune_wantrack_openvid_sparse_1p3b_merged_bias.yaml (wandb run
|
||||
# openvid_bidir_1p3b_merged_bias). Hparams are copied EXACTLY from that run — lr 1e-5,
|
||||
# global bs 128, 4800 steps, 480x832, 121f, 24fps, flow_shift 6, bf16 DiT — the only changes
|
||||
# are the ones 14B geometry forces (hsdp dims, grad_accum to keep bs 128 at 32 GPUs).
|
||||
#
|
||||
# init_from is the step-C MERGED init: a fresh 14B base whose track_encoder AND
|
||||
# patch_embedding.weight[:, 36:] both come from the 14B-native overfit (steps A+B), so the two
|
||||
# are co-adapted. This is the whole point of the pipeline: gradient-SNR analysis showed a
|
||||
# random 14B init gives per-sample track_encoder.proj grad ~0.011 (the same regime that made
|
||||
# 1.3B-random FAIL), while 1.3B-merged sat at ~0.502 and succeeded.
|
||||
#
|
||||
# Launched with WANTRACK_FREEZE_HEAD=1 (track_encoder frozen) + TRACKWAN_TRACK_BIAS=1.
|
||||
# 8 nodes x 4 GPU = 32 GPUs; grad_accum 4 -> 1 * 4 * 32 = 128 global.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_merged_from_fixed2000_bias
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_merged_from_fixed2000_bias
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by launcher to nodes*4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard the 14B within each node over NVLink (8*4 = 32)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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
|
||||
gradient_accumulation_steps: 8 # 1 * 8 * 16 gpus = 128 global (4 nodes)
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_14b_stage1_out
|
||||
training_state_checkpointing_steps: 100 # ~backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-14b-stage1-fixed2000
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # 2 openvid + car + dog
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
+124
@@ -0,0 +1,124 @@
|
||||
# STEP D — MotionStream stage-1 teacher training on the full OpenVid-1M bidir set, Wan2.1-14B.
|
||||
#
|
||||
# Direct 14B analog of finetune_wantrack_openvid_sparse_1p3b_merged_bias.yaml (wandb run
|
||||
# openvid_bidir_1p3b_merged_bias). Hparams are copied EXACTLY from that run — lr 1e-5,
|
||||
# global bs 128, 4800 steps, 480x832, 121f, 24fps, flow_shift 6, bf16 DiT — the only changes
|
||||
# are the ones 14B geometry forces (hsdp dims, grad_accum to keep bs 128 at 32 GPUs).
|
||||
#
|
||||
# init_from is the step-C MERGED init: a fresh 14B base whose track_encoder AND
|
||||
# patch_embedding.weight[:, 36:] both come from the 14B-native overfit (steps A+B), so the two
|
||||
# are co-adapted. This is the whole point of the pipeline: gradient-SNR analysis showed a
|
||||
# random 14B init gives per-sample track_encoder.proj grad ~0.011 (the same regime that made
|
||||
# 1.3B-random FAIL), while 1.3B-merged sat at ~0.502 and succeeded.
|
||||
#
|
||||
# Launch with WANTRACK_FREEZE_HEAD=0 (track_encoder TRAINABLE) + TRACKWAN_TRACK_BIAS=1.
|
||||
#
|
||||
# CORRECTED 2026-07-31 — this header previously said FREEZE_HEAD=1, which contradicted the
|
||||
# 1.3B run it claims to copy exactly:
|
||||
# * run_openvid_bidir_held.sh (the 1.3B stage-1 launcher) never sets WANTRACK_FREEZE_HEAD,
|
||||
# and wantrack.py defaults it to "0" -> the head was TRAINABLE.
|
||||
# * the three 1.3B run logs (openvid_bidir_1p3b_merged_bias_hold_{507,511,513}.out) contain
|
||||
# ZERO "FROZE track_encoder" lines, confirming it.
|
||||
# * every real freeze in this repo is stage-2/3: run_openvid_stage2*_slurm.sh and
|
||||
# run_openvid_stage3_slurm.sh ("head is converged" — converged BECAUSE stage-1 trained it).
|
||||
# * wantrack.py:78 itself calls it a "MotionStream stage-2 knob ... after initial training";
|
||||
# stage-1 bidir finetuning IS the initial training.
|
||||
# * that same comment notes the head only plateaus by step ~4700 of this 4800-step stage,
|
||||
# so freezing at step 0 removes essentially the whole trajectory, not a redundant tail.
|
||||
# * no 14B stage-1 run had ever been launched, so FREEZE_HEAD=1 here was untested.
|
||||
# (The wantrack.py comment attributes "The track head remains frozen after initial training as
|
||||
# it already operates chunk-wise" to 2511.01266; that sentence is not in the paper — its
|
||||
# "chunk-wise" usages refer to KV-cache chunking.)
|
||||
#
|
||||
# 16 nodes x 4 GPU = 64 GPUs; grad_accum 2 -> 1 * 2 * 64 = 128 global.
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_merged_from_bs16_1200_bias
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_14b_i2v_d64_merged_from_bs16_1200_bias
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 64 # 16 nodes x 4
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 16 # replicate across 16 nodes
|
||||
hsdp_shard_dim: 4 # shard 14B within each node (16x4=64)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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
|
||||
gradient_accumulation_steps: 2 # 64 gpu x 1 x 2 = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_14b_stage1_bs16merge_out
|
||||
training_state_checkpointing_steps: 100 # ~backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-14b-stage1-bs16merge
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # 2 openvid + car + dog
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
# PROVENANCE of flow_shift 6 (traced 2026-07-31, it is NOT inherited from anywhere):
|
||||
# - NOT a Wan default: Wan2.1-I2V-14B-720P's own scheduler_config.json says 5.0, the
|
||||
# Fun-1.3B-InP says 3.0 — and this model's init checkpoint declares 5.0.
|
||||
# - NOT from MotionStream: the paper (2511.01266) specifies no shift value at all, and
|
||||
# trained Wan2.1-1.3B @832x480 / Wan2.2-5B @1280x704.
|
||||
# - NOT from matrixgame: matrixgame2.py / matrixgame3.py both default to 5.0.
|
||||
# - FastVideo's own causal recipe uses 5: step1_kd.yaml's t_list [995,937,833,625] is the
|
||||
# shift-5 warp of the Self-Forcing grid [1000,750,500,250].
|
||||
# It was set by hand in ac919b1c (2026-07-14) with the first openvid config, then copied.
|
||||
# KEEP IT ANYWAY: the downstream stack is calibrated to it — the CD/SF/KD student schedule
|
||||
# t_list [1000, 947, 857, 667] IS the shift-6 warp of [1000,750,500,250] (see
|
||||
# ablation/wantrack_causal_joint/kd_ode_init.yaml), so moving the teacher to 5 would desync
|
||||
# teacher and student. The stake is small either way: shift scales SNR by exactly 1/shift^2
|
||||
# at EVERY timestep, so 5->6 is a uniform 1.58 dB log-SNR translation, against the 4.44 dB
|
||||
# Wan itself applies between 480P (3) and 720P (5).
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,85 @@
|
||||
# MotionStream-style INITIAL teacher training (stage 1) on the full OpenVid-1M bidir set.
|
||||
# Recipe: sparse point conditioning (1-per-SAM-object + extras) from run c0xyfg57
|
||||
# (synth-sparse-random, d64_nobias init) but with MotionStream stage-1 hparams:
|
||||
# batch size 128, lr 1e-5, 480p (480x832), 24fps, 121f. NO stochastic track masking,
|
||||
# NO fixed/overfit sampling (set via WANTRACK_* env at launch — see run_openvid_bidir_slurm.sh).
|
||||
# num_gpus + hsdp dims below assume 4 nodes x 4 GPU = 16 GPUs (grad_accum 8 -> global bs 128).
|
||||
# Runs CONCURRENTLY with the 14B run (8 nodes) = 12 nodes total <= 14 usable. The small 1.3B
|
||||
# gets fewer nodes; grad_accum keeps global bs at 128 regardless of GPU count.
|
||||
# DiT runs bf16 (fp32 full-scale would be ~70h; verified 1-GPU fp32 ~11s/it in smoke).
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_nobias_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_nobias_init
|
||||
dit_precision: bf16 # fp32 was ~11s/it on 1 GPU (=~70h full); bf16 for throughput
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard within each node over NVLink (4*4 = 16)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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 # MotionStream stage-1 (=4.8K); ~2.4 epochs over 259k at bs128; resumable
|
||||
gradient_accumulation_steps: 8 # 1 (per-gpu) * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_1p3b_out
|
||||
training_state_checkpointing_steps: 100 # ~40min backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-1p3b-sparse-stage1
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,85 @@
|
||||
# MotionStream-style INITIAL teacher training (stage 1) on the full OpenVid-1M bidir set.
|
||||
# Recipe: sparse point conditioning (1-per-SAM-object + extras) from run c0xyfg57
|
||||
# (synth-sparse-random, d64_nobias init) but with MotionStream stage-1 hparams:
|
||||
# batch size 128, lr 1e-5, 480p (480x832), 24fps, 121f. NO stochastic track masking,
|
||||
# NO fixed/overfit sampling (set via WANTRACK_* env at launch — see run_openvid_bidir_slurm.sh).
|
||||
# num_gpus + hsdp dims below assume 4 nodes x 4 GPU = 16 GPUs (grad_accum 8 -> global bs 128).
|
||||
# Runs CONCURRENTLY with the 14B run (8 nodes) = 12 nodes total <= 14 usable. The small 1.3B
|
||||
# gets fewer nodes; grad_accum keeps global bs at 128 regardless of GPU count.
|
||||
# DiT runs bf16 (fp32 full-scale would be ~70h; verified 1-GPU fp32 ~11s/it in smoke).
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_merged_from_overfit
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_merged_from_overfit
|
||||
dit_precision: bf16 # fp32 was ~11s/it on 1 GPU (=~70h full); bf16 for throughput
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard within each node over NVLink (4*4 = 16)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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 # MotionStream stage-1 (=4.8K); ~2.4 epochs over 259k at bs128; resumable
|
||||
gradient_accumulation_steps: 8 # 1 (per-gpu) * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_1p3b_merged_out
|
||||
training_state_checkpointing_steps: 100 # ~40min backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-1p3b-merged-init
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
+85
@@ -0,0 +1,85 @@
|
||||
# MotionStream-style INITIAL teacher training (stage 1) on the full OpenVid-1M bidir set.
|
||||
# Recipe: sparse point conditioning (1-per-SAM-object + extras) from run c0xyfg57
|
||||
# (synth-sparse-random, d64_nobias init) but with MotionStream stage-1 hparams:
|
||||
# batch size 128, lr 1e-5, 480p (480x832), 24fps, 121f. NO stochastic track masking,
|
||||
# NO fixed/overfit sampling (set via WANTRACK_* env at launch — see run_openvid_bidir_slurm.sh).
|
||||
# num_gpus + hsdp dims below assume 4 nodes x 4 GPU = 16 GPUs (grad_accum 8 -> global bs 128).
|
||||
# Runs CONCURRENTLY with the 14B run (8 nodes) = 12 nodes total <= 14 usable. The small 1.3B
|
||||
# gets fewer nodes; grad_accum keeps global bs at 128 regardless of GPU count.
|
||||
# DiT runs bf16 (fp32 full-scale would be ~70h; verified 1-GPU fp32 ~11s/it in smoke).
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_merged_from_overfit_bias
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_merged_from_overfit_bias
|
||||
dit_precision: bf16 # fp32 was ~11s/it on 1 GPU (=~70h full); bf16 for throughput
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard within each node over NVLink (4*4 = 16)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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 # MotionStream stage-1 (=4.8K); ~2.4 epochs over 259k at bs128; resumable
|
||||
gradient_accumulation_steps: 8 # 1 (per-gpu) * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_1p3b_merged_bias_out
|
||||
training_state_checkpointing_steps: 100 # ~40min backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-1p3b-merged-init-bias
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,85 @@
|
||||
# MotionStream-style INITIAL teacher training (stage 1) on the full OpenVid-1M bidir set.
|
||||
# Recipe: sparse point conditioning (1-per-SAM-object + extras) from run c0xyfg57
|
||||
# (synth-sparse-random, d64_nobias init) but with MotionStream stage-1 hparams:
|
||||
# batch size 128, lr 1e-5, 480p (480x832), 24fps, 121f. NO stochastic track masking,
|
||||
# NO fixed/overfit sampling (set via WANTRACK_* env at launch — see run_openvid_bidir_slurm.sh).
|
||||
# num_gpus + hsdp dims below assume 4 nodes x 4 GPU = 16 GPUs (grad_accum 8 -> global bs 128).
|
||||
# Runs CONCURRENTLY with the 14B run (8 nodes) = 12 nodes total <= 14 usable. The small 1.3B
|
||||
# gets fewer nodes; grad_accum keeps global bs at 128 regardless of GPU count.
|
||||
# DiT runs bf16 (fp32 full-scale would be ~70h; verified 1-GPU fp32 ~11s/it in smoke).
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_nobias_random_init
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/models/trackwan_1.3b_i2v_d64_nobias_random_init
|
||||
dit_precision: bf16 # fp32 was ~11s/it on 1 GPU (=~70h full); bf16 for throughput
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard within each node over NVLink (4*4 = 16)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
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 # MotionStream stage-1 (=4.8K); ~2.4 epochs over 259k at bs128; resumable
|
||||
gradient_accumulation_steps: 8 # 1 (per-gpu) * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_1p3b_random_out
|
||||
training_state_checkpointing_steps: 100 # ~40min backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-1p3b-random-init
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,85 @@
|
||||
# MotionStream-style INITIAL teacher training (stage 1) on the full OpenVid-1M bidir set.
|
||||
# Recipe: sparse point conditioning (1-per-SAM-object + extras) from run c0xyfg57
|
||||
# (synth-sparse-random, d64_nobias init) but with MotionStream stage-1 hparams:
|
||||
# batch size 128, lr 1e-5, 480p (480x832), 24fps, 121f. NO stochastic track masking,
|
||||
# NO fixed/overfit sampling (set via WANTRACK_* env at launch — see run_openvid_bidir_slurm.sh).
|
||||
# num_gpus + hsdp dims below assume 4 nodes x 4 GPU = 16 GPUs (grad_accum 8 -> global bs 128).
|
||||
# Runs CONCURRENTLY with the 14B run (8 nodes) = 12 nodes total <= 14 usable. The small 1.3B
|
||||
# gets fewer nodes; grad_accum keeps global bs at 128 regardless of GPU count.
|
||||
# DiT runs bf16 (fp32 full-scale would be ~70h; verified 1-GPU fp32 ~11s/it in smoke).
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800
|
||||
dit_precision: bf16 # fp32 was ~11s/it on 1 GPU (=~70h full); bf16 for throughput
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard within each node over NVLink (4*4 = 16)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-6 # MotionStream teacher stage-2
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 800 # MotionStream stage-2 (6:1 ratio); ~0.4 epoch on openvid
|
||||
gradient_accumulation_steps: 8 # 1 (per-gpu) * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_1p3b_stage2_frozen_out
|
||||
training_state_checkpointing_steps: 100 # ~40min backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: null
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-1p3b-stage2-frozen
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
@@ -0,0 +1,85 @@
|
||||
# MotionStream-style INITIAL teacher training (stage 1) on the full OpenVid-1M bidir set.
|
||||
# Recipe: sparse point conditioning (1-per-SAM-object + extras) from run c0xyfg57
|
||||
# (synth-sparse-random, d64_nobias init) but with MotionStream stage-1 hparams:
|
||||
# batch size 128, lr 1e-5, 480p (480x832), 24fps, 121f. NO stochastic track masking,
|
||||
# NO fixed/overfit sampling (set via WANTRACK_* env at launch — see run_openvid_bidir_slurm.sh).
|
||||
# num_gpus + hsdp dims below assume 4 nodes x 4 GPU = 16 GPUs (grad_accum 8 -> global bs 128).
|
||||
# Runs CONCURRENTLY with the 14B run (8 nodes) = 12 nodes total <= 14 usable. The small 1.3B
|
||||
# gets fewer nodes; grad_accum keeps global bs at 128 regardless of GPU count.
|
||||
# DiT runs bf16 (fp32 full-scale would be ~70h; verified 1-GPU fp32 ~11s/it in smoke).
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wantrack.wantrack.WanTrackModel
|
||||
init_from: /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
model_path: /mnt/lustre/vlm-s4duan/exports/merged_bias_ckpt4800
|
||||
dit_precision: bf16 # fp32 was ~11s/it on 1 GPU (=~70h full); bf16 for throughput
|
||||
|
||||
distributed:
|
||||
num_gpus: 16 # overridden by run_slurm.sh to nodes*NUM_GPUS
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 4 # replicate across the 4 nodes
|
||||
hsdp_shard_dim: 4 # shard within each node over NVLink (4*4 = 16)
|
||||
|
||||
data:
|
||||
data_path: /mnt/lustre/vlm-s4duan/openvid_1m/combined_parquet_dataset
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 31
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 121
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-6 # MotionStream teacher stage-2
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 800 # MotionStream stage-2 (6:1 ratio); ~0.4 epoch on openvid
|
||||
gradient_accumulation_steps: 8 # 1 (per-gpu) * 8 * 16 gpus = 128 global
|
||||
|
||||
checkpoint:
|
||||
output_dir: /mnt/lustre/vlm-s4duan/openvid_bidir_1p3b_stage2_unfrozen_out
|
||||
training_state_checkpointing_steps: 100 # ~40min backstop vs the ~80min cordon cycle
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: null
|
||||
|
||||
tracker:
|
||||
trackers: [wandb]
|
||||
project_name: wantrack-bidir
|
||||
run_name: openvid-bidir-1p3b-stage2-unfrozen
|
||||
|
||||
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: 250
|
||||
num_val_samples: 4
|
||||
val_data_path: /mnt/lustre/vlm-s4duan/val_examples_mixed/combined_parquet_dataset # fixed car+dog synth clips
|
||||
num_inference_steps: 30
|
||||
guidance_scale: 3.0
|
||||
motion_guidance_scale: 1.5
|
||||
fps: 24
|
||||
grid_stride: 3
|
||||
tail: 12
|
||||
validate_at_start: true
|
||||
seed: 1000
|
||||
|
||||
pipeline:
|
||||
flow_shift: 6
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user