sra-trajectory-code / MID /main_ethucy_mid.py
po03087's picture
SRA: MID/LED/MoFlow code + RUNNING.md instructions (code only, no data/ckpts)
d4cbafd verified
Raw
History Blame Contribute Delete
12.3 kB
"""
main_ethucy_mid.py — MID baseline on ETH/UCY (leave-one-out, variable A).
Uses the original scene-level pickle files at
`MoFlow/data/eth_ucy/original/{scene}/{scene}_{train,test}.pkl` which store:
traj [N_total, 20, 2] (all ped trajectories, concatenated)
seq_start_end [N_scenes, 2] (start,end into traj per scene window)
num_peds_in_seq [N_scenes] (A per scene — variable, >=1)
frame_list [N_scenes]
Leave-one-out: {scene}_train.pkl is the union of the other four ETH/UCY
subsets; {scene}_test.pkl is the held-out scene. Standard 8 past + 12 future
frames @ 2.5 Hz (4.8 s horizon).
Each "sample" is one scene window with variable A agents. We use
batch_size=1 so scene-batching is just stacking along A; the social
transformer attends across all A agents within the scene.
Usage:
python main_ethucy_mid.py --scene univ --gpu 3
"""
import os, sys, time, pickle, logging, argparse
import numpy as np
import torch
import torch.nn as nn
from torch import optim
from torch.utils.data import Dataset, DataLoader
from torch.utils.tensorboard import SummaryWriter # tbX-broken
from tqdm.auto import tqdm
from models.diffusion import DiffusionTraj, VarianceSchedule, TransformerConcatLinear
OBS_LEN = 8
PRED_LEN = 12
K_EVAL = 20
HORIZONS_FULL = {'1.6s': 4, '3.2s': 8, '4.8s': 12}
DATA_ROOT = '/mnt/jaewoo4tb/srtp/MoFlow/data/eth_ucy/original'
class ETHUCYDataset(Dataset):
"""Scene-window dataset: each item is one scene with variable A agents."""
def __init__(self, scene, split='train'):
super().__init__()
path = os.path.join(DATA_ROOT, scene, f'{scene}_{split}.pkl')
with open(path, 'rb') as f:
d = pickle.load(f)
traj = d['traj'].astype(np.float32) # [N_total, 20, 2]
sse = d['seq_start_end'] # [N_scenes, 2]
assert traj.shape[1] == OBS_LEN + PRED_LEN
self.scenes = []
for s, e in sse:
self.scenes.append(torch.from_numpy(traj[s:e])) # [A, 20, 2]
a_counts = np.array([len(x) for x in self.scenes])
print(f'[ETHUCYDataset] {scene} {split}: {len(self.scenes)} scenes, '
f'A min/mean/max = {a_counts.min()}/{a_counts.mean():.1f}/{a_counts.max()}')
def __len__(self): return len(self.scenes)
def __getitem__(self, i):
x = self.scenes[i] # [A, 20, 2]
return x[:, :OBS_LEN, :], x[:, OBS_LEN:, :]
def collate_bs1(batch):
assert len(batch) == 1, 'batch_size must be 1 (variable-A scenes)'
return batch[0] # (pre[A,8,2], fut[A,12,2])
def preprocess_scene(pre, fut, device):
"""Per-agent last-obs-relative normalization for one scene (A agents).
Returns past_6ch [A,8,6], fut_rel [A,12,2], mask [A,A] (all-zeros, full social
attention within the scene), last_obs [A,1,2]."""
pre = pre.to(device)
fut = fut.to(device)
last_obs = pre[:, -1:, :] # [A, 1, 2]
abs_xy = pre - last_obs
rel_xy = abs_xy
vel_xy = torch.cat([rel_xy[:, 1:] - rel_xy[:, :-1],
torch.zeros_like(rel_xy[:, :1])], dim=1)
past_6ch = torch.cat([abs_xy, rel_xy, vel_xy], dim=-1) # [A, 8, 6]
fut_rel = fut - last_obs # [A, 12, 2]
A = pre.size(0)
mask = torch.zeros(A, A, device=device) # full intra-scene attention
return past_6ch, fut_rel, mask, last_obs
class _STEncoder(nn.Module):
def __init__(self, in_channels=6, hidden=256):
super().__init__()
self.conv = nn.Conv1d(in_channels, 32, kernel_size=3, stride=1, padding=1)
self.relu = nn.ReLU()
self.gru = nn.GRU(32, hidden, num_layers=1, batch_first=True)
nn.init.kaiming_normal_(self.conv.weight)
nn.init.kaiming_normal_(self.gru.weight_ih_l0)
nn.init.kaiming_normal_(self.gru.weight_hh_l0)
nn.init.zeros_(self.conv.bias)
nn.init.zeros_(self.gru.bias_ih_l0)
nn.init.zeros_(self.gru.bias_hh_l0)
def forward(self, x):
h = self.relu(self.conv(x.transpose(1, 2)))
_, state = self.gru(h.transpose(1, 2))
return state.squeeze(0)
class _SocialTransformer(nn.Module):
def __init__(self, past_len=OBS_LEN, hidden=256):
super().__init__()
self.proj = nn.Linear(past_len * 6, hidden, bias=False)
layer = nn.TransformerEncoderLayer(
d_model=hidden, nhead=2, dim_feedforward=hidden, batch_first=False)
self.encoder = nn.TransformerEncoder(layer, num_layers=2)
def forward(self, x_flat, mask):
h = self.proj(x_flat).unsqueeze(1)
h = h + self.encoder(h, mask)
return h.squeeze(1)
class ETHUCYEncoder(nn.Module):
def __init__(self, encoder_dim=256, past_len=OBS_LEN):
super().__init__()
self.ego_encoder = _STEncoder(in_channels=6, hidden=256)
self.social_encoder = _SocialTransformer(past_len=past_len, hidden=256)
self.fusion = nn.Linear(512, encoder_dim)
def forward(self, past_6ch, social_mask):
ego = self.ego_encoder(past_6ch)
social = self.social_encoder(past_6ch.reshape(past_6ch.size(0), -1), social_mask)
return self.fusion(torch.cat([ego, social], dim=-1))
class Trainer:
def __init__(self, args):
self.args = args
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
self._build_dirs()
self._build_data()
self._build_model()
self._build_optimizer()
def _build_dirs(self):
self.exp_dir = os.path.join('experiments', self.args.exp_name)
os.makedirs(self.exp_dir, exist_ok=True)
self.tb_log = SummaryWriter(log_dir=self.exp_dir)
log_path = os.path.join(self.exp_dir,
f'{self.args.scene}_{time.strftime("%Y-%m-%d-%H-%M")}.log')
self.log = logging.getLogger(self.args.exp_name)
self.log.setLevel(logging.INFO)
self.log.addHandler(logging.FileHandler(log_path))
self.log.addHandler(logging.StreamHandler(sys.stdout))
self.log.info(f'Args: {self.args}')
def _build_data(self):
train_dset = ETHUCYDataset(self.args.scene, split='train')
test_dset = ETHUCYDataset(self.args.scene, split='test')
self.train_loader = DataLoader(
train_dset, batch_size=1, shuffle=True,
num_workers=2, collate_fn=collate_bs1, pin_memory=False)
self.test_loader = DataLoader(
test_dset, batch_size=1, shuffle=False,
num_workers=2, collate_fn=collate_bs1, pin_memory=False)
self.log.info(f'Scene={self.args.scene} Train={len(train_dset)} Test={len(test_dset)}')
def _build_model(self):
self.encoder = ETHUCYEncoder(encoder_dim=self.args.encoder_dim, past_len=OBS_LEN).to(self.device)
net = TransformerConcatLinear(point_dim=2, context_dim=self.args.encoder_dim,
tf_layer=self.args.tf_layer, residual=False)
self.diffusion = DiffusionTraj(
net=net,
var_sched=VarianceSchedule(num_steps=100, beta_T=5e-2, mode='linear'),
).to(self.device)
n_enc = sum(p.numel() for p in self.encoder.parameters())
n_diff = sum(p.numel() for p in self.diffusion.parameters())
self.log.info(f'Encoder: {n_enc:,} Diffusion: {n_diff:,}')
def _build_optimizer(self):
params = list(self.encoder.parameters()) + list(self.diffusion.parameters())
self.optimizer = optim.Adam(params, lr=self.args.lr)
self.scheduler = optim.lr_scheduler.ExponentialLR(self.optimizer, gamma=0.98)
def _run_step(self, pre, fut, grad_accum_every):
past_6ch, fut_rel, mask, _ = preprocess_scene(pre, fut, self.device)
context = self.encoder(past_6ch, mask)
loss = self.diffusion.get_loss(fut_rel, context)
return loss
def train(self):
best_ade = float('inf')
accum = self.args.grad_accum
for epoch in range(1, self.args.epochs + 1):
self.encoder.train(); self.diffusion.train()
total_loss, count = 0.0, 0
self.optimizer.zero_grad()
for i, (pre, fut) in enumerate(tqdm(self.train_loader, ncols=90, desc=f'E{epoch}')):
if pre.size(0) < 1: continue
loss = self._run_step(pre, fut, accum)
(loss / accum).backward()
if (i + 1) % accum == 0:
nn.utils.clip_grad_norm_(
list(self.encoder.parameters()) + list(self.diffusion.parameters()), 1.0)
self.optimizer.step()
self.optimizer.zero_grad()
total_loss += loss.item(); count += 1
self.optimizer.step(); self.optimizer.zero_grad()
self.scheduler.step()
avg = total_loss / max(count, 1)
self.tb_log.add_scalar('loss/train', avg, epoch)
self.log.info(f'Epoch {epoch} train_loss={avg:.4f}')
if epoch % self.args.eval_every == 0:
m = self.evaluate()
for k, v in m.items():
self.tb_log.add_scalar(f'metric/{k}', v, epoch)
self.log.info(
f'Epoch {epoch} ADE(4.8s)={m["ADE_4.8s"]:.4f} FDE(4.8s)={m["FDE_4.8s"]:.4f}'
f' ADE(1.6s)={m["ADE_1.6s"]:.4f} ADE(3.2s)={m["ADE_3.2s"]:.4f}')
ade = m['ADE_4.8s']; fde = m['FDE_4.8s']
if ade < best_ade:
best_ade = ade
torch.save({'encoder': self.encoder.state_dict(),
'diffusion': self.diffusion.state_dict(),
'epoch': epoch, 'metrics': m},
os.path.join(self.exp_dir, 'best.pt'))
self.log.info(f' ** New best ADE(4.8s)={ade:.4f} FDE(4.8s)={fde:.4f}')
@torch.no_grad()
def evaluate(self):
self.encoder.eval(); self.diffusion.eval()
sums = {f'{k}_{h}': 0.0 for h in HORIZONS_FULL for k in ('ADE', 'FDE')}
n_agents = 0
for pre, fut in tqdm(self.test_loader, ncols=90, desc='Eval'):
A = pre.size(0)
if A < 1: continue
past_6ch, _, mask, last_obs = preprocess_scene(pre, fut, self.device)
context = self.encoder(past_6ch, mask)
pred_rel = self.diffusion.sample(
num_points=PRED_LEN, context=context, sample=K_EVAL,
bestof=True, sampling=self.args.sampling, step=self.args.sampling_step)
pred_abs = pred_rel + last_obs.unsqueeze(0) # [K, A, T, 2]
fut_abs = fut.to(self.device) # [A, T, 2]
dist = (pred_abs - fut_abs.unsqueeze(0)).norm(dim=-1) # [K, A, T]
for h, end in HORIZONS_FULL.items():
sums[f'ADE_{h}'] += dist[:, :, :end].mean(dim=-1).min(dim=0).values.sum().item()
sums[f'FDE_{h}'] += dist[:, :, end - 1].min(dim=0).values.sum().item()
n_agents += A
return {k: v / n_agents for k, v in sums.items()}
def parse_args():
p = argparse.ArgumentParser()
p.add_argument('--scene', type=str, required=True,
choices=['eth', 'hotel', 'univ', 'zara1', 'zara2'])
p.add_argument('--exp_name', type=str, default=None)
p.add_argument('--gpu', type=int, default=0)
p.add_argument('--epochs', type=int, default=100)
p.add_argument('--grad_accum', type=int, default=32,
help='Gradient accumulation (since batch_size=1).')
p.add_argument('--lr', type=float, default=1e-3)
p.add_argument('--eval_every', type=int, default=1)
p.add_argument('--encoder_dim', type=int, default=256)
p.add_argument('--tf_layer', type=int, default=3)
p.add_argument('--sampling', type=str, default='ddim', choices=['ddpm', 'ddim'])
p.add_argument('--sampling_step', type=int, default=10)
args = p.parse_args()
if args.exp_name is None:
args.exp_name = f'mid_ethucy_baseline_{args.scene}'
return args
if __name__ == '__main__':
args = parse_args()
trainer = Trainer(args)
trainer.train()