| #!/bin/bash |
| |
| |
| 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} |
| GA=${GA:-1} |
|
|
| STEPS=${STEPS:-300} |
| 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 |
|
|