(args, data, flows_pre)
| 6 | import time |
| 7 | |
| 8 | def compute_all_loss(args, data, flows_pre): |
| 9 | # loss_AB means compute loss using flow_ab |
| 10 | losses = {'sed_BA_loss' : 0, |
| 11 | 'sed_AB_loss' : 0, |
| 12 | 'tnf_BB1_loss': 0, |
| 13 | 'tnf_B1B_loss': 0, |
| 14 | 'total_loss' : 0} |
| 15 | |
| 16 | flow_BA = flows_pre['output_BA'] |
| 17 | flow_AB = flows_pre['output_AB'] |
| 18 | flow_BB1 = flows_pre['output_BB1'] |
| 19 | flow_B1B = flows_pre['output_B1B'] |
| 20 | |
| 21 | N = args.iters |
| 22 | B, _, H, W = data['im1'].shape |
| 23 | Fm = data['fundamental_matrix'].to(data['im1'].device) |
| 24 | |
| 25 | for i in range(N): |
| 26 | i_weight = args.gamma ** (N - i - 1) |
| 27 | |
| 28 | if args.sed_loss: |
| 29 | losses['sed_BA_loss'] += i_weight * compute_symmetrical_epipolar_distance_loss(args, flow_BA[i], Fm.permute(0, 2, 1)) |
| 30 | losses['sed_AB_loss'] += i_weight * compute_symmetrical_epipolar_distance_loss(args, flow_AB[i], Fm) |
| 31 | if args.tnf_loss: |
| 32 | losses['tnf_BB1_loss'] += i_weight * compute_transformation_loss(args, data, flow_BB1[i], type='BB1') |
| 33 | losses['tnf_B1B_loss'] += i_weight * compute_transformation_loss(args, data, flow_B1B[i], type='B1B') |
| 34 | |
| 35 | losses['total_loss'] = losses['sed_BA_loss'] + losses['sed_AB_loss'] + losses['tnf_B1B_loss'] + losses['tnf_BB1_loss'] |
| 36 | |
| 37 | return losses |
| 38 | |
| 39 | def compute_symmetrical_epipolar_distance_loss(args, flow_ab, Fm): |
| 40 | B, _, H, W = flow_ab.shape |
no test coverage detected