sra-trajectory-code / LED /models /future_interaction_graph_v6.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
7.37 kB
"""
V6-style future-trajectory interaction graph for LED.
Re-uses MoFlow's `FutureInteractionGraphV6` verbatim (imported via sys.path)
so that LED, MID, and MoFlow share the identical neighbor-scoring and
message-passing component:
Per-node: q_i = W_q([y0_emb, sigma_mean]), k_j = W_k([y0_emb, sigma_mean])
Per-edge: semantic = (q_i · k_j)/sqrt(D_s), geo_bias = geo_mlp(mean_rel, std_rel, min_dist, heading_diff_mean)
Top-N by (semantic + geo_bias) → RelTrajEncoder([rel_pos, heading_diff], sigma_bias)
GNN layers → gated residual → refined node embedding
This file is a thin adapter: it produces the V6 inputs from LED's native call
signature and decodes the refined node embeddings into a per-agent, per-mode
residual trajectory that LED adds to its frozen core denoiser's epsilon
output (or to the implied y_0 estimate, depending on --residual_on).
The decoder's final layer is zero-initialized so the overall residual starts
as a no-op; V6's own gated residual inside is additionally safe at init
because its gate_proj is randomly initialized but small.
"""
import os, sys
import torch
import torch.nn as nn
# Import MoFlow's V6 graph module (and dependencies) directly.
_MOFLOW_ROOT = os.path.abspath(
os.path.join(os.path.dirname(__file__), '..', '..', 'MoFlow'))
if _MOFLOW_ROOT not in sys.path:
sys.path.insert(0, _MOFLOW_ROOT)
from models.graph_interaction_nba_v6 import FutureInteractionGraphV6
from models.interaction_baselines import build_interaction_module # E4: SRA_MODULE factory
class FutureInteractionGraphV6Wrapper(nn.Module):
"""Adapter that exposes MoFlow's V6 graph to LED's call signature.
LED call:
delta = graph(y0_hat, past, step_idx, sigma=sigma)
y0_hat [B*A, K, T, 2], past [B*A, T_h, 6], step_idx int, sigma [B*A, 1] or None
delta [B*A, K, T, 2]
V6 call (underneath):
refined = future_graph(y_emb, y_abs, t_emb, tau, sigma_agent)
y_emb [B, K, A, D], y_abs [B, K, A, T, 2], t_emb [B, D], tau [B], sigma_agent [B, K, A, T] or None
"""
def __init__(
self,
num_agents: int = 11,
future_steps: int = 20,
past_steps: int = 10,
past_channels: int = 6,
node_dim: int = 128,
top_n: int = 5, # matches MoFlow default
num_denoise_steps: int = 5,
num_gnn_layers: int = 2,
rel_traj_hidden: int = 32,
y0_score_dim: int = 32,
num_heads: int = 4,
dropout: float = 0.1,
edge_mode: str = 'full',
neighbor_mode: str = 'rag',
):
super().__init__()
self.A = num_agents
self.T = future_steps
self.D = node_dim
# ---- Input encoders for the node embedding fed to V6 ----
self.past_proj = nn.Sequential(
nn.Linear(past_steps * past_channels, node_dim),
nn.ReLU(inplace=True),
nn.Linear(node_dim, node_dim),
)
self.y0_proj = nn.Sequential(
nn.Linear(future_steps * 2, node_dim),
nn.ReLU(inplace=True),
nn.Linear(node_dim, node_dim),
)
self.step_emb = nn.Embedding(num_denoise_steps, node_dim)
# ---- MoFlow V6 graph (scoring, edge encoding, GNN, gated residual) ----
self.future_graph = build_interaction_module( # E4: sra|gameformer|c2f via env SRA_MODULE
embed_dim = node_dim,
future_steps = future_steps,
num_agents = num_agents,
num_heads = num_heads,
dropout = dropout,
num_gnn_layers = num_gnn_layers,
time_dim = node_dim,
top_n_neighbors = min(top_n, num_agents - 1),
rel_traj_hidden = rel_traj_hidden,
y0_score_dim = y0_score_dim,
edge_mode = edge_mode,
neighbor_mode = neighbor_mode,
)
# ---- Decoder: refined node embedding -> per-timestep trajectory residual ----
self.decode = nn.Sequential(
nn.Linear(node_dim, node_dim),
nn.ReLU(inplace=True),
nn.Linear(node_dim, future_steps * 2),
)
nn.init.zeros_(self.decode[-1].weight)
nn.init.zeros_(self.decode[-1].bias)
def forward(
self,
y0_hat: torch.Tensor, # [B*A, K, T, 2]
past: torch.Tensor, # [B*A, T_h, 6]
step_idx: int,
sigma: torch.Tensor = None, # [B*A, 1] (per-agent scalar uncertainty, optional)
A_override: int = None,
) -> torch.Tensor: # [B*A, K, T, 2]
BA, K, T, _ = y0_hat.shape
# If A_override is given (variable-A scenes), use it; otherwise use init A.
if A_override is not None:
A = A_override
else:
A = self.A
if BA % A != 0:
# Fall back to treating whole batch as a single scene (B=1, A=BA)
A = BA
B = BA // A
device = y0_hat.device
# If A < 2, graph is meaningless — return zero residual
if A < 2:
return torch.zeros_like(y0_hat)
# Rebuild graph's edge index for this A
fg = self.future_graph
if fg.num_agents != A:
fg.num_agents = A
fg._E0 = A * (A - 1)
fg.top_n = max(1, min(fg.top_n, A - 1))
src, dst = [], []
for i in range(A):
for j in range(A):
if i != j: src.append(j); dst.append(i)
fg._single_edge_index = torch.tensor([src, dst], dtype=torch.long, device=device)
# --- Build y_emb [B, K, A, D] from past + y0 + step ---
past_feat = self.past_proj(past.reshape(BA, -1)) # [BA, D]
past_feat = past_feat.unsqueeze(1).expand(BA, K, self.D) # [BA, K, D]
y0_feat = self.y0_proj(y0_hat.reshape(BA, K, T * 2)) # [BA, K, D]
step = self.step_emb(torch.tensor(step_idx, device=device)) # [D]
nodes_bka = past_feat + y0_feat + step # [BA, K, D]
y_emb = nodes_bka.view(B, A, K, self.D).permute(0, 2, 1, 3).contiguous() # [B, K, A, D]
# --- y_abs [B, K, A, T, 2] (unnormalized future positions) ---
y_abs = y0_hat.view(B, A, K, T, 2).permute(0, 2, 1, 3, 4).contiguous()
# --- t_emb [B, D] and tau [B] ---
t_emb = self.step_emb(torch.tensor(step_idx, device=device)).unsqueeze(0).expand(B, -1) # [B, D]
tau = torch.full((B,), float(step_idx), dtype=torch.float32, device=device)
# --- sigma_agent [B, K, A, T] or None ---
# LED's variance_estimation is per-agent scalar [BA, 1]; broadcast over K and T.
sigma_agent = None
if sigma is not None:
sigma_ba = sigma.view(B, A) # [B, A]
sigma_agent = sigma_ba.view(B, 1, A, 1).expand(B, K, A, T)
# --- Run V6 graph: scoring, sparse edges, RelTrajEncoder, GNN, gated residual ---
y_emb_refined = self.future_graph(y_emb, y_abs, t_emb, tau, sigma_agent=sigma_agent)
# y_emb_refined: [B, K, A, D]
# --- Decode to per-mode per-agent trajectory residual ---
refined_bka = y_emb_refined.permute(0, 2, 1, 3).contiguous().view(BA, K, self.D)
residual = self.decode(refined_bka).view(BA, K, T, 2)
return residual