| """ |
| MotionTransformerGraphV7 — single-pass training with GT graph, two-pass inference. |
| |
| Difference from V6: |
| Training: Single forward pass. Graph edge features built from GT future |
| trajectory (teacher forcing). No wasted no_grad pass 1. |
| Inference: Two-pass. Pass 1 uses y_0_prev from the previous sampling step |
| (or skips graph at step 0). Pass 2 refines with y_0_hat from |
| pass 1 + sigma-weighted message passing. |
| """ |
|
|
| import torch |
| from models.backbone_graph_v6 import MotionTransformerGraphV6 |
|
|
|
|
| class MotionTransformerGraphV7(MotionTransformerGraphV6): |
| """V6 graph module + single-pass GT training / two-pass chained inference.""" |
|
|
| def forward(self, y, time, x_data, y_0_prev=None): |
| if self.training: |
| |
| K = self.model_cfg.NUM_PROPOSED_QUERY |
| gt_fut = x_data['fut_traj'] |
| y_0_gt = (gt_fut.unsqueeze(1) |
| .expand(-1, K, -1, -1, -1) |
| .reshape(gt_fut.shape[0], K, self.A, -1)) |
| return self._forward_impl( |
| y, time, x_data, |
| y_0_for_graph=y_0_gt, |
| sigma_for_graph=None, |
| ) |
| else: |
| |
| |
| with torch.no_grad(): |
| y_0_hat, _, logvar_ = self._forward_impl( |
| y, time, x_data, |
| y_0_for_graph=y_0_prev, |
| sigma_for_graph=None, |
| skip_graph=(y_0_prev is None), |
| ) |
|
|
| |
| B = y_0_hat.shape[0] |
| K = self.model_cfg.NUM_PROPOSED_QUERY |
| A = self.A |
| T = self.T_future |
| sigma_ = (logvar_.detach() |
| .view(B, K, A, T, 2) |
| .clamp(-10, 10) |
| .exp() |
| .mean(dim=-1) |
| .sqrt()) |
|
|
| |
| return self._forward_impl( |
| y, time, x_data, |
| y_0_for_graph=y_0_hat.detach(), |
| sigma_for_graph=sigma_, |
| ) |
|
|