File size: 5,921 Bytes
d4cbafd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 | """
FutureInteractionGraphV8 — V6 + kinematic edge channels.
Instead of [rel_pos(2), heading_diff(1)] = 3 channels per timestep,
provides [rel_pos(2), rel_vel(2), closing_rate(1), rel_speed(1),
heading_diff(1), intensity(1)] = 8 channels.
Kinematic features are computed analytically (zero neural cost).
Only change: RelTrajEncoder.in_channels = 8 instead of 3.
"""
import torch
import torch.nn as nn
from models.graph_interaction_nba_v6 import FutureInteractionGraphV6
from models.graph_interaction_nba_v3 import RelTrajEncoder
from models.graph_interaction_nba_v5 import _heading_diff
def _compute_kinematic_features(pos_i, pos_j):
"""Compute kinematic pairwise features from predicted positions.
Args:
pos_i, pos_j: [E, T, 2] absolute future positions
Returns:
[E, T, 8] kinematic features
"""
rel_pos = pos_j - pos_i # [E, T, 2]
# Relative velocity via finite differences
rel_vel = torch.cat([rel_pos[:, 1:] - rel_pos[:, :-1],
torch.zeros_like(rel_pos[:, :1])], dim=1) # [E, T, 2]
# Relative speed (scalar)
rel_speed = rel_vel.norm(dim=-1, keepdim=True) # [E, T, 1]
# Closing rate: dot(rel_vel, rel_pos_unit) — negative = converging
rel_dist = rel_pos.norm(dim=-1, keepdim=True).clamp(min=1e-4)
rel_pos_unit = rel_pos / rel_dist
closing_rate = (rel_vel * rel_pos_unit).sum(dim=-1, keepdim=True) # [E, T, 1]
# Heading difference
heading_diff = _heading_diff(pos_i, pos_j) # [E, T, 1]
# Interaction intensity: exp(-dist/temperature)
intensity = torch.exp(-rel_dist / 5.0) # [E, T, 1]
return torch.cat([rel_pos, rel_vel, closing_rate,
rel_speed, heading_diff, intensity], dim=-1) # [E, T, 8]
class FutureInteractionGraphV8(FutureInteractionGraphV6):
"""V6 + kinematic edge channels (8 instead of 3)."""
def __init__(self, embed_dim, future_steps, num_agents,
num_heads=4, dropout=0.1, num_gnn_layers=2,
time_dim=128, top_n_neighbors=5, rel_traj_hidden=32,
y0_score_dim=32):
super().__init__(
embed_dim=embed_dim, future_steps=future_steps,
num_agents=num_agents, num_heads=num_heads,
dropout=dropout, num_gnn_layers=num_gnn_layers,
time_dim=time_dim, top_n_neighbors=top_n_neighbors,
rel_traj_hidden=rel_traj_hidden, y0_score_dim=y0_score_dim)
# Replace RelTrajEncoder with 8 input channels
self.rel_traj_encoder = RelTrajEncoder(
out_dim=embed_dim, T=future_steps,
D_hidden=rel_traj_hidden, num_heads=4,
in_channels=8) # was 3
def forward(self, y_emb, y_abs, t_emb, tau, sigma_agent=None):
B, K, A, D = y_emb.shape
T = y_abs.shape[3]
E0 = self._E0
N = self.top_n
# ---- Scoring (same as V6) ----
y0_flat = y_abs.reshape(B * K * A, T * 2)
y0_emb = self.y0_score_proj(y0_flat)
if sigma_agent is not None:
sigma_mean = sigma_agent.mean(dim=-1).reshape(B * K * A, 1)
tau_bka = sigma_mean.squeeze(-1)
else:
sigma_mean = torch.zeros(B * K * A, 1, device=y_abs.device)
tau_bka = (tau.unsqueeze(1).unsqueeze(2)
.expand(-1, K, A).reshape(B * K * A))
node_feat = torch.cat([y0_emb, sigma_mean], dim=-1)
q_bka = self.W_q(node_feat)
k_bka = self.W_k(node_feat)
pos_bk = y_abs.reshape(B * K * A, T, 2)
edge_index_bk = self._make_batched_edge_index(B * K)
pos_i_t = pos_bk[edge_index_bk[1]]
pos_j_t = pos_bk[edge_index_bk[0]]
rel_pos_t = pos_j_t - pos_i_t
mean_rel = rel_pos_t.mean(dim=1)
std_rel = rel_pos_t.std(dim=1)
min_dist = rel_pos_t.norm(dim=-1).min(dim=1).values.unsqueeze(-1)
heading_full = _heading_diff(pos_i_t, pos_j_t)
heading_mean = heading_full.mean(dim=1)
q_i = q_bka[edge_index_bk[1]]
k_j = k_bka[edge_index_bk[0]]
semantic_score = (q_i * k_j).sum(dim=-1) * self.scale
geo_feat = torch.cat([mean_rel, std_rel, min_dist, heading_mean], dim=-1)
geo_bias = self.geo_mlp(geo_feat).squeeze(-1)
scores = semantic_score + geo_bias
# ---- Top-N selection ----
scores_grouped = scores.view(B * K * A, A - 1)
_, top_idx = scores_grouped.topk(N, dim=-1, sorted=False)
mask = torch.zeros(B * K * A, A - 1, device=scores.device, dtype=torch.bool)
mask.scatter_(1, top_idx, True)
mask_flat = mask.view(-1)
# ---- Kinematic features (NEW: 8 channels instead of 3) ----
kinematic_sparse = _compute_kinematic_features(
pos_i_t[mask_flat], pos_j_t[mask_flat]) # [E_sel, T, 8]
if sigma_agent is not None:
sigma_full = sigma_agent.reshape(B * K * A, T)
sigma_i_t = sigma_full[edge_index_bk[1][mask_flat]]
sigma_j_t = sigma_full[edge_index_bk[0][mask_flat]]
sigma_bias = sigma_i_t - sigma_j_t
else:
sigma_bias = None
edge_attr_sparse = self.rel_traj_encoder(kinematic_sparse, sigma_bias)
# ---- GNN + gated residual (same as V6) ----
edge_index_sparse = edge_index_bk[:, mask_flat]
temb_bka = (t_emb.unsqueeze(1).unsqueeze(2)
.expand(-1, K, A, -1).reshape(B * K * A, D))
nodes = y_emb.reshape(B * K * A, D)
for layer in self.gnn_layers:
nodes = layer(nodes, edge_index_sparse, edge_attr_sparse,
temb_agent=temb_bka, tau=tau_bka)
orig = y_emb.reshape(B * K * A, D)
gate = self.gate_proj(torch.cat([orig, nodes], dim=-1))
out = orig + gate * self.out_proj(nodes)
return out.view(B, K, A, D)
|