sra-trajectory-code / MID /mid_sdd_graph_v4.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
17.3 kB
"""
MID + Graph on SDD — v4 (output-level).
Combines v3's correct fixes with v2's stable training recipe:
From v3 (bug fixes):
1. Velocity→position: y_pos = cumsum(vel)*dt + init_pos
2. Position normalization: y_abs / pos_scale (SDD pixels ~±500, scale=100)
3. Trajectory-aware nodes: node_proj([context, y_vel_flat])
4. No sigma/logvar (removed untrained uncertainty head)
5. Cosine LR schedule
From v2 (stability):
1. Per-agent independent diffusion t (not shared — MID batches mix agents
from different timesteps, shared t degrades base training)
2. Two-pass training: pass1 base-only → y_0_hat, pass2 base+graph with
y_0_hat edges. Avoids GT train/eval gap.
3. delta_scale=0.1, gate_init=0.1 (moderate — v2's 0.05 too weak, v3's 0.3 too strong)
4. graph_warmup=5 epochs
Target: beat previous graph best ADE=8.53 (baseline=8.27).
"""
import os, sys, time, logging, argparse, math, random
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import optim
from torch.utils.tensorboard import SummaryWriter # tbX-broken
import dill
from dataset import EnvironmentDataset, collate, get_timesteps_data, restore
from models.autoencoder import AutoEncoder
from models.trajectron import Trajectron
from utils.model_registrar import ModelRegistrar
from utils.trajectron_hypers import get_traj_hypers
from models.diffusion import DiffusionTraj, VarianceSchedule, TransformerConcatLinear
import evaluation
MOFLOW_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'MoFlow'))
sys.path.insert(0, MOFLOW_ROOT)
from models.graph_interaction_nba_v6 import FutureInteractionGraphV6
from models.context_encoder.mtr_encoder import SinusoidalPosEmb
class GraphDenoiserWrapperV4(nn.Module):
def __init__(self, base_net, encoder_dim=256, pred_len=12,
graph_hidden=128, top_n=5, num_gnn_layers=2,
graph_dropout=0.1, dt=0.4, pos_scale=100.0):
super().__init__()
self.base_net = base_net
self.pred_len = pred_len
self.graph_hidden = graph_hidden
self.max_top_n = top_n
self.dt = dt
self.pos_scale = pos_scale
self.node_proj = nn.Sequential(
nn.Linear(encoder_dim + pred_len * 2, graph_hidden),
nn.ReLU(inplace=True))
self.time_mlp = nn.Sequential(
SinusoidalPosEmb(graph_hidden),
nn.Linear(graph_hidden, graph_hidden), nn.ReLU(),
nn.Linear(graph_hidden, graph_hidden))
self.future_graph = FutureInteractionGraphV6(
embed_dim=graph_hidden, future_steps=pred_len, num_agents=64,
num_heads=4, dropout=graph_dropout, num_gnn_layers=num_gnn_layers,
time_dim=graph_hidden, top_n_neighbors=min(top_n, 63),
rel_traj_hidden=32, y0_score_dim=32)
self.graph_out_proj = nn.Sequential(
nn.Linear(graph_hidden, graph_hidden), nn.ReLU(inplace=True),
nn.Linear(graph_hidden, pred_len * 2),
nn.Tanh())
nn.init.zeros_(self.graph_out_proj[-2].weight)
nn.init.zeros_(self.graph_out_proj[-2].bias)
gi = 0.1
self.raw_gate = nn.Parameter(torch.tensor(math.log(gi / (1.0 - gi))))
self.register_buffer('delta_scale', torch.tensor(0.1))
def forward(self, x_t, beta, context,
y_vel=None, init_pos=None, skip_graph=False):
"""
x_t: [N, T, 2] noisy trajectory (velocity space)
beta: [N] diffusion beta
context: [N, enc_dim] encoder output
y_vel: [N, T, 2] velocity for graph edge features (y_0_hat from pass1)
init_pos: [N, 2] last observed absolute position
"""
eps_pred = self.base_net(x_t, beta=beta, context=context)
N = x_t.size(0)
T = self.pred_len
if (not skip_graph) and y_vel is not None and init_pos is not None and N >= 2:
D = self.graph_hidden
y_vel_flat = y_vel.view(N, T * 2)
node_input = torch.cat([context, y_vel_flat], dim=-1)
node_emb = self.node_proj(node_input).view(1, 1, N, D)
y_pos = (torch.cumsum(y_vel.view(N, T, 2), dim=1) * self.dt
+ init_pos.unsqueeze(1))
y_abs = (y_pos / self.pos_scale).view(1, 1, N, T, 2)
beta_scene = beta.mean().unsqueeze(0)
tau = (beta_scene / 0.05).clamp(0, 1)
t_emb = self.time_mlp(beta_scene)
y_emb_out = self.future_graph(
node_emb, y_abs, t_emb, tau, sigma_agent=None)
graph_out = self.graph_out_proj(
y_emb_out.squeeze(1).squeeze(0))
delta = graph_out.view(N, T, 2) * self.delta_scale
gate = torch.sigmoid(self.raw_gate)
eps_pred = eps_pred + gate * delta
return eps_pred
class MIDGraphV4:
def __init__(self, config):
self.config = config
torch.backends.cudnn.benchmark = True
self._build()
def _build(self):
self._skip_graph_override = False
self.model_dir = os.path.join("./experiments", self.config.exp_name)
self.log_writer = SummaryWriter(log_dir=self.model_dir)
os.makedirs(self.model_dir, exist_ok=True)
log_name = f"sdd_{time.strftime('%Y-%m-%d-%H-%M')}.log"
self.log = logging.getLogger(self.config.exp_name)
self.log.setLevel(logging.INFO)
self.log.addHandler(logging.FileHandler(os.path.join(self.model_dir, log_name)))
self.log.addHandler(logging.StreamHandler())
self.log.info(f"Config: {self.config}")
self.train_data_path = os.path.join(self.config.data_dir, "sdd_train.pkl")
self.eval_data_path = os.path.join(self.config.data_dir, "sdd_test.pkl")
self.hyperparams = get_traj_hypers()
self.hyperparams['enc_rnn_dim_edge'] = self.config.encoder_dim // 2
self.hyperparams['enc_rnn_dim_edge_influence'] = self.config.encoder_dim // 2
self.hyperparams['enc_rnn_dim_history'] = self.config.encoder_dim // 2
self.hyperparams['enc_rnn_dim_future'] = self.config.encoder_dim // 2
self.registrar = ModelRegistrar(self.model_dir, "cuda")
with open(self.train_data_path, 'rb') as f:
self.train_env = dill.load(f, encoding='latin1')
with open(self.eval_data_path, 'rb') as f:
self.eval_env = dill.load(f, encoding='latin1')
self.encoder = Trajectron(self.registrar, self.hyperparams, "cuda")
self.encoder.set_environment(self.train_env)
self.encoder.set_annealing_params()
base_net = TransformerConcatLinear(
point_dim=2, context_dim=self.config.encoder_dim,
tf_layer=self.config.tf_layer, residual=False)
self.graph_net = GraphDenoiserWrapperV4(
base_net, encoder_dim=self.config.encoder_dim,
pred_len=12, graph_hidden=128,
top_n=self.config.top_n_neighbors,
num_gnn_layers=self.config.graph_gnn_layers,
graph_dropout=self.config.graph_dropout,
dt=self.config.dt,
pos_scale=self.config.pos_scale).cuda()
if hasattr(self.config, 'graph_gate_init') and self.config.graph_gate_init is not None:
gi = float(max(min(self.config.graph_gate_init, 0.999), 1e-4))
with torch.no_grad():
self.graph_net.raw_gate.fill_(math.log(gi / (1.0 - gi)))
self.var_sched = VarianceSchedule(num_steps=100, beta_T=5e-2, mode='linear')
graph_keys = ('future_graph', 'node_proj', 'time_mlp',
'graph_out_proj', 'raw_gate')
self._graph_params = [p for n, p in self.graph_net.named_parameters()
if any(k in n for k in graph_keys)]
self._base_params = [p for n, p in self.graph_net.named_parameters()
if not any(k in n for k in graph_keys)]
self.optimizer = optim.Adam([
{'params': self.registrar.get_all_but_name_match('map_encoder').parameters()},
{'params': self.graph_net.parameters()},
], lr=self.config.lr)
warm = max(1, int(getattr(self.config, 'lr_warmup_epochs', 2)))
from torch.optim.lr_scheduler import LambdaLR, CosineAnnealingLR, SequentialLR
warm_sched = LambdaLR(self.optimizer,
lr_lambda=lambda e: min(1.0, (e + 1) / warm))
cosine_sched = CosineAnnealingLR(
self.optimizer, T_max=self.config.epochs - warm, eta_min=1e-5)
self.scheduler = SequentialLR(
self.optimizer, schedulers=[warm_sched, cosine_sched], milestones=[warm])
self.train_scenes = self.train_env.scenes
self.eval_scenes = self.eval_env.scenes
self.log.info(f"Train scenes: {len(self.train_scenes)}, "
f"Eval scenes: {len(self.eval_scenes)}")
def _get_loss(self, batch, node_type):
(first_history_index, x_t_raw, y_t, x_st_t, y_st_t,
neighbors_data_st, neighbors_edge_value,
robot_traj_st_t, map_) = batch
context = self.encoder.get_latent(batch, node_type)
y_0 = y_t.cuda() # [N, 12, 2] velocities
N = y_0.size(0)
init_pos = x_t_raw[:, -1, 0:2].cuda() # [N, 2] absolute position
# Per-agent independent diffusion t (v2 style — stable for MID)
t = self.var_sched.uniform_sample_t(N)
alpha_bar = self.var_sched.alpha_bars[t].cuda()
beta = self.var_sched.betas[t].cuda()
c0 = alpha_bar.sqrt().view(N, 1, 1)
c1 = (1 - alpha_bar).sqrt().view(N, 1, 1)
e_rand = torch.randn_like(y_0)
x_noisy = c0 * y_0 + c1 * e_rand
if self._skip_graph_override or N < 2:
eps_pred = self.graph_net(x_noisy, beta, context, skip_graph=True)
else:
# Two-pass: pass1 base-only → y_0_hat, pass2 base+graph
with torch.no_grad():
eps_base = self.graph_net(x_noisy, beta, context, skip_graph=True)
y_0_hat = (x_noisy - c1 * eps_base) / c0 # [N, T, 2] velocity
eps_pred = self.graph_net(x_noisy, beta, context,
y_vel=y_0_hat, init_pos=init_pos)
return F.mse_loss(eps_pred.reshape(-1, 2), e_rand.reshape(-1, 2))
def train(self):
node_type = "PEDESTRIAN"
ph = self.hyperparams['prediction_horizon']
max_hl = self.hyperparams['maximum_history_length']
graph_warm = int(getattr(self.config, 'graph_warmup_epochs', 5))
for epoch in range(1, self.config.epochs + 1):
self.graph_net.train()
total_loss, n_batches = 0.0, 0
self._skip_graph_override = (epoch <= graph_warm)
for scene in self.train_scenes:
for t_start in range(0, scene.timesteps, 10):
timesteps = np.arange(t_start, t_start + 10)
batch = get_timesteps_data(
env=self.train_env, scene=scene, t=timesteps,
node_type=node_type, state=self.hyperparams['state'],
pred_state=self.hyperparams['pred_state'],
edge_types=self.train_env.get_edge_types(),
min_ht=1, max_ht=max_hl, min_ft=12, max_ft=12,
hyperparams=self.hyperparams)
if batch is None:
continue
loss = self._get_loss(batch[0], node_type)
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm_(self._graph_params, 0.1)
nn.utils.clip_grad_norm_(self._base_params, 1.0)
self.optimizer.step()
total_loss += loss.item()
n_batches += 1
self.scheduler.step()
avg = total_loss / max(1, n_batches)
self.log.info(f"Epoch {epoch} train_loss={avg:.4f}")
self.log_writer.add_scalar('loss/train', avg, epoch)
if epoch % self.config.eval_every == 0:
ade, fde = self._eval(node_type, ph, max_hl)
ade *= 50; fde *= 50
self.log.info(f"Epoch {epoch} Best Of 20: "
f"ADE: {ade:.4f} FDE: {fde:.4f}")
self.log_writer.add_scalar('metric/ADE', ade, epoch)
self.log_writer.add_scalar('metric/FDE', fde, epoch)
torch.save({
'encoder': self.registrar.model_dict,
'graph_net': self.graph_net.state_dict(),
}, os.path.join(self.model_dir, f"sdd_epoch{epoch}.pt"))
@torch.no_grad()
def _eval(self, node_type, ph, max_hl):
self.graph_net.eval()
ade_errors, fde_errors = [], []
for scene in self.eval_scenes:
for t_start in range(0, scene.timesteps, 10):
timesteps = np.arange(t_start, t_start + 10)
batch = get_timesteps_data(
env=self.eval_env, scene=scene, t=timesteps,
node_type=node_type, state=self.hyperparams['state'],
pred_state=self.hyperparams['pred_state'],
edge_types=self.eval_env.get_edge_types(),
min_ht=7, max_ht=max_hl, min_ft=12, max_ft=12,
hyperparams=self.hyperparams)
if batch is None:
continue
test_batch, nodes, timesteps_o = batch
context = self.encoder.get_latent(test_batch, node_type)
dynamics = self.encoder.node_models_dict[node_type].dynamic
N = context.size(0)
_, x_t_raw, *_ = test_batch
init_pos = x_t_raw[:, -1, 0:2].cuda()
preds = self._sample_with_graph(
context, N, init_pos, num_points=12, K=20)
predicted_y_pos = dynamics.integrate_samples(preds)
predictions = predicted_y_pos.cpu().numpy()
predictions_dict = {}
for i, ts in enumerate(timesteps_o):
if ts not in predictions_dict:
predictions_dict[ts] = {}
predictions_dict[ts][nodes[i]] = np.transpose(
predictions[:, [i]], (1, 0, 2, 3))
batch_error = evaluation.compute_batch_statistics(
predictions_dict, scene.dt, max_hl=max_hl, ph=ph,
node_type_enum=self.eval_env.NodeType, kde=False,
map=None, best_of=True, prune_ph_to_future=True)
ade_errors = np.hstack(
(ade_errors, batch_error[node_type]['ade']))
fde_errors = np.hstack(
(fde_errors, batch_error[node_type]['fde']))
return np.mean(ade_errors), np.mean(fde_errors)
def _sample_with_graph(self, context, N, init_pos,
num_points=12, K=20):
traj_list = []
stride = 5
for _ in range(K):
x_t = torch.randn(N, num_points, 2, device=context.device)
y_0_prev = None
for t in range(self.var_sched.num_steps, 0, -stride):
alpha_bar = self.var_sched.alpha_bars[t]
alpha_bar_next = self.var_sched.alpha_bars[t - stride]
beta = self.var_sched.betas[[t] * N].cuda()
if y_0_prev is not None and N >= 2:
eps = self.graph_net(x_t, beta, context,
y_vel=y_0_prev, init_pos=init_pos)
else:
eps = self.graph_net(x_t, beta, context, skip_graph=True)
x0_pred = ((x_t - (1 - alpha_bar).sqrt() * eps)
/ alpha_bar.sqrt())
y_0_prev = x0_pred
x_t = (alpha_bar_next.sqrt() * x0_pred
+ (1 - alpha_bar_next).sqrt() * eps)
traj_list.append(x_t)
return torch.stack(traj_list)
def main():
p = argparse.ArgumentParser()
p.add_argument('--data_dir', default='processed_data')
p.add_argument('--exp_name', default='mid_sdd_graph_v4')
p.add_argument('--gpu', type=int, default=0)
p.add_argument('--epochs', type=int, default=100)
p.add_argument('--lr', type=float, default=1e-3)
p.add_argument('--eval_every', type=int, default=3)
p.add_argument('--encoder_dim', type=int, default=256)
p.add_argument('--tf_layer', type=int, default=3)
p.add_argument('--top_n_neighbors', type=int, default=5)
p.add_argument('--graph_gnn_layers', type=int, default=2)
p.add_argument('--graph_dropout', type=float, default=0.1)
p.add_argument('--graph_gate_init', type=float, default=0.1)
p.add_argument('--graph_warmup_epochs', type=int, default=5)
p.add_argument('--lr_warmup_epochs', type=int, default=2)
p.add_argument('--dt', type=float, default=0.4,
help='Scene timestep (SDD: 0.4s at 2.5Hz)')
p.add_argument('--pos_scale', type=float, default=100.0,
help='Divide positions by this before graph')
config = p.parse_args()
torch.cuda.set_device(config.gpu)
MIDGraphV4(config).train()
if __name__ == '__main__':
main()