(args, model, step=None, split=None)
| 47 | |
| 48 | @torch.no_grad() |
| 49 | def visualization(args, model, step=None, split=None): |
| 50 | model.eval() |
| 51 | dataset_val = ValidateData(args, split) |
| 52 | for val_id in range(len(dataset_val)): |
| 53 | im1, im2, flow_gt, valid_mask = dataset_val[val_id] |
| 54 | output = model(im1[None].cuda(), im2[None].cuda(), iters=args.iters, test_mode=True) |
| 55 | flow_pr = output[1] |
| 56 | |
| 57 | im1_vis = im1.permute([1,2,0]).cpu().numpy().astype(np.uint8) |
| 58 | im2_vis = im2.permute([1,2,0]).cpu().numpy().astype(np.uint8) |
| 59 | im2_warp_vis = image_flow_warp(im2_vis, flow_pr[0].permute([1,2,0])) |
| 60 | |
| 61 | im_all = np.concatenate([im1_vis, im2_vis, im2_warp_vis], axis=1)[:,:,::-1] |
| 62 | |
| 63 | im_all = wandb.Image(im_all, caption='step: {:d}'.format(step)) |
| 64 | wandb.log({'vis/{:d}'.format(val_id): im_all}) |
| 65 | |
| 66 | |
| 67 | if __name__ == '__main__': |
nothing calls this directly
no test coverage detected