-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy patheval.py
More file actions
72 lines (65 loc) · 2.64 KB
/
Copy patheval.py
File metadata and controls
72 lines (65 loc) · 2.64 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
import torch
from torchvision.utils import make_grid, save_image
import numpy as np
import argparse
import skimage
from model import IntroVAE
from main import DB, colormap
def main(args):
print(args)
device = torch.device('cuda')
torch.set_grad_enabled(False)
args.alpha, args.beta, args.margin, args.lr = 0, 0, 0, 0
vae = IntroVAE(args).to(device)
vae.load_state_dict(torch.load(args.load))
print('load ckpt from:', args.load)
args.root = '/dev/null'
args.data_aug = False
db = DB(args)
db.images = args.input
imgs = [img for img in db]
x = torch.stack(imgs, dim=0).to(device)
mu, logvar = vae.encoder(x)
z = torch.nn.functional.interpolate(mu.permute(1,0).unsqueeze(0), \
size=args.n_interp, mode='linear', align_corners=True \
).squeeze(0).permute(1,0)
xr = vae.decoder(z)
xr, x = xr.cpu(), x.cpu()
if len(args.output) == 1:
# save pretty-formated result
if args.num_classes >= 0:
x, xr = [colormap[img.argmax(1)].permute(0,3,1,2) \
for img in (x, xr)]
max_len = max(x.shape[0], xr.shape[0])
output = torch.cat( \
[x[:]]+[torch.zeros_like(x[0:1])]*(max_len - x.shape[0]) + \
[xr[:]]+[torch.zeros_like(xr[0:1])]*(max_len - xr.shape[0]), \
dim=0)
save_image(output, args.output[0], nrow=max_len, range=(0,1))
print('write output to :', args.output[0])
else:
# save raw output
assert len(args.output) == args.n_interp
if args.num_classes >= 0:
xr = xr.argmax(1).to(torch.uint8)
for img, output in zip(xr, args.output):
skimage.io.imsave(output, img.numpy())
print('write outputs to :', args.output)
if __name__ == '__main__':
argparser = argparse.ArgumentParser()
argparser.add_argument('--imgsz', type=int, default=128, \
help='imgsz')
argparser.add_argument('--z_dim', type=int, default=256, \
help='hidden latent z dim')
argparser.add_argument('--n_interp', type=int, default=3, \
help='number of images to be interpolated')
argparser.add_argument('--load', type=str, required=True, \
help='checkpoint to load')
argparser.add_argument('--input', type=str, required=True, nargs='*', \
help='checkpoint to load')
argparser.add_argument('--output', type=str, required=True, nargs='*', \
help='output path')
argparser.add_argument('--num_classes', type=int, default=-1, \
help='set to positive value to model shapes (e.g. segmentation)')
args = argparser.parse_args()
main(args)