return bb_min, bb_max, and mean for each track (B, 3) over entire trajectory :param verts (B, T, V, 3) :param vis_mask (B, T)
(verts, vis_mask)
| 110 | |
| 111 | |
| 112 | def get_bboxes(verts, vis_mask): |
| 113 | """ |
| 114 | return bb_min, bb_max, and mean for each track (B, 3) over entire trajectory |
| 115 | :param verts (B, T, V, 3) |
| 116 | :param vis_mask (B, T) |
| 117 | """ |
| 118 | B, T, *_ = verts.shape |
| 119 | bb_min, bb_max, mean = [], [], [] |
| 120 | for b in range(B): |
| 121 | v = verts[b, vis_mask[b, :T]] # (Tb, V, 3) |
| 122 | bb_min.append(v.amin(dim=(0, 1))) |
| 123 | bb_max.append(v.amax(dim=(0, 1))) |
| 124 | mean.append(v.mean(dim=(0, 1))) |
| 125 | bb_min = torch.stack(bb_min, dim=0) |
| 126 | bb_max = torch.stack(bb_max, dim=0) |
| 127 | mean = torch.stack(mean, dim=0) |
| 128 | # point to a track that's long and close to the camera |
| 129 | zs = mean[:, 2] |
| 130 | counts = vis_mask[:, :T].sum(dim=-1) # (B,) |
| 131 | mask = counts < 0.8 * T |
| 132 | zs[mask] = torch.inf |
| 133 | sel = torch.argmin(zs) |
| 134 | return bb_min.amin(dim=0), bb_max.amax(dim=0), mean[sel] |
| 135 | |
| 136 | |
| 137 | def track_to_colors(track_ids): |
no test coverage detected