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