AI 中文总结
WAM-Diff2通过三阶段分层蒸馏策略,将预训练自回归VLA模型转化为高效扩散模型,缓解暴露偏差、实现与基线相当性能,解码速度提升2.8倍,结合系统级优化后达15.1倍加速。
AI 中文摘要
视觉-语言-动作(VLA)模型已成为自动驾驶端到端的突出范式,但其高效部署受限于高计算延迟和顺序自回归解码产生的暴露偏差。专用扩散策略虽能实现低延迟并行执行,但从头训练通常得到狭窄的单任务架构,缺乏整体视觉-语言推理能力。将预训练的自回归通用模型转化为并行扩散模型,可结合多任务认知智能与执行效率,但因注意力模式(因果与双向)不匹配、优化目标不同,该转换存在巨大架构挑战。为解决此问题,我们提出WAM-Diff2,这是一个基于三阶段分层蒸馏策略的多任务离散扩散VLA框架。通过分阶段的块级适配、块级蒸馏和模型级跨尺度蒸馏来构建架构转换,WAM-Diff2在保留基础模型底层语义基础的同时加快推理速度。在驾驶理解、感知和规划基准上的大量评估表明,WAM-Diff2可有效缓解暴露偏差,达到与自回归基线相当的性能。关键的是,自回归到扩散的转换实现了2.8倍的解码速度提升,结合FlashInfer和CUDA Graphs等系统级优化后,最终加速比可达15.1倍。
英文摘要
Vision-Language-Action (VLA) models have emerged as a prominent paradigm for end-to-end autonomous driving; however, their efficient deployment is severely constrained by high computational latency and exposure bias arising from sequential autoregressive decoding. Conversely, while specialized diffusion policies enable low-latency, parallel execution, training them from scratch typically yields narrow, single-task architectures that lack holistic visual-linguistic reasoning. Successfully transforming pre-trained autoregressive generalists into parallel diffusion models could combine multi-task cognitive intelligence with execution efficiency, yet this transition presents a formidable architectural challenge due to mismatched attention patterns (causal versus bidirectional) and divergent optimization objectives. To bridge this divide, we introduce WAM-Diff2, a multi-task discrete diffusion VLA framework powered by a three-stage hierarchical distillation strategy. By structuring the architectural shift through progressive block-wise adaptation, block-wise distillation, and model-wise cross-scale distillation, WAM-Diff2 preserves the underlying semantic foundations of the base model while accelerating inference. Extensive evaluations across driving understanding, perception, and planning benchmarks demonstrate that WAM-Diff2 effectively mitigates exposure bias and achieves performance parity with autoregressive baselines. Crucially, the autoregressive-to-diffusion transition yields a 2.8x decoding speedup, which scales to an ultimate 15.1x acceleration when combined with system-level optimizations including FlashInfer and CUDA Graphs.