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

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`

TRAIN_PATH=/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="/app/suno/checkpoints" \
    --data_dir="/app/suno/data/chirp_v3_ft_classical_4min" \
    \
    --learning_rate=6e-5 \
    --max_iters=10_000 \
    --warmup_iters=1_000 \
    --eval_interval=500 \
    --eval_iters=25 \
    --step_save_iters=1_000 \
    --step_save_infer=False \
    \
    --block_size=8184 \
    --t_text=2048 \
    --t_memmap=6008 \
    --t_audio=6136 \
    \
    --n_layer=32 \
    --n_head=32 \
    --d_head=128 \
    --n_kv_head=4 \
    --use_rotary_pos_emb=False \
    --attention_type="tao" \
    --attention_sliding_window_size=2048 \
    \
    --gradient_accumulation_steps=1 \
    --batch_size=4 \
    --last_codebook_weight=0.1 \
    \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --grad_checkpointing=True \
    \
    --wandb_log=True \
    --wandb_project="chirp-v2_6" \
    --wandb_run_name="classical_ft_4min_sliding_win_reweight_norope"

#    --data_dir="/app/suno/data/chirp_v4_test" \
#    --data_dir="/mnt/localdisk/data/test_base" \
#    --debug_val_only=True \
#    --eval_iters=50






