#!/bin/bash # AV + AR SFT for the gemma-4-26B-A4B NLA, concurrent on GPU 0 / GPU 1 (bf16, LoRA). # Dedicated pod => pinning CUDA_VISIBLE_DEVICES per process is allowed. set -u set -a; source /root/.nla_env; set +a export PYTHONUNBUFFERED=1 export WANDB_ENTITY=MATS10-CS-JB export WANDB_RUN_GROUP=g4a export TOKENIZERS_PARALLELISM=false DATA=${DATA:-/data/gemma4_nla/or} CK=/ckpts/gemma4_nla BASE=google/gemma-4-26B-A4B-it BATCH=${BATCH:-64} # bf16 on 183GB B200; drop to 32 (+GA 2) if OOM GA=${GA:-1} STEPS=${STEPS:-300} # ~1 epoch on the ~20k/side subset (avoid overfit) SAVE=$(( STEPS / 25 )); [ $SAVE -lt 1 ] && SAVE=1 COMMON="--use-lora --lora-r 128 --lora-alpha 16 \ --lr 3e-5 --min-lr 3e-6 --lr-warmup-steps 20 --max-grad-norm 1.0 \ --num-steps $STEPS --batch-size $BATCH --gradient-accumulation-steps $GA \ --save-every $SAVE --wandb-project cot-oracle --seed 0" echo "=== AV SFT (GPU0) + AR SFT (GPU1), batch=$BATCH ga=$GA ===" CUDA_VISIBLE_DEVICES=0 python3 -m nla.train_sft --mode av --base-ckpt "$BASE" \ --parquet $DATA/av_sft_shuf.parquet --sidecar $DATA/av_sft_shuf.parquet \ --save-dir $CK/av $COMMON --wandb-name g4a_av_sft \ > /root/av_sft.log 2>&1 & AVPID=$! echo "AV pid=$AVPID" CUDA_VISIBLE_DEVICES=1 python3 -m nla.train_sft --mode ar --base-ckpt "$BASE" \ --parquet $DATA/ar_sft_shuf.parquet --sidecar $DATA/ar_sft_shuf.parquet \ --save-dir $CK/ar --ar-num-layers 21 \ --heldout-parquet $DATA/av_sft_shuf.parquet --heldout-every $SAVE \ $COMMON --wandb-name g4a_ar_sft \ > /root/ar_sft.log 2>&1 & ARPID=$! echo "AR pid=$ARPID" FAIL=0 wait $AVPID || { echo "AV FAILED"; FAIL=1; } wait $ARPID || { echo "AR FAILED"; FAIL=1; } echo "=== AV+AR done (fail=$FAIL) ===" exit $FAIL