| #!/usr/bin/env bash |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| 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; } |
|
|
| |
| 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" |
|
|