vla-sft-suite / scripts /bootstrap_cluster.sh
poet70's picture
Update cluster delivery docs and verified 6x20 workflow
5ff102d verified
Raw
History Blame Contribute Delete
15.7 kB
#!/usr/bin/env bash
set -euo pipefail
SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
PROJECT_ROOT=${PROJECT_ROOT:-/data/vla_sft}
DATA_ROOT=${DATA_ROOT:-$PROJECT_ROOT/data}
CKPT_ROOT=${CKPT_ROOT:-$PROJECT_ROOT/checkpoints}
CODE_ROOT=${CODE_ROOT:-$PROJECT_ROOT/baselines}
VENV_ROOT=${VENV_ROOT:-$PROJECT_ROOT/envs}
OUTPUT_ROOT=${OUTPUT_ROOT:-$PROJECT_ROOT/outputs}
SMOKE_STEPS=${SMOKE_STEPS:-10}
FULL_STEPS=${FULL_STEPS:-200000}
GPUS_PER_NODE=${GPUS_PER_NODE:-8}
DREAMZERO_OPTIM=${DREAMZERO_OPTIM:-adamw_bnb_8bit}
DREAMZERO_DATALOADER_NUM_WORKERS=${DREAMZERO_DATALOADER_NUM_WORKERS:-0}
DREAMZERO_SAVE_STEPS=${DREAMZERO_SAVE_STEPS:-10000}
DREAMZERO_SAVE_TOTAL_LIMIT=${DREAMZERO_SAVE_TOTAL_LIMIT:-2}
DREAMZERO_SKIP_FINAL_SAVE=${DREAMZERO_SKIP_FINAL_SAVE:-}
MOTUS_SKIP_FINAL_SAVE=${MOTUS_SKIP_FINAL_SAVE:-}
DRY_RUN=0
MODE=smoke
PROFILE=""
BENCH=""
TMUX_SESSION=""
LOG_FILE=""
CHILD_ARGS=()
while [[ $# -gt 0 ]]; do
case "$1" in
--profile|--baseline)
PROFILE="$2"; CHILD_ARGS+=("--profile" "$2"); shift 2 ;;
--bench|--benchmark)
BENCH="$2"; CHILD_ARGS+=("--bench" "$2"); shift 2 ;;
--smoke)
MODE=smoke; CHILD_ARGS+=("--smoke"); shift ;;
--full)
MODE=full; CHILD_ARGS+=("--full"); shift ;;
--dry-run)
DRY_RUN=1; CHILD_ARGS+=("--dry-run"); shift ;;
--tmux)
TMUX_SESSION="$2"; shift 2 ;;
--log)
LOG_FILE="$2"; CHILD_ARGS+=("--log" "$2"); shift 2 ;;
*)
echo "unknown arg: $1" >&2; exit 2 ;;
esac
done
if [[ -z "$PROFILE" || -z "$BENCH" ]]; then
echo "usage: $0 --profile dreamzero|motus|dreamtacvla --bench manifeel|univtac [--smoke|--full] [--dry-run] [--tmux SESSION] [--log FILE]" >&2
echo "note: libero/robotwin are non-core legacy benchmarks and are only kept for explicit compatibility." >&2
exit 2
fi
run() {
echo "+ $*"
if [[ "$DRY_RUN" -eq 0 ]]; then
"$@"
fi
}
steps="$SMOKE_STEPS"
if [[ "$MODE" == "full" ]]; then
steps="$FULL_STEPS"
fi
patch_dreamzero_verify_runtime() {
[[ "$PROFILE" == "dreamzero" && "$DRY_RUN" -eq 0 ]] || return 0
local dreamzero_root="$CODE_ROOT/dreamzero"
[[ -d "$dreamzero_root" ]] || return 0
"${PYTHON_BIN:-python3}" - "$dreamzero_root" <<'PYFIX'
from pathlib import Path
import sys
root = Path(sys.argv[1])
for rel in ("groot/vla/data/dataset/manifeel.py", "groot/vla/data/dataset/univtac.py"):
path = root / rel
if not path.is_file():
continue
text = path.read_text()
old = ''' frame = root["data/wrist"][fi]
frame = (frame * 255).clip(0, 255).astype(np.uint8)'''
new = ''' frame = root["data/wrist"][fi]
if frame.dtype == np.uint8:
frame = frame.astype(np.uint8, copy=False)
else:
frame = (frame * 255).clip(0, 255).astype(np.uint8)'''
if old in text:
path.write_text(text.replace(old, new))
print(f"[patch] DreamZero frame dtype updated: {path}")
base_py = root / "groot/vla/experiment/base.py"
if base_py.is_file():
text = base_py.read_text()
old = ''' self.trainer.save_state()
safe_save_model_for_hf_trainer(
trainer=self.trainer,
output_dir=self.training_args.output_dir,
)
'''
new = ''' if __import__("os").environ.get("DREAMZERO_SKIP_FINAL_SAVE", "0") == "1":
print("[verify] DREAMZERO_SKIP_FINAL_SAVE=1, skip final save_state/save_model", flush=True)
return
self.trainer.save_state()
safe_save_model_for_hf_trainer(
trainer=self.trainer,
output_dir=self.training_args.output_dir,
)
'''
if "DREAMZERO_SKIP_FINAL_SAVE" not in text and old in text:
if "import os" not in text.splitlines()[:40]:
text = text.replace("import json\n", "import json\nimport os\n", 1)
base_py.write_text(text.replace(old, new))
print(f"[patch] DreamZero final save skip updated: {base_py}")
PYFIX
}
patch_motus_verify_runtime() {
[[ "$PROFILE" == "motus" && "$DRY_RUN" -eq 0 ]] || return 0
local motus_root="$CODE_ROOT/motus"
local dataset_py="$motus_root/data/manifeel/manifeel_zarr_dataset.py"
local train_py="$motus_root/train/train.py"
local cfg_dir="$motus_root/configs"
[[ -f "$dataset_py" && -f "$train_py" && -d "$cfg_dir" ]] || return 0
"${PYTHON_BIN:-python3}" - "$dataset_py" "$train_py" "$cfg_dir/manifeel_sft.yaml" "$cfg_dir/univtac_sft.yaml" <<'PYFIX'
from pathlib import Path
import sys
dataset_py = Path(sys.argv[1])
train_py = Path(sys.argv[2])
manifeel_cfg = Path(sys.argv[3])
univtac_cfg = Path(sys.argv[4])
text = dataset_py.read_text()
old = ''' frame = root["data/wrist"][fi]
frame = (frame * 255).clip(0, 255).astype(np.uint8)'''
new = ''' frame = root["data/wrist"][fi]
if frame.dtype == np.uint8:
frame = frame.astype(np.uint8, copy=False)
else:
frame = (frame * 255).clip(0, 255).astype(np.uint8)'''
if old in text:
text = text.replace(old, new)
dataset_py.write_text(text)
print(f"[patch] Motus Zarr frame dtype updated: {dataset_py}")
text = dataset_py.read_text()
old = ''' action_seq = np.zeros((self.action_chunk_size, self.action_dim), dtype=np.float32)
for i in range(self.action_chunk_size):
raw_act = root["data/action"][min(abs_t + i, len(root["data/action"]) - 1)]
action_seq[i, : len(raw_act)] = raw_act
frames = []
w, h = self.video_size
for i in range(self.num_video_frames):
fi = min(abs_t + i, len(root["data/wrist"]) - 1)
frame = root["data/wrist"][fi]
'''
new = ''' action_arr = root["data/action"]
action_seq = np.zeros((self.action_chunk_size, self.action_dim), dtype=np.float32)
for i in range(self.action_chunk_size):
raw_act = action_arr[min(abs_t + i, action_arr.shape[0] - 1)]
action_seq[i, : len(raw_act)] = raw_act
wrist_arr = root["data/wrist"]
frames = []
w, h = self.video_size
for i in range(self.num_video_frames):
fi = min(abs_t + i, wrist_arr.shape[0] - 1)
frame = wrist_arr[fi]
'''
if old in text:
dataset_py.write_text(text.replace(old, new))
print(f"[patch] Motus Zarr length handling updated: {dataset_py}")
elif 'action_arr = root["data/action"]' in text and 'wrist_arr = root["data/wrist"]' in text:
print(f"[patch] Motus Zarr length handling already patched: {dataset_py}")
text = train_py.read_text()
old = ''' if self.rank == 0:
logger.info(f"UniDiffuser training completed in {total_time:.2f}s ({self.global_step} steps)")
self.save_checkpoint()
'''
new = ''' if self.rank == 0:
logger.info(f"UniDiffuser training completed in {total_time:.2f}s ({self.global_step} steps)")
if __import__("os").environ.get("MOTUS_SKIP_FINAL_SAVE", "0") == "1":
logger.info("[verify] MOTUS_SKIP_FINAL_SAVE=1, skip final checkpoint save")
else:
self.save_checkpoint()
'''
if old in text:
if "import os" not in text.splitlines()[:40]:
text = text.replace("import logging\n", "import logging\nimport os\n", 1)
train_py.write_text(text.replace(old, new))
print(f"[patch] Motus final save skip updated: {train_py}")
elif "MOTUS_SKIP_FINAL_SAVE" in text:
print(f"[patch] Motus final save skip already patched: {train_py}")
if manifeel_cfg.is_file():
cfg = manifeel_cfg.read_text()
cfg = cfg.replace("dataset_dir: /root/autodl-tmp/tmp/manifeel_extracted", "dataset_dir: /root/autodl-tmp/vla_sft/data/univtac_zarr")
cfg = cfg.replace("wandb_project: motus_manifeel", "wandb_project: motus_univtac")
univtac_cfg.write_text(cfg)
print(f"[patch] Motus UniVTac config written: {univtac_cfg}")
PYFIX
}
run mkdir -p "$PROJECT_ROOT" "$DATA_ROOT" "$CKPT_ROOT" "$CODE_ROOT" "$VENV_ROOT" "$OUTPUT_ROOT"
if [[ -n "$TMUX_SESSION" ]]; then
command -v tmux >/dev/null 2>&1 || {
echo "tmux not found. Install tmux, or run the same command without --tmux." >&2
exit 4
}
if tmux has-session -t "$TMUX_SESSION" 2>/dev/null; then
echo "tmux session already exists: $TMUX_SESSION" >&2
echo "attach with: tmux attach -t $TMUX_SESSION" >&2
exit 4
fi
if [[ -z "$LOG_FILE" ]]; then
LOG_FILE="$OUTPUT_ROOT/deploy_${PROFILE}_${BENCH}_${MODE}.log"
fi
mkdir -p "$(dirname "$LOG_FILE")"
tmux_cmd=(
env
"PROJECT_ROOT=$PROJECT_ROOT"
"DATA_ROOT=$DATA_ROOT"
"CKPT_ROOT=$CKPT_ROOT"
"CODE_ROOT=$CODE_ROOT"
"VENV_ROOT=$VENV_ROOT"
"OUTPUT_ROOT=$OUTPUT_ROOT"
"SMOKE_STEPS=$SMOKE_STEPS"
"FULL_STEPS=$FULL_STEPS"
"GPUS_PER_NODE=$GPUS_PER_NODE"
"DREAMZERO_OPTIM=$DREAMZERO_OPTIM"
"DREAMZERO_DATALOADER_NUM_WORKERS=$DREAMZERO_DATALOADER_NUM_WORKERS"
"DREAMZERO_SAVE_STEPS=$DREAMZERO_SAVE_STEPS"
"DREAMZERO_SAVE_TOTAL_LIMIT=$DREAMZERO_SAVE_TOTAL_LIMIT"
"DREAMZERO_SKIP_FINAL_SAVE=${DREAMZERO_SKIP_FINAL_SAVE:-}"
"MOTUS_SKIP_FINAL_SAVE=${MOTUS_SKIP_FINAL_SAVE:-}"
"HF_HUB_DISABLE_XET=${HF_HUB_DISABLE_XET:-1}"
"HF_ARCHIVE_DOWNLOAD_BACKEND=${HF_ARCHIVE_DOWNLOAD_BACKEND:-hf}"
"HF_ARCHIVE_DELETE_AFTER_EXTRACT=${HF_ARCHIVE_DELETE_AFTER_EXTRACT:-0}"
"ARTIFACT_FORCE_DOWNLOAD=${ARTIFACT_FORCE_DOWNLOAD:-0}"
"USE_NETWORK_TURBO=${USE_NETWORK_TURBO:-auto}"
"HF_ENDPOINT=${HF_ENDPOINT:-}"
"PIP_INDEX_URL=${PIP_INDEX_URL-https://pypi.tuna.tsinghua.edu.cn/simple}"
"PIP_TRUSTED_HOST=${PIP_TRUSTED_HOST-pypi.tuna.tsinghua.edu.cn}"
"PIP_FALLBACK_INDEX_URL=${PIP_FALLBACK_INDEX_URL-https://pypi.org/simple}"
"PIP_DISABLE_FALLBACK=${PIP_DISABLE_FALLBACK:-0}"
"TORCH_INDEX_URL=${TORCH_INDEX_URL-}"
"TORCH_FALLBACK_INDEX_URL=${TORCH_FALLBACK_INDEX_URL-}"
"PYTHON_BIN=${PYTHON_BIN:-python3}"
"DREAMTACVLA_PYTHON_BIN=${DREAMTACVLA_PYTHON_BIN:-}"
bash "$0" "${CHILD_ARGS[@]}"
)
quoted_cmd=$(printf ' %q' "${tmux_cmd[@]}")
quoted_log=$(printf '%q' "$LOG_FILE")
child_script="set -euo pipefail;${quoted_cmd} 2>&1 | tee -a ${quoted_log}"
tmux new -d -s "$TMUX_SESSION" "bash -lc $(printf '%q' "$child_script")"
echo "[tmux] started: $TMUX_SESSION"
echo "[tmux] log: $LOG_FILE"
echo "[tmux] attach: tmux attach -t $TMUX_SESSION"
exit 0
fi
bash "$SCRIPT_DIR/check_system.sh"
download_args=(
--baseline "$PROFILE" \
--bench "$BENCH" \
--data-root "$DATA_ROOT" \
--ckpt-root "$CKPT_ROOT" \
--code-root "$CODE_ROOT"
)
if [[ "$DRY_RUN" -eq 1 ]]; then
download_args+=(--dry-run)
fi
bash "$SCRIPT_DIR/download_artifacts.sh" "${download_args[@]}"
patch_dreamzero_verify_runtime
patch_motus_verify_runtime
install_args=(
--profile "$PROFILE" \
--code-root "$CODE_ROOT" \
--venv-root "$VENV_ROOT"
)
if [[ "$DRY_RUN" -eq 1 ]]; then
install_args+=(--dry-run)
fi
bash "$SCRIPT_DIR/install_env.sh" "${install_args[@]}"
VENV_DIR="$VENV_ROOT/$PROFILE"
OUT_DIR="$OUTPUT_ROOT/${PROFILE}_${BENCH}_${MODE}"
case "$PROFILE:$BENCH" in
dreamtacvla:univtac)
TRAIN_CMD=(
python "$CODE_ROOT/dreamtacvla/ModelTrain/model_train.py"
--ckpt_dir "$OUT_DIR"
--policy_class ACTJEPAAdapter
--task_name univtac_all
--batch_size 8
--seed 42
--num_steps "$steps"
--lr 1e-5
--save_every 100
--enable_hsa
--freeze_clip
--vit_ckpt_path "$CKPT_ROOT/dreamtacvla_jepa/vitl_manifeel_base.pt"
)
PRELUDE=("cd" "$CODE_ROOT/dreamtacvla")
;;
dreamtacvla:manifeel)
TRAIN_CMD=(
python "$CODE_ROOT/dreamtacvla/ModelTrain/model_train.py"
--ckpt_dir "$OUT_DIR"
--policy_class ACTJEPAAdapter
--task_name manifeel_all
--batch_size 8
--seed 42
--num_steps "$steps"
--lr 1e-5
--save_every 100
--enable_hsa
--freeze_clip
--vit_ckpt_path "$CKPT_ROOT/dreamtacvla_jepa/vitl_manifeel_base.pt"
)
PRELUDE=("cd" "$CODE_ROOT/dreamtacvla")
;;
dreamzero:libero|dreamzero:manifeel|dreamzero:robotwin|dreamzero:univtac)
TRAIN_CMD=(
torchrun --standalone --nproc_per_node="$GPUS_PER_NODE"
"$CODE_ROOT/dreamzero/groot/vla/experiment/experiment.py"
report_to=none
"data=dreamzero/$BENCH"
model=dreamzero/vla
model/dreamzero/action_head=wan_flow_matching_action_tf_wan22
model/dreamzero/transform=dreamzero_cotrain
train_architecture=full
"max_steps=$steps"
"optim=$DREAMZERO_OPTIM"
"dataloader_num_workers=$DREAMZERO_DATALOADER_NUM_WORKERS"
dataloader_persistent_workers=false
"save_steps=$DREAMZERO_SAVE_STEPS"
"save_total_limit=$DREAMZERO_SAVE_TOTAL_LIMIT"
"training_args.save_steps=$DREAMZERO_SAVE_STEPS"
"training_args.save_total_limit=$DREAMZERO_SAVE_TOTAL_LIMIT"
training_args.dataloader_persistent_workers=false
training_args.deepspeed=groot/vla/configs/deepspeed/zero2_offload.json
per_device_train_batch_size=1
training_args.learning_rate=1e-5
training_args.bf16=true
training_args.tf32=true
wandb_project=vla_sft_verify
"output_dir=$OUT_DIR"
"dit_version=$CKPT_ROOT/Wan2.2-TI2V-5B"
"text_encoder_pretrained_path=$CKPT_ROOT/Wan2.2-TI2V-5B/models_t5_umt5-xxl-enc-bf16.pth"
"image_encoder_pretrained_path=$CKPT_ROOT/clip-encoder/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"
"vae_pretrained_path=$CKPT_ROOT/Wan2.2-TI2V-5B/Wan2.2_VAE.pth"
"tokenizer_path=$CKPT_ROOT/umt5-xxl"
"+train_dataset.tokenizer_path=$CKPT_ROOT/umt5-xxl"
)
if [[ "$steps" -le 20 ]]; then
TRAIN_CMD+=(
save_strategy=no
training_args.save_strategy=no
)
export DREAMZERO_SKIP_FINAL_SAVE=${DREAMZERO_SKIP_FINAL_SAVE:-1}
fi
if [[ "$BENCH" == "manifeel" ]]; then
TRAIN_CMD+=("manifeel_dataset_dir=$DATA_ROOT/manifeel_extracted")
fi
if [[ "$BENCH" == "univtac" ]]; then
TRAIN_CMD+=("univtac_dataset_dir=$DATA_ROOT/univtac_zarr")
fi
PRELUDE=("cd" "$CODE_ROOT/dreamzero")
;;
motus:libero|motus:manifeel|motus:robotwin|motus:univtac)
cfg="$BENCH"
if [[ "$BENCH" == "manifeel" ]]; then cfg="manifeel_sft"; fi
if [[ "$BENCH" == "robotwin" ]]; then cfg="robotwin_sft"; fi
if [[ "$BENCH" == "univtac" ]]; then cfg="univtac_sft"; fi
MOTUS_RENDERED_CFG="$OUTPUT_ROOT/rendered_configs/motus_${BENCH}_${MODE}.yaml"
run python3 "$SCRIPT_DIR/render_motus_config.py" \
--source "$CODE_ROOT/motus/configs/$cfg.yaml" \
--output "$MOTUS_RENDERED_CFG" \
--bench "$BENCH" \
--data-root "$DATA_ROOT" \
--ckpt-root "$CKPT_ROOT" \
--output-root "$OUTPUT_ROOT" \
--steps "$steps"
TRAIN_CMD=(
torchrun --standalone --nproc_per_node="$GPUS_PER_NODE"
"$CODE_ROOT/motus/train/train.py"
--config "$MOTUS_RENDERED_CFG"
--deepspeed "$CODE_ROOT/motus/configs/zero2.json"
)
if [[ "$steps" -le 20 ]]; then
export MOTUS_SKIP_FINAL_SAVE=${MOTUS_SKIP_FINAL_SAVE:-1}
fi
PRELUDE=("cd" "$CODE_ROOT/motus")
;;
*)
echo "unsupported experiment: $PROFILE + $BENCH" >&2
exit 3
;;
esac
echo "[train] mode: $MODE"
echo "[train] output: $OUT_DIR"
echo "[train] command:"
printf ' %q' "${TRAIN_CMD[@]}"
echo
if [[ "$DRY_RUN" -eq 1 ]]; then
echo "[dry-run] skip training"
exit 0
fi
# shellcheck disable=SC1091
source "$VENV_DIR/bin/activate"
export DOBOT_DATA_DIR="$DATA_ROOT"
export CHECKPOINT_DIR="$CKPT_ROOT"
export DREAMZERO_DATA_ROOT="$DATA_ROOT"
export MOTUS_DATA_ROOT="$DATA_ROOT"
export MOTUS_CKPT_ROOT="$CKPT_ROOT"
export OUTPUT_ROOT="$OUTPUT_ROOT"
run mkdir -p "$OUT_DIR"
"${PRELUDE[@]}"
"${TRAIN_CMD[@]}"