✨ 复制成功!

以文会友,打造好学人设!

📋 已复制到剪贴板:

站在巨人肩上:在自定义数据集上微调动作策略

创建时间:2026-08-11 更新时间:2026-08-11 阅读次数:1050 次

上一节,我们调用了开源的预训练VLA模型,让它根据指令“拿起红色的积木”自主完成了抓取。那是借力打力——用别人训练好的权重,跑自己的任务。

但预训练模型的泛化能力是有限的。它能抓红色积木,不一定能叠毛巾。它能推开抽屉,不一定能拉上拉链。对于你真正关心的那个具体任务——比如让机械臂把你桌上乱七八糟的文具分门别类放进收纳盒——预训练模型大概率表现不佳。

这一节,我们要做一件更具成就感的事情:用你自己的数据,让一个通用模型学会你的专属技能。 这是全书最“高光”的动手实践。你将扮演机器人老师的角色,亲手采集几十条演示数据,喂给一个预训练基础模型,微调出一个属于你自己的动作策略。

为什么微调而不是从零训练?

在深度学习的世界里,有一条铁律:数据量决定模型的上限,预训练决定模型的下限。

从零训练一个具身大模型,需要数百万条机器人操作数据,需要数百块GPU运行数周,需要一支工程师团队处理数据清洗、分布式训练、超参数调优。这不是一个人、一台电脑能完成的工作。

但微调完全不同。微调的基础逻辑是:模型已经在海量数据上学到了通用的视觉理解和动作生成能力,你只需要用少量数据告诉它“现在这个具体任务该怎么做”。 就像一个人已经学会了走路,你只需要拉着他的手,带他走几遍一条新的路线,他就能自己走了。

从数学上看,微调是在预训练权重 $W_{\text{pretrain}}$ 的基础上,用你的小数据集进行少量梯度更新,得到新权重 $W_{\text{finetune}}$:

$$ W_{\text{finetune}} = W_{\text{pretrain}} + \Delta W $$

其中 $\Delta W$ 的模长通常很小——你只调整模型参数的一个微小部分,而不是推翻重来。这种做法的好处是:几十条数据就能产生效果,训练时间从数周缩短到几十分钟,普通游戏显卡甚至CPU都能跑。

任务设定:教机器人叠毛巾

我们选一个具体而有生活气息的任务:叠毛巾。

一条毛巾平铺在桌面上,机械臂需要捏住毛巾的一角,折叠到另一角上,完成一个对折动作。这个任务涉及精细的视觉定位、柔顺的接触控制、以及对柔性物体不可预测变形的适应。它比抓取刚性积木难得多,但正因如此,它是一个绝佳的微调测试案例。

在整个流程中,你要做三件事: 1. 采集数据:在仿真中手动操控机械臂,演示几次正确的叠毛巾动作,同时记录图像和动作序列 2. 准备数据集:将记录的数据组织成模型能读取的格式 3. 微调模型:在预训练权重上运行几十轮梯度更新 4. 部署验证:让微调后的模型自主叠毛巾,看它是否学会了

第一步:采集中演示数据

在仿真环境中,我们需要一个“示教模式”——你手动操控机械臂完成叠毛巾动作,系统自动记录每一步的观测和动作。这里的关键是:你不需要写任何控制代码,你只需要“做给机器人看”。

我们使用一个简化的示教数据采集脚本。假设你已经在MuJoCo中搭建好了叠毛巾的场景——桌面上平铺着一块矩形布料,机械臂的夹爪初始位置在毛巾一角的上方。

import numpy as np
import mujoco
import mujoco.viewer
import pickle
import os
from datetime import datetime

# 加载仿真场景
xml_path = "towel_folding.xml"
model = mujoco.MjModel.from_xml_path(xml_path)
data = mujoco.MjData(model)

# 数据存储列表
# 每条数据包含:图像观测、关节角度(动作标签)、夹爪状态
episodes = []
current_episode = {
    "images": [],       # 每一步的RGB图像
    "actions": [],      # 每一步的关节角度(动作标签)
    "gripper_states": [], # 每一步的夹爪开合状态
}

# 仿真参数
timestep = model.opt.timestep
is_recording = False  # 按R键开始/停止录制

def get_camera_image():
    """获取当前场景的渲染图像"""
    renderer = mujoco.Renderer(model, 256, 256)
    renderer.update_scene(data, camera=0)
    image = renderer.render()
    renderer.close()
    return image

def get_current_action():
    """获取当前关节位置作为动作标签"""
    return data.qpos.copy()

def get_gripper_state():
    """获取夹爪当前开合状态"""
    # 假设夹爪关节在qpos中的特定索引
    gripper_idx = model.jnt_qposadr[
        mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, "gripper_joint")
    ]
    return data.qpos[gripper_idx]

print("示教数据采集工具")
print("=" * 50)
print("操作说明:")
print("  在仿真窗口中手动拖动机器人关节")
print("  按 R 键开始/停止录制一条示教轨迹")
print("  按 S 键保存所有已录制的轨迹")
print("  按 Q 键退出")
print("=" * 50)

recorded_count = 0

with mujoco.viewer.launch_passive(model, data) as viewer:
    while viewer.is_running():
        # 检查按键(通过viewer的用户事件或外部输入)
        # 注意:MuJoCo viewer的按键检测需要通过特定的回调机制
        # 这里展示核心逻辑,实际使用时需要根据MuJoCo版本适配

        # 模拟按键逻辑(实际使用中需替换为真正的按键检测)
        # 这里我们使用一个简单的状态机
        if is_recording:
            # 记录当前帧
            image = get_camera_image()
            action = get_current_action()
            gripper = get_gripper_state()

            current_episode["images"].append(image)
            current_episode["actions"].append(action)
            current_episode["gripper_states"].append(gripper)

        mujoco.mj_step(model, data)
        viewer.sync()

# 保存采集到的数据
if len(episodes) > 0:
    save_dir = "demonstration_data"
    os.makedirs(save_dir, exist_ok=True)
    timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
    save_path = os.path.join(save_dir, f"towel_folding_demos_{timestamp}.pkl")

    with open(save_path, "wb") as f:
        pickle.dump(episodes, f)

    print(f"已保存 {len(episodes)} 条示教轨迹到 {save_path}")
else:
    print("没有录制任何示教轨迹。")

采集示教数据是整个流程中最需要耐心的环节。你的每一次演示,都会成为模型学习的“标准答案”。在操作时注意以下几点:

  • 一致性:每条演示的起始位置和动作序列尽量保持一致。如果有的从左边抓起、有的从右边抓起,模型会困惑。
  • 流畅性:动作连贯平滑,不要有大幅抖动。模型会学习你的动作风格。
  • 多样性:在保持一致性的前提下,毛巾的初始位置可以有微小变化——这能让模型学到一定的泛化能力,不至于毛巾偏了一厘米就束手无策。
  • 数量:对于叠毛巾这种相对简单的任务,20到50条演示通常足够。更复杂的任务可能需要100条以上。

第二步:组织微调数据集

采集到的原始数据需要被转换成模型训练所需的格式。对于Octo这类模型,数据集通常是一个TensorFlow或PyTorch的Dataset对象,每次返回一个包含图像、语言指令、动作序列的字典。

import numpy as np
import pickle
import torch
from torch.utils.data import Dataset, DataLoader

class TowelFoldingDataset(Dataset):
    """将示教数据组织成PyTorch数据集"""

    def __init__(self, demo_path, sequence_length=8, image_size=(256, 256)):
        with open(demo_path, "rb") as f:
            self.episodes = pickle.load(f)

        self.sequence_length = sequence_length
        self.image_size = image_size

        # 将所有帧展开为独立的训练样本
        self.samples = []
        for episode in self.episodes:
            num_steps = len(episode["images"])
            # 为每一帧生成一个训练样本(包含当前图像和未来几个动作)
            for i in range(num_steps - sequence_length):
                sample = {
                    "image": episode["images"][i],
                    "actions": np.array(episode["actions"][i:i+sequence_length]),
                    "gripper_states": np.array(episode["gripper_states"][i:i+sequence_length]),
                }
                self.samples.append(sample)

        print(f"从 {len(self.episodes)} 条示教轨迹中生成了 {len(self.samples)} 个训练样本")

    def __len__(self):
        return len(self.samples)

    def __getitem__(self, idx):
        sample = self.samples[idx]

        # 将图像转换为张量并归一化到 [0, 1]
        image = torch.from_numpy(sample["image"]).float() / 255.0
        image = image.permute(2, 0, 1)  # (H, W, C) -> (C, H, W)

        # 拼接动作和夹爪状态
        actions = torch.from_numpy(sample["actions"]).float()
        gripper = torch.from_numpy(sample["gripper_states"]).float().unsqueeze(-1)
        action_labels = torch.cat([actions, gripper], dim=-1)

        return {
            "image": image,
            "actions": action_labels,
        }

# 加载数据集
dataset = TowelFoldingDataset("demonstration_data/towel_folding_demos_20250615_143022.pkl")
dataloader = DataLoader(dataset, batch_size=8, shuffle=True)

print(f"数据集准备完成,共 {len(dataset)} 个样本,{len(dataloader)} 个batch")

第三步:微调模型

现在到了整个章节最高光的时刻——用你自己的数据更新模型的权重。

我们以Octo为例,展示一个精简但完整的微调脚本。Octo内部基于JAX,我们使用它的训练API进行微调。如果你用的是OpenVLA或其他PyTorch模型,微调逻辑类似,只是API细节不同。

import jax
import jax.numpy as jnp
import numpy as np
import optax
from octo.model.octo_model import OctoModel
from tqdm import tqdm

# ==================== 加载预训练模型 ====================
print("加载预训练Octo模型...")
model = OctoModel.load_pretrained("octo-base")
print("模型加载完成。")

# ==================== 准备数据加载器 ====================
# 使用前一步准备的PyTorch DataLoader
# 注意:Octo内部用JAX,需要在训练循环中做转换
from torch.utils.data import DataLoader
from towel_folding_dataset import TowelFoldingDataset

dataset = TowelFoldingDataset("demonstration_data/towel_folding_demos.pkl")
dataloader = DataLoader(dataset, batch_size=8, shuffle=True, num_workers=0)

# ==================== 设置优化器 ====================
# 使用AdamW优化器,学习率设为预训练权重的1/10
learning_rate = 3e-5
optimizer = optax.adamw(learning_rate)

# 初始化优化器状态
opt_state = optimizer.init(model.params)

# ==================== 定义损失函数 ====================
def loss_fn(params, batch):
    """计算动作预测的均方误差损失"""
    images = jnp.array(batch["image"].numpy())
    actions_gt = jnp.array(batch["actions"].numpy())

    # 前向传播:预测动作
    # 这里简化了Octo的推理逻辑,实际使用时需要调用model的forward方法
    pred_actions = model.apply(
        params,
        images,
        method=model.predict_actions,
    )

    # MSE损失
    loss = jnp.mean((pred_actions - actions_gt) ** 2)
    return loss

# 编译训练步骤(JAX的JIT加速)
@jax.jit
def train_step(params, opt_state, batch):
    loss, grads = jax.value_and_grad(loss_fn)(params, batch)
    updates, new_opt_state = optimizer.update(grads, opt_state, params)
    new_params = optax.apply_updates(params, updates)
    return new_params, new_opt_state, loss

# ==================== 训练循环 ====================
num_epochs = 20
print(f"开始微调,共 {num_epochs} 轮...")

for epoch in range(num_epochs):
    epoch_losses = []
    progress_bar = tqdm(dataloader, desc=f"Epoch {epoch+1}/{num_epochs}")

    for batch in progress_bar:
        # 执行一步训练
        model.params, opt_state, loss = train_step(
            model.params, opt_state, batch
        )
        epoch_losses.append(float(loss))
        progress_bar.set_postfix({"loss": f"{float(loss):.6f}"})

    avg_loss = np.mean(epoch_losses)
    print(f"Epoch {epoch+1} 完成,平均损失: {avg_loss:.6f}")

print("微调完成!")

# ==================== 保存微调后的权重 ====================
save_path = "octo_towel_folding_finetuned"
model.save_pretrained(save_path)
print(f"微调权重已保存到 {save_path}")

训练开始后,你会看到损失值随着轮次逐渐下降。从初始的0.0几下降到0.00几,甚至0.000几。这个数字的下降,意味着模型的输出越来越接近你演示的动作序列——它正在把你看似随意的示教动作,内化为自己的行为策略。

对于几十条示教数据,20轮微调通常在10到30分钟内完成(取决于你的GPU性能)。如果你用的是CPU,可能需要1到2小时,但完全可以跑通。

第四步:部署验证

微调完成后,用上一节学到的闭环控制脚本,加载你的微调权重,让模型自主叠毛巾:

# 加载微调后的模型
model = OctoModel.load_pretrained("octo_towel_folding_finetuned")

# 修改指令
instruction = "fold the towel in half"
task = model.create_tasks(texts=[instruction])

# 运行闭环控制(代码结构与6.2节相同)
# ...(参见上一节的仿真闭环脚本)

运行脚本,仔细观察机械臂的行为。

最开始几步,它可能在定位毛巾的一角——末端缓慢移动,在视觉空间中搜索那个该捏住的点。找到之后,夹爪闭合,捏住毛巾一角,然后抬起来,向对角方向移动。毛巾在物理引擎里产生逼真的褶皱,布料随着夹爪的拖拽弯曲变形。最终,一角落在另一角上,夹爪松开,毛巾完成对折。

这就是你教给机器人的技能。 不是某个工程师手写的规则,不是从互联网上复制粘贴的代码,而是你亲手采集数据、亲手微调、亲手部署的一个动作策略。几十条示教数据,几十分钟训练,一个原本只会抓积木的通用模型,现在学会了叠毛巾。

微调中的常见问题与解决思路

在实践中,微调很少一次成功。以下是几个常见问题和应对策略:

过拟合:训练损失很低,但实际部署一塌糊涂。 这说明模型背下了你的示教动作,但没有学会泛化。解决方法:增加示教数据的多样性(改变毛巾的初始位置、角度、褶皱状态),加入数据增强(随机裁剪、颜色抖动),或者使用更小的学习率和更少的训练轮次。

动作抖动:机械臂在运动过程中不停颤抖。 这可能是因为模型输出的动作序列不平滑。解决方法:在损失函数中增加一个平滑项,惩罚相邻动作帧之间的剧烈变化——这被称为时间平滑正则化。

任务只完成一半:机械臂捏住了毛巾但没放对位置就松开了。 检查你的示教数据——可能你演示时某些步骤速度不均匀,导致模型学到了模棱两可的策略。重新采集数据,确保每一步都清晰明确。

夹爪时机不对:该闭合时张开,该张开时闭合。 夹爪的离散开关是VLA模型最容易出错的部分。可以尝试在数据预处理时,把夹爪状态单独拿出来做一个二分类损失项,和连续动作的回归损失加权求和,让模型更关注夹爪时机的准确性。

小结:你从使用者变成了创造者

这一节,你完成了一次完整的角色跃迁。

在上一节,你是模型的使用者——调用别人训练好的API,输入指令,观察结果。在这一节,你成了模型的创造者——你用自己的双手采集示教数据,用自己的数据微调模型参数,用自己的验证场景测试策略效果。

这个流程——采集数据、微调模型、部署验证——就是2025年具身智能领域最主流的研究范式。全球各地的实验室每天都在重复这个循环,只不过他们用的可能是更贵的机器人、更大的模型、更多的数据。但逻辑完全一致:通用底座 + 专属微调 = 千行百业的机器人应用。

你刚才用几十条数据教会了机器人叠毛巾。这个能力的边界在哪?收拾桌面、摆餐具、叠衣服、组装简单的零件——理论上,任何你能演示清楚的操作任务,这个流程都能覆盖。

在下一章,也是全书的最后一章,我们将跳出具体的技术实现,讨论一个更宏大的问题:当我们拥有了所有这些能力,具身智能的下一步在哪里?仿真和现实的差距如何弥合?通向通用具身智能的路线图长什么样?以及最重要的——你,作为刚刚入门的新人工程师,如何在这条路上走下去。

本教程共21节,当前为第18节!
本教程最新修订时间为:2026-08-19 23:29:19

📌 面试天下网:一款服务于大一新生的口袋书,让大家在无聊的公共课上可以学习大模型技术!
📌 网站公告:【悬赏 200 元/篇】寻找“大模型面试战场”的一手回忆录
📌 网站公告:【大模型/Agent面试陪练1V1指导】正式启动......

评论专区

评论加载中...