Download code/scripts/train_baseline.sh from teawhite/ActionRoPE: direct link, hf CLI and curl.
- Browser
- Download file 2.28 kB
-
https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/scripts/train_baseline.sh
- Command line
-
hf download hf://teawhite/ActionRoPE/code/scripts/train_baseline.sh
-
curl -L -o train_baseline.sh https://huggingface.co/teawhite/ActionRoPE/resolve/main/code/scripts/train_baseline.sh
2.28 kB
| # baseline 动作臂(linear / xattn / prompt / adaln,baseline/SPEC.md)正式训练:与 train_arope.sh 完全同一配方 | |
| # —— 8 卡 ZeRO-2、GA 默认 2、lr 1e-5、warmup 500、每 1000 步存权重(+ --save_state 两槽轮转)、每 500 步验证、 | |
| # --min_rt_ratio 0.8 / --val_min_rt_ratio 0.9、seed 0;唯一的差别是 --arm(普通 RoPE、无 mask、权重全 1、 | |
| # 文本按臂选:prompt 带动作从句,其余剥掉)。5k 步实验照 arope_5k_ga2 的做法在末尾覆盖: | |
| # bash scripts/train_baseline.sh linear --max_steps 5000 --warmup_steps 200 | |
| # 用法:bash scripts/train_baseline.sh <arm> [额外参数...] | |
| # RUN=xxx 改输出名(默认 <arm>_40k);GA=1 用 accelerate_zero2.yaml;ARM_KWARGS='{"enable_mouse": false, "window_frames": 1}' | |
| # 给 xattn 复刻旧报告的 98.6M 口径;崩溃续训:bash scripts/train_baseline.sh xattn --resume outputs/xattn_40k/state | |
| set -euo pipefail | |
| ROOT=/opt/dlami/nvme/zhiyangdeng/ActionRoPE | |
| export DIFFSYNTH_SKIP_DOWNLOAD=True | |
| export PYTHONPATH="$ROOT" | |
| export CUDA_VISIBLE_DEVICES="${CUDA_VISIBLE_DEVICES:-0,1,2,3,4,5,6,7}" | |
| ARM="${1:?用法: bash scripts/train_baseline.sh <linear|xattn|prompt|adaln|plain> [额外参数...]}"; shift | |
| case "$ARM" in linear|xattn|prompt|adaln|plain) ;; *) echo "未知臂 $ARM(arope 请用 scripts/train_arope.sh)" >&2; exit 2;; esac | |
| GA="${GA:-2}" # 梯度累积(全局 batch = 8 卡 × GA);yaml 随之选 accelerate_zero2_ga${GA}.yaml(GA=1 用 accelerate_zero2.yaml) | |
| CFG=$([ "$GA" = 1 ] && echo accelerate_zero2.yaml || echo "accelerate_zero2_ga${GA}.yaml") | |
| RUN="${RUN:-${ARM}_40k}" | |
| ARM_KWARGS="${ARM_KWARGS:-}" | |
| KW_ARG=() | |
| [ -n "$ARM_KWARGS" ] && KW_ARG=(--arm_kwargs "$ARM_KWARGS") | |
| mkdir -p "$ROOT/outputs/$RUN" | |
| "$ROOT/.venv/bin/accelerate" launch --config_file "$ROOT/configs/$CFG" \ | |
| "$ROOT/actionrope/train.py" \ | |
| --arm "$ARM" "${KW_ARG[@]}" --scene_dropout 0.1 \ | |
| --lr 1e-5 --weight_decay 0.01 --warmup_steps 500 \ | |
| --max_steps 40000 --save_every 1000 --save_state --val_every 500 --val_n 16 --val_timesteps 200,500,800 \ | |
| --min_rt_ratio 0.8 --val_min_rt_ratio 0.9 \ | |
| --grad_accum "$GA" --num_workers 2 --seed 0 \ | |
| --output "$ROOT/outputs/$RUN" "$@" 2>&1 | tee -a "$ROOT/outputs/$RUN/train.log" | |