跳转至

Social-Mamba: Socially-Aware Trajectory Forecasting with State-Space Models

会议: ECCV2026
论文: ECCV 原文
代码: https://github.com/vita-epfl/Social-Mamba
领域: 时间序列
关键词: 人体轨迹预测、社交交互、状态空间模型、连续双向扫描、多假设预测

一句话总结

Social-Mamba 将无序邻居关系组织成可扫描的社交网格,用 Cycle Mamba 支撑时间、自我和目标三路交互,在 NBA-Full 上以 1.9M 参数获得 0.72/0.92 米 ADE/FDE,但低计算量不等于所有场景下延迟最低。

研究背景与动机

预测一个人的下一段行走轨迹,不只是延长他刚才的速度方向。 在拥挤通道里,迎面来人的位置会改变避让动作;在篮球场上,队友和防守者又会改变切入路线。 Social-LSTM 用邻居池化引入这种影响,图网络通过消息传递显式表示人与人的关系,Transformer 则让不同智能体之间直接交换信息。 这些方法的共同问题是,预测既要保留每个人的运动历史,又要理解其他人如何影响当前被预测对象。 尤其是全连接注意力,智能体数量增加时,成对交互的计算与存储成本按平方增长。 因此,交互表示的组织方式不仅影响误差,也关系到拥挤场景下能否实时执行。

Mamba 的选择性状态空间模型可以沿序列累计信息,提供线性序列扫描的替代方案。 但人的集合没有天然的先后顺序,把邻居随意排成一列,会让索引顺序被误当成社会关系。 另一个障碍是信息方向:普通单向扫描只能利用此前处理过的 token,而当前行为往往需要结合整段已知历史才能理解。 既有轨迹方法因此经常只用 Mamba 编码时间维度,再把智能体交互交给其他结构。 即使采用两个方向的 Mamba,若两路状态独立、最后才相加,方向之间也没有递归状态的直接传递。 本文要解决的不是把注意力算得更快,而是让社交关系本身变成适合状态空间模型处理的序列。

作者选择以一个目标智能体为中心整理附近轨迹,再把交互拆成个体运动、当前位置影响和预测终点影响。 这样,同一邻居历史可以在不同语义锚点下被读取,而不必构造完整的全连接注意力矩阵。 需要注意,这里的目标中心交互不是额外输入真实目的地,而是围绕预测时域末端的表示进行信息汇聚。 核心 idea:用具有语义锚点的三路序列替代无结构的邻居集合,并让正反方向共享一条连续状态流,再按场景动态融合各路信息。

方法详解

整体框架

输入是目标智能体及其邻居的二维观测轨迹,输出只针对该目标智能体,而不是一次性联合生成所有人的未来。 社交网格首先补齐预测时域的空位置;Cycle Mamba 与三路交互随后并行提取不同关系;门控融合与解码最后跨智能体汇总并输出多条候选轨迹。 为避免原文用同一字母表示邻居数量和特征维度,本笔记将保留的智能体数量记为 \(N\)、总时域记为 \(T=T_{obs}+T_{pred}\)、嵌入维度记为 \(d\)。 编码后的网格形状为 \(N\times T\times d\),三个交互分支最终都恢复到这一形状,才可以逐位置融合。 邻居的预测时域在此是潜在表示槽位,不是已经观测到的未来,也不意味着模型已为邻居提供真实轨迹。 图中的虚线表示训练监督;推理过程不访问未来真值。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400, 'subGraphTitleMargin': {'top': 8, 'bottom': 16}}}}%%
flowchart TD
        Input["目标与邻居的<br/>二维观测轨迹"] --> Grid["社交网格"]
        subgraph Triplet["Cycle Mamba 与三路交互"]
                direction TB
                Temporal["时间扫描"]
                Ego["自我中心扫描"]
                Goal["目标中心扫描"]
        end
        Grid --> Temporal
        Grid --> Ego
        Grid --> Goal
        Temporal --> Fusion["门控融合与解码"]
        Ego --> Fusion
        Goal --> Fusion
        Fusion --> Output["目标智能体的<br/>K 条候选轨迹"]
        Output -.-> Loss["训练:best-of-K MSE"]
        Truth["未来真值<br/>仅训练时使用"] -.-> Loss

关键设计

1. 社交网格:先规定谁参与交互、每个 token 代表什么

模型先在最后一个观测时刻围绕目标智能体选取空间邻居,正文给出的局部半径示例是 10 米。 这是当前时刻的筛选条件,不是对整个预测过程持续更新的邻居图,也不是已确认所有数据集都采用同一个阈值。 每个保留的智能体占据一行,列方向依次排列观测时刻和待预测时刻。 观测槽位填入二维坐标,未来槽位以零向量初始化,再由 MLP 投影成特征。 由此形成的“网格”是智能体与时间的张量布局,而不是把场景离散成带占用标签的二维地图。 统一布局让时间扫描可以逐行工作,也让最后的全局扫描能够沿智能体轴交换信息。 目标身份则通过后续插入的当前位置 token 与终点 token 显式进入邻居序列。

论文图 3 将这一准备过程称为经过排序的自我中心网格,但正文未交代具体排序键及并列情况的处理。 因此不能进一步断言它按距离、方位角或球员角色排序,也不能把模型视为已证明的置换不变集合网络。 能够确定的是,它不再把邻居当作完全缺少目标参照的普通序列,而是提供局部筛选和明确的目标锚点。 这种设计的适用范围也由输入决定:本文主要使用坐标轨迹,没有将语义地图或场地边界作为显式输入。 预测终点对应的初始 token 来自零填充槽位的编码,不代表外部给定的导航意图。 区分“终点位置的表示”与“已知终点坐标”,是理解下一步目标中心交互的关键。

2. Cycle Mamba 与三路交互:连续状态流读取不同语义锚点

Cycle Mamba 的输入是一段特征序列,它先构造逆序副本,再把原序列接在后面,交给同一个 Mamba 连续处理。 下面沿用第 4.1 节式 (4) 的序列顺序,用清晰的记法表示该构造:

\[ S_{cycle}=[\operatorname{reverse}(S);S]=(s_L,\ldots,s_1,s_1,\ldots,s_L). \]

关键不在于简单重复输入,而在于处理到连接点时不重置隐藏状态。 因此,正序段的第一个 token 已经接收到逆序段对整段输入的压缩记忆。 扫描完成后,将前半段输出翻回原顺序,再与后半段对齐合并;正文给出的合并方式是逐元素相加。 常规双向 Mamba 的两路只在输出处碰面,这里则在递归状态层面已经发生了信息传递。 扫描长度变为 \(2L\),但使用一套扫描参数;参数共享并不等于算力或延迟自动减半。 而所谓双向上下文是输入序列内部的上下文,观测之后的槽位依然没有未来真值。 原文引言和图 2 图注对前后方向的表述与第 4.1 节存在差异,本笔记采用式 (4) 及其后关于正序段继承状态的解释。

三路交互并行接收同一份初始网格,各自用 Cycle Mamba 扫描,不是把一条分支的输出顺次喂给另一条。 时间扫描逐智能体读取整个时间轴,保留其自身运动模式,并让观测信息影响预测槽位。 自我中心扫描把目标在 \(T_{obs}\) 的当前位置 token 插入每个邻居的观测段与预测段之间。 这个插入位置使邻居过去的运动与目标当前状态在同一条扫描中建立联系,而不是事后简单拼接两个独立向量。 扫描后,模型用可学习的加权求和聚合相关位置和邻居的信息,再去掉额外插入的 token,恢复原来的时间长度。 目标中心扫描则把目标在 \(T\) 的终点槽位 token 附加到每个邻居序列末尾,再做同样的扫描和聚合。 它让邻居信息围绕预测终点表示组织起来,补充当前位置锚点无法充分表达的终点相关影响。 在基础模型的输入端,这仍是潜在终点表示,不是先由另一个网络估计出来、再作为坐标条件输入的目的地。 两个社会交互分支最终与时间分支同形,分别保存不同上下文,而不在这一阶段强制融合。 缓存中式 (12)、式 (13) 的索引与运算符抽取受损,无法可靠复原精确聚合实现;这里依据正文说明解释其加权汇聚与形状恢复,不猜补作者公式。

3. 门控融合与解码:先选择交互来源,再跨智能体整合

三个分支不是等权相加:模型先拼接它们的表示,用 MLP 和 Softmax 计算对应的三路权重。 随后将时间、自我中心、目标中心表示按权重逐元素加权求和,得到融合网格。 这些权重依赖输入上下文,而不是给整个数据集设定三个固定超参数。 这样,独立运动时可以更多依赖个体历史,交互密集时则可以提高社会关系分支的贡献。 这也是并行分支的意义:融合器能够同时比较多种解释,而不必接收已经被后续分支改写过的单一表示。 作者将这一机制称为 social gate;门控的价值由有无可学习权重的消融验证,但正文没有给出逐场景权重的定量解释性评测。

门控之后还要沿智能体轴进行一次全局 Mamba 扫描,补充各行之间的信息交换。 时间维度上的三路扫描与智能体维度上的全局扫描承担不同职责,不能把前者视为已经等价完成所有成对消息传递。 最后只抽取目标智能体的融合表示,送入由双向 Mamba 和 \(K\) 个 MLP 投影头构成的解码器。 这些投影头输出 \(K\times T_{pred}\times 2\) 的候选坐标序列,不是仅预测一个终点再线性插值。 全局扫描仍采用顺序状态传递,因此线性扫描成本并不构成对任意邻居排列都等价的保证。 若要预测场景里所有人,需要针对目标选择重新组织计算或另行设计共享机制;正文没有给出全场联合预测的复杂度结论。

一个完整示例

以 NBA-Full 协议为例,模型看到 10 帧历史,需要预测接下来 20 帧,分别对应 2.0 秒与 4.0 秒。 以下只解释信息流,不虚构某次实验中的坐标、预测概率或门控权重。 假设当前要预测一名切入球员,附近防守者在其运动方向上形成潜在阻挡。 社交网格先为保留的球员及其他智能体建立 30 个时间槽位,后 20 个槽位填零。 时间分支读取每个人的运动趋势;自我中心分支在邻居的历史和未来槽位之间插入目标的当前 token。 目标中心分支在邻居序列末端再放入目标的终点槽位 token,使终点相关表示可以汇聚邻居影响。 Cycle Mamba 的连续双向处理发生在这些已构造好的序列上,而不是偷看比赛后续录像。 门控融合后,全局扫描跨智能体传播信息,解码器输出 20 条候选未来;这里的 20 是默认假设数,恰好也等于 NBA 的预测帧数。 训练时只有最接近真实未来的一条候选主导该样本的误差;测试时没有真值帮助模型挑选最终路线。

损失函数 / 训练策略

模型采用 best-of-\(K\) 均方误差:分别计算每条候选与真实未来之间的平均平方距离,只对最小者施加该样本的训练损失。 下面是按第 4.3 节文字定义重述的式 (17),用预测区间内的局部时间索引简化记号,并非照录受损公式:

\[ \mathcal{L}_e=\min_{k\in\{1,\ldots,K\}}\frac{1}{T_{pred}}\sum_{\tau=1}^{T_{pred}}\left\|\hat{y}_{e,k,\tau}-y_{e,\tau}\right\|_2^2. \]

批次目标是各样本损失的平均值,默认 \(K=20\)。 它鼓励至少存在一条接近真实未来的轨迹,但没有单独保证所有候选都可行、互不重复或具有校准概率。 因此“多条输出”与“完整刻画未来分布”不能直接等同,best-of-\(K\) 评测也不能替代部署时的候选选择策略。 正文说明 NBA-Full 上从头训练,不依赖对比方法使用的大规模外部预训练;效率测试采用单张 NVIDIA A100,训练批大小为 128,推理批大小为 1。 作者还将其作为 MoFlow 的社交编码器替换项,用于流匹配框架,但没有因此把基础 best-of-\(K\) 模型变成流匹配模型。 相关实现细节指向附录;本次全文缓存止于第 18 页参考文献,没有附录,无法核验优化器、完整训练日程和 MoFlow 接口细节。

实验关键数据

主实验

默认指标为 minADE20 和 minFDE20,单位均为米、越低越好;前者对预测时域的欧氏距离取平均后选择最佳候选,后者只看终点并选择最佳候选。 两种指标的最佳候选不必是同一条,训练用平方距离也与这两个评测量不同。 下表摘自原文表 1(第 10 页),使用 NBA-Full 的 10 帧观测、20 帧预测协议。 星号表示使用大规模外部轨迹数据预训练,因此表内同时包含架构与训练数据条件不同的方法。

方法 ADE ↓(米) FDE ↓(米) 参数量 ↓(M) GFLOPs ↓
Social-Transmotion 0.78 1.01 2.0 0.87
Multi-Transmotion* 0.75 0.97 5.7 0.87
OmniTraj* 0.73 0.94 7.5 1.45
Social-Mamba 0.72 0.92 1.9 0.66

Social-Mamba 相比 OmniTraj 的绝对 ADE 优势为 0.01 米,参数量和 GFLOPs 的差异则明显更大,不能把精度提升描述成数量级变化。 其他场景中,表 3(第 11 页)的 SDD 结果为 0.25/0.38,NMRF 为 0.25/0.39;协议是 8 帧观测、12 帧预测,采用米而非像素。 表 4(第 11 页)的 JRDB 在 4.8 秒预测时域上为 0.13/0.21,NMRF 为 0.15/0.23;其输入是 9 帧观测、输出是 12 帧预测。 原文“五个基准”的概括涉及多个 NBA 设置,不应理解成五个完全独立的数据来源。

消融实验

下表摘自表 8(第 14 页)的交互模块消融,采用正文 NBA-Full 消融设置,报告默认 \(K=20\) 的 ADE/FDE(米)。 时间扫描与全局扫描始终保留,所以它检验的是另外两个社会交互分支的增益,而不是证明前两者不可缺少。

配置 时间扫描 自我中心扫描 目标中心扫描 全局扫描 ADE ↓ FDE ↓
基础交互 保留 移除 移除 保留 0.735 0.939
加入当前位置锚点 保留 保留 移除 保留 0.729 0.928
加入终点锚点 保留 移除 保留 保留 0.727 0.928
完整模型 保留 保留 保留 保留 0.719 0.919

两路单独加入都有效,同时保留时最好,支持它们在此设置中的互补性;表内没有置信区间,不能据此宣称统计显著。 表 11(第 14 页)中,等权相加为 0.744/0.954,可学习权重为 0.719/0.919,说明融合方式也是整体收益的重要来源。 表 9(第 14 页)中,常规双向 Mamba 为 0.741/0.948、2.4M 参数,Cycle Mamba 为 0.719/0.919、1.9M 参数;整模型参数减少并非减半。

关键发现

效率不能只用参数量代替;下表摘自表 7(第 12 页),测量条件为单张 NVIDIA A100、推理批大小为 1。 “模型内存”沿用论文表头,不是完整的峰值运行显存。

模型 推理时间 ↓(ms) 模型内存 ↓(MB)
Social-Transmotion 1.8 7.6
Multi-Transmotion 7.3 21.8
Social-Mamba 3.4 7.3
  • Social-Mamba 比 Multi-Transmotion 更快、更小,但比 Social-Transmotion 更慢;短序列下注意力的并行实现仍有优势。
  • 表 5(第 11 页)在 NBA-LED 的 MoFlow 编码器替换中报告 0.71/0.87 → 0.70/0.85,参数量为 1.3M → 0.5M;正文“2.3 倍更小”与这些四舍五入数值的比值不一致,宜分别保留,不强行修正。
  • NBA-Full 的表 1 与消融表使用不同精度展示同一量级结果,0.72/0.92 与 0.719/0.919 不应误认成不同性能结论。

亮点与洞察

  • 将“插入谁的 token、插在什么位置”变成社交建模手段,比直接把邻居压成一个池化向量更有结构。当前位置与终点槽位提供两种不同的关系参照。
  • Cycle Mamba 改变的是状态流的连接方式,而不只是多跑一个方向。其可迁移价值在于需要整段输入上下文的编码任务,但在线因果任务不能不加修改地照搬。
  • 三路分支保留各自表示后再门控,给模型留下按场景选择信息来源的空间。消融比单独的总榜结果更能说明这一设计为何值得保留。

局限与展望

  • 作者在 NBA 定性结果中承认偶发越界,图 4(i) 给出例子;没有显式地图或边界约束时,低平均误差并不保证路径可执行。
  • 从评测角度看,minADE/minFDE 只评价候选集合中的最佳项,没有充分验证碰撞率、候选概率校准或多人联合预测的一致性。这些是本笔记提出的后续验证方向。
  • 固定局部邻居筛选可能忽略当前较远、稍后进入交互范围的对象;可研究动态邻居选择,但这不是论文已经验证的改进。
  • 原文未充分说明邻居排序规则,且缓存缺少附录、部分公式抽取损坏;复现前应核对源码中的排序、token 聚合与训练配置,而不是把本文解释当成完整实现规范。

相关工作与启发

  • vs Social-LSTM / Trajectron++:前者采用社交池化,后者使用图结构建模动态交互;Social-Mamba 则通过序列锚点和状态扫描表达关系。它避免全连接注意力,但不是无损复现任意图上的显式消息传递。
  • vs U2Diff / Sports-Traj:原文将这些方法的 Mamba 使用概括为主要处理时间依赖;本文把 Mamba 的作用进一步扩展到带目标语义的邻居交互,而不只是替换时间编码器。
  • vs MambaPTP:根据原文第 3 页讨论,其社会交互主要位于解码阶段并采用通用邻居扫描;本文在编码阶段构造三路交互。这里是作者的比较定位,未另行核验该方法全文。
  • 与 MoFlow 的关系:流匹配负责生成机制,Social-Mamba 提供社交条件表示,两者可以组合。一个值得测试的方向是加入地图条件后,检查更好的编码是否同时降低最佳误差与越界率。

评分

  • 新颖性: 4/5。连续双向状态流与语义锚点扫描结合紧密,不只是将 Transformer 层替换为 Mamba。
  • 实验充分度: 4/5。覆盖多种数据设置、模块消融与效率比较,但缺少排序鲁棒性、概率校准和安全性指标。
  • 写作质量: 3/5。总体架构清晰,但扫描方向表述不一致,排序细节不足,当前缓存还存在公式抽取损坏。
  • 价值: 4/5。为轻量社交轨迹编码提供可复用方案,实际部署仍需核验延迟条件、候选选择与场景约束。