- x : (B, T, D) - map_grid_feat : (B, C, H, W) - raster_from_agent: (B, 3, 3)
(self, x, map_grid_feat, raster_from_agent)
| 479 | return aux_info |
| 480 | |
| 481 | def query_map_feats(self, x, map_grid_feat, raster_from_agent): |
| 482 | ''' |
| 483 | - x : (B, T, D) |
| 484 | - map_grid_feat : (B, C, H, W) |
| 485 | - raster_from_agent: (B, 3, 3) |
| 486 | ''' |
| 487 | B, T, _ = x.size() |
| 488 | _, C, Hfeat, Wfeat = map_grid_feat.size() |
| 489 | |
| 490 | # unscale to agent coords |
| 491 | pos_traj = self.descale_traj(x.detach())[:,:,:2] |
| 492 | # convert to raster frame |
| 493 | raster_pos_traj = transform_points_tensor(pos_traj, raster_from_agent) |
| 494 | |
| 495 | # scale to the feature map size |
| 496 | _, H, W = self.input_image_shape |
| 497 | xscale = Wfeat / W |
| 498 | yscale = Hfeat / H |
| 499 | raster_pos_traj[:,:,0] = raster_pos_traj[:,:,0] * xscale |
| 500 | raster_pos_traj[:,:,1] = raster_pos_traj[:,:,1] * yscale |
| 501 | |
| 502 | # interpolate into feature grid |
| 503 | feats_out = query_feature_grid( |
| 504 | raster_pos_traj, |
| 505 | map_grid_feat |
| 506 | ) |
| 507 | feats_out = feats_out.reshape((B, T, -1)) |
| 508 | return feats_out |
| 509 | |
| 510 | def get_state_and_action_from_data_batch(self, data_batch, chosen_inds=[]): |
| 511 | ''' |
no test coverage detected