#!/bin/bash
export CUDA_LAUNCH_BLOCKING=1
# NCCL: only shout when something's wrong; still fail fast
export NCCL_DEBUG=WARN                 # INFO spams per-collective; WARN is quiet
export NCCL_DEBUG_SUBSYS=INIT          # INIT only; skip COLL/P2P chatter
export TORCH_NCCL_ASYNC_ERROR_HANDLING=1
export TORCH_NCCL_BLOCKING_WAIT=1
export NCCL_TIMEOUT=600

# PyTorch distributed: minimal breadcrumbs, full stack only on crash
export TORCH_SHOW_CPP_STACKTRACES=1    # only prints on error

# Allocator: no periodic dumps, but better chance to avoid fragmentation
export PYTORCH_CUDA_ALLOC_CONF=garbage_collection_threshold:0.6,max_split_size_mb:512

export OMP_NUM_THREADS=1
export MASTER_ADDR=$(scontrol show hostnames "$SLURM_JOB_NODELIST" | head -n 1)
export MASTER_PORT=12835
export TRITON_CACHE_DIR=/mnt/localdisk/.triton_cache_$USER


TRAIN_PATH=/home/tony/Work/neon/sunoGPT
echo working from $TRAIN_PATH
cd $TRAIN_PATH
pkill -f 'spawn_main'
rm -rf /mnt/localdisk/tmp_tony
mkdir -p /mnt/localdisk/tmp_tony
chmod -R 777 /mnt/localdisk/tmp_tony
# https://wandb.ai/suno/chirp-v4_dev2/runs/335lzmnj/overview?nw=nwusertonytongsuno
/home/tony/anaconda3/envs/gpt_n/bin/python -u \
    train_dpo.py \
    --master_addr=$MASTER_ADDR \
    --master_port=$MASTER_PORT \
    \
    --out_dir="/app2/suno/checkpoints" \
    --data_dir="/app2/suno/data/dpo/crow_t1_v55" \
    \
    --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" \
    \
    --coarse_n_codebooks=0 \
    --block_size=14_080 \
    --t_text=2048 \
    --t_audio=12_000 \
    \
    --learning_rate=5e-6 \
    --min_lr=1e-7 \
    --do_ipo=True \
    --dpo_beta=20.0 \
    --sft_loss_scale=0.0 \
    --semantic_codebook_weight=4.0 \
    --last_codebook_weight=0.5 \
    --warmup_iters=50 \
    --max_iters=546 \
    \
    --grad_clip=0.1 \
    --eval_interval=1000 \
    --eval_iters=25 \
    --step_save_iters=4_000 \
    \
    --t_memmap=6016 \
    --t_data_memmap=12_000 \
    --data_coarse_n_codebooks=0 \
    --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=32 \
    --n_head=32 \
    --d_head=128 \
    --n_kv_head=4 \
    --attention_type="tao" \
    \
    --gradient_accumulation_steps=1 \
    --batch_size=8 \
    --eval_loss_batch_size=16 \
    \
    --fsdp=True \
    --sharding_strategy="full_shard" \
    --grad_checkpointing=True \
    --shuffle_data=False \
    --custom_seed_offset=830 \
    --local_shuffle_data=False \
    \
    --preload_checkpoint="/app2/suno/checkpoints/2025-11-02_08-24-35/last_ckpt_infer.pt" \
    --model_cache_loss_name="dodo_0828_crow_v54_sft005_n16_r4_stem" \
    --preload_strict=False \
    --local_cache_dir="/mnt/localdisk/tmp_tony" \
    \
    --wandb_log=True \
    --wandb_project="chirp-dodo-dpo" \
    --wandb_run_name="dodo_0828_crow_v55_n16_r5"
    

rm -rf /mnt/localdisk/tmp_tony
echo working from $TRAIN_PATH
echo DONE