0622_franka_v4_qwen3OFT_sandwich / run_franka_v4_finetune_qwenoft.sh
Vincent2311's picture
Add files using upload-large-folder tool
888a910 verified
Raw
History Blame Contribute Delete
3.62 kB
#!/bin/bash
set -eo pipefail
# Activate the starVLA conda environment.
CONDA_ENV=${CONDA_ENV:-/gpfs/wangzixuan/conda_envs/starVLA}
# shellcheck disable=SC1091
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:-}"
# v4 single-arm Franka datasets (make_sandwich + move_place_gaizi, co-trained).
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}
# Leave empty to train from Qwen3-VL base only (action head from scratch).
pretrained_checkpoint=${pretrained_checkpoint:-}
base_vlm=${base_vlm:-/gpfs/wangzixuan/Octopus/workspace/MODEL/Qwen3-VL-4B-Instruct}
# reload_modules=${reload_modules:-'qwen_vl_interface'}
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}"
# --trainer.reload_modules "${reload_modules}"
--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[@]}"