#!/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