(batch: dict, overwrite_nan=True)
| 344 | return d |
| 345 | |
| 346 | def parse_node_centric(batch: dict, overwrite_nan=True): |
| 347 | maybe_pad_neighbor(batch) |
| 348 | fut_pos, fut_yaw, _, fut_mask = trajdata2posyawspeed(batch["agent_fut"], nan_to_zero=overwrite_nan) |
| 349 | hist_pos, hist_yaw, hist_speed, hist_mask = trajdata2posyawspeed(batch["agent_hist"], nan_to_zero=overwrite_nan) |
| 350 | curr_speed = hist_speed[..., -1] |
| 351 | curr_state = batch["curr_agent_state"] |
| 352 | assert isinstance(curr_state,StateTensor) or isinstance(curr_state,StateArray) |
| 353 | h1, h2 = curr_state[:, -1], curr_state.heading[...,0] |
| 354 | p1, p2 = curr_state[:, :2], curr_state.position |
| 355 | assert torch.all(h1[~torch.isnan(h1)] == h2[~torch.isnan(h2)]) |
| 356 | assert torch.all(p1[~torch.isnan(p1)] == p2[~torch.isnan(p2)]) |
| 357 | curr_yaw = curr_state.heading[...,0] |
| 358 | curr_pos = curr_state.position |
| 359 | |
| 360 | # convert nuscenes types to l5kit types |
| 361 | agent_type = batch["agent_type"] |
| 362 | agent_type = convert_nusc_type_to_lyft_type(agent_type) |
| 363 | |
| 364 | # mask out invalid extents |
| 365 | agent_hist_extent = batch["agent_hist_extent"] |
| 366 | agent_hist_extent[torch.isnan(agent_hist_extent)] = 0. |
| 367 | neigh_indices = batch["neigh_indices"] |
| 368 | neigh_hist_pos, neigh_hist_yaw, neigh_hist_speed, neigh_hist_mask = trajdata2posyawspeed(batch["neigh_hist"], nan_to_zero=overwrite_nan) |
| 369 | neigh_fut_pos, neigh_fut_yaw, _, neigh_fut_mask = trajdata2posyawspeed(batch["neigh_fut"], nan_to_zero=overwrite_nan) |
| 370 | neigh_curr_speed = neigh_hist_speed[..., -1] |
| 371 | neigh_types = batch["neigh_types"] |
| 372 | # convert nuscenes types to l5kit types |
| 373 | neigh_types = convert_nusc_type_to_lyft_type(neigh_types) |
| 374 | # mask out invalid extents |
| 375 | neigh_hist_extents = batch["neigh_hist_extents"] |
| 376 | neigh_hist_extents[torch.isnan(neigh_hist_extents)] = 0. |
| 377 | |
| 378 | world_from_agents = torch.inverse(batch["agents_from_world_tf"]) |
| 379 | |
| 380 | raster_cfg = BATCH_RASTER_CFG |
| 381 | map_res = 1.0 / raster_cfg["pixel_size"] # convert to pixels/meter |
| 382 | h = w = raster_cfg["raster_size"] |
| 383 | ego_cent = raster_cfg["ego_center"] |
| 384 | |
| 385 | raster_from_agent = torch.Tensor([ |
| 386 | [map_res, 0, ((1.0 + ego_cent[0])/2.0) * w], |
| 387 | [0, map_res, ((1.0 + ego_cent[1])/2.0) * h], |
| 388 | [0, 0, 1] |
| 389 | ]).to(curr_state.device) |
| 390 | |
| 391 | bsize = batch["agents_from_world_tf"].shape[0] |
| 392 | agent_from_raster = torch.inverse(raster_from_agent) |
| 393 | raster_from_agent = TensorUtils.unsqueeze_expand_at(raster_from_agent, size=bsize, dim=0) |
| 394 | agent_from_raster = TensorUtils.unsqueeze_expand_at(agent_from_raster, size=bsize, dim=0) |
| 395 | raster_from_world = torch.bmm(raster_from_agent, batch["agents_from_world_tf"]) |
| 396 | |
| 397 | all_hist_pos = torch.cat((hist_pos[:, None], neigh_hist_pos.to(hist_pos.device)), dim=1) |
| 398 | all_hist_yaw = torch.cat((hist_yaw[:, None], neigh_hist_yaw.to(hist_pos.device)), dim=1) |
| 399 | all_hist_mask = torch.cat((hist_mask[:, None], neigh_hist_mask.to(hist_pos.device)), dim=1) |
| 400 | |
| 401 | maps_rasterize_in = batch["maps"] |
| 402 | if maps_rasterize_in is None and BATCH_RASTER_CFG["include_hist"]: |
| 403 | maps_rasterize_in = torch.empty((bsize, 0, h, w)).to(all_hist_pos.device) |
no test coverage detected