(args, data, flow_est, type='BB1')
| 113 | return (out + eps).sqrt() |
| 114 | |
| 115 | def compute_transformation_loss(args, data, flow_est, type='BB1'): |
| 116 | device = flow_est.device |
| 117 | B, C, H, W = flow_est.shape |
| 118 | gt_map, valid_mask = None, None |
| 119 | if type == 'BB1': |
| 120 | gt_map = data['forward_map'].to(device) |
| 121 | valid_mask = data['forward_mask'] |
| 122 | elif type == 'B1B': |
| 123 | gt_map = data['backward_map'].to(device) |
| 124 | valid_mask = data['backward_mask'] |
| 125 | |
| 126 | est_map = (coords_grid(B, H, W, device) + flow_est) |
| 127 | epe = torch.norm(gt_map - est_map, dim=1, p=1) |
| 128 | return epe[valid_mask].mean() |
no test coverage detected