| 11 | from model import resnet34 |
| 12 | |
| 13 | def main(): |
| 14 | device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") |
| 15 | |
| 16 | data_transform = transforms.Compose( |
| 17 | [transforms.Resize(256), |
| 18 | transforms.CenterCrop(224), |
| 19 | transforms.ToTensor(), |
| 20 | transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) |
| 21 | |
| 22 | # load image |
| 23 | img_path = '/Users/WH/Desktop/Deep-Learning-for-image-processing/data_set/tulip.jpg' |
| 24 | assert os.path.exists(img_path), "file: '{}' does not exist.".format(img_path) |
| 25 | img = Image.open(img_path) |
| 26 | plt.imshow(img) |
| 27 | # [N, C, H, W] N为batch_size,此处应该等于1 |
| 28 | img = data_transform(img) |
| 29 | # expand batch dimension |
| 30 | img = torch.unsqueeze(img, dim=0) |
| 31 | |
| 32 | # read class_indict |
| 33 | json_path = '/Users/WH/Desktop/Deep-Learning-for-image-processing/data_set/class_indices.json' |
| 34 | assert os.path.exists(json_path), "file: '{}' does not exist.".format(json_path) |
| 35 | |
| 36 | with open(json_path, "r") as f: |
| 37 | class_indict = json.load(f) |
| 38 | |
| 39 | # create model |
| 40 | model = resnet34(num_classes=5).to(device) |
| 41 | |
| 42 | # load model weights |
| 43 | weigths_path = "/Users/WH/Desktop/Deep-Learning-for-image-processing/Pytorch_classification/ResNet/ResNet34_retrain.pth" |
| 44 | assert os.path.exists(weigths_path), "file: '{}' does not exist.".format(weigths_path) |
| 45 | model.load_state_dict(torch.load(weigths_path, map_location=device)) |
| 46 | |
| 47 | # prediction |
| 48 | model.eval() |
| 49 | with torch.no_grad(): |
| 50 | # predict class |
| 51 | output = torch.squeeze(model(img.to(device))).cpu() |
| 52 | predict = torch.softmax(output, dim=0) |
| 53 | predict_cla = torch.argmax(predict).numpy() |
| 54 | |
| 55 | print_res = "class: {} prob: {:.3}".format(class_indict[str(predict_cla)], |
| 56 | predict[predict_cla].numpy()) |
| 57 | plt.title(print_res) |
| 58 | for i in range(len(predict)): |
| 59 | print("class: {:10} prob: {:.3}".format(class_indict[str(i)], |
| 60 | predict[i].numpy())) |
| 61 | |
| 62 | if __name__ == '__main__': |
| 63 | main() |