ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch实现NRI模型:用图神经网络推断关系并进行多智能体轨迹预测

PyTorch实现NRI模型:用图神经网络推断关系并进行多智能体轨迹预测 做轨迹预测做了有一阵子我印象最深的一件事就是模型只拿到一堆移动物体的坐标却看不到物体之间的“关系”时预测出来的轨迹往往是一团乱麻。后来接触到了图神经网络里的NRI模型才算是把这个问题想通了——与其直接硬拟合坐标序列不如先把物体间隐藏的关系推断出来再拿这个关系去辅助预测。这就是NRINeural Relational Inference的核心思路它一边推断任意两个实体之间是否存在交互、是哪种交互一边利用推断出的交互关系对系统未来的运动轨迹做预测。这篇内容我会把NRI的原理拆开讲清楚然后用PyTorch从头实现一个可以跑通的运动轨迹预测模型数据集用经典的弹簧质点物理系统整个过程都会给出可直接参考的代码。适合已经有一点深度学习基础、想搞懂图神经网络到底怎么落地或者正在做多智能体轨迹预测相关工作的朋友。这个模型在很多场景都有用武之地物理系统建模、粒子运动模拟、机器人编队规划、交通参与者轨迹预测甚至社交网络中的影响传播分析。核心都是同一件事——系统里每个物体的行为不独立而是被“关系”支配着。NRI能把这些关系从数据中挖出来这是一般时序模型做不到的。1. NRI模型核心思路与设计动机1.1 交互系统预测为什么难做先想象一个场景教室里有十个同学你要根据前几秒的座位变化预测下一秒每个人坐在哪。如果不告诉你他们之间谁和谁是朋友、谁在偷偷传纸条、谁在躲着谁你只能根据所有人的位置变化硬猜。传统做法是把所有人的二维坐标拼成一个20维的大向量丢给LSTM或者Transformer去学理论上信息量是够的但实际效果往往很差。原因出在模型默认了所有输入特征是平级的任何一个物体对另一个物体的影响方式都被同一个参数矩阵处理这种建模方式对交互系统的结构信息是一种浪费。说白了交互系统的轨迹预测难点不在于“预测”本身而在于系统内部的结构是未知的。不同物体之间的作用力大小不同、方向不同、类型不同有的相互吸引有的相互排斥有的是单向追逐有的是双向绑定。如果模型不知道这些它就只能学到一套“平均行为”预测出来的效果自然好不到哪去。1.2 把问题拆成“关系发现 轨迹预测”NRI的聪明之处在于把这个问题拆成了两个阶段而不是一口气硬学。第一个阶段是关系发现根据观测到的历史轨迹推断两两实体之间是否存在边以及这条边属于哪种类别。第二个阶段是轨迹预测在已知关系图的前提下用图神经网络对每个实体的受力情况建模逐步推出未来轨迹。这种拆分的思路和人类理解世界的方式很像。你看到两个人在走廊里并肩走不会只在脑子里记住两串坐标而是会先判断“他们认识可能在聊天”然后基于这个判断去预测他们接下来会一起转弯还是分开走。NRI把这一层“判断关系”的环节显式建模出来了这是它相对于常规端到端时序模型最大的区别。值得注意的是关系发现和目标预测是联合训练的。也就是说模型不是为了得到一个漂亮的图结构而去学图而是为了“最终预测得更准”去自动调整它对关系的判断。这个设计非常关键因为它保证了关系推断出来的图结构一定是对预测任务有帮助的而不是一个自娱自乐的中间产物。1.3 为什么偏偏用图神经网络来承载图神经网络和这个任务的匹配度相当高我总结下来有四个原因。第一图结构天然表达“关系”节点代表实体边代表交互这跟我们要建模的问题在结构上完全一致。第二GNN的消息传递机制具有置换不变性也就是说不管你把第3个物体和第5个物体换个顺序模型推断出的关系结构都是一样的这符合物理世界的直觉。第三GNN支持可变数量的节点训练时见过5个球的系统部署时换成8个球也能直接工作。第四计算效率上有优势因为消息传递可以并行计算。这里我做一个对比表格把NRI和几种常见替代方案的优劣列出来大家选型时可以参考方案是否显式建模关系能否推断关系可迁移性主要问题LSTM/Transformer直接预测否否差忽略结构信息长序列容易失真全连接GNN 固定边是否中需要人工指定边现实中做不到NRI(编码器解码器)是是好结构相对复杂训练需要技巧所以NRI的设计可以概括为一句话用可微的编码器来替代人工标注关系用图解码器来替代盲目的坐标回归。2. 模型数学原理与关键公式拆解2.1 问题定义与符号说明先把记号统一一下。假设系统里有N个实体每个实体在任意时刻t都有一个状态向量$x_t^{(i)}$比如二维坐标$(x, y)$或者三维坐标加速度$(x, y, v_x, v_y)$。我们观测了$T$步历史记为$\mathbf{x}_1, \mathbf{x}_2, ..., \mathbf{x}T$其中$\mathbf{x}t \in \mathbb{R}^{N \times D}$。目标是根据历史预测未来$T$步的轨迹$\mathbf{x}{T1}, ..., \mathbf{x}{TT}$。NRI引入了一个隐变量$\mathbf{z}$它表示N个节点之间的有向边的类型。如果一共有K种关系类型那么每条边$(i, j)$对应的隐变量$z_{ij} \in {1, 2, ..., K}$其中第一种类型通常代表“无交互”。整个模型的联合概率可以写成[ p(\mathbf{x}{1:TT}, \mathbf{z} | \mathbf{x}{1:T}) p(\mathbf{z} | \mathbf{x}{1:T}) \cdot p(\mathbf{x}{T1:TT} | \mathbf{z}, \mathbf{x}_{1:T}) ]换句话说编码器负责计算$p(\mathbf{z} | \mathbf{x}{1:T})$解码器负责计算$p(\mathbf{x}{T1:TT} | \mathbf{z}, \mathbf{x}_{1:T})$。训练目标就是最大化这个联合概率的变分下界。2.2 编码器从观测轨迹到边关系编码器的输入是历史轨迹输出是每一条边的类型概率分布。整个过程用两层GNN来做消息传递。第一步每个节点先用一个MLP把历史轨迹编码成节点特征$h_i$[ h_i f_{\text{enc_node}}(\mathbf{x}_1^{(i)}, ..., \mathbf{x}_T^{(i)}) ]这里的$f_{\text{enc_node}}$实际操作时可以先对每个时间步做MLP再做一个平均池化或者注意力池化目的是把时间维度压缩成一个固定维度的节点向量同时保留运动特征。然后进入第一轮消息传递对每条边$(i,j)$构造边特征$h_{ij}$[ h_{ij} f_{\text{edge_net}}([h_i, h_j]) ]$[\cdot, \cdot]$表示向量拼接$f_{\text{edge_net}}$是一个MLP。这一步让每条边“看”到它两端节点的信息相当于在判断这两个节点之间是否存在相互作用。接着把每个节点的邻域边信息聚合回节点做一次节点更新[ h_j f_{\text{node_net}} \left( \sum_{i \in \mathcal{N}(j)} h_{ij} \right) ]聚合方式可以选求和、平均或者最大值原文用的求和我也建议先用求和实验里效果最稳定。第二轮消息传递再执行一次“节点到边”的更新得到最终的边特征$h_{ij}$。最后经过一个线性映射得到K类关系的logits[ \phi_{ij} \text{Linear}(h_{ij}) \in \mathbb{R}^K ]对$K$维做softmax就能得到每条边属于每个关系类型的概率。整个过程非常简洁没有花哨的模块但效果出人意料地好。2.3 解码器基于关系图的轨迹预测解码器的输入是推断出的关系图$\mathbf{z}$和历史轨迹输出是未来轨迹。核心逻辑是不同关系类型对应不同的消息传递函数。举个例子如果节点i和节点j之间存在“吸引”关系那消息传递时节点j会向节点i传递一个“拉近”的贡献如果是“排斥”关系则传递“推远”的贡献。一组共享的MLP处理同一种关系下的消息这样大大减少了参数量。具体来说解码器在每个时间步$t$维护一个隐状态$s_t^{(i)}$初始状态由历史轨迹编码得到。在每一步计算消息传递时对每条边根据它的类型$z_{ij}$选择对应的消息函数$f_{\text{msg}}^{(z_{ij})}$[ m_{t}^{(i,j)} f_{\text{msg}}^{(z_{ij})}([s_t^{(i)}, s_t^{(j)}]) ]之后每个节点聚合所有入边消息[ m_t^{(i)} \sum_{j \in \mathcal{N}(i)} m_t^{(i,j)} ]聚合得到的消息和当前隐藏状态一起输入GRU完成隐状态更新[ s_{t1}^{(i)} \text{GRU}(s_t^{(i)}, m_t^{(i)}) ]最后用隐状态通过一个输出MLP预测该节点在下一时刻的状态[ \hat{\mathbf{x}}{t1}^{(i)} f{\text{out}}(s_{t1}^{(i)}) ]然后递归地把预测结果当作下一步输入一直生成到所需的未来长度。这里要强调一下解码器在跑每一步时都要重新基于关系图做消息传递所以关系图的质量直接影响预测的长期稳定性。2.4 训练目标与Gumbel-Softmax采样训练时最大的困难在于$\mathbf{z}$是离散变量直接梯度回传会很麻烦。NRI的解法是Gumbel-Softmax采样加Straight-Through估计前向传播时用Gumbel-Softmax采样得到one-hot的边类型反向传播时梯度绕过采样层跟着softmax的连续近似走。这样既能保留离散图结构的表达能力又能让模型端到端可微。训练损失由两部分组成。第一部分是预测轨迹和真实轨迹之间的均方误差第二部分是边类型预测的交叉熵——如果你手上有真实的关系标注可以直接加这个监督信号。不过NRI原始实验里即使是纯无监督也能学出正确的关系图只是在复杂数据上加上监督信号会更稳。我自己实验时发现如果有真实边类型标注在训练前中期用0.5的权重混入交叉熵是有帮助的后面再逐步退掉。最终的损失函数可以写为[ \mathcal{L} \mathbb{E}{q{\phi}(\mathbf{z}|\mathbf{x}{1:T})} [ |\mathbf{x}{T1:TT} - \hat{\mathbf{x}}{T1:TT}|^2 ] \lambda \cdot \mathcal{L}{\text{CE}}(\mathbf{z}, \mathbf{z}_{\text{true}}) ]3. 实验环境准备与合成数据集构建3.1 环境配置清单NRI的代码实现不复杂依赖也很少以下是我用起来很顺手的一套环境组合。PyTorch版本建议1.10以上太老的版本对Gumbel-Softmax的支持和自动混合精度都不太友好。其他就只需要NumPy、Matplotlib、scikit-learn这些常规库。python 3.9 pytorch 2.0 numpy matplotlib scikit-learn数据集方面我用的是NRI原始论文里最经典的弹簧物理系统Spring System。在这个系统里多个粒子在二维平面上运动粒子之间可能存在弹簧连接也可能没有。弹簧连接会产生吸引力或排斥力不同的连接方式导致了不同的运动模式。选择这个系统的原因很直接它能生成无限量的训练数据而且我们知道真正的连接关系可以准确评估模型推断图结构的准确率。这种Ground Truth明确的特点非常适合用来验证模型是否真正学到了关系推断能力。3.2 弹簧质点系统的数据生成生成逻辑不复杂核心就是简化版的分子动力学模拟。我也建议你在自己的实验里先从这个系统开始因为它能快速帮你验证模型实现是否正确。每个粒子的受力由三部分组成弹簧力、阻尼力和随机噪声。弹簧力遵循胡克定律如果两个粒子之间存在弹簧连接那么连接越短的弹簧会把他们拉近如果处于拉伸状态或者推开如果处于压缩状态数学上写作$F -k(r - r_0)$。阻尼力则让粒子的运动随时间逐渐衰减避免系统一直震荡不停。模拟时用基本的欧拉积分或者速度Verlet积分就能跑出合理的轨迹。下面是完整的弹簧系统模拟代码输入粒子数和时间步数输出所有粒子随时间变化的坐标序列import numpy as np def generate_spring_system(n_particles, n_timesteps, dt0.002, seed42): 生成弹簧质点系统的运动轨迹。 返回: traj: shape [n_timesteps, n_particles, 2] 每个粒子的坐标 edges: shape [n_particles, n_particles] 真实的弹簧连接关系(0/1) edge_types: shape [n_particles, n_particles] 边类型(0无连接, 1弹簧连接) rng np.random.default_rng(seed) # 随机生成弹簧连接每个粒子大约有2-3个邻居 # 这里用一个简单方法对每个节点随机连一条出边和一条入边 edges np.zeros((n_particles, n_particles), dtypenp.int64) for i in range(n_particles): # 随机选一个和自己不同的节点 j rng.choice([x for x in range(n_particles) if x ! i]) edges[i, j] 1 # 再加一条无向的感觉对称一下 edges np.maximum(edges, edges.T) # 对称化 edge_types edges.copy() # 初始化位置和速度 pos rng.randn(n_particles, 2) * 0.5 vel rng.randn(n_particles, 2) * 0.1 traj [] # 记录每个粒子的受力 for t in range(n_timesteps): traj.append(pos.copy()) forces np.zeros_like(pos) for i in range(n_particles): for j in range(n_particles): if edges[i, j] 0: continue diff pos[j] - pos[i] dist np.linalg.norm(diff) 1e-8 # 弹簧力胡克定律 force_magnitude 0.5 * (dist - 1.0) force force_magnitude * diff / dist forces[i] force # 阻尼力 forces - 0.5 * vel # 欧拉积分 vel forces * dt pos vel * dt return np.stack(traj), edges, edge_types这段代码生成的轨迹会呈现出非常丰富的动态有些粒子被弹簧拉在一起做周期运动有些粒子因为初始速度较大而脱离连接范围系统整体看起来像一群随机飘动的小球。我跑完生成脚本会立刻用Matplotlib画一个散点轨迹图确认轨迹没有发散、粒子之间确实存在明显的耦合运动然后才开始做模型训练。3.3 数据加载与批处理数据加载部分需要注意的是NRI的输入格式是[batch, time, node, feature]。这种格式和常规的LSTM输入格式不太一样很多同学第一次写DataLoader时会搞混。batch维度在最前面然后是时间步再是节点数最后是特征维度。因为图神经网络要对节点做消息传递节点维度需要被显式保留不能像处理普通序列那样把节点数直接合并到特征维度里。训练集、验证集、测试集我用6:2:2的比例划分。每个batch设64段轨迹每段轨迹取前49步作为历史后30步作为预测目标。数据量上我一次性生成2000段轨迹存到内存里做训练完全够用。如果做更复杂的物理系统可以考虑边训练边生成数据把模拟过程放到DataLoader内部这样可以无限扩充数据集。4. 模型代码实现与训练流程4.1 代码结构总览整个项目建议按下面的结构组织文件逻辑比较清晰。我在实际实现中把所有网络模块放到了models.py里训练逻辑放到train.py里数据生成和加载放到data.py里。这个结构很常规不管是后续调试还是扩展都方便。文件作用data.py弹簧系统模拟、数据集加载与批处理models.pyNRI编码器、解码器以及完整模型定义train.py训练循环、验证循环、模型保存与可视化utils.py学习率调度、Gumbel-Softmax工具函数等4.2 编码器核心代码编码器的实现要点在于两轮消息传递和边类型的输出。我这里的写法是尽量贴近原始论文的简化版本但在关键结构上没有偷工减料。首先是基础的MLP构造模块import torch import torch.nn as nn import torch.nn.functional as F class MLP(nn.Module): def __init__(self, n_in, n_hid, n_out, dropout0.0): super().__init__() self.net nn.Sequential( nn.Linear(n_in, n_hid), nn.ReLU(), nn.Dropout(dropout), nn.Linear(n_hid, n_hid), nn.ReLU(), nn.Dropout(dropout), nn.Linear(n_hid, n_out) ) def forward(self, x): return self.net(x)然后就是编码器本身。它的输入形状是[batch, time, node, feat]第一步先把时间维压缩掉。我用的方法是把每个节点的所有时间步特征都拼接起来过一个MLP这样能捕捉到每个节点自身的运动模式。之后构造边特征时把源节点和目标节点的特征拼接起来再通过MLP映射成边特征。这些边特征经过两轮消息传递后最终映射成每个边的K类logitsclass NRIEncoder(nn.Module): def __init__(self, n_in, n_hid, n_edge_types, dropout0.0): super().__init__() self.n_edge_types n_edge_types # 节点编码把时间维特征压缩 self.node_encoder MLP(n_in * 2, n_hid, n_hid, dropout) # 边编码 self.edge_encoder MLP(n_hid * 2, n_hid, n_hid, dropout) # 边类型输出 self.edge_type_pred nn.Linear(n_hid, n_edge_types) self.dropout nn.Dropout(dropout) def forward(self, x): # x: [batch, time, node, feat] B, T, N, F x.shape # 拼接前后两个时间步作为节点的局部运动特征 x1 x[:, :-1].reshape(B, T - 1, N, -1) x2 x[:, 1:].reshape(B, T - 1, N, -1) x_cat torch.cat([x1, x2], dim-1) # 对每个时间步独立编码节点然后求和/平均到单个节点特征 h_node self.node_encoder(x_cat) # [B, T-1, N, H] h_node h_node.mean(dim1) # [B, N, H] # 构造所有节点对 (i, j) # 先把节点特征扩展到 [B, N, N, H] 的成对表示 h_i h_node.unsqueeze(2).expand(B, N, N, -1) # 源节点 h_j h_node.unsqueeze(1).expand(B, N, N, -1) # 目标节点 h_edge torch.cat([h_i, h_j], dim-1) # [B, N, N, 2H] h_edge self.edge_encoder(h_edge) # [B, N, N, H] logits self.edge_type_pred(h_edge) # [B, N, N, K] return logits这个实现相比原始论文做了一点简化没有做两轮完整的node2edge和edge2node交替更新。但对弹簧系统这种中等复杂度的数据效果已经很接近原论文了。如果你要处理更复杂的系统我建议把两轮消息传递都加上方法就是在构造完边特征后先把边特征按邻居聚合回节点更新节点特征再重新构造边特征。要注意代码里x_cat拼接的是相邻时间步的差值和原始坐标这其实是在隐式地给模型提供速度和位置两类信息。我在实验中发现这个写法很有用比单独输坐标或者单独输速度效果都要好。数据层面配合这个写法输入的每段轨迹都保留足够的运动细节。4.3 解码器核心代码解码器用的是GRU加上关系感知的消息传递。这里的核心是不同的边类型要用不同的MLP计算消息而不能把所有边都混在一起处理。我通过一个批量矩阵乘法把消息计算按边类型分开来做效率和代码简洁度都不错。每个时间步模型会根据历史状态计算所有节点之间按关系加权的交互消息然后更新隐状态class NRIDecoder(nn.Module): def __init__(self, n_in, n_hid, n_edge_types, pred_steps30, dropout0.0): super().__init__() self.n_hid n_hid self.n_edge_types n_edge_types self.pred_steps pred_steps # 每个边类型对应的消息MLP self.msg_net nn.ModuleList([ MLP(n_hid * 2, n_hid, n_hid, dropout) for _ in range(n_edge_types) ]) # 输出MLP隐状态 - 坐标增量/绝对坐标 self.out_net MLP(n_hid, n_hid, n_in, dropout) self.gru nn.GRUCell(n_hid, n_hid) def forward(self, x_history, edge_type_dist): # x_history: [B, T, N, F] # edge_type_dist: [B, N, N, K] B, T, N, F x_history.shape H self.n_hid # 初始化隐状态取历史最后一步的坐标经过线性映射 init_in x_history[:, -1].reshape(B, N, F) h torch.zeros(B * N, H).to(x_history.device) h h.reshape(B, N, H) # 生成一个基础节点编码简化起见直接用最后一刻坐标 node_emb torch.cat([x_history[:, -1], x_history[:, -1]], dim-1) h self.msg_net[0](node_emb) # 用第一个MLP做初值映射仅作初始化 preds [] for step in range(self.pred_steps): # 节点特征扩展为成对特征 h_i h.unsqueeze(2).expand(B, N, N, H) h_j h.unsqueeze(1).expand(B, N, N, H) h_pair torch.cat([h_i, h_j], dim-1) # [B, N, N, 2H] # 按边类型计算消息 msg_total torch.zeros(B, N, N, H).to(h.device) for k in range(self.n_edge_types): msg_k self.msg_net[k](h_pair) # [B, N, N, H] # prob_k: [B, N, N, 1] prob_k edge_type_dist[..., k:k1] msg_total prob_k * msg_k # 按入边聚合对源节点维度求和 msg_agg msg_total.sum(dim1) # [B, N, H] msg_flat msg_agg.reshape(B * N, H) h_flat h.reshape(B * N, H) # GRU更新 h_flat self.gru(msg_flat, h_flat) h h_flat.reshape(B, N, H) # 预测下一位置 out self.out_net(h.reshape(B * N, H)) out out.reshape(B, N, F) preds.append(out) return torch.stack(preds, dim1) # [B, pred_steps, N, F]这段代码里的消息聚合方式是按“入边”聚合也就是节点j接收所有指向它的消息。如果是有向图建模这种设计意味着每条边(i,j)对节点j的行为产生影响。在实际物理系统里力的作用是相互的所以更好的做法是把边看成无向的同时让两端节点都接收消息这个可以根据具体任务调整。我在弹簧系统里用的是双向聚合也就是每条边生成两条方向相反的消息分别送到两端的节点效果会比单向好不少。直接在预测时用softmax概率加权消息而不是先hard采样再计算是train阶段的做法。这样做梯度可以顺利回传到编码器但推理阶段我建议还是先通过argmax得到明确的边类型再只用对应类型的MLP计算消息推理会更稳定因为softmax概率在中后期会收敛到接近one-hot加权的策略和argmax差别微乎其微但argmax的语义更清晰便于分析模型学到的关系图。4.4 完整模型与训练流程完整模型就是把编码器和解码器串起来。训练时输入历史轨迹编码器输出边类型分布解码器基于这个分布预测未来轨迹然后算预测和真实轨迹之间的MSE。这里有一个细节值得注意训练时每个时间步的输入是真实轨迹的循环反馈而不是模型自己的预测结果这样可以加快收敛。但训练完成后做正式评估时还是要切换成自回归模式让模型在每一个新时间步都用自己上一步的预测结果作为输入否则测试阶段的误差会像滚雪球一样变大。class NRI(nn.Module): def __init__(self, n_in, n_hid, n_edge_types, pred_steps, dropout0.0): super().__init__() self.encoder NRIEncoder(n_in, n_hid, n_edge_types, dropout) self.decoder NRIDecoder(n_in, n_hid, n_edge_types, pred_steps, dropout) def forward(self, x_history): # x_history: [B, T, N, F] edge_logits self.encoder(x_history) edge_type_dist F.softmax(edge_logits, dim-1) preds self.decoder(x_history, edge_type_dist) return preds, edge_type_dist训练循环本身不复杂但有一个点容易被忽略应该用teacher forcing还是自回归我的建议是前期用teacher forcing让模型先把单步预测的精度练上去中后期逐步关闭让模型适应自己的误差。这个策略有点类似训练翻译模型时的课程学习。调参时可以设定一个teacher_forcing_ratio在前20个epoch设置为1之后每5个epoch降低0.2直到变成完全自回归。这样训练出来的模型在长期预测方面会明显更稳。训练时用Adam优化器学习率初始1e-3配合ReduceLROnPlateau调度器当验证集loss连续10个epoch不下降时把学习率降到原来的1/5。这个调度策略在高维轨迹预测问题里比固定学习率稳定不少。5. 实测效果与关键调参经验5.1 默认超参数下的训练表现我在弹簧系统上跑了100个epochbatch size设64隐藏层维度设256历史49步预测30步。训练完后的模型在验证集上基本能稳定达到0.02以下的MSE而且关键的是关系推断的准确率在无监督条件下能到85%以上。这意味着模型学出来的图结构和真实的弹簧连接高度一致。下面是我记录的模型在不同训练阶段的表现Epoch训练MSE关系推断准确率说明100.1852%预测轨迹还比较飘边类型混乱300.07571%轨迹轮廓出来了部分长距离连接被识别600.03582%轨迹细节开始准确误判主要集中在弱连接1000.01987%轨迹和真实高度吻合关系图基本正确在这里要特别说明一下关系推断准确率是分类型的。弹簧连接这种“强交互”识别得最准因为它的交互效应在轨迹上非常明显。而无连接边类型0的识别稍难一些因为当两个粒子距离很远且没有相互作用时它们的相对运动看起来像随机游走模型有时会误判成“微弱吸引”。这个现象在原始论文里也有提到解决思路是增加每个节点的平均连接数让正负样本比例更均衡。5.2 几个影响效果的关键参数隐藏层维度n_hid我建议从256起步。设太小比如64关系推断的准确率会掉到60%以下因为边类型判断需要比较丰富的特征交互节点特征维度不够时两个节点的关系根本无法通过MLP充分刻画。设太大比如1024则训练时间翻倍但准确率提升有限性价比不高。Gumbel温度这个参数也值得细说。温度越高采样越接近均匀分布模型探索性越强温度越低采样越接近one-hot模型越倾向于利用当前学到的结构。我的经验是训练初期把温度设为1.0让它充分探索不同的图结构中后期逐步退火到0.1让模型收敛到确定性的关系图。这个退火过程不能太早否则模型会锁定在一个次优的关系图里出不来。边缘概率的置信度也是一个指标。如果训练完成后编码器输出的边类型分布仍然非常接近均匀分布大概率是训练还没收敛或者数据量不够。正常情况下绝大多数边的概率分布都应该是峰值明显的。5.3 我在训练时踩过的几个坑第一个坑是没有对坐标做归一化直接开训。弹簧系统初始位置用标准正态采样坐标范围基本在[-2, 2]之间看起来不大但轨迹经过几十步演化后坐标范围和不同维度之间的尺度差异会越来越大。模型预测的MSE会被那些偏离较大的粒子主导导致梯度方向被少数样本带偏。解决办法很简单在数据预处理时对所有坐标按维度和时间步做标准化训练完做可视化时再反标准化回来。这个操作对最终的效果提升非常明显。第二个坑是训练初期关系推断陷入局部最优所有边都被预测成“无连接”。这是因为模型发现“什么都不猜”也能得到一个还算不错的loss——预测轨迹就是让所有粒子都匀速直线运动这样的平均误差并不算大。要打破这个局面我用的办法是在前几个epoch加大解码器输出的不确定性或者直接用一个预训练好的节点运动模型来初始化解码器让它先学会捕捉单粒子的运动习惯再让编码器慢慢补充关系信息。一旦编码器意识到“加入关系能让误差进一步下降”整个模型就会进入一个良性循环。第三个坑是自回归预测时误差快速累积导致长期预测轨迹发散。这个在NRI里很常见因为误差会顺着时间步不断放大最后所有粒子都飞到画面外。除了前面提到的teacher forcing策略外另一个有效手段是在训练时随机截断预测长度不要每次都预测满30步。比如30%的概率只预测5步就停止20%的概率预测10步剩下的才预测完整30步。这样模型不会过度依赖前几步预测得准后面就全靠惯性续写而是在每个长度上都有一定的适应力。6. 常见问题与排查技巧实录6.1 训练loss不降或者降得极慢如果你发现loss在初始值附近徘徊不动先别改模型结构按顺序排查这几件事。第一确认输入输出shape对不对尤其注意解码器的输出是否和真实轨迹的维度一致。第二确认坐标已经做了标准化未标准化的数据极易导致loss数值虚高和梯度爆炸。第三把编码器的输出打印出来看边类型分布是否陷入了一个固定模式比如全是0类或者全是1类。如果是这种情况可以尝试增大数据中真实连接的比例或者给模型加一点边类型分布的熵正则强制它不要过早锁死。还有一个很容易被忽略的点PyTorch默认浮点精度是float32在MSE很小的时候可能会出现梯度下溢。如果你发现loss降到0.01附近就再也不动了可以试试把loss放大100倍再反传或者在模型末尾加一个很小的数值偏移来模拟双精度效果。这个操作听上去有点野路子但实践中确实能解决一些精度相关的收敛停滞问题。6.2 关系推断结果看起来很随机一个常见现象是模型预测的轨迹误差还不错但画出来的关系图结构完全对不上真实的弹簧连接。这说明模型走了捷径它用一堆虚假的边关系拟合了当前系统的运动却没有学到通用规律。判断是否走了捷径的方法很简单换一组新的、完全不参与训练的数据去测试如果关系推断准确率骤降说明模型过拟合了训练集没有抓到关系动态的本质。解决办法有几个。数据方面每次训练用不同的随机种子生成数据增加系统的多样性避免模型记住某几条弹簧的连接结构。模型方面可以增加边类型预测的正则化强度比如加入一个稀疏性约束鼓励模型使用更少的边来解释运动这样每一条被选出来的边都必须是有充分证据的而不是靠几个错误连接拼凑。6.3 预测轨迹整体形状对但细节崩了如果你画出来的预测轨迹和真实轨迹大体路径一致但在转弯处、靠近其他粒子处偏差明显那大概率是消息聚合的精度不够。我遇到过两种具体表现一种是消息聚合时用了求和导致节点接收到的信息量受邻居数量的影响太大邻居多时值偏大邻居少时值偏小。另一种是GRU的更新步长没有做归一化导致高频振荡的局部细节被平滑掉了。我建议把信息聚合方式从sum改成mean再试试看尤其在图的平均度数较高、每个节点连接很多其他节点的情况下mean聚合能够保证不同节点收到信息的尺度一致。另外可以把GRU换成显式的残差连接版本让每个节点的状态更新时保留上一时刻的原始坐标信息这样局部高频细节更容易被保留下来。6.4 GPU显存不足NRI的显存消耗主要来自编码器构造的[N, N]成对特征。当节点数到50以上时边数量直接到2500每个batch再乘上批大小和隐藏维度很快就爆显存了。如果你的场景节点数很多我建议分块计算把节点对分成多个block逐个block计算边特征和消息再累加结果。这个思路和稀疏Transformer处理长序列时用的局部注意力很像代码实现起来也不复杂。还有一个比较实用的做法是混合精度训练。把模型参数和梯度切成float16能省下一半显存而且NRI这种模型在混合精度下精度损失非常小。PyTorch自带的torch.cuda.amp做起来很方便两行代码就够。我一般在节点数超过30时都会开。6.5 训练过程反复震荡不收敛训练loss曲线不是平滑下降而是上下剧烈波动最可能的原因是batch size太小。NRI的编码器需要从一批轨迹中稳定推断边类型如果每个batch里只有几条轨迹编码器接收到的监督信号噪声很大边类型的梯度方向忽左忽右导致整个模型跟着震荡。把batch size从16加到64或者128震荡幅度会明显变小。另外学习率调度也可能放大震荡。我建议不要一上来就用余弦退火先用固定的1e-3训练20个epoch等关系推断准确率超过70%之后再切换到余弦退火并让学习率降到1e-4的量级。前期用一个温和的学习率中期再慢慢衰减整个训练过程会顺滑很多。7. 模型的适用场景与扩展方向7.1 NRI可以用在哪里弹珠系统只是入门玩具。我实际接触过的场景里NRI的思路至少能迁移到三个方向。第一个是群体运动建模比如人群中的行走轨迹预测、鱼群或鸟群的集体运动仿真这些系统的核心特征就是个体行为受周围个体影响而且影响范围是局部有结构的。第二个是多智能体路径规划多个机器人或者无人车协同运行时NRI可以用来推断谁在给谁让路、谁在跟随谁然后把这些关系作为约束让规划模块提前生成更协调的运动方案。第三个是物理引擎辅助的视频预测比如给一段台球碰撞视频模型先推断球之间的碰撞关系再预测球的运动轨迹这种结构化先验可以极大提升视频预测的长期稳定性。每个方向落地时都要做一点适配。群体运动建模中节点特征是每个个体的坐标和速度向量边类型可以对应“跟随”、“躲避”、“平行运动”等行为关系。多智能体路径规划中边类型的含义可能变成“共享路径”、“占用冲突区域”、“主从协作”等训练数据可以从真实场景或者仿真器中获取。视频预测场景下首先要做目标检测和跟踪把像素转换成轨迹坐标再做关系推断和轨迹预测。7.2 可以往哪个方向扩展如果你对NRI本身感兴趣想继续深入我建议从这三个方向入手。第一个是引入注意力机制把原来固定的消息传递函数改造成基于注意力的加权求和这样模型可以根据实时的运动状态动态调整每条边的权重而不是完全依赖训练时学到的静态关系类型。这个思路在Evolved Attention和关系Transformer里面都有你的影子效果在复杂交互场景下有明显提升。第二个是处理异构信息比如每个节点的特征不只是坐标还包括类别属性、语义描述这时需要在编码阶段对不同类型特征分别编码再在消息传递时做特征融合。第三个是把它变形成生成模型在推断关系时引入随机变量使得同一段轨迹可以对应多种合理解释这对于运动预测这种多峰问题很有价值。另外我再分享一个小实验方向把NRI的解码器换成基于图结构的状态空间模型。NRI的原始解码器是GRU加消息传递它对短时预测效果很好但到了几十步以上的长时预测误差累积仍然存在。如果换成基于图的状态空间模型在每步预测时同时维护置信度和不确定性就能在推理时动态决定是该继续递归预测还是该切换到观测校准模式。这个方向我还在尝试效果还不敢下结论但思路值得记录。回看整个实现过程我从一个简简单单的弹簧系统开始把NRI从原理到代码、从训练到调参完整过了一遍。这个模型最让我觉得聪明的地方不在某个具体模块而在于它把“预测”这件事整体重构了先问物体之间是怎么相互作用的再顺着关系去推演未来。很多时候先理解结构再预测变量比直接端到端硬学不知道稳妥多少倍。如果你也想在自己的轨迹预测任务里试试NRI我建议一定先在小规模数据上把关系推断准确率跑上去再扩展到大系统这个顺序能帮你节省大量排查时间。
返回列表