Object-Intelligence-Backend / backend /ml /embedding_recognizer.py
muhammadpriv001's picture
Backend
9cfe107
Raw History Blame Contribute Delete
5.37 kB
import logging
import cv2
import torch
import torchvision.transforms as T
import torchvision.models as models
from PIL import Image
import numpy as np
from typing import Optional, List, Dict, Any, Tuple
from backend.config import DEVICE, SPECIFIC_SIMILARITY_THRESHOLD
from backend.database.storage import db
logger = logging.getLogger(__name__)
class EmbeddingRecognizer:
def __init__(self, model_name: str = "resnet18"):
self.model_name = model_name
self.model = None
self.transform = T.Compose([
T.Resize((224, 224)),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
self.load_model()
def load_model(self):
try:
logger.info("Loading PyTorch vision backbone for Specific Object Recognition...")
backbone = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)
# Remove classification layer to get 512-dim feature embedding vector
self.model = torch.nn.Sequential(*list(backbone.children())[:-1])
self.model.eval()
self.model.to(DEVICE)
logger.info("Specific Object Embedding backbone loaded successfully.")
except Exception as e:
logger.error(f"Error loading embedding backbone model: {e}")
self.model = None
def extract_embedding_from_crop(self, crop_np: np.ndarray) -> Optional[np.ndarray]:
if self.model is None or crop_np is None or crop_np.size == 0:
return None
try:
pil_img = Image.fromarray(crop_np).convert("RGB")
tensor = self.transform(pil_img).unsqueeze(0).to(DEVICE)
with torch.no_grad():
feat = self.model(tensor)
feat = feat.squeeze().cpu().numpy()
# L2 normalize vector
norm = np.linalg.norm(feat)
if norm > 0:
feat = feat / norm
return feat
except Exception as e:
logger.error(f"Failed to extract embedding from crop: {e}")
return None
def extract_embedding_from_image(self, img_np: np.ndarray) -> Optional[np.ndarray]:
if self.model is None or img_np is None or img_np.size == 0:
return None
try:
rgb_img = cv2.cvtColor(img_np, cv2.COLOR_BGR2RGB) if len(img_np.shape) == 3 and img_np.shape[2] == 3 else img_np
pil_img = Image.fromarray(rgb_img).convert("RGB")
tensor = self.transform(pil_img).unsqueeze(0).to(DEVICE)
with torch.no_grad():
feat = self.model(tensor)
feat = feat.squeeze().cpu().numpy()
norm = np.linalg.norm(feat)
if norm > 0:
feat = feat / norm
return feat
except Exception as e:
logger.error(f"Failed to extract embedding from image: {e}")
return None
def extract_embedding_from_image_path(self, image_path: str) -> Optional[np.ndarray]:
try:
pil_img = Image.open(image_path).convert("RGB")
tensor = self.transform(pil_img).unsqueeze(0).to(DEVICE)
with torch.no_grad():
feat = self.model(tensor)
feat = feat.squeeze().cpu().numpy()
norm = np.linalg.norm(feat)
if norm > 0:
feat = feat / norm
return feat
except Exception as e:
logger.error(f"Failed to extract embedding from file {image_path}: {e}")
return None
def recognize_crop(self, crop_np: np.ndarray, threshold: float = SPECIFIC_SIMILARITY_THRESHOLD) -> Tuple[Optional[str], float]:
emb = self.extract_embedding_from_crop(crop_np)
if emb is None:
return None, 0.0
matched_name, similarity = db.find_best_matching_object(emb, min_threshold=threshold)
return matched_name, similarity
def refine_detections_with_specific_objects(
self, image_np: np.ndarray, candidate_detections: List[Dict[str, Any]], threshold: float = SPECIFIC_SIMILARITY_THRESHOLD
) -> List[Dict[str, Any]]:
"""Scans candidate bounding boxes and overrides label if matched to a specific user-taught object."""
if image_np is None or not candidate_detections:
return candidate_detections
h, w = image_np.shape[:2]
refined = []
for det in candidate_detections:
bbox = det["bbox"]
x1, y1, x2, y2 = max(0, bbox[0]), max(0, bbox[1]), min(w, bbox[2]), min(h, bbox[3])
if (x2 - x1) > 10 and (y2 - y1) > 10:
crop = image_np[y1:y2, x1:x2]
matched_name, similarity = self.recognize_crop(crop, threshold=threshold)
if matched_name is not None:
# Specific Object Identity Matched!
det_copy = dict(det)
det_copy["label"] = matched_name
det_copy["specific_identity"] = matched_name
det_copy["type"] = "specific"
det_copy["source"] = "embedding"
det_copy["confidence"] = round(similarity, 4)
refined.append(det_copy)
continue
refined.append(det)
return refined