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

export TRITON_CACHE_DIR=/tmp/.triton_cache
export NCCL_P2P_DISABLE=1

TRAIN_PATH=/mnt/round-surf/home/tony/Work/glockenspiel/sunoGPT
echo working from $TRAIN_PATH
cd $TRAIN_PATH

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.py \
    \
    --out_dir="/mnt/round-surf/checkpoints" \
    --data_dir="/mnt/round-surf/data/chirp_v4/multi" \
    --local_data_shard_dir="/tmp/data_shards" \
    --allow_data_shard_reuse=True \
    --weights_multiplier="youtube_music:1;youtube_music_lyrics_foreign:3;youtube_music_lyrics:2;genius_hq_lyrics:2;genius_hq_lyrics_foreign:3;jamendo:1;imslp:1;pond5_music:0.5;deezer_lyrics_foreign:3;deezer_lyrics:2;deezer:1;ytm_tagged:1;discogs_lyrics:1;discogs:0.5;discogs_covers:5" \
    --preload_checkpoint="/mnt/round-surf/checkpoints/13b/model.pt" \
    --preload_optimizer=False \
    --local_cache_dir="/tmp/checkpoint" \
    --custom_seed_offset=2 \
    \
    --learning_rate=1e-4 \
    --min_lr=1e-5 \
    --max_iters=20_000 \
    --warmup_iters=2_000 \
    --eval_interval=2_000 \
    --checkpoint_save_old_format=True \
    \
    --block_size=9728 \
    --t_memmap=6016 \
    --t_audio=7168 \
    --t_text=2560 \
    --use_rotary_pos_emb=True \
    --mask_padding=True \
    --pack=True \
    --activation_f="gelu" \
    --use_qk_norm=False \
    --suffix_first=True \
    --dropout_semantic=True \
    --artist_condition=True \
    --cover_condition=True \
    --semantic_shift_factor=250 \
    --coarse_shift_factor=25 \
    \
    --n_layer=40 \
    --n_head=40 \
    --d_head=128 \
    --n_kv_head=8 \
    \
    --gradient_accumulation_steps=1 \
    --batch_size=2 \
    --attention_type="tao" \
    \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --grad_checkpointing=True \
    \
    --wandb_log=True \
    --wandb_project="chirp-v4_dev2" \
    --wandb_run_name="13b_ft"



