首页 / 资讯中心 / 文章详情

MeanFuser:单步扩散蒸馏,实现434FPS自动驾驶轨迹规划

MeanFuser:单步扩散蒸馏,实现434FPS自动驾驶轨迹规划 ★ FEATURED ARTICLE
自动驾驶的轨迹预测与生成我一直觉得是个“能跑起来不难、跑得快才要命”的活。特别是需要输出多条未来轨迹假设多模态的时候以扩散模型为代表的一类生成式方法效果确实好但几十步迭代去噪的推理延迟放到真车上就是灾难性的。所以当看到自动化所和小米合作的 MeanFuser 这个工作——单步生成轨迹、纯规划 434 FPS、还中了 CVPR——第一反应是“终于有人正面刚这个问题了”。这篇就把 MeanFuser 为什么能这么快、它的“Mean 方向约束 蒸馏”思路到底怎么落地、以及单步轨迹生成在训练和部署里那些没人明文写的坑一次性掰开说清楚。1. 为什么轨迹生成必须“极速”1.1 多模态轨迹生成到底在解决什么问题先对齐一下任务背景。在城市道路或者高速场景里自车周围每辆车的未来行为都不是唯一的。举个例子你准备变道旁车可能减速让你、可能加速顶上来、也可能保持当前速度。作为决策层规控系统需要的不只是一条“最可能轨迹”而是一组覆盖多种行为的候选轨迹并且最好还带置信度。这就是“多模态轨迹生成”的由来给定目标车历史轨迹、自车状态、地图上下文模型要输出 K 条未来轨迹比如 K6彼此之间差异要足够大、又都要合理。这事以前常用 CVAE、GAN后来 Transformer 加 DETR 风格的目标查询learnable query也能做效果不错。但真正把“多样性和真实性”同时拉满的是扩散模型。扩散在生成的每一步都在逐步“去噪”天然能覆盖多种模式生成的轨迹也更细腻不会像 CVAE 那样容易崩塌成平均轨迹。代价也明摆着慢。1.2 434FPS 是什么概念纯规划帧率怎么读先拆一下 FPS。434 FPS 意味着单次前向推理大概 2.3 毫秒1000/434 ≈ 2.3ms。这个数字乍看不稀奇关键在“纯规划”三个字。日常我们说一个感知模型多少 FPS往往包含了预处理、后处理、多模型串行整个链路。纯规划帧率指的是从特征进来到轨迹头输出候选轨迹整个 DNN 模块的推理吞吐。也就是说它在真机上给其他环节留出了极大量的算力富余。可以横向对比一下感受量级一般基于 Transformer 的轨迹预测模型单次推理在 10ms~30ms 左右扩散模型要迭代几十步哪怕用 DDIM 压缩到 8 步、每步 2ms整体也要 16ms 起步而且这还没算每步之间不可避免的调度开销。MeanFuser 能把纯规划做到 2ms 级别说明它不只是把步数压到了 1 步连网络本身的体量也做了收敛设计。这个帧率放在整个自动驾驶计算平台里意味着规划模块几乎不占预算完全可以把 GPU 算力让给视觉感知、激光雷达点云处理甚至一个更大的世界模型。1.3 扩散模型“慢”的根源在哪扩散过程可以理解为先对一条标注好的轨迹逐步加噪直到变成纯噪声生成时反向操作从纯噪声出发通过一连串去噪步骤还原轨迹。每一步去噪都是一个网络前向推理。训练时步数多一些没关系可以开大 batch、充分采样但推理时每多一点步数延迟就线性上涨。这也是所有想用扩散做轨迹的人们头痛的地方DDPM 要 1000 步、DDIM 压缩到 20 步、DPM-Solver 压到 4~8 步模型已经从最初的 Unet 换成了更轻的 Transformer 或者 MLP但始终没法突破“必须迭代多次”的瓶颈。MeanFuser 选择了一条更激进的路不做 N 步去噪而是用蒸馏把“多步去噪的期望效果”压缩成单步前向直接从带噪轨迹映射到干净轨迹。下面细看它到底怎么做到。2. MeanFuser 的核心思路拆解2.1 “Mean” 这个名字到底指的是什么MeanFuser 这个名称里Mean 不是“均值函数”这么简单我理解它包含了两层意思。第一层它关注的是轨迹分布里的“均值方向”mean direction。对于决策规划来说自动驾驶系统最希望拿到的是可执行的、稳定的轨迹。扩散采样天然带有随机性两次采样的结果可能抖得厉害。直接把随机采样轨迹给到控制模块车身会晃。与其去纠结单个采样轨迹的噪声不如去学“所有可行行为里最具有代表性的那条期望轨迹”以及围绕这个期望的多模态分支。第二层是网络结构上的“Mean 特征融合”。轨迹预测通常要融合三类信息目标车历史观测、周边障碍物的交互关系、高精地图的拓扑约束。简单做法是把三类特征拼在一起过 MLP但这种粗糙融合会丢失模态间的对齐关系。MeanFuser 的做法是建立一个“融合中心”Fuser把历史轨迹编码后的 token、地图编码后的 token 和交互特征 token 在一个统一的特征空间里做多次交叉注意力cross-attention让轨迹的均值表征在多层之间被反复精炼。这个融合后的向量就是后续单步生成条件的基础。2.2 单步生成如何用蒸馏把几十步压成一步这一步是整套方法的重心。如果只是训练一个从噪声直接映射到轨迹的网络学出来会很粗糙。MeanFuser 走的是“Teacher-Student 蒸馏”框架先训练一个标准的扩散轨迹生成模型 Teacher步数可以设得比较大方比如 50 步去噪任务是把带噪轨迹还原成干净的未来轨迹损失函数是预测噪声与真实噪声之间的 MSE。然后训练一个 Student 网络输入同样带噪的初始状态但只前向一次目标是让输出轨迹尽量等于 Teacher 多步迭代后得到的最终轨迹。这里蒸馏损失用的是轨迹空间的距离而不是噪声空间的因为系统最终要的是轨迹不是中间噪声。这个思路和图像生成里的一致性模型Consistency Model很像。一致性模型最关键的一点是经过蒸馏后的单步模型输出应该满足“自洽性”——即任意时刻的带噪状态经单步映射后得到的干净状态应该一致。MeanFuser 把同样的思想搬到轨迹上但做了两个改动一是直接蒸馏到完整轨迹序列而不是逐帧映射保证时间维度的连贯性二是在蒸馏损失里叠加了动力学约束和地图约束防止单步输出越过车道边界或者产生急转弯。第三步也是容易被忽略的是蒸馏后的“校正阶段”。只做蒸馏的话学生网络在某些 corner case 上可能学得不够好所以还需要用真实轨迹数据再做一轮监督微调重点加权那些急刹、变道、路口转弯等困难样本。这一步能显著压低碰撞率和离群值。2.3 多模态怎么保得住模式坍缩是单步生成的头号敌人单步蒸馏模型最大的毛病就是“懒”既然只能输出一个结果网络倾向学一个平均轨迹把所有场景都往中庸方向输出。这在视觉上就像用一步 GAN 生成的图片会糊成一片轨迹里就是各种模式被平均成一条歪歪扭扭的线。MeanFuser 防坍缩的方法是“多查询multi-query 分配损失”。模型不是只输出一条轨迹而是并行输出 K 条候选轨迹每条轨迹都由独立的可学习查询向量引导这些查询向量在网络里会和融合后的场景特征做交叉注意力迫使不同查询关注不同行为模式。比如其中一个查询关注“旁车让行”、另一个关注“旁车加速”。训练时计算每个查询输出和当前真实轨迹之间的距离用二分图匹配或 Top-K 匹配把轨迹样本分配给离它最近的查询分支去监督。这样每个查询专注自己的模式不会大家挤在同一个模式里。为了进一步保证分支之间足够“分得开”训练损失里还加了一个“模式间距惩罚”如果两条候选轨迹之间的平均距离小于阈值就在损失函数里引入惩罚项强迫网络把模式拉开。这里经验值是阈值大概取 1~1.5 米在 8 秒预测时域下太小了没作用太大了会让有些分支硬生生拐出去。2.4 Fuser 到底“融”了什么除了上面说的特征融合Fuser 还承担了一个很重要的职责候选轨迹与最终执行轨迹之间的融合与选择。纯规划模式下系统其实不需要 K 条轨迹都送去控制最终只能选一条执行。MeanFuser 在 K 个候选轨迹后面接了一个轻量打分头综合评估每条轨迹的动力学可行性、与目标点的接近程度、碰撞风险输出一个分数。执行轨迹不是简单选最高分那条而是拿 Top-3 轨迹加权融合权重就是分数经过 softmax 之后的值这样做的目的是平滑波动。相邻两个控制周期最高分轨迹可能跳变但加权融合后的轨迹会平滑过渡这对下游控制模块是非常重要的细节。3. 训练实现与工程配置照着做能少踩一半坑3.1 输入输出与网络骨架建议这部分我不贴完整源码但把关键配置和形状定义写清楚你拿去对接自己的数据格式会非常顺利。以车辆轨迹预测为例历史观测目标车过去 3 秒、10Hz也就是 30 帧轨迹点x, y, v, heading, 加速度等序列长度 T_obs 30特征维度 D_obs 6。地图信息当前车道中心线采样点、左右车道边界点、车道连接关系这些先编码成矢量片段再通过地图编码器用轻量 MLP 注意力输出 M 个地图 token特征维度 D_map 32。交互上下文周边 10 辆车的相对位置和运动状态用一个小型交互图网络编码成 D_inter 32 的向量。输出未来 6 秒、10Hz共 T_fut 60 个轨迹点候选条数 K 6。Backbone 我建议直接沿用视觉领域验证过的那套组合历史轨迹用 GRU 或轻量因果卷积编码地图和交互用两个独立的浅层 Transformer encoder每个 encoder 两到三层、隐藏维度 128 就够了。别把 backbone 堆太大MeanFuser 的帧率优势一半靠算法、一半靠“网络真的不大”。3.2 三阶段蒸馏训练流程我建议把整个训练流程严格拆成三个阶段阶段之间不要跳级Teacher 扩散模型训练阶段用真实的未来轨迹作为 target训练一个标准的 DDPM/DDIM 模型步数建议在 1000 步加噪、50 步采样。Teacher 不需要做任何加速优化它的任务是“尽可能准”因为它是学生模型的上限。这个阶段大概占全部训练资源的 40%。Student 单步蒸馏阶段冻结 Teacher 参数。Student 输入同一个带噪初始轨迹可以用 20 步左右的加噪程度一步映射到轨迹。损失函数是Student 输出轨迹与 Teacher 最终输出轨迹之间的均方误差 轨迹一阶差分速度一致性惩罚。这个阶段占 30%。联合精调阶段解冻 Student加入真实轨迹监督、碰撞约束、候选分支匹配损失。这个阶段最重要的超参是分支匹配距离阈值和模式间距惩罚权重。这个阶段占 30%。给一段简化的伪代码意思一下单步蒸馏的核心逻辑# 单步蒸馏核心示意 for batch in dataloader: obs_feat encoder(obs_history) # 历史观测编码 map_feat map_encoder(map_data) # 地图编码 ctx_feat cross_attention(obs_feat, map_feat) # 用加噪到 t 步的轨迹作为输入 noise_trj add_noise(gt_future, t20) # Teacher 多步去噪得到参考轨迹 with torch.no_grad(): teacher_out teacher_sampling(noise_trj, ctx_feat, steps50) # Student 单步输出 student_out student(noise_trj, ctx_feat, query_bank) # (K, T_fut, 2) loss_distill MSELoss(student_out, teacher_out) loss_temp MSELoss(student_out[:, 1:] - student_out[:, :-1], teacher_out[:, 1:] - teacher_out[:, :-1]) total_loss loss_distill 0.5 * loss_temp total_loss.backward()如果你不用 Teacher 蒸馏也可以直接从最原始的“一步生成 真实轨迹监督”开始训练但那样生成质量会差不少特别是预测时域大于 4 秒以后轨迹会发飘。蒸馏最大的价值不是“学到的知识”而是“单步模型的分布起点更接近真实分布”。3.3 关键训练超参数速查表下面这些参数是我基于同类轨迹扩散模型和蒸馏模型经验给的起步值具体数值要按你的数据集微调但量级基本差不太多。参数建议值备注初始学习率1e-4AdamWwarmup 3000 步Batch size128单卡不够用梯度累积Teacher 采样步数50蒸馏阶段固定不参与梯度加噪程度 t学生输入20 步太接近纯噪声学生学不动太接近干净轨迹蒸馏没意义候选分支数 K6高速场景 4 就够城区建议 6~8模式间距阈值1.2 米超这个阈值开始惩罚分支过近EMA 衰减0.999学生模型用 EMA 版本推理更稳这里特别提一下加噪程度 t 的选择。我试过从 5 到 50 步的区间太小的 t 学生网络几乎只需要复制输入没有学到真正轨迹生成的能力太大的 t 又变成一步从噪声猜轨迹损失太大收敛不动。20 步左右是目前看下来最平衡的点蒸馏损失收敛速度和最终轨迹精度都比较好。4. 模型评估与踩坑实录4.1 指标怎么看从 ADE/FDE 到模式覆盖轨迹预测的标准指标大家应该熟悉了。平均位移误差ADE是预测轨迹和真实轨迹逐帧的平均距离最终位移误差FDE只看最后一帧的距离。多模态场景下通常报告 minADE / minFDE即 K 条候选轨迹里距离真实轨迹最近的那一条的误差。这套指标的优点是直观缺点也明显它不管其余 5 条轨迹合不合理。所以实际项目里还要看模式覆盖率和碰撞指数——即 K 条轨迹的多样性是否够、其中有多少条会和地图边界或障碍物发生碰撞。MeanFuser 这类方法在 minADE / minFDE 上肯定不如迭代 50 步的 Teacher 模型这是蒸馏的固有损失。但纯规划落地叠加上真车体验2ms 的延迟和 30ms 的延迟带来的控制品质差异远比一点点指标差距更明显。这不是加快速度弥补质量而是在轨道规划这种场景里“低延迟、平滑、可执行”本身就是质量的一部分。拿一个类似规模数据集的典型结果做参照的话可以看下面这个趋势表数值是类比量级不代表论文原表方法推理延迟minADE (m)minFDE (m)模式覆盖率Teacher DDPM 50 步~100ms0.821.40高Teacher DDIM 8 步~16ms0.881.52中高MeanFuser 单步~2.3ms0.941.65中高差距大约在 10%~15% 的精度损失换来 40 倍以上的延迟下降。对于规划模块而言只要碰撞率压得住这个买卖非常划算。4.2 单步轨迹生成最常见的三个坑坑一模式坍缩。训练中期最常见的表现是 K 条候选轨迹越来越像最后几乎完全重叠。这时候别急着调损失函数先去看模式间距惩罚的梯度方向对不对很多时候是因为分支匹配分配得太平均导致每个查询都覆盖了所有样本没有明确分工。建议先把分配损失改成“一对一硬分配”每个轨迹样本只喂给距离它最近的查询分支强制每个查询专注一片区域。坑二轨迹抖动。单步生成的轨迹在空间上可能平滑但速度轮廓上很毛糙一阶差分跳变明显。这个问题我排查了很久最后发现根因在蒸馏损失里只约束了位置没约束速度。后来在损失里同时加上了一阶差分项抖动立刻缓解了大半。如果要更细还可以加 jerk二阶差分约束控制体验会更好。坑三过拟合历史轨迹。如果训练集里多数样本是直行模型容易遇到弯道也强行直行。解决思路是在训练时对地图和交互特征做随机 dropout让模型不能只依赖其中单一强特征逼迫它综合上下文信息。我加了一个概率为 0.1 的地图 token mask直行过拟合的现象明显下降。4.3 部署阶段把帧率坐实TensorRT 与算子选择模型训练完之后真正上车的还有一道部署优化工序。434 FPS 不是单靠小模型就能在真机上轻松拿到的工程侧至少要做三件事换 TensorRT 并开 FP16。轨迹预测网络全是矩阵乘加和注意力TensorRT 优化后通常能比 PyTorch GPU 推理快 30%~60%。地图编码做静态化缓存。高精地图在一个控制周期内是不变的没必要每帧重新编码。把地图 encoder 的输出缓存起来每帧只算动态变化的交互特征。这个优化能把每次推理省下大概 0.5ms。如果 Batch 是 1还要注意把动态形状dynamic shape配置好。很多加速框架在 batch1 时因为 shape 推断耗费额外时间反而比 batch4 更慢这点要用 profiling 工具实测别想当然。说白了纯规划 FPS 是个系统工程算法层面把步数压到 1 是前提工程层面每一毫秒都得抠出来。4.4 MeanFuser 后续能往哪儿延伸这类方法的价值不止于轨迹预测。我看到的一个很自然的延伸方向是把它接在端到端大模型后面作为“规划头”。如今很多自动驾驶大模型语言模型、世界模型输出的并不直接是轨迹而是语义意图和场景描述MeanFuser 这类轻量规划头正好可以把语义条件编码成查询向量单步输出多条可执行轨迹。这样大模型负责场景理解和意图决策MeanFuser 负责把意图变成可控的、平滑的、带多模态备选的轨迹两边算力分工明确部署的时候规划部分依然是极速模块。另外如果把 MeanFuser 的蒸馏对象从“扩散模型”换成“复杂交互搜索”或者“博弈求解器”它也能成为一个通用的“策略压缩器”把一些慢但准确的离线决策算法压成一步前向的近似策略。这个思路在机器人操作、无人机编队、甚至游戏 AI 里都是通用的。按照我个人实际做轨迹预测落地的体会单步生成最大的价值不是省了推理时间本身而是把规划模块的延迟降到了一个可以让下游控制直接使用、不再需要复杂补偿的水平。以前我们做 20 步扩散预测控制端总在抱怨轨迹出来得太晚、太碎得加好多滤波器去平滑。MeanFuser 的思路换了赛道既然迭代一步是散射状的噪声到轨迹映射那就用干净明确的均值方向约束整个生成让网络一次输出就满足执行需求。这个取舍我认为比单纯刷 FPS 数字要有价值得多。将来想再接大模型规划头或者做机器人动作生成的朋友不妨就从这套“蒸馏 单步 均值约束”的打法开始试试很多问题是相通的。
阅读完成 · 觉得有帮助?
咨询建站