#!/bin/bash # Visual Prompt Training Script for SimplerEnv (OXE datasets) # export CUDA_VISIBLE_DEVICES=2,3,4,5,6,7 export NCCL_SOCKET_IFNAME=bond0 export NCCL_IB_HCA=mlx5_2,mlx5_3 # used for check save when communication export NCCL_BLOCKING_WAIT=1 export NCCL_ASYNC_ERROR_HANDLING=1 export TORCH_NCCL_BLOCKING_WAIT=1 export TORCH_NCCL_ASYNC_ERROR_HANDLING=1 # Timeout settings for distributed operations (in seconds) export NCCL_TIMEOUT=3600 export TORCH_DISTRIBUTED_DEBUG=DETAIL ########################################################################################### # === Please modify the following paths according to your environment === Framework_name=QwenOFT base_vlm=/gpfs/wangzixuan/visual_prompting/starVLA_robocasa/playground/Pretrained_models/Qwen3-VL-4B-Instruct freeze_module_list='' DIT_TYPE="DiT-B" # Data paths data_root_dir=/gpfs/wangzixuan/visual_prompting/data_process/datasets visual_prompt_dir=/gpfs/wangzixuan/visual_prompting/starVLA_robocasa/simplerenv_process/visual_prompts_output extracted_frames_dir=/gpfs/wangzixuan/visual_prompting/data_process/extracted_frames_simplerenv data_mix=bridge_rt_1_vp # Config: choose between separate VP dataloader or inline VP # For separate VP dataloader mode: config_yaml=./examples/SimplerEnv/train_files/starvla_cotrain_oxe_visual_prompt.yaml # For inline VP mode (uncomment to use): # config_yaml=./examples/SimplerEnv/train_files/starvla_cotrain_oxe_visual_prompt_inline_vp.yaml # Output run_root_dir=/gpfs/wangzixuan/visual_prompting/starVLA_robocasa/playground/Checkpoints run_id=oxe_${data_mix}_visual_prompt_training_${Framework_name}_both_image_no_subtask # === End of environment variable configuration === ########################################################################################### output_dir=${run_root_dir}/${run_id} mkdir -p ${output_dir} # Save this script to the output dir cp $0 ${output_dir}/ accelerate launch \ --config_file starVLA/config/deepseeds/deepspeed_zero2.yaml \ --num_processes 8 \ starVLA/training/train_starvla_visual_prompt.py \ --config_yaml ${config_yaml} \ --framework.name ${Framework_name} \ --framework.qwenvl.base_vlm ${base_vlm} \ --framework.action_model.action_model_type ${DIT_TYPE} \ --datasets.vla_data.data_root_dir ${data_root_dir} \ --datasets.vla_data.visual_prompt_dir ${visual_prompt_dir} \ --datasets.vla_data.data_mix ${data_mix} \ --datasets.vla_data.per_device_batch_size 32 \ --datasets.vla_data.video_backend pyav \ --datasets.vp_data.visual_prompt_dir ${visual_prompt_dir} \ --datasets.vp_data.extracted_frames_dir ${extracted_frames_dir} \ --datasets.vp_data.per_device_batch_size 8 \ --trainer.freeze_modules "${freeze_module_list}" \ --trainer.max_train_steps 100000 \ --trainer.save_interval 10000 \ --trainer.logging_frequency 10 \ --trainer.eval_interval 100 \ --trainer.learning_rate.base 3e-5 \ --trainer.learning_rate.qwen_vl_interface 1e-5 \ --trainer.loss_scale.visual_prompt 0.1 \ --datasets.vla_data.use_subtask false \ --datasets.vla_data.feed_both_images true \ --datasets.vp_data.feed_both_images false \ --run_root_dir ${run_root_dir} \ --run_id ${run_id} \ --wandb_project starVLA_simplerEnv_visual_prompt \ --wandb_entity zwanggk # --is_debug True