Forward function of SPyNet. This function computes the optical flow from ref to supp. Args: ref (Tensor): Reference image with shape of (n, 3, h, w). supp (Tensor): Supporting image with shape of (n, 3, h, w). Returns: Tensor: Estimated
(self, ref, supp)
| 316 | |
| 317 | return flow |
| 318 | |
| 319 | def forward(self, ref, supp): |
| 320 | """Forward function of SPyNet. |
| 321 | |
| 322 | This function computes the optical flow from ref to supp. |
| 323 | |
| 324 | Args: |
| 325 | ref (Tensor): Reference image with shape of (n, 3, h, w). |
| 326 | supp (Tensor): Supporting image with shape of (n, 3, h, w). |
| 327 | |
| 328 | Returns: |
| 329 | Tensor: Estimated optical flow: (n, 2, h, w). |
| 330 | """ |
| 331 | |
| 332 | # upsize to a multiple of 32 |
| 333 | h, w = ref.shape[2:4] |
| 334 | w_up = w if (w % 32) == 0 else 32 * (w // 32 + 1) |
| 335 | h_up = h if (h % 32) == 0 else 32 * (h // 32 + 1) |
| 336 | ref = F.interpolate( |
| 337 | input=ref, size=(h_up, w_up), mode='bilinear', align_corners=False) |
| 338 | supp = F.interpolate( |
| 339 | input=supp, |
| 340 | size=(h_up, w_up), |
| 341 | mode='bilinear', |
| 342 | align_corners=False) |
| 343 | |
| 344 | # compute flow, and resize back to the original resolution |
| 345 | flow = F.interpolate( |
| 346 | input=self.compute_flow(ref, supp), |
| 347 | size=(h, w), |
| 348 | mode='bilinear', |
| 349 | align_corners=False) |
| 350 | |
| 351 | # adjust the flow values |
| 352 | flow[:, 0, :, :] *= float(w) / float(w_up) |
| 353 | flow[:, 1, :, :] *= float(h) / float(h_up) |
| 354 | |
| 355 | return flow |
| 356 | |
| 357 |
nothing calls this directly
no test coverage detected