Classification-Test / core /predict.py
Shaik Ahamad
Application files added
050025b unverified
Raw
History Blame Contribute Delete
3.96 kB
import torch
import torch.nn as nn
import torchvision.transforms as transforms
from PIL import Image
import cv2
import os
# define CNN architecture
class CustomCNNModel(nn.Module):
def __init__(self, input_dim, num_classes):
super(CustomCNNModel,self).__init__()
self.input_dim = input_dim
self.num_classes = num_classes
self.conv_layers = nn.Sequential(
nn.Conv2d(3, 32, kernel_size=3, padding =1, stride=1),
nn.BatchNorm2d(32),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2, stride=2),
nn.Conv2d(32, 64, kernel_size=3, padding =1, stride=1),
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2, stride=2),
nn.Conv2d(64, 128, kernel_size=3, padding =1, stride=1),
nn.BatchNorm2d(128),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2, stride=2),
nn.Conv2d(128, 256, kernel_size=3, padding =1, stride=1),
nn.BatchNorm2d(256),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2, stride=2),
)
self._to_linear = None
self._get_conv_output(self.input_dim)
self.fc_layers = nn.Sequential(
nn.Linear(self._to_linear, 512),
nn.ReLU(),
nn.Linear(512, 128),
nn.ReLU(),
nn.Linear(128,self.num_classes)
)
def forward(self, x):
x = self.conv_layers(x)
x = x.view(x.size(0),-1)
x = self.fc_layers(x)
return x
def _get_conv_output(self, input_dim=128):
with torch.no_grad():
dummy_input = torch.zeros(1, 3,input_dim, input_dim) #batch size, no of channels, input
output = self.conv_layers(dummy_input)
self._to_linear = output.view(1, -1).size(1)
# load CNN architecture trained weights
# index to label map
# transformation
class ImageClassifier():
def __init__(self, model_path, class_names=None):
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
self.model = CustomCNNModel(input_dim=128, num_classes=3).to(self.device)
self.model.load_state_dict(torch.load(model_path, map_location = self.device))
self.model.eval()
if class_names is None:
self.class_names = {0: 'person', 1: 'Dog', 2: 'Cat'}
else:
self.class_names = class_names
self.transform = transforms.Compose([
transforms.Resize((128, 128)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std = [0.5, 0.5, 0.5])
])
# load image with pillow(pytorch expects pillow)
# prediction
# label map --> class
# opencv, write text on our input image
# return class, output image
def predict(self, image_path):
image = Image.open(image_path).convert("RGB")
image_tensor = self.transform(image).unsqueeze(0).to(self.device) # adds batch size to the image unsqueeze
with torch.no_grad():
output = self.model(image_tensor)
## [[0.3, 0.7, 0.9]] -> 2D tensor
_, predicted = torch.max(output, 1)
label = self.class_names[predicted.item()]
img = cv2.imread(image_path)
cv2.putText(img, label, (10,30), cv2.FONT_HERSHEY_SIMPLEX, 1, (0,0,255), 2)
output_path = "output_image.jpg"
cv2.imwrite("output_image.jpg",img)
cwd = os.getcwd()
output_path = os.path.join(cwd, output_path)
return label, output_path
## CNN Architecture
## CNN() model instance
## CNN.load_dict()