CoolFace
Apppublic

ARTeLab/DTM_Estimation_SRandD

sourceHugging Faceupdated 4y agoView on Hugging Face
1likes
test.py59 linesDownload Raw Back to root
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')