Files
graphene/Examples/pytorch/pytorchexample.py
T
Dayeol Lee e651c1d262 [Examples] Change PyTorch example to load model from a pickle file
- This example will be compared with another example for loading
encrypted models and inputs using PFS. Thus, load the model from a file.
- Updated Makefile and README.md
2020-06-03 12:11:27 -07:00

49 lines
1.3 KiB
Python

# This PyTorch image classification example is based off
# https://www.learnopencv.com/pytorch-for-beginners-image-classification-using-pre-trained-models/
from torchvision import models
import torch
# Load the model from a file
alexnet = torch.load("alexnet-pretrained.pt")
# Prepare a transform to get the input image into a format (e.g., x,y dimensions) the classifier
# expects.
from torchvision import transforms
transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)])
# Load the image.
from PIL import Image
img = Image.open("input.jpg")
# Apply the transform to the image.
img_t = transform(img)
# Magic (not sure what this does).
batch_t = torch.unsqueeze(img_t, 0)
# Prepare the model and run the classifier.
alexnet.eval()
out = alexnet(batch_t)
# Load the classes from disk.
with open('classes.txt') as f:
classes = [line.strip() for line in f.readlines()]
# Sort the predictions.
_, indices = torch.sort(out, descending=True)
# Convert into percentages.
percentage = torch.nn.functional.softmax(out, dim=1)[0] * 100
# Print the 5 most likely predictions.
with open("result.txt", "w") as outfile:
outfile.write(str([(classes[idx], percentage[idx].item()) for idx in indices[0][:5]]))