VIPRA-reproduce / scripts /full_run.sbatch
yqi19's picture
VIPER reproduction: code, mini VIPER-19K pipeline, stage-1 checkpoint, qualitative comparisons and eval
e0db531 verified
Raw
History Blame Contribute Delete
1.63 kB
#!/bin/bash
#SBATCH --account=healthcareeng_computervision
#SBATCH --partition=backfill_singlenode
#SBATCH --nodes=1
#SBATCH --gpus-per-node=8
#SBATCH --exclusive
#SBATCH --time=24:00:00
#SBATCH --job-name=viper_train
#SBATCH --output=logs/full_run_%j.log
#
# Three-stage VIPER training, then batch inference + eval on the held-out pairs.
#
# Step counts are scaled DOWN from the paper's 3K/6K/6K. The paper trains on
# VIPER-19K; our mini dataset is 523 clips / 64 pairs, so the paper's schedule
# would be many hundreds of epochs over the pair set and would just memorise it.
# We keep the stage ratio (1 : 2 : 2 relative weight) and size stage 1 to a few
# epochs over the clip set.
#
# Throughput reference: 35.8 s/optimizer-step measured on one A100-80GB at
# 81x480x832 (seq_len 32760). With 8 ranks that is 8 samples per step.
set -e
cd "$SLURM_SUBMIT_DIR"
source env.sh
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
export PYTHONUNBUFFERED=1
STEPS=400 scripts/train.sh 1
STEPS=200 scripts/train.sh 2
STEPS=200 scripts/train.sh 3
echo "=== training complete: running validation inference (VIPER vs no-physics) ==="
srun --ntasks=4 scripts/run_infer.sh checkpoints/stage3/final.pt
cat results/manifest.jsonl.* > results/manifest.jsonl
echo "=== evaluating generated videos ==="
.venv/bin/python -m viper.eval --manifest results/manifest.jsonl \
--out results/eval_viper.json
echo "=== paired conditioning ablation (reference vs zeroed physics tokens) ==="
for STAGE in 1 2 3; do
.venv/bin/python -m viper.eval_loss --ckpt checkpoints/stage$STAGE/final.pt \
--out results/eval_loss_stage$STAGE.json
done