japhba's picture
Upload folder using huggingface_hub
74d1994 verified
Raw
History Blame Contribute Delete
1.75 kB
#!/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