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

Function plot_agent_batch_dict

tbsim/utils/trajdata_utils.py:1524–1695  ·  view source on GitHub ↗
(
    batch_original: dict,
    batch_idx: int,
    ax: Optional[Axes] = None,
    legend: bool = True,
    show: bool = True,
    close: bool = True,
)

Source from the content-addressed store, hash-verified

1522 )
1523
1524def plot_agent_batch_dict(
1525 batch_original: dict,
1526 batch_idx: int,
1527 ax: Optional[Axes] = None,
1528 legend: bool = True,
1529 show: bool = True,
1530 close: bool = True,
1531) -> None:
1532 # In order to visualize during rollout
1533 if len(batch_original['history_positions'].shape) == 4:
1534 batch = {}
1535 for k in batch_original:
1536 if k == 'extras':
1537 batch['extras'] = {}
1538 for k2 in batch_original['extras']:
1539 batch['extras'][k2] = batch_original['extras'][k2][0].detach().cpu()
1540 elif k in ['agent_name', 'map_names']:
1541 batch[k] = batch_original[k][0]
1542 else:
1543 batch[k] = batch_original[k][0].detach().cpu()
1544 else:
1545 batch = batch_original
1546
1547 if ax is None:
1548 _, ax = plt.subplots()
1549
1550 if 'agent_name' in batch:
1551 agent_name: str = batch['agent_name'][batch_idx]
1552 elif 'agent_names' in batch:
1553 agent_name: str = batch['agent_names'][batch_idx]
1554 else:
1555 raise ValueError("No agent name found in batch")
1556
1557 agent_type = batch['type']
1558 agent_type = convert_lyft_type_to_nusc_type(agent_type)
1559 agent_type: AgentType = AgentType(agent_type[batch_idx].item())
1560
1561 current_state = np.concatenate([batch['centroid'][batch_idx].numpy(), batch['yaw'][:, None][batch_idx].numpy()], axis=0)
1562 ax.set_title(
1563 f"{str(agent_type)}/{agent_name}\nat x={current_state[0]:.2f},y={current_state[1]:.2f},h={current_state[-1]:.2f}"
1564 )
1565
1566 agent_from_world_tf: Tensor = batch['agent_from_world'][batch_idx].cpu()
1567
1568 if batch['maps'] is not None:
1569 world_from_raster_tf: Tensor = torch.linalg.inv(
1570 batch['raster_from_world'][batch_idx].cpu()
1571 )
1572
1573 agent_from_raster_tf: Tensor = agent_from_world_tf @ world_from_raster_tf
1574
1575 draw_map(ax, batch['maps'][batch_idx], agent_from_raster_tf, alpha=1.0)
1576
1577 agent_hist = StateTensor.from_array(torch.cat([batch['history_positions'][batch_idx].cpu(), batch['history_yaws'][batch_idx].cpu()], axis=-1), format='x,y,h')
1578
1579 agent_fut = StateTensor.from_array(torch.cat([batch['target_positions'][batch_idx].cpu(), batch['target_yaws'][batch_idx].cpu()], axis=-1), format='x,y,h')
1580
1581 agent_extent = batch['extent'][batch_idx].cpu()

Callers 1

get_actionMethod · 0.90

Calls 3

plot_vec_map_lanesFunction · 0.85
closeMethod · 0.80

Tested by

no test coverage detected