#!/bin/bash
export CUDA_LAUNCH_BLOCKING=0
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=12832
export COUNT_NODE=`scontrol show hostnames "$SLURM_JOB_NODELIST" | wc -l`

export NCCL_DEBUG=WARN
export NCCL_CROSS_NIC=2

export TRITON_CACHE_DIR=/mnt/localdisk/.triton_cache_tony

TRAIN_PATH=/home/tony/Work/neon/sunoGPT
# ls /mnt/localdisk/tmp_tony_ckpt/
# rm /mnt/localdisk/tmp_tony_ckpt/*.pt
# ls /mnt/localdisk/tmp_tony_ckpt/
pkill -f 'spawn_main'
echo working from $TRAIN_PATH
cd $TRAIN_PATH

# https://wandb.ai/suno/chirp-v4_dev2/runs/7g9yttzu/overview 
# 30b pretrain ckpt: /app/suno/checkpoints/2024-06-14_16-56-28/last_ckpt_infer.pt
torchrun \
    --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.py \
    \
    --out_dir="/app/suno/checkpoints" \
    --data_dir="/app/suno/data/dpo/chirp_v4_multi_ft_classical_t1" \
    --local_data_shard_dir="/mnt/localdisk/tmp" \
    --allow_data_shard_reuse=False \
    --checkpoint_save_old_format=True \
    --local_cache_dir="/mnt/localdisk/tmp_tony_ckpt_classical_ft" \
    --preload_checkpoint="/app/suno/data/dpo/models/model_30b_ft_t3.pt" \
    --preload_optimizer=False \
    --preload_strict=False \
    --custom_seed_offset=88 \
    --layer_init=False \
    \
    --learning_rate=1e-5 \
    --min_lr=1e-6 \
    --max_iters=5_000 \
    --warmup_iters=300 \
    --eval_interval=2_500 \
    --step_save_iters=2_500 \
    --last_codebook_weight=0.5 \
    \
    --block_size=8832 \
    --t_memmap=6016 \
    --t_audio=6272 \
    --t_text=2560 \
    --use_rotary_pos_emb=True \
    --mask_padding=True \
    --pack=False \
    --embed_scale_factor=10.0 \
    --suffix_first=False \
    --dropout_semantic=False \
    --artist_condition=False \
    --cover_condition=False \
    \
    --n_layer=60 \
    --n_head=56 \
    --d_head=128 \
    --n_kv_head=4 \
    \
    --gradient_accumulation_steps=1 \
    --batch_size=2 \
    --global_every_n_layers=1 \
    \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --grad_checkpointing=True \
    \
    --wandb_log=True \
    --wandb_project="chirp-v4_classical" \
    --wandb_run_name="30b_multi_ft_t3_classical_v1"

#    --preload_checkpoint="/mnt/localdisk/ckpt.pt" \
#    --preload_optimizer=False \
#    --custom_seed_offset=2 \
#    --data_dir="/tmp/data" \
#    --data_dir="/mnt/round-surf/data/chirp_v4_test/base" \
#    --debug_val_only=True \
#    --eval_iters=50

