自定义数据集在单张RTX 4090上通过LoRA微调机械臂PI0模型的完整指南
基于Mujoco仿真平台robosuite采集的自定义数据集在单张RTX 4090上通过LoRA微调PI0模型的完整指南
重要说明
LoRA微调仅支持JAX训练器,不支持PyTorch训练器。
RTX 4090 (24GB显存)完全满足LoRA微调的内存要求(>22.5GB)。
1 环境安装
ubuntu 20.04系统
python =3.11
1.1 克隆仓库并初始化子模块
conda create -n openpi python=3.11
conda activate openpi
git clone --recurse-submodules https://github.com/Physical-Intelligence/openpi.git
cd openpi
1.2 安装uv包管理器
pip install uv
1.3 安装依赖
GIT_LFS_SKIP_SMUDGE=1 uv sync
GIT_LFS_SKIP_SMUDGE=1 uv pip install -e .
步骤二:数据转换
2 数据转换和下载预训练权重
2.1 采集自定义数据集并转换为lerobot格式
(1)步骤一:安装robosuite和robomimic
conda create -n robomimic python=3.10.0 -y
conda activate robomimic
git clone https://github.com/ARISE-Initiative/robosuite.git
cd robosuite
# (可选)切到v1.5.1以获得与robomimic示例最匹配的版本:
git checkout v1.5.1
# 安装依赖
pip install -r requirements.txt
# 安装 robosuite(可选)
pip install -e .
git clone https://github.com/ARISE-Initiative/robomimic.git
cd robomimic
pip install -e .
(2)步骤二:从robosuite采集
使用 collect_human_demonstrations.py 脚本来收集人类演示数据并生成 demo.hdf5 文件
python robosuite/scripts/collect_human_demonstrations.py \
--environment PickPlace \
--robots Kinova3 \
--device keyboard \
--directory ./my_demonstrations
运行脚本文件,会自动打开mujoco的仿真平台,通过键盘、或游戏手柄操控机械臂,生成演示数据。
手柄操作机械臂教程:在Robosuite中如何使用Xbox游戏手柄操控mujoco仿真中的机械臂?
(3)步骤三: 转换robosuite数据集格式
首先需要将原始的demo.hdf5文件转换为robomimic兼容格式:
python robomimic/scripts/conversion/convert_robosuite.py --dataset /path/to/demo.hdf5
此步骤会就地修改demo.hdf5文件,使其包含robomimic所需的元数据结构 。转换后的文件包含states和actions,但缺少observations、rewards和dones 。
(4)步骤四: 提取观测数据
使用dataset_states_to_obs.py从MuJoCo状态提取观测 ,包含图像观测:
python dataset_states_to_obs.py --dataset /path/to/demo.hdf5 \
--output_name image.hdf5 --done_mode 2 \
--camera_names agentview robot0_eye_in_hand \
--camera_height 84 --camera_width 84
必须要包含图像观测,因为RT_1模型是VLA模型,需要图片作为输入。
生成的image.hdf5文件就是完整的数据集。
(5)将hdf5格式文件转为lerobot文件
参考文章:基于Robosuite和Robomimic采集mujoco平台的机械臂数据微调预训练PI0模型,实现快速训练机械臂任务
安装lerobot的环境后执行脚本,这将在./kinova_data目录下创建LeRobot格式的数据集。
python convert_hdf5_to_lerobot.py
2.2 下载预训练权重
(1)方法一
# 安装 gsutil (如果尚未安装)
pip install gsutil
# 下载 params 目录
gsutil -m cp -r gs://openpi-assets/checkpoints/pi0_base/params /home/user/openpi/hub/pi0_base/
(2)方法二
打开科学上网后,运行以下脚本下载11.2G文件
import fsspec
import time
import concurrent.futures
import pathlib
import tqdm
url = "gs://openpi-assets/checkpoints/pi0_base"
save_path = "/home/user/openpi/hub/pi0_base/ "
local_path = pathlib.Path(save_path)
fs, _ = fsspec.core.url_to_fs(url)
info = fs.info(url)
# Folders are represented by 0-byte objects with a trailing forward slash.
if is_dir := (info["type"] == "directory" or (info["size"] == 0 and info["name"].endswith("/"))):
total_size = fs.du(url)
else:
total_size = info["size"]
with tqdm.tqdm(total=total_size, unit="iB", unit_scale=True, unit_divisor=1024) as pbar:
executor = concurrent.futures.ThreadPoolExecutor(max_workers=1)
future = executor.submit(fs.get, url, local_path, recursive=is_dir)
while not future.done():
current_size = sum(f.stat().st_size for f in [*local_path.rglob("*"), local_path] if f.is_file())
pbar.update(current_size - pbar.n)
time.sleep(1)
pbar.update(total_size - pbar.n)
3 创建自定义训练配置
3.1 创建Kinova数据配置类
创建文件src/openpi/policies/kinova_policy.py:
# src/openpi/policies/kinova_policy.py
import openpi.transforms as transforms
from openpi.models.model import ModelType
class KinovaInputs(transforms.DataTransformFn):
"""将Kinova数据集格式转换为模型输入格式"""
def __init__(self, model_type: ModelType = ModelType.PI0):
self.model_type = model_type
def __call__(self, x: dict) -> dict:
# 解析图像为uint8 HWC格式
images = {}
if "observation/images/agentview" in x:
images["base_0_rgb"] = transforms.parse_image(x["observation/images/agentview"])
if "observation/images/wrist" in x:
images["wrist_0_rgb"] = transforms.parse_image(x["observation/images/wrist"])
# 状态数据 (21维: 7关节位置 + 7关节速度 + 3末端位置 + 4末端四元数)
state = x["observation/state"]
# 动作数据 (7维: Kinova Gen3机械臂)
actions = x["actions"]
return {
"images": images,
"state": state,
"actions": actions,
}
class KinovaOutputs(transforms.DataTransformFn):
"""将模型输出转换为Kinova动作格式"""
def __call__(self, x: dict) -> dict:
return x # 动作已经是正确格式
3.2 在src/openpi/training/config.py中添加配置
在_CONFIGS列表中添加以下配置(在文件末尾的_CONFIGS列表中):
# 在config.py文件中添加以下DataConfigFactory
@dataclasses.dataclass(frozen=True)
class LeRobotKinovaDataConfig(DataConfigFactory):
"""Kinova Gen3数据配置"""
@override
def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig:
from openpi.policies import kinova_policy # 导入您创建的文件
repack_transform = _transforms.Group(
inputs=[
_transforms.RepackTransform({
"observation/images/agentview": "observation/images/agentview",
"observation/images/wrist": "observation/images/wrist",
"observation/state": "observation/state",
"action": "actions",
})
]
)
data_transforms = _transforms.Group(
inputs=[kinova_policy.KinovaInputs(model_type=model_config.model_type)],
outputs=[kinova_policy.KinovaOutputs()],
)
# Kinova动作已经是绝对值,需要转换为delta动作
delta_action_mask = _transforms.make_bool_mask(7)
data_transforms = data_transforms.push(
inputs=[_transforms.DeltaActions(delta_action_mask)],
outputs=[_transforms.AbsoluteActions(delta_action_mask)],
)
model_transforms = ModelTransformFactory(default_prompt="pick and place milk")(model_config)
return dataclasses.replace(
self.create_base_config(assets_dirs, model_config),
repack_transforms=repack_transform,
data_transforms=data_transforms,
model_transforms=model_transforms,
)
# 然后在_CONFIGS列表中添加LoRA微调配置:
TrainConfig(
name="pi0_kinova_lora",
# 使用LoRA变体进行低显存微调
model=pi0_config.Pi0Config(
action_dim=7, # Kinova Gen3有7个自由度
action_horizon=10, # 动作块长度
paligemma_variant="gemma_2b_lora", # 使用LoRA变体
action_expert_variant="gemma_300m_lora" # 动作专家也使用LoRA
),
data=LeRobotKinovaDataConfig(
repo_id="kinova_data", # /yourpath/openpi/kinova_data数据集位置
base_config=DataConfig(prompt_from_task=True),
),
# 加载PI0基础模型权重
# 修改这里: 使用本地路径
weight_loader=weight_loaders.CheckpointWeightLoader(
"/home/user/openpi/pi0_base/params"
),
# 设置freeze filter以冻结非LoRA参数
freeze_filter=pi0_config.Pi0Config(
paligemma_variant="gemma_2b_lora",
action_expert_variant="gemma_300m_lora"
).get_freeze_filter(),
# LoRA训练不使用EMA
ema_decay=None,
num_train_steps=20_000,
batch_size=32,
save_interval=1000,
),
freeze_filter的作用: 定义哪些参数在训练时被冻结。对于LoRA微调,只有LoRA适配器层会被训练,其他参数会被冻结。
4 计算归一化统计信息
在训练前必须计算数据集的归一化统计信息:
# 将当前openpi的路径添加
export HF_LEROBOT_HOME="/yourpath/openpi"
export HF_HOME="/yourpath/openpi"
uv run scripts/compute_norm_stats.py --config-name pi0_kinova_lora
# 或
uv run python -u scripts/compute_norm_stats.py --config-name pi0_kinova_lora
# -u 参数会禁用 Python 的输出缓冲,让您立即看到错误信息。
这会分析自定义的数据集并将统计信息保存到./assets/pi0_kinova_lora/local/kinova_pickplace/目录。
5 步骤五:运行LoRA微调
5.1 设置环境变量以最大化GPU内存使用
export XLA_PYTHON_CLIENT_MEM_FRACTION=0.9
这允许JAX使用最多90%的GPU内存。
5.2 启动训练
XLA_PYTHON_CLIENT_MEM_FRACTION=0.9 uv run scripts/train.py pi0_kinova_lora --exp-name=kinova_1 --overwrite
训练过程会:
- 在控制台输出训练进度
- 保存checkpoint到
./checkpoints/pi0_kinova_lora/kinova_milk_lora/ - 上传指标到Weights & Biases (如果启用)
6 运行推理
训练完成后,启动策略服务器:
uv run scripts/serve_policy.py policy:checkpoint \
--policy.config=pi0_kinova_lora \
--policy.dir=checkpoints/pi0_kinova_lora/kinova_milk_lora/20000
6.1 LoRA变体选择
PI0支持两种LoRA变体:
gemma_2b_lora: PaliGemma视觉-语言模型使用rank-16 LoRA适配器gemma_300m_lora: 动作专家使用rank-32 LoRA适配器
6.2 内存优化
- LoRA微调: >22.5GB (适合RTX 4090)
- 完整微调: >70GB (需要A100/H100)
如果仍然内存不足,可以使用FSDP分片:
uv run scripts/train.py pi0_kinova_lora --exp-name=kinova_milk_lora --fsdp-devices 1
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐
所有评论(0)