(self, data)
| 125 | nn.init.normal_(m.weight, mean=0.0, std=0.02) |
| 126 | |
| 127 | def forward(self, data): |
| 128 | agent_pos = data["agent"]["position"][:, :, self.history_steps - 1] |
| 129 | agent_heading = data["agent"]["heading"][:, :, self.history_steps - 1] |
| 130 | agent_mask = data["agent"]["valid_mask"][:, :, : self.history_steps] |
| 131 | polygon_center = data["map"]["polygon_center"] |
| 132 | polygon_mask = data["map"]["valid_mask"] |
| 133 | |
| 134 | bs, A = agent_pos.shape[0:2] |
| 135 | |
| 136 | position = torch.cat([agent_pos, polygon_center[..., :2]], dim=1) |
| 137 | angle = torch.cat([agent_heading, polygon_center[..., 2]], dim=1) |
| 138 | angle = (angle + math.pi) % (2 * math.pi) - math.pi |
| 139 | pos = torch.cat([position, angle.unsqueeze(-1)], dim=-1) |
| 140 | |
| 141 | agent_key_padding = ~(agent_mask.any(-1)) |
| 142 | polygon_key_padding = ~(polygon_mask.any(-1)) |
| 143 | key_padding_mask = torch.cat([agent_key_padding, polygon_key_padding], dim=-1) |
| 144 | |
| 145 | x_agent = self.agent_encoder(data) |
| 146 | x_polygon = self.map_encoder(data) |
| 147 | x_static, static_pos, static_key_padding = self.static_objects_encoder(data) |
| 148 | |
| 149 | x = torch.cat([x_agent, x_polygon, x_static], dim=1) |
| 150 | |
| 151 | pos = torch.cat([pos, static_pos], dim=1) |
| 152 | pos_embed = self.pos_emb(pos) |
| 153 | |
| 154 | key_padding_mask = torch.cat([key_padding_mask, static_key_padding], dim=-1) |
| 155 | x = x + pos_embed |
| 156 | |
| 157 | for blk in self.encoder_blocks: |
| 158 | x = blk(x, key_padding_mask=key_padding_mask, return_attn_weights=False) |
| 159 | x = self.norm(x) |
| 160 | |
| 161 | prediction = self.agent_predictor(x[:, 1:A]) |
| 162 | |
| 163 | ref_line_available = data["reference_line"]["position"].shape[1] > 0 |
| 164 | |
| 165 | if ref_line_available: |
| 166 | trajectory, probability = self.planning_decoder( |
| 167 | data, {"enc_emb": x, "enc_key_padding_mask": key_padding_mask} |
| 168 | ) |
| 169 | else: |
| 170 | trajectory, probability = None, None |
| 171 | |
| 172 | out = { |
| 173 | "trajectory": trajectory, |
| 174 | "probability": probability, # (bs, R, M) |
| 175 | "prediction": prediction, # (bs, A-1, T, 2) |
| 176 | } |
| 177 | |
| 178 | if self.use_hidden_proj: |
| 179 | out["hidden"] = self.hidden_proj(x[:, 0]) |
| 180 | |
| 181 | if self.ref_free_traj: |
| 182 | ref_free_traj = self.ref_free_decoder(x[:, 0]).reshape( |
| 183 | bs, self.future_steps, 4 |
| 184 | ) |
nothing calls this directly
no outgoing calls
no test coverage detected