#!/usr/bin/env bash
# SmolVLA 微调 · 优化配方与硬件档位分离
#
# 为什么有这个脚本(2026-08-24):
#   旧的 train_smolvla.sh 把 BATCH 当成一个普通环境变量,于是「换了张大卡就顺手把 batch
#   调大」会静默改掉学习问题本身 —— A100 那次 batch=64 / 只跑 6000 步就是这么来的:
#   命令里写 --steps=20000,调度器把整条 cosine 摊在 20000 步上,停在 6000 时学习率
#   还有 7.99e-05(峰值的 80%)。分析见 SMOLVLA-TRAINING-RECIPE.md。
#
# 这里把参数分成两组,并且【只允许改第二组】来适配硬件:
#   ① 优化配方 —— 决定学到什么:  B_EFF / LR / UPDATES / WARMUP
#   ② 硬件档位 —— 只决定快慢和显存: MICRO_BATCH / WORKERS
#   ACCUM 由 ① 和 ② 推出来,不手填。
#
# 用法:
#   ./train_smolvla_v2.sh                          # A100 默认档
#   MICRO_BATCH=16 ./train_smolvla_v2.sh           # 4090:自动 ACCUM=4,B_EFF 仍是 64
#   B_EFF=8 MICRO_BATCH=8 ./train_smolvla_v2.sh    # 想换配方就显式换
#   DRY=1 ./train_smolvla_v2.sh                    # 只打印,不跑
set -euo pipefail

# ── ① 优化配方(改这些 = 改学习问题,要写进实验记录)────────────────────────────
B_EFF="${B_EFF:-64}"          # effective batch = MICRO_BATCH × ACCUM × WORLD
LR="${LR:-1e-4}"              # LeRobot 为 SmolVLA 调好的默认值,配 B_EFF=64
UPDATES="${UPDATES:-20000}"   # 优化器真正更新的次数 U
WARMUP="${WARMUP:-1000}"      # 以【更新次数】计
# eval_split:默认 0.0 = 全部拿去训。见下面的守卫,以及 RECIPE §〇·七。
# 【不是随机 20%】LeRobot 留出的是「每个 task 的【最后】ceil(N×split) 集」
# (datasets/factory.py:156)。本数据集只有 1 个 task、50 集,所以 0.2 = 砍掉 ep40-49,
# 那正好是整个第五个杯位。
EVAL_SPLIT="${EVAL_SPLIT:-0.0}"
EVAL_STEPS="${EVAL_STEPS:-0}"     # 0 = 不做评估。>0 才会真的跑 eval
SEED="${SEED:-1000}"
AUG="${AUG:-0}"               # 1 = 开颜色增强(affine 恒为 0,见下)

# ── ② 硬件档位(只影响快慢/显存,不影响学到什么)──────────────────────────────
MICRO_BATCH="${MICRO_BATCH:-64}"
WORKERS="${WORKERS:-8}"
WORLD="${WORLD:-1}"

# ── 固定项 ────────────────────────────────────────────────────────────────
REPO_ID="${REPO_ID:-Suyang99/xlerobot-cup-grasp-20260820-0230}"
BASE="${BASE:-lerobot/smolvla_base}"
DEVICE="${DEVICE:-cuda}"
PY="${PY:-python}"
RENAME='{"observation.images.head":"observation.images.camera1","observation.images.right_arm_wrist":"observation.images.camera2","observation.images.left_arm_wrist":"observation.images.camera3"}'

# ── 守卫:不许悄悄浪费数据 ──────────────────────────────────────────────────
# 这正是 2026-08-20 那次踩的坑:eval_split=0.2 砍掉了 10 集,而 eval_steps=0 意味着
# 评估从来没跑过(lerobot_train.py:748 是 `cfg.eval_steps > 0 and ...`,
# 训练日志里 eval loss 行数为 0)。20% 的数据 + 一整个杯位,换来零信息。
if [[ "$EVAL_SPLIT" != "0.0" && "$EVAL_SPLIT" != "0" && "$EVAL_STEPS" == "0" ]]; then
  cat >&2 <<WARN
────────────────────────────────────────────────────────────────
⚠ EVAL_SPLIT=$EVAL_SPLIT 会把【每个 task 的最后 ceil(N×$EVAL_SPLIT) 集】从训练里拿掉,
  而 EVAL_STEPS=0 意味着【评估根本不会运行】—— 这些数据被丢掉且不产生任何信息。
  对 xlerobot-cup-grasp-20260820-0230(50 集 / 1 个 task):砍掉 ep40-49,
  也就是整个第五个杯位。

  想要什么就选什么:
    · 要最好的模型            -> EVAL_SPLIT=0.0        (默认)
    · 要真的有验证曲线         -> EVAL_SPLIT=0.2 EVAL_STEPS=1000
    · 要复刻 2026-08-20 那次   -> EVAL_SPLIT=0.2 EVAL_STEPS=0 ACK_WASTE=1
────────────────────────────────────────────────────────────────
WARN
  if [[ "${ACK_WASTE:-0}" != "1" ]]; then
    echo "已停止。确实要复刻历史设置就加 ACK_WASTE=1。" >&2; exit 1
  fi
  echo "(ACK_WASTE=1,明知故犯,继续)" >&2
fi

# ── 推导 ──────────────────────────────────────────────────────────────────
if (( B_EFF % (MICRO_BATCH * WORLD) != 0 )); then
  echo "错误: B_EFF=$B_EFF 不能被 MICRO_BATCH×WORLD=$((MICRO_BATCH*WORLD)) 整除。" >&2
  echo "      请把 MICRO_BATCH 调成 B_EFF 的因数,或显式改 B_EFF(那是在改配方)。" >&2
  exit 1
fi
ACCUM=$(( B_EFF / (MICRO_BATCH * WORLD) ))
STEPS=$(( UPDATES * ACCUM ))          # --steps 数的是 micro-batch,不是更新次数
SAVE_FREQ=$(( STEPS / 10 ))           # 10 个存档 -> S = 0.1 .. 1.0

# 关键:调度器长度用 cfg.steps(micro-batch 数)建,但只在更新时前进。
# 用了累积就必须显式把 decay 设成 UPDATES,否则曲线走不完:
# 自动缩放只会把曲线改短不会改长(schedulers.py:149),所以
#   曲线长度 = min(STEPS, policy 的 scheduler_decay_steps=30000)
#   S        = UPDATES / 曲线长度
# 例:micro16 accum4 U=20000 -> STEPS=80000,曲线长仍是 30000,S=0.667,终点 lr 2.69e-05。
# 显式设成 20000 之后:曲线长 20000 = 实际更新次数,S=1.000,终点 lr 2.50e-06。
# 【始终】显式传,不依赖自动缩放。原因:自动缩放只在 steps < decay_steps 时触发,
# 触发时还会把 warmup 一起按比例缩小(1000 -> 666)。显式传之后,ACCUM 是 1 还是 4,
# 曲线长度和 warmup 都完全一样,打印出来的就是真正生效的值。
#
# 注意:两次历史训练(batch2 和 batch64)因为走的是自动缩放,实际 warmup 是 666 不是 1000。
# 要严格复刻它们,设 WARMUP=666。
SCHED_ARGS=( --policy.scheduler_decay_steps="$UPDATES" --policy.scheduler_warmup_steps="$WARMUP" )

AUG_ARGS=()
if [[ "$AUG" == "1" ]]; then
  # 只开颜色类。affine 的随机平移 ±0.05 与相邻杯位在画面里的间距 0.06 同量级,
  # 会把唯一的空间线索抹掉 —— 必须置 0。见 RECIPE §8.1。
  AUG_ARGS+=( --dataset.image_transforms.enable=true
              --dataset.image_transforms.tfs.affine.weight=0 )
fi

OUT="${OUT:-$(dirname "$0")/runs/smolvla_b${B_EFF}_u${UPDATES}_$(date +%Y%m%d-%H%M)}"
JOB="${JOB:-smolvla_b${B_EFF}_u${UPDATES}}"

cat <<EOF
──────────────── 优化配方(改这些=改学习问题)────────────────
  effective batch B_eff   : $B_EFF
  learning rate           : $LR
  optimizer updates U     : $UPDATES
  warmup (更新次数)        : $WARMUP
  eval_split / eval_steps : $EVAL_SPLIT / $EVAL_STEPS$( [[ "$EVAL_SPLIT" == "0.0" || "$EVAL_SPLIT" == "0" ]] && echo "   (全部 50 集参与训练)" || echo "   (⚠ 砍掉每 task 最后 ceil(N×$EVAL_SPLIT) 集)" )
  seed                    : $SEED
  颜色增强                 : $([[ "$AUG" == "1" ]] && echo 开 || echo 关)   (affine 恒关)
──────────────── 硬件档位(只影响快慢/显存)──────────────────
  micro_batch             : $MICRO_BATCH
  accumulation (推导)      : $ACCUM
  world_size              : $WORLD
  workers                 : $WORKERS
──────────────── 传给 lerobot 的值 ────────────────────────
  --steps                 : $STEPS   (= U × ACCUM)
  --save_freq             : $SAVE_FREQ   (10 个存档, S = 0.1..1.0)
  scheduler decay horizon : $UPDATES (显式传,不走自动缩放)
  预期 S 终点              : 1.0
EOF

if [[ "${DRY:-0}" == "1" ]]; then echo "(DRY=1,不执行)"; exit 0; fi

exec "$PY" -m lerobot.scripts.lerobot_train \
  --dataset.repo_id="$REPO_ID" \
  --policy.path="$BASE" \
  --policy.device="$DEVICE" \
  --policy.push_to_hub=false \
  --policy.optimizer_lr="$LR" \
  --rename_map="$RENAME" \
  --dataset.eval_split="$EVAL_SPLIT" \
  --eval_steps="$EVAL_STEPS" \
  --batch_size="$MICRO_BATCH" \
  --accelerator.gradient_accumulation.steps="$ACCUM" \
  --num_workers="$WORKERS" \
  --steps="$STEPS" \
  --save_freq="$SAVE_FREQ" \
  --log_freq=100 \
  --seed="$SEED" \
  "${SCHED_ARGS[@]}" "${AUG_ARGS[@]}" \
  --output_dir="$OUT" \
  --job_name="$JOB" "$@"
