(
batch_original: dict,
batch_idx: int,
ax: Optional[Axes] = None,
legend: bool = True,
show: bool = True,
close: bool = True,
)
| 1522 | ) |
| 1523 | |
| 1524 | def 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() |
no test coverage detected