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

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

TRAIN_PATH=/mnt/round-surf/home/tony/Work/glockenspiel/sunoGPT
echo working from $TRAIN_PATH
cd $TRAIN_PATH
/mnt/round-surf/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" \
    train_dpo.py \
    \
    --out_dir="/mnt/round-surf/checkpoints" \
    --data_dir="/mnt/round-surf/data/dpo/2b_before_recode_v0" \
    --local_data_shard_dir="/tmp/data_shards_tony_2b" \
    --allow_data_shard_reuse=False \
    \
    --train_filename="data_tr.bin" \
    --train_metas_filename="meta_tr.jsonl" \
    --train_info_filename="info_tr.json" \
    --val_filename="data_val.bin" \
    --val_metas_filename="meta_val.jsonl" \
    --val_info_filename="info_val.json" \
    \
    --learning_rate=5e-7 \
    --min_lr=1e-8 \
    --do_ipo=True \
    --dpo_beta=5.0 \
    --semantic_codebook_weight=4.0 \
    --last_codebook_weight=0.5 \
    --warmup_iters=200 \
    --max_iters=1_000 \
    \
    --grad_clip=0.1 \
    --eval_interval=500 \
    --eval_iters=25 \
    --step_save_iters=2_000 \
    \
    --block_size=8704 \
    --t_text=2560 \
    --t_memmap=6016 \
    --t_audio=6144 \
    --use_rotary_pos_emb=True \
    --rope_theta=500_000 \
    --use_qk_norm=False \
    --activation_f="gelu" \
    \
    --n_layer=24 \
    --n_head=20 \
    --d_head=128 \
    --n_kv_head=4 \
    --attention_type="tao" \
    \
    --gradient_accumulation_steps=1 \
    --batch_size=4 \
    \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --grad_checkpointing=True \
    \
    --preload_checkpoint="/mnt/round-surf/data/dpo/models/2b_mod.pt" \
    --preload_strict=False \
    --preload_optimizer=False \
    --local_cache_dir="/tmp/checkpoint_tony" \
    \
    --wandb_log=True \
    --wandb_project="chirp-v4-dpo" \
    --wandb_run_name="dpo_2b_test"
    