#!/bin/bash

# Batch job submission script for reward model caching
# Usage: ./submit_cache_batch.sh v32 v33 v34 v35 v36 v37 v38 v39
# Or: ./submit_cache_batch.sh 32 33 34 35 36 37 38 39

if [ $# -eq 0 ]; then
    echo "Usage: $0 <version1> <version2> ... <versionN>"
    echo "Example: $0 v32 v33 v34 v35 v36 v37 v38 v39"
    echo "Example: $0 32 33 34 35 36 37 38 39"
    exit 1
fi

# Configuration - edit these as needed
CHECKPOINT="/app2/suno/checkpoints/2025-11-10_10-10-52/last_ckpt_infer.pt"
OUTPUT_PREFIX="reward_2025-11-10_10-10-52"
BATCH_SIZE=20
T_DATA_MEMMAP=12_000

echo "=================================================="
echo "Submitting ${#@} reward caching jobs..."
echo "Checkpoint: $CHECKPOINT"
echo "=================================================="

SUBMITTED_JOBS=()

for VERSION in "$@"; do
    # Strip 'v' prefix if present
    VERSION_NUM="${VERSION#v}"
    
    DATA_DIR="/app2/suno/data/dpo/crow_t1_v${VERSION_NUM}"
    OUTPUT_NAME="${OUTPUT_PREFIX}_v${VERSION_NUM}"
    
    # Create a temporary sbatch script for this version
    TEMP_SBATCH=$(mktemp /tmp/sbatch_cache_reward_v${VERSION_NUM}_XXXXXX.sh)
    
    cat > "$TEMP_SBATCH" << 'EOF'
#!/bin/bash
#SBATCH --job-name="cache_reward_vVERSION_NUM"
#SBATCH --nodes=4
#SBATCH --ntasks-per-node=8
#SBATCH --cpus-per-task=4
#SBATCH --gres=gpu:8
#SBATCH --output=/home/tony/slurm/logs/run_%x_%j.txt
#SBATCH --error=/home/tony/slurm/logs/run_%x_%j_err.txt

# Reward caching job - computes rewards for all samples in DPO dataset
# Dataset: DATA_DIR
# Output: OUTPUT_NAME

# ============================================================================
# Environment Configuration
# ============================================================================
export CUDA_LAUNCH_BLOCKING=1

# NCCL: only shout when something's wrong; still fail fast
export NCCL_DEBUG=WARN
export NCCL_DEBUG_SUBSYS=INIT
export TORCH_NCCL_ASYNC_ERROR_HANDLING=1
export TORCH_NCCL_BLOCKING_WAIT=1
export NCCL_TIMEOUT=600

# PyTorch distributed: minimal breadcrumbs, full stack only on crash
export TORCH_SHOW_CPP_STACKTRACES=1

# Allocator: no periodic dumps, but better chance to avoid fragmentation
export PYTORCH_CUDA_ALLOC_CONF=garbage_collection_threshold:0.6,max_split_size_mb:512

export OMP_NUM_THREADS=1
export MASTER_ADDR=$(scontrol show hostnames "$SLURM_JOB_NODELIST" | head -n 1)
export MASTER_PORT=12836
export TRITON_CACHE_DIR=/mnt/localdisk/.triton_cache_$USER

# ============================================================================
# Job Execution
# ============================================================================
TRAIN_PATH=/home/tony/Work/neon_2/sunoGPT
echo working from $TRAIN_PATH
cd $TRAIN_PATH

# Reward model caching - pre-compute rewards for all DPO samples
srun -K1 /home/tony/anaconda3/envs/gpt_n/bin/python -u \
    scripts/cache_reward.py \
    --checkpoint="CHECKPOINT" \
    --data_dir="DATA_DIR" \
    --output_name="OUTPUT_NAME" \
    --batch_size=BATCH_SIZE \
    --t_data_memmap=T_DATA_MEMMAP \
    --master_addr=$MASTER_ADDR \
    --master_port=$MASTER_PORT

echo working from $TRAIN_PATH
echo DONE
EOF
    
    # Replace placeholders with actual values
    sed -i "s|VERSION_NUM|${VERSION_NUM}|g" "$TEMP_SBATCH"
    sed -i "s|CHECKPOINT|${CHECKPOINT}|g" "$TEMP_SBATCH"
    sed -i "s|DATA_DIR|${DATA_DIR}|g" "$TEMP_SBATCH"
    sed -i "s|OUTPUT_NAME|${OUTPUT_NAME}|g" "$TEMP_SBATCH"
    sed -i "s|BATCH_SIZE|${BATCH_SIZE}|g" "$TEMP_SBATCH"
    sed -i "s|T_DATA_MEMMAP|${T_DATA_MEMMAP}|g" "$TEMP_SBATCH"
    
    # Submit the job
    JOB_OUTPUT=$(sbatch "$TEMP_SBATCH" 2>&1)
    JOB_ID=$(echo "$JOB_OUTPUT" | grep -oP 'Submitted batch job \K\d+')
    
    if [ -n "$JOB_ID" ]; then
        echo "✓ v${VERSION_NUM}: Job $JOB_ID submitted (data: $DATA_DIR)"
        SUBMITTED_JOBS+=("v${VERSION_NUM}:${JOB_ID}")
    else
        echo "✗ v${VERSION_NUM}: Submission failed - $JOB_OUTPUT"
    fi
    
    # Clean up temp file
    rm "$TEMP_SBATCH"
    
    # Small delay between submissions to avoid overwhelming scheduler
    sleep 0.5
done

echo "=================================================="
echo "Submission complete!"
echo "Submitted ${#SUBMITTED_JOBS[@]} jobs:"
for job in "${SUBMITTED_JOBS[@]}"; do
    echo "  - $job"
done
echo "=================================================="
echo "Monitor with: squeue -u $USER"
echo "Cancel all with: scancel ${SUBMITTED_JOBS[@]##*:}"

