#!/bin/bash
#SBATCH --job-name=ucf101_r3d18_array
#SBATCH --partition=gpu-nodes
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=8
#SBATCH --mem=32G
#SBATCH --time=48:00:00
#SBATCH --array=1-9
#SBATCH --output=slurm_%A_%a.out

# ==========================
# DDP Environment Variables
# ==========================
export MASTER_ADDR=$(hostname)
export MASTER_PORT=12345
export OMP_NUM_THREADS=8

echo "MASTER_ADDR=$MASTER_ADDR"
echo "MASTER_PORT=$MASTER_PORT"
echo "TASK ID = $SLURM_ARRAY_TASK_ID"

# =========================================================
# EXPERIMENT GRID (9 runs)
# splits × budgets
# =========================================================

# Splits
SPLITS=(
1 1 1
2 2 2
3 3 3
)

# Energy budgets
BUDGETS=(
0.001572 0.001048 0.000524
0.001572 0.001048 0.000524
0.001572 0.001048 0.000524
)

IDX=$((SLURM_ARRAY_TASK_ID - 1))

SPLIT=${SPLITS[$IDX]}
BUDGET=${BUDGETS[$IDX]}

echo "Running: split=$SPLIT  budget=$BUDGET"

# =========================================================
# RUN TRAINING
# =========================================================

python3 train.py \
   --dataset ucf101 \
   --model r3d18 \
   --split $SPLIT \
   --energy_budget $BUDGET \
   --optimizer adam \
   --epochs 80 \
   --batch_size 4
