(self, imgs)
| 171 | return torch.cat([final_output[0], final_output[1]], dim=-1), pos.reshape(B*N, hw, -1) |
| 172 | |
| 173 | def forward(self, imgs): |
| 174 | imgs = (imgs - self.image_mean) / self.image_std |
| 175 | |
| 176 | B, N, _, H, W = imgs.shape |
| 177 | patch_h, patch_w = H // 14, W // 14 |
| 178 | |
| 179 | # encode by dinov2 |
| 180 | imgs = imgs.reshape(B*N, _, H, W) |
| 181 | hidden = self.encoder(imgs, is_training=True) |
| 182 | |
| 183 | if isinstance(hidden, dict): |
| 184 | hidden = hidden["x_norm_patchtokens"] |
| 185 | |
| 186 | hidden, pos = self.decode(hidden, N, H, W) |
| 187 | |
| 188 | point_hidden = self.point_decoder(hidden, xpos=pos) |
| 189 | conf_hidden = self.conf_decoder(hidden, xpos=pos) |
| 190 | camera_hidden = self.camera_decoder(hidden, xpos=pos) |
| 191 | |
| 192 | with torch.amp.autocast(device_type='cuda', enabled=False): |
| 193 | # local points |
| 194 | point_hidden = point_hidden.float() |
| 195 | ret = self.point_head([point_hidden[:, self.patch_start_idx:]], (H, W)).reshape(B, N, H, W, -1) |
| 196 | xy, z = ret.split([2, 1], dim=-1) |
| 197 | z = torch.exp(z) |
| 198 | local_points = torch.cat([xy * z, z], dim=-1) |
| 199 | |
| 200 | # confidence |
| 201 | conf_hidden = conf_hidden.float() |
| 202 | conf = self.conf_head([conf_hidden[:, self.patch_start_idx:]], (H, W)).reshape(B, N, H, W, -1) |
| 203 | |
| 204 | # camera |
| 205 | camera_hidden = camera_hidden.float() |
| 206 | camera_poses = self.camera_head(camera_hidden[:, self.patch_start_idx:], patch_h, patch_w).reshape(B, N, 4, 4) |
| 207 | |
| 208 | # unproject local points using camera poses |
| 209 | points = torch.einsum('bnij, bnhwj -> bnhwi', camera_poses, homogenize_points(local_points))[..., :3] |
| 210 | |
| 211 | return dict( |
| 212 | points=points, |
| 213 | local_points=local_points, |
| 214 | conf=conf, |
| 215 | camera_poses=camera_poses, |
| 216 | ) |
nothing calls this directly
no test coverage detected