ARTeLab/DTM_Estimation_SRandD
1
1import torch2import torchvision3from torchvision import transforms4from PIL import Image5import matplotlib.pyplot as plt6import numpy as np7from models.modelNetA import Generator as GA8from models.modelNetB import Generator as GB9from models.modelNetC import Generator as GC10 11 12 13# DEVICE='cpu'14DEVICE='cuda'15model_type = 'model_c'16 17modeltype2path = {18 'model_a': 'DTM_exp_train10%_model_a/g-best.pth',19 'model_b': 'DTM_exp_train10%_model_b/g-best.pth',20 'model_c': 'DTM_exp_train10%_model_c/g-best.pth',21}22 23if model_type == 'model_a':24 generator = GA()25if model_type == 'model_b':26 generator = GB()27if model_type == 'model_c':28 generator = GC()29 30generator = torch.nn.DataParallel(generator)31state_dict_Gen = torch.load(modeltype2path[model_type], map_location=torch.device('cpu'))32generator.load_state_dict(state_dict_Gen)33generator = generator.module.to(DEVICE)34# generator.to(DEVICE)35generator.eval()36 37preprocess = transforms.Compose([38 transforms.Grayscale(),39 # transforms.Resize((128, 128)),40 transforms.ToTensor()41])42input_img = Image.open('demo_imgs/fake.jpg')43torch_img = preprocess(input_img).to(DEVICE).unsqueeze(0).to(DEVICE)44torch_img = (torch_img - torch.min(torch_img)) / (torch.max(torch_img) - torch.min(torch_img))45with torch.no_grad():46 output = generator(torch_img)47sr, sr_dem_selected = output[0], output[1]48sr = sr.squeeze(0).cpu()49 50print(sr.shape)51torchvision.utils.save_image(sr, 'sr.png')52# sr = Image.fromarray(sr.squeeze(0).detach().numpy() * 255, 'L')53# sr.save('sr2.png')54 55sr_dem_selected = sr_dem_selected.squeeze().cpu().detach().numpy()56print(sr_dem_selected.shape)57plt.imshow(sr_dem_selected, cmap='jet', vmin=0, vmax=np.max(sr_dem_selected))58plt.colorbar()59plt.savefig('test.png')