#!/bin/bash
#SBATCH --job-name="cache_reward"
#SBATCH --nodes=4
#SBATCH --ntasks-per-node=8
#SBATCH --cpus-per-task=4           # Number of cores per tasks
#SBATCH --gres=gpu:8                 # Number of gpus
#SBATCH --output=/home/tony/slurm/logs/run_%x_%j.txt   # Set this dir where you want slurm outs to go
#SBATCH --error=/home/tony/slurm/logs/run_%x_%j_err.txt    # Set this dir where you want slurm outs to go

# Reward caching job - computes rewards for all samples in DPO dataset
# This only needs to run once per reward model checkpoint

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

# NCCL: only shout when something's wrong; still fail fast
export NCCL_DEBUG=WARN                 # INFO spams per-collective; WARN is quiet
export NCCL_DEBUG_SUBSYS=INIT          # INIT only; skip COLL/P2P chatter
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    # only prints on error

# 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
# two models:
# 2025-11-09_22-27-28 -- just auk sft v0
# 2025-11-10_10-10-52 -- full reward
srun -K1 /home/tony/anaconda3/envs/gpt_n/bin/python -u \
    scripts/cache_reward.py \
    --checkpoint="/app2/suno/checkpoints/2025-11-09_22-27-28/last_ckpt_infer.pt" \
    --data_dir="/app2/suno/data/dpo/crow_t1_v58" \
    --output_name="reward_2025-11-09_22-27-28" \
    --batch_size=20 \
    --t_data_memmap=12_000 \
    --master_addr=$MASTER_ADDR \
    --master_port=$MASTER_PORT

echo working from $TRAIN_PATH
echo DONE
