-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtest.py
More file actions
101 lines (90 loc) · 3.65 KB
/
Copy pathtest.py
File metadata and controls
101 lines (90 loc) · 3.65 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
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
import os
import argparse
import time
import numpy as np
os.environ['CUDA_VISIBLE_DEVICES'] = '1'
import torch
import torch.nn.functional as F
from torch.nn.parallel import DataParallel
from torch.autograd import Variable
from torch.utils.data import DataLoader
from torch.autograd import Variable
from dataset import fusion_dataset
from model import backbone
from PIL import Image
import time
from rgb2ycbcr import RGB2YCrCb, YCrCb2RGB
def main():
fusion_model_path = './model/Fusion/model_fusion_final.pth'
fusionmodel = backbone(output=1)
# 多卡
fusionmodel = DataParallel(fusionmodel, device_ids=[0])
device = torch.device("cuda:{}".format(args.gpu) if torch.cuda.is_available() else "cpu")
if args.gpu >= 0:
fusionmodel.to(device)
fusionmodel.load_state_dict(torch.load(fusion_model_path))
print('Loading fusion model...')
# ir_path = './test_images/ir'
# vi_path = './test_images/vis'
# ir_path = './test_images/flair'
# vi_path = './test_images/t1'
# ir_path = './test_images/flair'
# vi_path = './test_images/t1ce'
ir_path = './test_images/flair'
vi_path = './test_images/t2'
test_dataset = fusion_dataset('val', ir_path=ir_path, vi_path=vi_path)
# test_dataset = Fusion_dataset('val')
test_loader = DataLoader(
dataset=test_dataset,
batch_size=args.batch_size,
shuffle=False,
num_workers=args.num_workers,
pin_memory=True,
drop_last=False,
)
test_loader.n_iter = len(test_loader)
with torch.no_grad():
for it, (images_vis, images_ir,name) in enumerate(test_loader):
images_vis = Variable(images_vis)
images_ir = Variable(images_ir)
if args.gpu >= 0:
images_vis = images_vis.to(device)
images_ir = images_ir.to(device)
images_vis_ycrcb = RGB2YCrCb(images_vis)
logits = fusionmodel(images_vis_ycrcb, images_ir)
fusion_ycrcb = torch.cat(
(logits, images_vis_ycrcb[:, 1:2, :, :], images_vis_ycrcb[:, 2:, :, :]),
dim=1,
)
fusion_image = YCrCb2RGB(fusion_ycrcb)
ones = torch.ones_like(fusion_image)
zeros = torch.zeros_like(fusion_image)
fusion_image = torch.where(fusion_image > ones, ones, fusion_image)
fusion_image = torch.where(fusion_image < zeros, zeros, fusion_image)
fused_image = fusion_image.cpu().numpy()
fused_image = fused_image.transpose((0, 2, 3, 1))
fused_image = (fused_image - np.min(fused_image)) / (
np.max(fused_image) - np.min(fused_image)
)
fused_image = np.uint8(255.0 * fused_image)
st = time.time()
for k in range(len(name)):
image = fused_image[k, :, :, :]
image = Image.fromarray(image)
save_path = os.path.join(fused_dir, name[k])
image.save(save_path)
ed = time.time()
print('File name: {0}'.format(save_path))
print(ed - st)
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Test backbone')
parser.add_argument('--model_name', '-M', type=str, default='')
parser.add_argument('--batch_size', '-B', type=int, default=1)
parser.add_argument('--gpu', '-G', type=int, default=0)
parser.add_argument('--num_workers', '-j', type=int, default=8)
args = parser.parse_args()
n_class = 9
fused_dir = './results/flair_t2/'
os.makedirs(fused_dir, mode=0o777, exist_ok=True)
print('| testing %s on GPU #%d with pytorch' % (args.model_name, args.gpu))
main()