


























SERL 的核心使命是:在真实世界中,让机器人在 20-40 分钟内学会高精度的机械操作。它通过集成 SAC、RLPD、DrQ 和 VICE,将原本需要数百万次尝试的 RL,压缩到了人类演示水平的量级。
RLPD(Reinforcement Learning with Prior Data)是一种基于 off-policy 的 actor-critic强化学习算法,借鉴了 soft-actor-critic 等时序差分算法的成功经验,但为满足上述需求做出了一些关键修改,它其实就是 SAC + 之前的数据(Prior Data)+ 极高的更新频率(High UTD)。
注:
SERL 的总体流程图如下,其中:

在许多场景中,强化学习的优异表现依赖于与环境进行大量的在线交互,这通常通过使用模拟器来实现。然而,在实际问题中,常常面临样本获取成本高昂的情况。此外,奖励信号稀疏,且高维的状态和动作空间往往使这一问题更加严重。
一些先前的研究致力于通过预训练利用这些数据,而其他方法则在在线训练时引入约束,以应对分布转移问题。然而,每种方法都有其缺点,例如需要额外的训练时间和超参数,或者在行为策略之外的提升有限。
RLPD关注的是,是否可以在 在线学习时,直接应用现有的离策略方法以充分利用离线数据。在每一步训练中,RLPD 在先验(离线)数据和on-policy数据之间等概率采样,以形成一个训练批次,即“对称采样”,即每个批次有50%的数据来自(在线)回放缓冲区,另外50%来自离线数据缓冲区(先验数据)。
我们对比如下:
原生:纯靠自己试
SERL:在 SAC 的 Batch 里塞进人类演示数据,相当于给 SAC 考试时递了一张带有参考答案的纸,让它不用从头瞎猜
那"RLPD"这四个字母贵在哪里?秘密不在"怎么算"(Loss Function),而在"喂什么"(Batch Composition)。RLPD 50/50 抽样 强行保证每个 Batch 都有 128 条 Demo。这相当于给机器人装了一个"强制记忆模块",让它每一秒钟都在看正确答案。
Prior Data 策略决定了演示数据与在线数据的黄金采样配比。
真机强化学习最怕"冷启动"— 机器人像没头苍蝇一样乱撞。SERL 引入了 RLPD(Reinforcement Learning with Prior Data)机制:
数据混洗:在训练的每一个 Batch 中,系统会强制性地混合:
算法价值:这种混合采样确保了模型在进化的每一秒,都在不断对比"正确答案"与"自己的尝试"。它解决了强化学习初期的探索困境,让机器人即使在完全没拿高分的情况下,也能通过模仿专家数据来迅速建立起任务的初步认知。
我们可以这样理解:
这也是 SERL 样本效率高的关键之一。它不是让机器人从零开始乱试,而是在 demonstrations 的引导下进行强化学习微调。论文中明确写到,每次更新使用 sample-based approximation,其中 half of the samples drawn from prior data,half drawn from replay buffer。
下面是根据 rlpd.py 的代码逻辑整理的 RLPD 训练流程逻辑图。这个图展示了从数据采样到网络更新的完整路径, 特别标注了 RLPD 相对于普通 SAC 的核心改进点 (如 BC Loss 和 Pessimistic Backup)。

RLPD 关键组件说明
lax.scan 允许在一个硬件循环内执行 20 次上述流程。结合 rlpd.py 代码的深度解读:
关于 rho 的计算:
next_qs = self.network.select('target_critic')(batch['next_observations'][..., -1, :], next_actions)
next_q = next_qs.mean(axis=0) - next_qs.std(axis=0) * self.config["rho"]
这是 RLPD 的精髓。普通的 SAC 是 min(Q1, Q2), 而这里是用标准差来量化"不确定性"。如果 10 个 Q 网络对某个状态动作意见不统一 (std 大), next_q 就会被压得很低。
关于 BC Loss:
bc_loss = -(dist.log_prob(jnp.clip(batch_actions, -1 + 1e-5, 1 - 1e-5)) * batch["valid"][..., -1]).mean() * self.config["bc_alpha"]
这行代码在告诉 Actor: "不管 Q 值怎么说, 你输出的动作最好和 Buffer 里的真实动作 (演示数据) 接近一些"。这对机器人任务极其关键, 因为它防止了机器人在训练初期因为乱甩而撞坏硬件。
关于 High UTD:
@jax.jit
def batch_update(self, batch):
agent, infos = jax.lax.scan(self._update, self, batch)
这里使用了 JAX 的 scan 原语。这比 Python 的 for 循环快得多, 它能把 20 次更新编译成一个高效的 GPU 算子。
演示数据(Prior Data):
SERL 所基于的 RLPD 算法中,作者发现最简单、最有效的办法就是一视同仁,比如:
"看未来"与"看现在"的逻辑
rlpd.py 中计算 next_q 用的是 target_critic, 而计算 actor_loss 时用的是 critic (当前网络)。
Target Critic (看未来): 用于计算 r + γ Q_{target}。由于 Q_{target} 更新得很慢 (Soft Update), 它提供了一个稳定的地基, 防止 Q 值计算产生正反馈螺旋 (即自己把自己估高)。
Critic (看现在): 用于 Actor 的更新。Actor 问: "我现在的动作好吗? "。由于 Critic 正在被最快地训练, 它能给 Actor 提供最及时的反馈。
矛盾解决: 这就是"评估要稳 (Target), 改进要快 (Current)"的权衡。
在 SERL 的复现中, 通常的步骤是:
在 SERL 的完整流程中,bc.py 的代码非常关键, 它揭示了 SERL 系统中 Behavioral Cloning (BC) 环节是如何运作的。
BCAgent 负责预训练冷启动,是纯监督学习的行为克隆实现,架构最为简洁。SERL先用 BC 模仿 Demo,让机器人学会"手往哪放",再开启 RL 寻找"怎么抓取"。
核心优势:BCAgent 的简洁性使其成为从演示到强化学习的理想桥梁,通过监督学习快速获得可用的策略,然后可以在此基础上进行RL微调。
BCAgent 的特性如下:
| 特性 | BCAgent | 说明 |
|---|---|---|
| 网络数量 | 仅1个Policy网络 | 无Critic,无Temperature |
| tanh_squash | False | 不使用tanh压缩 |
| 输出分布 | MultivariateNormalDiag | 标准高斯分布 |
| 训练目标 | 最小化MSE + 负对数似然 | 监督学习 |
唯一的 Policy 网络
network_kwargs["activate_final"] = True
networks = {
"actor": Policy(
encoder_def, # 视觉编码器
MLP(**network_kwargs), # 默认 [256, 256]
action_dim=actions.shape[-1],
tanh_squash_distribution=False, # 关键差异
)
}
| 组件 | 输入 | 网络结构 | 输出 | 特点 |
|---|---|---|---|---|
| Policy | 图像观测 | 编码器+MLP[256,256] | 动作分布(μ,σ) | 纯监督学习 |
"small" 编码器:
encoders = {
image_key: SmallEncoder(
features=(32, 64, 128, 256),
kernel_sizes=(3, 3, 3, 3),
strides=(2, 2, 2, 2),
padding="VALID",
pool_method="avg",
bottleneck_dim=256,
spatial_block_size=8,
)
}
"resnet" 编码器:
encoders = {
image_key: resnetv1_configs["resnetv1-10"](
pooling_method="spatial_learned_embeddings",
num_spatial_blocks=8,
bottleneck_dim=256,
)
}
"resnet-pretrained" 编码器:
pretrained_encoder = resnetv1_configs["resnetv1-10-frozen"](
pre_pooling=True,
)
encoders = {
image_key: PreTrainedResNetEncoder(
pooling_method="spatial_learned_embeddings",
num_spatial_blocks=8,
bottleneck_dim=256,
pretrained_encoder=pretrained_encoder,
)
}
def loss_fn(params, rng):
# 前向传播
dist = self.state.apply_fn(
{"params": params},
batch["observations"],
temperature=1.0,
train=True,
rngs={"dropout": key},
name="actor",
)
pi_actions = dist.mode() # 预测动作
log_probs = dist.log_prob(batch["actions"]) # 对数概率
# 多重损失
mse = ((pi_actions - batch["actions"]) ** 2).sum(-1) # MSE损失
actor_loss = -(log_probs).mean() # 负对数似然
return actor_loss, {
"actor_loss": actor_loss,
"mse": mse.mean(),
}
def sample_actions(self, observations, seed=None, temperature=1.0, argmax=False):
dist = self.state.apply_fn(
{"params": self.state.params},
observations,
temperature=temperature,
name="actor",
)
if argmax:
actions = dist.mode() # 确定性采样
else:
actions = dist.sample(seed=seed) # 随机采样
return actions
BC Agent (模仿学习) 核心流程图如下。BC 关键组件说明:

核心逻辑: update 函数
视觉处理: 数据增强 (Data Augmentation) data_augmentation_fn:
网络架构: ResNet-10
elif encoder_type == "resnet":
encoders = {
image_key: resnetv1_configs["resnetv1-10"](...)
}
为什么会有 mse 却不用它更新?
EncodingWrapper 的作用
这个包装器能把视觉图像和机械臂自身的状态 (关节角度、末端坐标) 揉在一起。这意味着机器人不仅知道自己"看到了什么", 还知道自己"现在手在哪"。
encoder_def = EncodingWrapper(..., use_proprio=use_proprio, enable_stacking=True, ...)
冷启动与热切换:如果复现SERL,一般会把 bc.py 练出来的模型会作为 RLPDAgent.create 时的初始参数 (或权重)。这相当于把原本需要几百万次尝试才能学会的动作, 压缩成了几千步的模仿。
bc_loss 的数据生效为(padding="VALID")。这是因为不能对在线数据做 BC Loss?
High UTD 的意义:它强迫神经网络在极短的时间内"吃透"每一张图片。
UTD(Update-to-Data Ratio)表示每采集一条环境数据,算法进行多少次梯度更新。
传统 RL 常用 UTD=1:采一步,训一步。
SERL / RLPD 使用更高 UTD(通常为 20 甚至更高):采集一条昂贵的真机数据后,learner 会多次从 buffer 中采样并更新网络。
我们可以把 High UTD 理解成:真机数据太贵,所以每一帧都要反复研读,不能看一遍就扔。
没有 UTD 的后果:普通的 SAC 每采样一个数据才更新一次。对于机器人这种高维度(ResNet 图像)且数据量极小(只有 2.5 小时数据)的任务:
High UTD 将数据的价值榨取到了极致:
@partial(jax.jit, static_argnames=("utd_ratio", "pmap_axis"))
def update_high_utd(
self,
batch: Batch,
*,
utd_ratio: int,
pmap_axis: Optional[str] = None,
) -> Tuple["SACAgent", dict]:
"""
Fast JITted high-UTD version of `.update`.
Splits the batch into minibatches, performs `utd_ratio` critic
(and target) updates, and then one actor/temperature update.
Batch dimension must be divisible by `utd_ratio`.
"""
batch_size = batch["rewards"].shape[0]
assert (
batch_size % utd_ratio == 0
), f"Batch size {batch_size} must be divisible by UTD ratio {utd_ratio}"
minibatch_size = batch_size // utd_ratio
chex.assert_tree_shape_prefix(batch, (batch_size,))
def scan_body(carry: Tuple[SACAgent], data: Tuple[Batch]):
(agent,) = carry
(minibatch,) = data
agent, info = agent.update(
minibatch, pmap_axis=pmap_axis, networks_to_update=frozenset({"critic"})
)
return (agent,), info
def make_minibatch(data: jnp.ndarray):
return jnp.reshape(data, (utd_ratio, minibatch_size) + data.shape[1:])
minibatches = jax.tree_map(make_minibatch, batch)
(agent,), critic_infos = jax.lax.scan(scan_body, (self,), (minibatches,))
critic_infos = jax.tree_map(lambda x: jnp.mean(x, axis=0), critic_infos)
del critic_infos["actor"]
del critic_infos["temperature"]
# Take one gradient descent step on the actor and temperature
agent, actor_temp_infos = agent.update(
batch,
pmap_axis=pmap_axis,
networks_to_update=frozenset({"actor", "temperature"}),
)
del actor_temp_infos["critic"]
infos = {**critic_infos, **actor_temp_infos}
return agent, infos
结论:如果把 UTD 降为 1,效果会大幅变差,甚至完全学不会。
把 cta_ratio(UTD 比率)从 20 降到 1,导致效果变差的原理主要有三点:
但 High UTD 也有副作用。对同一批数据反复训练,critic 容易过拟合和过估计,最终让策略崩溃。因此,SERL 还需要配套的稳定性机制。
High UTD 是发动机,但发动机太猛就需要刹车系统。SERL 通过多种机制的协同,实现了在极高样本效率下的稳定训练。
整体稳定性保障:
这些机制不是孤立工作的,而是协同配合。SERL 的工程价值在于不是单独实现某个技巧,而是把一整套相互配合的稳定性机制整合起来,使得高 UTD 这种"激进"的训练策略能够在真机上稳定运行,形成一套可工作的系统。SAC 的巧妙之处恰恰在于它如何利用"不确定性"来获得最终的稳定。
我们接下来选择部分机制进行解读。
论文中提到,regularizing the critic with layer normalization allows for higher UTD ratios and thus more efficient training。也就是说,SERL 并不是单纯把 UTD 拉高,而是通过 critic 正则化让高频更新不至于数值失控。
即,为了抗住 20 倍的更新强度而不崩盘,SERL 在 Critic 网络中引入了(LayerNorm)。这在传统 SAC 中是不常见的,但在高 UTD 的 RLPD 算法中至关重要。
从 MLP可以看到已有的层归一化支持:
class MLP(nn.Module):
use_layer_norm: bool = False # 层归归一化开关
@nn.compact
def __call__(self, x: jnp.ndarray, train: bool = False) -> jnp.ndarray:
for i, size in enumerate(self.hidden_dims):
x = nn.Dense(size, kernel_init=default_init())(x) # 线性变换
if i + 1 < len(self.hidden_dims) or self.activate_final:
# 正则化层(可选)
if self.dropout_rate is not None and self.dropout_rate > 0:
x = nn.Dropout(rate=self.dropout_rate)(x, deterministic=not train)
if self.use_layer_norm: # 关键:层归一化应用
x = nn.LayerNorm()(x) # 标准化层输出
x = activations(x) # 激活函数
return x
Critic 网络的特点:
层归一化的具体好处:
SACAgent 创建时启用层归一化.
critic_network_kwargs={
"activations": nn.tanh,
"use_layer_norm": True,
"hidden_dims": [256, 256],
},
policy_network_kwargs={
"activations": nn.tanh,
"use_layer_norm": True,
"hidden_dims": [256, 256],
},
针对 DrQAgent 的实现
critic_network_kwargs={
"activations": nn.tanh,
"use_layer_norm": True,
"hidden_dims": [256, 256],
},
policy_network_kwargs={
"activations": nn.tanh,
"use_layer_norm": True,
"hidden_dims": [256, 256],
},
VICE 中的层归一化(已实现)
critic_network_kwargs={
"activations": nn.tanh,
"use_layer_norm": True,
"hidden_dims": [256, 256],
},
vice_network_kwargs={
"activations": nn.leaky_relu,
"use_layer_norm": True,
"hidden_dims": [
256,
],
"dropout_rate": 0.1,
},
policy_network_kwargs={
"activations": nn.tanh,
"use_layer_norm": True,
"hidden_dims": [256, 256],
},
对 Critic 进行层归一化正则化的关键是:
critic_network_kwargs 中设置 use_layer_norm: True在 SERL 框架中,这种实现方式既保持了代码的简洁性,又充分利用了 Flax/JAX 的模块化优势,是提高 Critic 网络训练稳定性和性能的有效手段。
Soft Update 让目标网络始终缓慢追踪当前 Q 值,保持贝尔曼目标的平稳性。在 REDQ 的高 UTD 场景下尤为重要。
在机器人控制中,动作的连续性决定了硬件的寿命。SERL 坚持使用 Soft Update(软更新)维护目标网络:
平滑公式:θ(target) = τ θ_online + (1−τ) θ{target}。其中 τ 通常设为极其微小的 0.005。
硬件意义:与直接拷贝权重的 Hard Update 不同,Soft Update 让目标值(Target)以一种近乎流体的方式缓慢漂移。这反映到机器人身上,就是动作的进化是"渐进"的,不会因为模型权重的突跳导致机械臂产生瞬时的冲击电流或抖动。
Soft update的核心实现如下:
def target_update(self, tau: float) -> "JaxRLTrainState":
"""
Performs an update of the target params via polyak averaging. The new
target params are given by:
new_target_params = tau * params + (1 - tau) * target_params
"""
new_target_params = jax.tree_map(
lambda p, tp: p * tau + tp * (1 - tau), self.params, self.target_params
)
return self.replace(target_params=new_target_params)
这个方法在 SACAgent.update 中被调用:
# Update target network (if requested)
if "critic" in networks_to_update:
new_state = new_state.target_update(self.config["soft_target_update_rate"])
原理分析: Soft Update采用Polyak averaging的方式缓慢更新目标网络。这种方法的核心思想是让目标网络以平滑的方式跟踪主网络,而不是周期性地完全复制。这种平滑跟踪有助于:
REDQ 模式:支持把 Q 网络增加到 10 个以上,并从中随机抽 2 个来计算 Target。这是另一种对抗高估问题的强力方法。
High UTD 虽然能加速学习,但会带来致命副作用:Q 值过估计(Overestimation Bias)。模型会因为反复研读少量样本而变得极端自信,单Q网络容易高估未见过的状态一动作对的价值,最终导致策略崩溃。
SERL 引入了 REDQ(Randomized Ensembled Double Q-learning)风格的机制来解决这个问题。我们可以将其理解为一种"陪审团机制":
通俗地说:如果十个裁判里随机抽出的几个裁判中,有一个觉得这个动作危险,那我们就保守一点。这种"悲观主义"巧妙地抵消了 High UTD 带来的"狂热乐观",使训练在极高强度下依然稳如磐石。
算法如下:

训练时的行为:
这种设计既保持了ensemble的容量优势,又通过子采样降低了计算成本和过拟合风险。
REDQ 论文证明:min(2 from 10) 的效果接近 min(10),但计算量减少 5 倍。

# drq.py:124-125 和 launcher.py:165-166
critic_ensemble_size=10, # 10 个独立 Q 网络
critic_subsample_size=2, # 计算 target 时只随机选 2 个
def critic_loss_fn(self, batch, params: Params, rng: PRNGKey):
# ...前期准备代码...
# 1. 计算所有ensemble成员的Q值
target_next_qs = self.forward_target_critic(
batch["next_observations"],
next_actions,
rng=rng,
) # shape: (critic_ensemble_size, batch_size)
# 2. 如果配置了子采样,则随机选择指定数量的网络
if self.config["critic_subsample_size"] is not None:
rng, subsample_key = jax.random.split(rng)
subsample_idcs = jax.random.randint(
subsample_key,
(self.config["critic_subsample_size"],), # 通常是2
0,
self.config["critic_ensemble_size"], # 通常是10
)
target_next_qs = target_next_qs[subsample_idcs] # 只保留选中的网络
# 3. 在(子采样后的)ensemble成员中取最小值
target_next_min_q = target_next_qs.min(axis=0) # shape: (batch_size,)
# ...后续使用target_next_min_q计算TD目标...
SERL 先计算所有10个网络的Q值,然后随机选择2个网络进行子采样,最后在这2个中选择最小值。
详细执行流程:
subsample_idcs)target_next_qs[subsample_idcs] 只保留这2个网络的Q值REDQ为什么这样设计:
这种设计在保持计算效率的同时,充分利用了ensemble的多样性,是REDQ算法的核心创新点之一。
代码中完全没有观测 - 动作的时间戳对齐机制。但这是刻意的设计取舍:
实际的时序设计:
step(action) 被调用
├─1. 计算目标位姿 + 安全裁剪 ~0.1ms
├─2. _send_gripper_command() ~600ms (夹爪动作时)
├─3. _send_pos_command() ~5ms (HTTP POST)
├─4. time.sleep(1/hz - elapsed) 补齐到 100ms 控制周期
├─5. _update_curipos() ~5ms (HTTP POST /getstate)
└─6. _get_obs()
├ get_im() ~10~30ms (从 VideoCapture 队列取最新)
└ 组装 state (用步骤5的数据)
为什么不做精确对齐?
SERL 在系统设计上做出了取舍,专注于整体架构的简洁性和可维护性,而非追求极致的硬件同步精度。
RLPD 的 Critic 更新和 SAC 没有区别。
对于 Batch 里的每一个样本(无论是来自 Demo 还是 Online),Critic 都在做同一件事:Loss = (Q(s, a) - Target)²,其中 Target = r_{vice} + γ Q(s', a')。
注意:
这里的"特殊处理"不在公式里,而是在数据的质量上。RLPD 并没有给 Critic 写一个"针对 Demo 的特殊公式",它的特殊在于强行喂给 Critic 50% 的高质量样本。秘密不在"怎么算"(Loss Function),而在"喂什么"(Batch Composition)。
a_{random}),它对应的 r_{vice} 大概率是 0。a_{demo}),它对应的 r_{vice} 大概率是 1(或者接近 1 的高分)。所有数据(无论谁做的)都必须经过 VICE 的安检。VICE 说它是成功,它才是成功。
想象你和机器人都在学数学:

此内容由惯性聚合(RSS阅读器)自动聚合整理,仅供阅读参考。 原文来自 — 版权归原作者所有。