#!/bin/bash
export CUDA_LAUNCH_BLOCKING=0
export NCCL_DEBUG=WARN
export TORCH_DISTRIBUTED_DEBUG=OFF

export OMP_NUM_THREADS=4
export HOSTNAMES=`scontrol show hostnames "$SLURM_JOB_NODELIST"`
export MASTER_ADDR=$(scontrol show hostnames "$SLURM_JOB_NODELIST" | head -n 1)
export MASTER_PORT=12835
export COUNT_NODE=`scontrol show hostnames "$SLURM_JOB_NODELIST" | wc -l`
export TRITON_CACHE_DIR=/mnt/localdisk/.triton_cache_tony

TRAIN_PATH=/home/tony/Work/neon/sunoGPT
LOCA_DATA_SHARD_DIR=/mnt/localdisk/tmp_tony_data_shards
pkill -f 'spawn_main'
echo "Clear cache data shard"
result=$(find $LOCA_DATA_SHARD_DIR)
echo "The contents of cache dir is: $result"
# rm $LOCA_DATA_SHARD_DIR/*
# result=$(find $LOCA_DATA_SHARD_DIR)
# echo "The contents of cache dir is: $result"
echo working from $TRAIN_PATH
cd $TRAIN_PATH

/home/tony/anaconda3/envs/suno_env/bin/torchrun \
    --master_addr=$MASTER_ADDR \
    --master_port=$MASTER_PORT \
    --nnodes=$SLURM_JOB_NUM_NODES \
    --nproc_per_node=$SLURM_GPUS_ON_NODE \
    --rdzv_id $SLURM_JOB_ID \
    --rdzv_backend c10d \
    --rdzv_endpoint "$MASTER_ADDR:$MASTER_PORT" \
    eval_loss.py \
    --preload_checkpoint="/app/suno/checkpoints/2024-07-04_04-14-58/2b_55k_infer.pt" \
    --preload_strict=True \
    --local_cache_dir="/mnt/localdisk/tmp" \
    --local_data_shard_dir="/mnt/localdisk/tmp" \
    --allow_data_shard_reuse=False \
    --compile=False \
    \
    --out_dir="/app/suno/data/dpo/chirp_v4_multi" \
    --out_sub_dir="eval_2b_raw_bt28_val_2node" \
    --data_dir="/app/suno/data/chirp_v4/multi" \
    --val_filename="data_val.bin" \
    --val_metas_filename="metas_val.jsonl" \
    --val_info_filename="info_val.json" \
    \
    --block_size=8832 \
    --t_memmap=6016 \
    --t_audio=6272 \
    --t_text=2560 \
    --use_rotary_pos_emb=True \
    --rope_theta=500_000 \
    --use_qk_norm=True \
    --embed_scale_factor=10.0 \
    --activation_f="silu" \
    --attention_sliding_window_size=1024 \
    --global_every_n_layers=1 \
    \
    --n_layer=24 \
    --n_head=20 \
    --d_head=128 \
    --n_kv_head=4 \
    --attention_type="tao" \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --batch_size=28 \
    --wandb_run_name="2b_val"