
这次我们来看一篇偏算法研究方向的论文Learning When to Stop: Prefix-Optimal Dynamic Diffusion Policies for Continuous Control。先说它解决什么问题。扩散策略Diffusion Policies在连续控制任务上的效果确实能打模仿学习、离线强化学习里都能拿到不错的分。但代价是推理慢每次生成动作都要完整跑一条噪声去除链常见设置要跑 20 步甚至 100 步。问题是很多状态下并不需要这么长的链扩散模型在前几步就已经把动作的大致方向定下来了后续步骤更多是在做细节修整。于是一个很自然的问题就出现了能不能让策略自己决定这次推理到底跑多少步论文标题里的三个关键词已经把思路写得很明确Prefix-Optimal、Dynamic Diffusion Policies、Early Stopping。它并不是简单地把去噪步数从 20 改成 5而是把去噪轨迹看作一个可以截断的序列让策略针对当前状态选择一个“最优前缀”只跑到这个前缀就停。这篇文章会做四件事把论文核心概念拆开讲清楚给出复现实验之前需要准备的环境整理一套训练、推理、批量评估的通用流程把容易踩的坑和排查方向列成清单。如果你在研究扩散策略、做连续控制方向的复现或者正在想办法降低策略推理延迟这篇可以直接往下看。1. 核心能力速览由于这是一篇学术论文而不是开箱即用的工具很多参数需要以论文原文和作者放出的示例代码为准。先给一张速览表方便快速判断这个方向和你是否相关。能力项说明项目类型学术研究论文 / 算法框架非桌面工具或 Web 服务核心机制动态早期停止根据状态决定扩散去噪的停止步数核心概念Prefix-Optimal前缀最优、Early Stopping提前停止目标任务连续控制常见测试环境为 MuJoCo 与 D4RL benchmark推理效率通过缩短平均去噪步数降低推理延迟具体提升幅度需以论文实验为准运行环境Python PyTorch MuJoCo D4RL 这类典型 RL 复现环境显存需求依赖具体模型和 batch size需按实际环境测试是否支持 CPU理论可运行但扩散策略训练和批量推理在 GPU 上更实际是否提供 API论文本身不提供 HTTP 接口接口化需要自行封装是否支持批量任务实验侧支持批量 episode 评估和批量推理适合读者强化学习、扩散模型、机器人控制方向的研究者与工程师从表里可以看到这篇论文的“交付物”不是一键启动包而是一套算法思路和对应的实验验证方式。所以下面各章节的内容都以复现实验为中心展开。2. 为什么连续控制里的扩散策略需要“知道何时停止”扩散策略的基本流程很好理解策略并不直接输出动作而是从一个随机噪声张量出发经过多次去噪最终得到动作向量。这个去噪过程可以看作一条马尔可夫链从完全噪声的 $x_T$ 逐步逼近数据分布中的 $x_0$。在行为克隆和离线强化学习中扩散策略的优势在于能够表示多峰分布避免传统高斯策略把多个可行动作平均成一个不可执行的动作。但代价是推理阶段的计算量很大。每个控制周期都要执行一次完整的去噪链步数越多策略延迟越高。对于机器人控制这类场景策略延迟直接影响闭环控制的稳定性。如果动作要 20 步去噪才能出来那控制频率就受限于这一步生成时间。固定步数的问题在于不是所有状态都需要相同的去噪步数。有些状态很清晰目标动作也很明确去噪 3 到 5 步就已经输出了一个合理的动作。有些状态比较复杂或者处于决策的关键时刻确实需要更多步去噪来区分不同的动作模式。用同一个步数处理所有状态必然导致一部分状态计算浪费。所以这里的关键词是Dynamic。论文不是提出一个全局的固定步数而是让策略动态决定本次去噪应该在哪个点停止。从扩散生成过程来看早期去噪步骤已经决定了动作的“骨架”大致方向、大致量级、属于哪个动作模式。后续步骤则是不断精修把动作映射到更精确的训练数据分布点上。连续控制并不要求每一步动作都达到像素级完美只要动作位于可行区域、能完成当前任务目标就可以执行。这给提前停止留出了很大的优化空间。3. Prefix-Optimal 与 Early Stopping论文核心概念拆解理解这篇论文关键是理解“前缀”和“最优前缀”。3.1 去噪轨迹的前缀视角我们把一次完整的扩散采样过程看成$$x_T \rightarrow x_{T-1} \rightarrow x_{T-2} \rightarrow \dots \rightarrow x_0$$这里的 $x_t$ 是去噪过程中第 $t$ 步的中间结果。完整生成需要 $T$ 步但如果我们在第 $k$ 步就停止那么实际使用的就是序列$$x_T \rightarrow x_{T-1} \rightarrow \dots \rightarrow x_k$$这个序列是完整序列的一个前缀。传统方法等于强制使用完整长度的前缀即整条序列。论文要做的是针对当前状态找到一条长度更短但依然能满足动作质量要求的前缀。3.2 前缀最优怎么理解“Prefix-Optimal”可以拆成两层意思。第一层是前缀层面的截断优化。扩散采样的每一步都会改变动作内容。如果停止得太早动作可能还带有明显噪声控制效果会下降如果停止得太晚又浪费算力。因此存在一个效率和质量的权衡问题。前缀最优就是在给定状态下找到一个长度 $k$使得前 $k$ 步去噪后的动作在某个评价标准下已经足够好。第二层是每个状态都有自己的最优前缀。不同状态下去噪收敛速度差异很大。有些状态一步就能去噪到可执行范围有些状态需要更多步。前缀最优强调的就是这种逐状态粒度而不是全任务固定一个步数。3.3 动态提前停止的通用流程虽然论文的具体实现细节需要以原文为准但从标题思路可以提炼出下面这套通用流程图。这个流程也适合作为复现时的参考框架。# dynamic early stopping for diffusion policy # 伪代码具体实现按论文和仓库调整 def generate_action(state, denoise_model, stop_model, max_steps20, threshold0.6): # 初始化噪声 x torch.randn_like(noise_template) # 从最大步数往回做去噪 for t in range(max_steps, 0, -1): x denoise_model(x, t, state) # 每一步都判断当前前缀是否已经足够好 stop_score stop_model(x, t, state) if stop_score threshold: break return x这个流程里有两个关键组件去噪网络负责从噪声中恢复动作也就是扩散策略本身。停止判断器负责评估当前前缀是否已经足够好。它可以是一个价值网络、一个置信度网络也可以是通过强化学习训练的 stopping policy。从论文标题看停止判断器的输出会直接影响采样过程因此这个模块的训练目标设计很重要。它不能只追求“快”否则会把动作截断成噪声也不能只追求“准”否则会退化成固定长采样。它需要学习效率和动作质量之间的最优平衡。4. 停止判断器可能的训练方向论文没有公开训练细节时我们只能做合理推断。下面这几个方向是扩散策略相关工作中比较常见的构造方式可以作为复现论文思想时的参考。第一种基于价值或回报的评分。训练一个价值网络或 Q 网络对于当前状态和当前前缀动作评估“如果在这里停止并执行这个动作预期回报是多少”。当预期回报超过某个阈值时停止去噪。这种做法的好处是可以和离线强化学习框架直接结合。第二种基于动作质量预测的置信度。训练一个二分类网络判断当前去噪结果是否已经接近真实数据分布。这个网络可以拿不同前缀长度的去噪结果作为训练样本标签是“该结果与最终完整结果的距离是否小于容忍范围”。第三种将停止决策本身建模成一个策略。把“继续去噪”和“停止并输出当前动作”当成两个动作用强化学习训练一个 meta-policy。奖励同时考虑控制回报和负的推理步数惩罚让模型自己权衡。不管采用哪种方式训练数据里都要包含不同状态、不同前缀长度下的动作质量信息。有一个可以预见的工程点是需要保存扩散采样过程中每个前缀步骤的中间动作用于构造停止判断器的训练集。这会增加一些存储开销但换取的是推理阶段的动态效率。5. 复现前环境准备复现这类论文环境准备比跑普通深度学习项目要麻烦一些。核心依赖包括 PyTorch、MuJoCo、D4RL 数据集、Gym/Gymnasium以及可能用到的离线强化学习工具库。以下是一套常见的 conda 环境准备命令具体版本号需要根据 PyTorch 和 MuJoCo 的最新兼容情况微调。conda create -n diffusion_policy python3.10 -y conda activate diffusion_policy # 安装 PyTorch按自己的 CUDA 版本选择 pip install torch --index-url https://download.pytorch.org/whl/cu121 # 安装环境与数据集工具 pip install gymnasium pip install mujoco pip install d4rl # 可选离线强化学习工具库 pip install d3rlpy # 其他常用 pip install numpy matplotlib tqdm wandb这里有三个容易出坑的地方第一d4rl 和 gymnasium 存在版本兼容问题。d4rl 是比较早的库设计时主要针对 gym 0.21 以下的接口。如果跑的是 D4RL 数据集可能需要单独创造一个 gym 0.21 的环境或者通过兼容层把 d4rl 的数据接口接到 gymnasium 上。更稳妥的做法是看论文作者给的 requirements.txt。第二MuJoCo 的安装方式变化较大。旧版本需要 mujoco-py 和独立的 license key新版本可以直接用pip install mujoco但 API 和渲染方式有差异。复现时建议先跑通环境中自带的一个简单示例确认 MuJoCo 环境能正常 reset 和 step。第三D4RL 数据集下载依赖网络。如果下载失败可以先检查网络连通性或者手动把数据集文件放到~/.d4rl/datasets/目录。6. 复现目录组织与训练流程复现论文时建议从一开始就把目录结构整理好。下面是一个可以参考的项目布局diffusion-policy-early-stop/ ├── config/ │ └── halfcheetah.yaml ├── datasets/ ├── models/ │ ├── diffusion_policy.py │ ├── stop_critic.py │ └── value_net.py ├── scripts/ │ ├── train_policy.py │ ├── train_stop_critic.py │ └── evaluate.py ├── checkpoints/ ├── results/ └── README.md训练流程可以分成两个阶段。第一阶段训练扩散策略。这一阶段先不要考虑提前停止重点是把扩散策略本身的控制能力训练到稳定水平。只有策略本身够强后续做动态停止的对比才有意义。# 训练流程模板先学扩散策略再训练停止判断器 # 需要按实际代码库替换 dataset 和模型类 import torch from torch.utils.data import DataLoader from diffusion_policy import DiffusionPolicy from dataset import load_d4rl_dataset # 加载 D4RL 数据集 data load_d4rl_dataset(halfcheetah-medium-expert-v2) loader DataLoader(data, batch_size256, shuffleTrue) policy DiffusionPolicy( observation_dim17, action_dim6, max_timesteps20, devicecuda ) optimizer torch.optim.Adam(policy.parameters(), lr3e-4) for epoch in range(200): for obs, action in loader: obs obs.cuda() action action.cuda() loss policy.compute_loss(obs, action) optimizer.zero_grad() loss.backward() optimizer.step() print(fepoch {epoch} loss {loss.item():.4f})第二阶段训练停止判断器。这一阶段需要使用扩散策略生成多个前缀步骤的中间结果并记录每个前缀步骤对应的动作质量再用这些数据训练停止判断器。# 训练停止判断器模板 # 思路生成不同前缀长度的动作标注质量训练停止评分器 from models.stop_critic import StopCritic stop_critic StopCritic( observation_dim17, action_dim6, hidden_dim256 ) for epoch in range(100): for obs, action in loader: # 对每个 obs 采样多条去噪轨迹 prefix_actions, prefix_steps policy.sample_prefixes(obs, max_steps20) # 计算每个前缀动作与真实动作的距离作为质量信号 quality -((prefix_actions - action[:, None, :]) ** 2).mean(dim-1) # 训练停止判断器 loss stop_critic.fit(obs, prefix_actions, prefix_steps, quality) print(fstop critic epoch {epoch} loss {loss.item():.4f})这里只是伪代码实际使用时需要根据论文的停止判断器设计来调整。7. 推理与批量评估如何量化效率提升论文实验的核心指标通常分成两套一套是控制性能指标比如 D4RL 分数或平均回报另一套是推理效率指标比如平均去噪步数、单步推理耗时、回合总耗时。如果只报告分数不报告步数节省情况就无法体现动态早期停止的价值。评估脚本里值得记录的数据项包括记录项说明D4RL score控制性能用于和固定步数基线对比平均提前停止步数每个 episode 里实际去的去噪步数均值被截断的步数占比反映停止判断器的活跃程度单个动作推理耗时包括停止判断和去噪总耗时额外计算开销停止判断器本身带来的延迟是否可忽略下面是一个批量评估的 shell 循环模板。用多个随机种子跑输出结果统一存成 JSON方便后续画曲线和对比表格。for seed in 101 202 303; do python evaluate.py \ --env halfcheetah-medium-expert-v2 \ --ckpt ./checkpoints/policy_${seed}.pt \ --stop_ckpt ./checkpoints/stop_critic_${seed}.pt \ --max_steps 20 \ --early_stop_threshold 0.6 \ --seed ${seed} \ --save_path ./results/seed_${seed}.json done评估时可以设置三个对比组固定 20 步完整采样代表最保守的基线。固定 5 步采样代表暴力压缩步数的基线。动态早期停止论文方法的核心对比组。如果动态早期停止的效果确实如论文所提它的性能应该接近甚至超过固定 20 步同时平均步数远低于 20 步。如果动态停止后的性能反而大幅低于固定 5 步就要怀疑停止判断器的训练信号或网络容量可能有问题。8. 资源占用与性能观察方法从工程角度看扩散策略的显存占用主要体现在去噪网络的前向计算上。MLP 结构的策略网络通常不大但扩散模型的训练会涉及多个时间步的噪声预测梯度计算峰值会比较明显。批量推理时显存占用还会受 batch size 影响。动态早期停止的主要收益在于降低总计算量而不是直接降低显存峰值。停止判断器本身是轻量网络前向计算的开销通常远小于一次去噪。如果观察到停止判断器的推理时间反而拖慢整体流程需要考虑是否它的网络过大或者没有放到 GPU 上运行。可以用下面这种方式观察每个 episode 的去噪步数分布python evaluate.py \ --env hopper-medium-v2 \ --ckpt ./checkpoints/policy.pt \ --stop_ckpt ./checkpoints/stop_critic.pt \ --log_steps_distribution日志中会出现类似“step_1: 12%, step_3: 30%, step_5: 25%, step_10: 5%”的分布信息。如果大部分状态集中在很短的步数说明提前停止机制生效如果几乎所有状态都跑到最大步数说明停止判断器没有学会“该停就停”可能需要调整判断阈值或训练信号。需要注意去噪步数减少并不等于控制频率一定上升。控制频率还取决于单步去噪的时间和停止判断的额外开销。所以评估时要直接测量“从输入状态到输出动作”的端到端时间。9. 常见问题与排查方法复现这类论文遇到的问题往往不在算法本身而在环境搭建和数据接口。下面整理一张排查表。问题现象可能原因排查方式解决方案安装 d4rl 后无法导入与 gymnasium 版本不兼容查看 import 报错信息创建 gym 0.21 环境或找兼容层环境创建成功但 reset 报错MuJoCo 版本与代码接口不匹配运行 MuJoCo 官方示例切换 mujoco 或 mujoco-py 版本D4RL 数据集下载失败网络或证书问题检查数据集目录是否生成文件手动下载并放到 ~/.d4rl/datasets/训练时显存不足batch size 过大或扩散步数过多观察训练时的显存占用降低 batch size或降低模型宽度提前停止后动作明显抖动停止判断器阈值过低输出不同阈值下的评估曲线提高阈值或重新训练停止判断器所有状态都跑到最大步数停止判断器未收敛检查训练 loss 和数据集标注增加训练轮次检查质量信号动态停止后性能低于固定短步数停止判断器本身引入了错误决策单独评估停止判断器准确率增加训练样本或简化判断网络其中“所有状态都跑到最大步数”是复现早期停止类工作最常见的问题。一个实用的调试方法是先对停止判断器做可视化采样一批状态绘制不同前缀步骤下的停止分数曲线。如果曲线几乎没有区分度说明训练数据中的质量信号构造有问题。10. 最佳实践与使用建议从复现和落地两个角度给出下面几条建议。先跑通固定步数基线。动态早期停止是建立在扩散策略本身已经稳定的前提上的。如果固定步数策略的 D4RL 分数都不对先解决策略训练问题再叠加停止判断器。保留最小可运行配置。固定步数 5 步和 20 步各保存一份 checkpoint动态早期停止再保存一份。这样后续对比实验可以随时回到任意一个基线。固定随机种子。强化学习环境本身有随机性评估时必须让多个种子在相同条件下对比。不同种子的结果差异可能比方法差异还大。记录每步数据。训练停止判断器需要保存扩散采样过程中的中间动作。建议同时保存的前缀步骤、中间动作和最终动作差距避免后续重新采样浪费时间。批量任务要加失败重试。批量评估时如果某个 episode 因为环境异常或 CUDA 错误中断脚本要能跳过并记录原因不能让整个评估流程卡死。注意数据许可和使用边界。如果使用 D4RL 等公开数据集注意数据集许可和引用规范。如果后续把这种方法迁移到真实机器人或者与人相关的数据上需要考虑数据来源授权、算法失败可能导致的安全风险以及部署前充分测试。接口化不是论文自带能力。如果想把动态早期停止策略封装成服务或接入现有控制框架需要自己写一层推理接口。建议先把推理脚本整理成函数输入状态、输出动作和实际步数再考虑 HTTP 或 gRPC 包装。11. 总结与下一步这篇论文最值得关注的点是把扩散模型的采样过程从“固定全长”改成了“前缀最优动态停止”。这个思路在连续控制里很有价值因为机器人控制对推理延迟敏感而扩散策略的推理慢一直是实际部署的痛点。复现这篇论文时最先应该验证的是停止判断器是否真的学到了数据分布里的“何时停止”而不是在固定步数上做无意义的步数缩短。可以从一个简单环境开始比如 HalfCheetah 或 Hopper先跑通固定步数基线再叠加动态早期停止对比平均步数和控制分数之间的权衡。最容易踩的坑不在算法而在环境和数据接口。d4rl 和 gymnasium 的版本冲突、MuJoCo 的版本变动、数据集下载失败这些前置问题会消耗大量时间。先把环境跑通再谈模型和算法。后续可以扩展的方向包括把动态停止策略和扩散采样器加速方法结合比如 DPM-Solver 这类加速采样器叠加早期停止把停止判断器做成更轻量的模块部署到真实机器人上或者在训练阶段就引入效率惩罚让扩散策略本身学会生成更适合早停的动作模式。如果要做基于扩散策略的连续控制方向这篇论文的思路值得收藏备查。下一步建议先把固定步数基线跑出来再加上动态停止用数据判断收益。