Cat-vs-Dog-Classifier / predict.py
Zaman01's picture
Upload 5 files
713f590 verified
Raw
History Blame Contribute Delete
1.56 kB
import os
import cv2
import numpy as np
import torch
from model import CatDogCNN
LABELS = {
0: "cat",
1: "dog"
}
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CatDogCNN().to(device)
# Load model from the same directory as this file
_model_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "cat_dog_model.pth")
model.load_state_dict(torch.load(_model_path, map_location=device))
model.eval()
def preprocess_image(image):
if image is None:
raise ValueError("No image provided")
if hasattr(image, "convert"):
image = np.array(image.convert("RGB"))
else:
image = np.asarray(image)
if image.ndim == 2:
image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)
elif image.shape[-1] == 4:
image = cv2.cvtColor(image, cv2.COLOR_RGBA2RGB)
elif image.shape[-1] == 1:
image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)
image = cv2.resize(image, (64, 64))
image = image.astype(np.float32) / 255.0
image = np.transpose(image, (2, 0, 1))
image = torch.from_numpy(image).float().unsqueeze(0)
return image
def predict(image):
if image is None:
return {"cat": 0.0, "dog": 0.0}
image = preprocess_image(image)
image = image.to(device)
with torch.no_grad():
logits = model(image)
probabilities = torch.softmax(logits, dim=1)[0]
return {
"cat": float(probabilities[0]),
"dog": float(probabilities[1])
}