-
Notifications
You must be signed in to change notification settings - Fork 31
Expand file tree
/
Copy pathsave_disp_rt.py
More file actions
85 lines (70 loc) · 4.07 KB
/
Copy pathsave_disp_rt.py
File metadata and controls
85 lines (70 loc) · 4.07 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
import sys
sys.path.append('core_rt')
import argparse
import glob
import numpy as np
import torch
from tqdm import tqdm
from pathlib import Path
from core_rt.rt_igev_stereo import IGEVStereo
from core_rt.utils.utils import InputPadder
from PIL import Image
from matplotlib import pyplot as plt
import os
import skimage.io
import cv2
DEVICE = 'cuda'
os.environ['CUDA_VISIBLE_DEVICES'] = '0'
def load_image(imfile):
img = np.array(Image.open(imfile)).astype(np.uint8)
img = torch.from_numpy(img).permute(2, 0, 1).float()
return img[None].to(DEVICE)
def demo(args):
model = torch.nn.DataParallel(IGEVStereo(args), device_ids=[0])
model.load_state_dict(torch.load(args.restore_ckpt))
model = model.module
model.to(DEVICE)
model.eval()
output_directory = Path(args.output_directory)
output_directory.mkdir(exist_ok=True)
with torch.no_grad():
left_images = sorted(glob.glob(args.left_imgs, recursive=True))
right_images = sorted(glob.glob(args.right_imgs, recursive=True))
print(f"Found {len(left_images)} images. Saving files to {output_directory}/")
for (imfile1, imfile2) in tqdm(list(zip(left_images, right_images))):
image1 = load_image(imfile1)
image2 = load_image(imfile2)
padder = InputPadder(image1.shape, divis_by=32)
image1, image2 = padder.pad(image1, image2)
disp = model(image1, image2, iters=args.valid_iters, test_mode=True)
disp = padder.unpad(disp)
file_stem = os.path.join(output_directory, imfile1.split('/')[-1])
disp = disp.cpu().numpy().squeeze()
if args.save_png:
disp_16 = np.round(disp * 256).astype(np.uint16)
skimage.io.imsave(file_stem, disp_16)
# plt.imsave(file_stem, disp, cmap='jet')
if args.save_numpy:
np.save(file_stem.replace('.png', '.npy'), disp)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--restore_ckpt', help="restore checkpoint", default='./pretrained_models/igev_rt/kitti.pth')
parser.add_argument('--save_png', action='store_true', default=True, help='save output as gray images')
parser.add_argument('--save_numpy', action='store_true', help='save output as numpy arrays')
parser.add_argument('-l', '--left_imgs', help="path to all first (left) frames", default="/data/StereoDatasets/kitti/2015/testing/image_2/*_10.png")
parser.add_argument('-r', '--right_imgs', help="path to all second (right) frames", default="/data/StereoDatasets/kitti/2015/testing/image_3/*_10.png")
# parser.add_argument('-l', '--left_imgs', help="path to all first (left) frames", default="/data/StereoDatasets/kitti/2012/testing/colored_0/*_10.png")
# parser.add_argument('-r', '--right_imgs', help="path to all second (right) frames", default="/data/StereoDatasets/kitti/2012/testing/colored_1/*_10.png")
parser.add_argument('--output_directory', help="directory to save output", default="output/kitti2015/disp_0")
parser.add_argument('--mixed_precision', action='store_true', help='use mixed precision')
parser.add_argument('--precision_dtype', default='float32', choices=['float16', 'bfloat16', 'float32'], help='Choose precision type: float16 or bfloat16 or float32')
parser.add_argument('--valid_iters', type=int, default=8, help='number of flow-field updates during forward pass')
# Architecture choices
parser.add_argument('--hidden_dim', nargs='+', type=int, default=96, help="hidden state and context dimensions")
parser.add_argument('--corr_levels', type=int, default=2, help="number of levels in the correlation pyramid")
parser.add_argument('--corr_radius', type=int, default=4, help="width of the correlation pyramid")
parser.add_argument('--n_downsample', type=int, default=2, help="resolution of the disparity field (1/2^K)")
parser.add_argument('--n_gru_layers', type=int, default=3, help="number of hidden GRU levels")
parser.add_argument('--max_disp', type=int, default=192, help="max disp range")
args = parser.parse_args()
demo(args)