Team Ai
Apppublic

GoodWin/Deep-Multi-scale

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
util.py141 linesDownload Raw Back to util
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