| #!/bin/bash |
| set -eo pipefail |
|
|
| |
| CONDA_ENV=${CONDA_ENV:-/gpfs/wangzixuan/conda_envs/starVLA} |
| |
| source "$(conda info --base)/etc/profile.d/conda.sh" |
| conda activate "${CONDA_ENV}" |
|
|
| export NCCL_SOCKET_IFNAME=${NCCL_SOCKET_IFNAME:-bond0} |
| export NCCL_IB_HCA=${NCCL_IB_HCA:-mlx5_2,mlx5_3} |
| export NCCL_BLOCKING_WAIT=${NCCL_BLOCKING_WAIT:-1} |
| export NCCL_ASYNC_ERROR_HANDLING=${NCCL_ASYNC_ERROR_HANDLING:-1} |
| export NCCL_TIMEOUT=${NCCL_TIMEOUT:-10000} |
|
|
| SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) |
| REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../../.." && pwd) |
| cd "${REPO_ROOT}" |
| export PYTHONPATH="${REPO_ROOT}:${PYTHONPATH:-}" |
|
|
| |
| config_yaml=${config_yaml:-"${REPO_ROOT}/examples/Franka/train_files/starvla_cotrain_franka_v4.yaml"} |
| franka_data_root=${franka_data_root:-/gpfs/wangzixuan/Octopus/workspace/data/lerobot/v4} |
| data_mix=${data_mix:-real_franka_v4} |
| |
| pretrained_checkpoint=${pretrained_checkpoint:-} |
| base_vlm=${base_vlm:-/gpfs/wangzixuan/Octopus/workspace/MODEL/Qwen3-VL-4B-Instruct} |
| |
|
|
| run_root_dir=${run_root_dir:-"${REPO_ROOT}/results/Checkpoints"} |
| run_id=${run_id:-0622_franka_v4_qwen3OFT_sandwich} |
| num_processes=${num_processes:-8} |
| per_device_batch_size=${per_device_batch_size:-16} |
| max_train_steps=${max_train_steps:-100000} |
| save_interval=${save_interval:-10000} |
| logging_frequency=${logging_frequency:-100} |
| eval_interval=${eval_interval:-1000} |
|
|
| if [[ ! -f "${config_yaml}" ]]; then |
| echo "Config yaml not found: ${config_yaml}" >&2 |
| exit 1 |
| fi |
|
|
| if [[ -n "${pretrained_checkpoint}" && ! -f "${pretrained_checkpoint}" ]]; then |
| echo "Pretrained checkpoint not found: ${pretrained_checkpoint}" >&2 |
| exit 1 |
| fi |
|
|
| output_dir="${run_root_dir}/${run_id}" |
| mkdir -p "${output_dir}" |
| cp "$0" "${output_dir}/" |
|
|
| echo "============================================" |
| echo "Conda env : ${CONDA_ENV}" |
| echo "Config yaml : ${config_yaml}" |
| echo "Base VLM : ${base_vlm}" |
| echo "Pretrained checkpoint : ${pretrained_checkpoint:-<none> (Qwen3-VL base only)}" |
| echo "Franka data root : ${franka_data_root}" |
| echo "Data mix : ${data_mix}" |
| echo "Batch size / process : ${per_device_batch_size}" |
| echo "Num processes : ${num_processes}" |
| echo "Max train steps : ${max_train_steps}" |
| echo "Run ID : ${run_id}" |
| echo "============================================" |
|
|
| launch_args=( |
| --config_file "${REPO_ROOT}/starVLA/config/deepseeds/deepspeed_zero2.yaml" |
| --num_processes "${num_processes}" |
| "${REPO_ROOT}/starVLA/training/train_starvla.py" |
| --config_yaml "${config_yaml}" |
| --framework.name QwenOFT |
| --framework.qwenvl.base_vlm "${base_vlm}" |
| --datasets.vla_data.data_root_dir "${franka_data_root}" |
| --datasets.vla_data.data_mix "${data_mix}" |
| --datasets.vla_data.per_device_batch_size "${per_device_batch_size}" |
| --datasets.vla_data.sequential_step_sampling false |
| --trainer.max_train_steps "${max_train_steps}" |
| --trainer.save_interval "${save_interval}" |
| --trainer.logging_frequency "${logging_frequency}" |
| --trainer.eval_interval "${eval_interval}" |
| |
| --run_root_dir "${run_root_dir}" |
| --run_id "${run_id}" |
| --wandb_project roboclaw |
| --wandb_entity zwanggk |
| ) |
|
|
| if [[ -n "${pretrained_checkpoint}" ]]; then |
| launch_args+=(--trainer.pretrained_checkpoint "${pretrained_checkpoint}") |
| fi |
|
|
| accelerate launch "${launch_args[@]}" |
|
|