MCPcopy Create free account
hub / github.com/NVlabs/CTG / parse_node_centric

Function parse_node_centric

tbsim/utils/trajdata_utils.py:346–475  ·  view source on GitHub ↗
(batch: dict, overwrite_nan=True)

Source from the content-addressed store, hash-verified

344 return d
345
346def 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)

Callers 1

parse_trajdata_batchFunction · 0.85

Calls 6

maybe_pad_neighborFunction · 0.85
trajdata2posyawspeedFunction · 0.85
verify_mapFunction · 0.85
rasterize_agentsFunction · 0.85
get_drivable_region_mapFunction · 0.70

Tested by

no test coverage detected