GoodWin/Deep-Multi-scale
0
1# -- coding: utf-8 --2from __future__ import print_function3import torch4import numpy as np5from PIL import Image6import os7import torchvision8 9# Converts a Tensor into an image array (numpy)10# |imtype|: the desired type of the converted numpy array11def tensor2im(input_image, norm=1, imtype=np.uint8):12 if isinstance(input_image, torch.Tensor):13 image_tensor = input_image.data14 else:15 return input_image16 if norm == 1: #for clamp -1 to 117 image_numpy = image_tensor[0].cpu().float().clamp_(-1,1).numpy()18 elif norm == 2: # for norm through max-min19 image_ = image_tensor[0].cpu().float()20 max_ = torch.max(image_)21 min_ = torch.min(image_)22 image_numpy = (image_ - min_)/(max_-min_)*2-123 image_numpy = image_numpy.numpy() 24 else:25 pass26 if image_numpy.shape[0] == 1:27 image_numpy = np.tile(image_numpy, (3, 1, 1))28 # print(image_numpy.shape)29 image_numpy = (np.transpose(image_numpy, (1, 2, 0)) + 1) / 2.0 * 255.030 # print(image_numpy.shape)31 return image_numpy.astype(imtype)32def tensor2im3Channels(input_image, imtype=np.uint8):33 if isinstance(input_image, torch.Tensor):34 image_tensor = input_image.data35 else:36 return input_image37 38 image_numpy = image_tensor.cpu().float().clamp_(-1,1).numpy()39 40 # print(image_numpy.shape)41 image_numpy = (np.transpose(image_numpy, (1, 2, 0)) + 1) / 2.0 * 255.042 # print(image_numpy.shape)43 return image_numpy.astype(imtype)44 45def diagnose_network(net, name='network'):46 mean = 0.047 count = 048 for param in net.parameters():49 if param.grad is not None:50 mean += torch.mean(torch.abs(param.grad.data))51 count += 152 if count > 0:53 mean = mean / count54 print(name)55 print(mean)56 57 58 59 60def print_numpy(x, val=True, shp=False):61 x = x.astype(np.float64)62 if shp:63 print('shape,', x.shape)64 if val:65 x = x.flatten()66 print('mean = %3.3f, min = %3.3f, max = %3.3f, median = %3.3f, std=%3.3f' % (67 np.mean(x), np.min(x), np.max(x), np.median(x), np.std(x)))68 69 70def mkdirs(paths):71 if isinstance(paths, list) and not isinstance(paths, str):72 for path in paths:73 mkdir(path)74 else:75 mkdir(paths)76 77 78def mkdir(path):79 if not os.path.exists(path):80 os.makedirs(path)81 82 83def print_current_losses(epoch, i, losses, t, t_data):84 message = '(epoch: %d, iters: %d, time: %.3f, data: %.3f) ' % (epoch, i, t, t_data)85 for k, v in losses.items():86 message += '%s: %.3f ' % (k, v)87 88 print(message)89 # with open('', "a") as log_file:90 # log_file.write('%s\n' % message)91 92def display_current_results(writer,visuals,losses,step,save_result):93 for label, images in visuals.items():94 if 'Mask' in label:# or 'Scale' in label:95 grid = torchvision.utils.make_grid(images,normalize=False, scale_each=True)96 # pass97 else:98 pass99 grid = torchvision.utils.make_grid(images,normalize=True, scale_each=True)100 writer.add_image(label,grid,step)101 for k,v in losses.items():102 writer.add_scalar(k,v,step)103 104def VisualFeature(input_feature, imtype=np.uint8):105 if isinstance(input_feature, torch.Tensor):106 image_tensor = input_feature.data107 else:108 return input_feature109 110 image_ = image_tensor.cpu().float()111 112 if image_.size(1) == 3:113 image_ = image_.permute(1,2,0)114 115 # assert(image_.size(1) == 1)116 117 118 119 #####norm 0 to 1120 max_ = torch.max(image_)121 min_ = torch.min(image_)122 image_numpy = (image_ - min_)/(max_-min_)*2-1123 image_numpy = image_numpy.numpy()124 image_numpy = (image_numpy + 1) / 2.0 * 255.0125 #####no norm126 # print((max_,min_))127 # image_numpy = image_.numpy()128 # image_numpy = image_numpy*255.0129 130 131 # print('wwwwwwwwwwwwww')132 # print(max_)133 # print(min_)134 # print(image_numpy.shape)135 return image_numpy.astype(imtype)136 137 138def save_image(image_numpy, image_path):139 image_pil = Image.fromarray(image_numpy)140 image_pil.save(image_path)141 