#!/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=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
# export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True

TRAIN_PATH=/home/tony/Work/neon/sunoGPT
echo working from $TRAIN_PATH
cd $TRAIN_PATH
# 30b 8 nodes can do batch 12 in forward loss
# 150k 30b: /app/suno/data/dpo/models/model_30b_150k.pt
# 30b fted: /app/suno/checkpoints/2024-07-08_15-35-27/last_ckpt_infer.pt
# 30b t1: /app/suno/data/dpo/models/model_30b_fix_ft2_20k.pt
# 30b t1 v12: /app/suno/checkpoints/2024-08-04_01-35-06/last_ckpt_infer.pt
# 30b t2 (t1 v12): /app/suno/data/dpo/models/model_30b_ft_t2.pt
# 30b t3 (t2 v10): /app/suno/data/dpo/models/model_30b_ft_t3.pt
# 30b t2 v17: /app/suno/checkpoints/2024-09-04_05-11-53/last_ckpt_infer.pt
# https://wandb.ai/suno/chirp-v4_dev2/runs/335lzmnj/overview?nw=nwusertonytongsuno
/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="/app/suno/checkpoints" \
    --data_dir="/app/suno/data/dpo/13b_s8_v3" \
    \
    --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-9 \
    --do_ipo=True \
    --dpo_beta=5.0 \
    --sft_loss_scale=0.0 \
    --semantic_codebook_weight=4.0 \
    --last_codebook_weight=0.5 \
    --warmup_iters=50 \
    --max_iters=1000 \
    \
    --grad_clip=0.1 \
    --eval_interval=4_000 \
    --eval_iters=25 \
    --step_save_iters=4_000 \
    \
    --block_size=8832 \
    --t_text=2560 \
    --t_memmap=6016 \
    --t_audio=6272 \
    --use_rotary_pos_emb=True \
    --rope_theta=500_000 \
    --use_qk_norm=True \
    --activation_f="silu" \
    --embed_scale_factor=10.0 \
    --global_every_n_layers=1 \
    \
    --n_layer=60 \
    --n_head=56 \
    --d_head=128 \
    --n_kv_head=4 \
    --attention_type="tao" \
    \
    --gradient_accumulation_steps=1 \
    --batch_size=2 \
    --eval_loss_batch_size=16 \
    \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --grad_checkpointing=True \
    \
    --preload_checkpoint="/app/suno/checkpoints/2024-09-14_17-41-31/last_ckpt_infer.pt" \
    --model_cache_loss_name="30b_tg_bt16" \
    --preload_strict=False \
    --local_cache_dir="/mnt/localdisk/tmp" \
    \
    --wandb_log=True \
    --wandb_project="chirp-v4-dpo" \
    --wandb_run_name="dpo_30b_tg_13b_s8_v3"
    