File size: 4,120 Bytes
c653378
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
#!/usr/bin/env bash
# DA3 single-task fine-tune: task 5 `setting_mousetraps`, 8 GPUs.
#
# Recipe as specified:
#   geometry (core = bank builder + injection, geom = ray MLP) trains from step 0
#       1e-4, cosine to 1e-6 by 20k
#   backbone (vlm) frozen for the first 2k steps, then ramps 1e-7 -> 5e-5 over steps 2k..5k,
#       then cosine to 1e-6 by 20k
#   20k steps, batch 256 if it fits else 128
#
# LR SCALING: at BS=128 every LR is multiplied by sqrt(128/256) = 0.70711 (square-root scaling),
# so the schedule SHAPE is preserved and only the magnitude tracks the batch.
#
# DA3 is frozen throughout -- it is the frozen extractor, never in the optimizer. "geometry
# modules" here means the trainable spatial branch, not the backbone.
set -uo pipefail

BS=${BS:-256}
GPUS="${1:-$(nvidia-smi --query-gpu=index --format=csv,noheader | paste -sd, -)}"
NG=$(echo "$GPUS" | tr ',' '\n' | grep -c .)
STEPS=${STEPS:-20000}
EXP=${EXP_OVERRIDE:-b1k_da3_task5_mousetraps}

[ $((BS % NG)) -eq 0 ] || { echo "ABORT: batch $BS not divisible by $NG GPUs"; exit 2; }

# sqrt(BS/256) applied to every LR in both groups
S=$(awk -v b="$BS" 'BEGIN{printf "%.6f", sqrt(b/256.0)}')
f() { awk -v v="$1" -v s="$S" 'BEGIN{printf "%.4e", v*s}'; }
GEOM_PEAK=$(f 1e-4);  GEOM_END=$(f 1e-6)
VLM_START=$(f 1e-7);  VLM_PEAK=$(f 5e-5);  VLM_END=$(f 1e-6)

echo "=============================================================="
echo " task 5  setting_mousetraps   (DA3, frozen extractor)"
echo " GPUs   : $GPUS (n=$NG)   batch $BS ($((BS/NG))/GPU)   fsdp off"
echo " LR scale vs BS=256 reference: ${S}x   (sqrt rule)"
echo " geom+core: $GEOM_PEAK from step 0, cosine -> $GEOM_END @${STEPS}"
echo " vlm      : 0 for 2k, ramp $VLM_START -> $VLM_PEAK over 2k..5k, cosine -> $VLM_END @${STEPS}"
echo " data     : behavior_2026_224_gop8 (224, GOP=8)   depth: GROUND TRUTH"
echo " assets   : assets_qvelfix (global qvel-corrected norm stats)"
echo " init     : behavior_50t_checkpoint"
echo "=============================================================="

cd /work/jack/behavior-1k-solution
LOG=/work/jack/behavior1k/run_logs/${EXP}_train.log
mkdir -p /work/jack/behavior1k/run_logs

CUDA_VISIBLE_DEVICES="$GPUS" \
EXP="$EXP" \
USE_DA3_FULL=1 DA3_LR_GROUPS=1 \
B1K_ACTIVITIES=setting_mousetraps \
B1K_2026_ROOT=/work/jack/behavior1k/data/behavior_2026_224_gop8 \
B1K_TASK_DATA_JSON=/work/jack/behavior1k/task_data.json \
B1K_INIT_PARAMS=/work/jack/behavior1k/checkpoints/behavior_50t_checkpoint/params \
B1K_ASSETS_BASE=/work/jack/behavior1k/assets_qvelfix \
B1K_TASK_SPACE=100 B1K_USE_GT_DEPTH=1 B1K_DECODE_RESIZE=0 \
B1K_MP_CONTEXT=fork B1K_HOST_PREFETCH=1 B1K_TORCH_COLLATE=1 \
B1K_DLPACK=${B1K_DLPACK:-1} \
B1K_EXTRACT_DEVICES=$(python3 -c "print(','.join(f'cuda:{i}' for i in range($NG)))") \
B1K_DA3_FWD_CHUNK=${B1K_DA3_FWD_CHUNK:-16} \
XLA_PYTHON_CLIENT_MEM_FRACTION=${XLA_PYTHON_CLIENT_MEM_FRACTION:-0.72} \
XLA_PYTHON_CLIENT_ALLOCATOR=default \
XLA_FLAGS="--xla_gpu_enable_latency_hiding_scheduler=true --xla_gpu_all_reduce_combine_threshold_bytes=8388608 --xla_gpu_enable_highest_priority_async_stream=true" \
JAX_COMPILATION_CACHE_DIR=/work/jack/jax_compile_cache \
HF_HOME=/work/jack/.cache/huggingface \
PYTHONUNBUFFERED=1 \
RAYON_NUM_THREADS=1 OMP_NUM_THREADS=1 OPENBLAS_NUM_THREADS=1 MKL_NUM_THREADS=1 \
NUMEXPR_NUM_THREADS=1 POLARS_MAX_THREADS=1 TOKENIZERS_PARALLELISM=false \
BS=$BS STEPS=$STEPS FSDP_DEVICES=1 NW=${NW:-32} SHUFFLE=1 \
LR_DECAY_STEPS=$STEPS \
LR_GEOM_DELAY=0 LR_GEOM_RAMP=0 LR_GEOM_PEAK=$GEOM_PEAK LR_GEOM_DECAY_START=1 LR_GEOM_END=$GEOM_END \
LR_CORE_DELAY=0 LR_CORE_RAMP=0 LR_CORE_PEAK=$GEOM_PEAK LR_CORE_DECAY_START=1 LR_CORE_END=$GEOM_END \
LR_VLM_DELAY=2000 LR_VLM_RAMP=3000 LR_VLM_RAMP_START=$VLM_START LR_VLM_PEAK=$VLM_PEAK \
LR_VLM_DECAY_START=5000 LR_VLM_END=$VLM_END \
DA3_SPATIAL_EPS=1e-16 DA3_SPATIAL_WD=1e-4 \
 DA3_BANK_CENTER=${DA3_BANK_CENTER:-0} \
SAVE_INTERVAL=${SAVE_INTERVAL:-1000} KEEP_PERIOD=${KEEP_PERIOD:-4000} LOG_INTERVAL=${LOG_INTERVAL:-25} \
OVERWRITE=${OVERWRITE:-1} RESUME=${RESUME:-0} \
  nohup .venv/bin/python scripts/train_2026_da3.py >> "$LOG" 2>&1 &
echo "LAUNCHED pid $! -> $LOG"