oculus-ui-detector / modules /detectors /rectangular_detector.py
codebanesr
Initial commit for HuggingFace Spaces deployment
b9e2109
Raw
History Blame Contribute Delete
3.44 kB
import logging
import numpy as np
from typing import List, Dict, Any
from modules.detectors.base_detector import ObjectDetector
from modules.config import UI_MODEL_PATH, PYTORCH_ENABLE_MPS_FALLBACK
logger = logging.getLogger(__name__)
class RectangleDetector(ObjectDetector):
"""Detector for rectangular UI elements like buttons, input fields, etc."""
def __init__(self):
self.model = None
# self.ui_classes = ['button', 'card', 'field', 'heading', 'icon', 'image', 'link', 'paragraph', 'text']
self.ui_classes = ['card', 'field', 'icon', 'image', 'paragraph']
async def initialize(self):
"""Initialize the YOLO model for UI element detection."""
from ultralytics import YOLO
import torch
import os
# Check for available hardware acceleration
if torch.cuda.is_available():
self.device = "cuda:0"
logger.info("CUDA is available, using GPU acceleration for UI detection")
elif hasattr(torch, 'mps') and torch.mps.is_available():
self.device = "mps"
logger.info("MPS is available, using Apple Silicon acceleration for UI detection")
# Enable MPS fallback for operations not supported by MPS
os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = PYTORCH_ENABLE_MPS_FALLBACK
else:
self.device = "cpu"
logger.warning("No hardware acceleration available, falling back to CPU for UI detection")
# Initialize UI element detection model with custom-trained weights
model_path = UI_MODEL_PATH
if os.path.exists(model_path):
self.model = YOLO(model_path)
logger.info(f"UI element detector initialized from {model_path} on {self.device}")
else:
logger.error(f"UI model file {model_path} not found")
raise FileNotFoundError(f"UI model file {model_path} not found")
def detect(self, image: np.ndarray) -> List[Dict]:
"""Detect UI elements in an image."""
if self.model is None:
raise RuntimeError("UI detection model not initialized")
# Ensure image has 3 channels (RGB)
if len(image.shape) == 3 and image.shape[2] == 4: # If RGBA, convert to RGB
image = image[:, :, :3]
# Run inference
results = self.model(image, conf=0.25, device=self.device)
# Process and format results
detected_elements = []
for result in results:
boxes = result.boxes
for box in boxes:
# Get bounding box coordinates
x1, y1, x2, y2 = box.xyxy[0].cpu().numpy().astype(int)
# Get class name
cls_id = int(box.cls[0].item())
cls_name = result.names[cls_id]
confidence = float(box.conf[0].item())
# Only include UI elements
# Note: This is a placeholder - with a proper UI element model,
# this filtering wouldn't be necessary
if cls_name.lower() in self.ui_classes:
detected_elements.append({
"label": cls_name.lower(),
"confidence": confidence,
"bbox": [x1, y1, x2, y2]
})
return detected_elements