配置训练环境and准备训练数据

第一步按照文档配置已经完成:

Train_Tiny_DexVLA_train.yml 文件有改动,由于版本冲突以及依赖包等问题,删除了以下这些包:(后面需要单独安装flash-attn==2.7.4.post1 )

      - flash-attn==2.7.4.post1
      - nvidia-cublas-cu12==12.4.5.8
      - nvidia-cuda-cupti-cu12==12.1.105
      - nvidia-cuda-nvrtc-cu12==12.1.105
      - nvidia-cuda-runtime-cu12==12.1.105
      - nvidia-cudnn-cu12==8.9.2.26
      - nvidia-cufft-cu12==11.0.2.54
      - nvidia-cufile-cu12==1.11.1.6
      - nvidia-curand-cu12==10.3.2.106
      - nvidia-cusolver-cu12==11.4.5.107
      - nvidia-cusparse-cu12==12.1.0.106
      - nvidia-cusparselt-cu12==0.6.3
      - nvidia-nccl-cu12==2.20.5
      - nvidia-nvjitlink-cu12==12.6.85
      - nvidia-nvtx-cu12==12.1.105
pip install flash-attn==2.7.4.post1 -i https://pypi.tuna.tsinghua.edu.cn/simple

可以参考:(1 封私信 / 26 条消息) flash-attn安装避坑 - 知乎

不能保证改动是否正确(真的不确定),对DexVLA文件夹下的process_data.py文件做了以下改动

第一个改动是为了增加新任务click_bell和beat_block_hammer;

第二个改动是由于文件路径报错(原代码读取的是形如episode_0的文件夹,而/RoboTwin文件夹下的data文件夹中,数据存在)。

得到的数据文件如下:

ps:这一步还没做

下载预训练权重Qwen2-VL-2B

主包这里是本地下载传到服务器上的,下载链接如下:

Qwen/Qwen2-VL-2B-Instruct at main

下载预训练权重 ScaleDP-H 

主包这里是本地下载传到服务器上的,下载链接如下:

lesjie/scale_dp_h at main

设置训练配置and训练模型

1、修改aloha_scripts/constant.py

这个 constants.py 文件的作用是作为一个任务注册表。它里面列出的,是这个项目开发者们自己用来训练和测试的任务,比如 folding_data_0609place_object_scale 等。

现在,我们需要做的是为自己的数据集,在这个文件里创建一个新的任务配置

第1步:找到自己的数据集路径

首先,需要知道用来训练的数据集存放在服务器的哪个位置。假设 click_bell 数据集被存放在:/mnt/home/xuhuixin/RoboTwin/policy/DexVLA/data/sim-click_bell/demo_clean-50

注意:请务必将这个路径替换为您自己数据集的真实绝对路径)。

第2步:理解任务配置的结构

我们来看一个文件里已有的、比较简单的例子来分析其结构:

    'folding_blue_shirt': { # for local debug
        'dataset_dir': [
            "/media/rl/HDD/data/data/aloha_data/4_cameras_aloha/folding_shirt"
        ],
        'episode_len': 1000,
        'camera_names': ['cam_high', 'cam_left_wrist', 'cam_right_wrist']
    },

这个结构包含三个关键部分:

  1. 'folding_blue_shirt': 这就是任务的名称 (TASKNAME)。您可以自己定义一个简单明了的名称。

  2. 'dataset_dir': 这是一个列表 (list),包含了所有用于这个任务的数据文件夹路径。即使只有一个路径,也要放在方括号 [] 里。

  3. 'episode_len': 每个数据片段的最大长度。训练时,数据会被截断或填充到这个长度。

  4. 'camera_names': 您的机器人数据中包含的摄像头名称。数据加载器会根据这个列表去读取对应的图像数据。

由于原constant.py中的task_name只有folding_data_0609、place_object_scale、folding_blue_shirt这些,主包想测试的任务是click_bell,所以需要按照代码的格式增加,例如:

第3步:在 constants.py 中添加任务
# --- 在这里添加您自己的 click_bell 任务 ---
'click_bell': {
    'dataset_dir': [
        # 这是包含您所有 hdf5 文件的文件夹路径
        "/mnt/home/xuhuixin/RoboTwin/policy/DexVLA/data/sim-click_bell/demo_clean-50/" 
    ],
    'episode_len': 400,  # 这是一个通用的起始值,如果您的任务动作序列很长或很短,可以适当调整
    'camera_names': ['cam_high', 'cam_left_wrist', 'cam_right_wrist'] # 假设这是标准的摄像头设置
},
# --- 任务添加完毕 ---

# ... 文件中其他的任务配置保持不变 ...
'folding_data_0609': {
    # ...

2、修改scripts/aloha/vla_stage2_train.sh

1. 修改 VLM (Qwen2-VL) 路径

脚本中使用 mnop 变量来指向 VLM 的路径。它原来的逻辑比较复杂,我们会用你的绝对路径直接替换它,这样最清晰。

修改为:

# 我们不再使用复杂的 if/else 结构,直接指定你的路径
# if [ "${LLM}" == "paligemma" ]; then
#   echo "Using PaliGemma"
#   mnop=${ROOT}/wjj/model_param/PaliGemma/paligemma/pixel_224/vla-paligemma-3b-pt-224
# else
#   mnop=${ROOT}/Qwen2-VL-${LLM_MODEL_SIZE}-Instruct # original qwen2vl
# fi

# 直接将 mnop 设置为你的 Qwen2-VL-2B 文件夹的绝对路径
mnop=/mnt/home/xuhuixin/RoboTwin/policy/DexVLA/Qwen2-VL-2B
2. 修改策略头 (ScaleDP-H) 路径

脚本中使用 DIT_PRETRAIN 变量指向策略头的权重文件。特别注意:这里需要的是具体 .ckpt 文件的路径,而不是文件夹路径。

找到这一行:

DIT_PRETRAIN=/data/private/policy_step_60000_2025-06-15_09-15-25.ckpt

修改为: 你需要先确认一下你下载的 ScaleDP-H 文件夹里的文件名。它很可能就叫 policy_step_...

# 将 DIT_PRETRAIN 设置为你的 ScaleDP-H 文件夹中 .ckpt 文件的绝对路径
# 请根据你文件夹内的实际文件名确认!
DIT_PRETRAIN=/mnt/home/xuhuixin/RoboTwin/policy/DexVLA/ScaleDP-H/policy_step_60000_2025-06-15_09-15-25.ckpt

操作:将原来的路径替换成你的 ScaleDP-H 文件夹里那个 .ckpt 文件的完整路径。

3. 修改任务名称 (TASKNAME)

这个变量定义了你的训练任务,会影响到数据加载和输出目录。你需要把它改成你自己的任务名。

找到这一行:

TASKNAME=folding_data_0609

修改为:

# 替换成你在 constant.py 文件中定义的你自己的任务名
TASKNAME="your_actual_task_name" 

操作:把 folding_data_0609 换成你自己的任务名(例如 aloha_towel_folding)。

4. 修改输出目录 (OUTPUT)

这个变量决定了你的训练结果(模型、日志等)保存在哪里。脚本原来的路径是基于 ROOT 变量的,我们同样建议使用一个简单明了的绝对路径。

找到这一行:

OUTPUT=${ROOT}/dex-checkpoints/stage2/${LLM}_${LLM_MODEL_SIZE}/${TASKNAME}_Stage2_DIT_H_Stage1_1_17_using_state_correct

修改为:

CURR_DATE=$(date +"%y%m%d") # 获取当前日期,例如 251015
BATCH_SIZE=24 # 定义设定的批次大小
OUTPUT="/mnt/home/xuhuixin/RoboTwin/policy/DexVLA/exp/qwen2_vl_2B/${TASKNAME}_s2_lora_dit_h_bs${BATCH_SIZE}_${CURR_DATE}"

操作:将原来的 OUTPUT 定义替换为上面这一行。这会把所有输出都保存在 DexVLA/exp/ 目录下,非常清晰。

5. 自定义超参数 (vla_stage2_train.sh)

拥有两张 5090 意味着您可以使用非常高的设置来加速训练。以下是针对您硬件的详细建议。

请打开 scripts/aloha/vla_stage2_train.sh 文件,找到 deepspeed 命令后面的参数部分进行修改。

  • --num_gpus: (必须修改) 这是最重要的设置。您的硬件是两张卡。

    • 建议值: --num_gpus=2

  • --per_device_train_batch_size: 这个参数定义了每张 GPU 处理的批次大小。5090 显存巨大,您可以用一个很高的值来充分利用硬件,这有助于稳定训练。

    • 建议值: 1632 之间。可以从 24 开始尝试。

    • 计算: 您的总批次大小 (Total Batch Size) 将是 (per_device_train_batch_size) * (num_gpus)。例如,24 * 2 = 48

    • 注意: 如果训练开始后遇到 “CUDA out of memory” 错误,只需把这个值调低一些(比如从 24 降到 16)再试一次。

  • --save_steps: 这个参数决定了每隔多少步保存一次模型检查点 (checkpoint)。

    • 权衡: 值越小,保存越频繁,占用的硬盘空间越多,但如果训练中断,您能从一个更近的节点恢复。值越大则相反。

    • 建议值: 对于您这种强大的硬件,训练速度会很快。一个很高的值(如脚本中默认的 10000)可能意味着很长时间才保存一次。建议设置为 20005000。这样既能频繁备份进度,又不至于过于频繁。

  • --max_steps: 这个参数是总的训练步数。脚本中默认的 60000 是针对一个非常大的混合数据集的。对于您目前这个单一的 click_bell 任务,这个值可能太大了,会导致模型在您的数据上过拟合 (overfitting)。

    • 建议值: 2000030000。您可以先设为 30000,然后通过 TensorBoard 观察损失 (loss) 曲线。如果损失在 20000 步左右就不再下降并趋于平稳,您就可以提前手动停止训练,此时的模型效果可能就是最佳的了

deepspeed --master_port 29604 --num_gpus=2 --num_nodes=1 ./train_vla.py \
  --deepspeed scripts/zero2.json \
  # ... 其他参数 ...
  --output_dir $OUTPUT \
  --max_steps 30000 \                 # <-- 修改后的总步数
  --per_device_train_batch_size 24 \   # <-- 修改后的批次大小
  --gradient_accumulation_steps 1 \
  --save_strategy "steps" \
  --save_steps 2000 \                  # <-- 修改后的保存频率
  --save_total_limit 10 \              # 可以适当减少保存的checkpoint总数,节省空间
  # ... 其他参数 ...

3、最后一步:开始训练

DexVLA 目录下运行脚本:

bash ./scripts/aloha/vla_stage2_train.sh

训练过程中遇到了许多小问题,首先是发现有一些模型没有成功下载,对照github又重新下载上传了,然后遇到了数据类型不匹配的问题(这个很好解决,简单修改了代码即可),之后是遇上了显存不足的问题,后面我将batch_size设置为1,gpu数量设置为2,lora_enable设置为Ture。

再次训练:

见上,训练失败的根本原因:数据加载器与 Qwen2-VL 模型的图像 Token 数量不匹配。

然而将batch_size设置为4之后,报错信息为:

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐