diff --git a/apps/server/api/detection.py b/apps/server/api/detection.py index 1565a0d..77cba43 100644 --- a/apps/server/api/detection.py +++ b/apps/server/api/detection.py @@ -21,11 +21,12 @@ async def detect_image( confidence: float = Query(0.5), iou: float = Query(0.45), algorithm_config: Optional[str] = Query(None, description="算法配置JSON字符串"), - composite: bool = Query(False, description="是否启用复合检测(火灾检测时同时检测火焰和烟雾)") + composite: bool = Query(False, description="是否启用复合检测(火灾检测时同时检测火焰和烟雾)"), + region_polygon: Optional[str] = Query(None, description="禁停区域多边形JSON,例如:[[x1,y1],[x2,y2],...]") ): """ 图片检测接口 - + Args: algorithm_config: 算法配置JSON,例如: { @@ -36,6 +37,7 @@ async def detect_image( "loitering_threshold": 300.0, "movement_threshold": 5.0 } + region_polygon: 禁停区域多边形坐标,用于违停检测 """ from main import model_service from services.detection_service import DetectionService @@ -49,19 +51,27 @@ async def detect_image( algo_config = json.loads(algorithm_config) except json.JSONDecodeError as e: logger.warning(f"算法配置解析失败: {e}") - + + # 解析禁停区域 + region = None + if region_polygon: + try: + region = json.loads(region_polygon) + except json.JSONDecodeError as e: + logger.warning(f"禁停区域解析失败: {e}") + try: contents = await file.read() nparr = np.frombuffer(contents, np.uint8) frame = cv2.imdecode(nparr, cv2.IMREAD_COLOR) - + if frame is None: return ImageDetectionResult( success=False, message="无法读取图片", data={} ) - + # 判断是否启用复合火灾检测 if composite and model_id == 'fire_detection': result = await detection_service.detect_fire_composite( @@ -69,7 +79,9 @@ async def detect_image( ) else: result = await detection_service.detect_image( - frame, model_id, confidence, iou, algorithm_config=algo_config + frame, model_id, confidence, iou, + algorithm_config=algo_config, + region_polygon=region ) if result['success']: diff --git a/apps/server/main.py b/apps/server/main.py index 638b6da..e0a1b16 100644 --- a/apps/server/main.py +++ b/apps/server/main.py @@ -97,7 +97,12 @@ app.add_middleware( allow_headers=["*"], ) -app.mount("/static", StaticFiles(directory="static"), name="static") +static_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "static") +os.makedirs(static_dir, exist_ok=True) +os.makedirs(os.path.join(static_dir, "uploads"), exist_ok=True) +os.makedirs(os.path.join(static_dir, "results"), exist_ok=True) +os.makedirs(os.path.join(static_dir, "temp"), exist_ok=True) +app.mount("/static", StaticFiles(directory=static_dir), name="static") docker_output_dir = os.path.join( os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), @@ -122,10 +127,6 @@ async def camera_websocket_endpoint(websocket: WebSocket): await camera_service.handle_connection(websocket) if __name__ == "__main__": - os.makedirs("static/uploads", exist_ok=True) - os.makedirs("static/results", exist_ok=True) - os.makedirs("static/temp", exist_ok=True) - # 设置信号处理器 setup_signal_handlers() diff --git a/apps/server/services/detection_service.py b/apps/server/services/detection_service.py index 1e197fc..4789081 100644 --- a/apps/server/services/detection_service.py +++ b/apps/server/services/detection_service.py @@ -48,7 +48,8 @@ class DetectionService: model_id: str, confidence: float = 0.5, iou: float = 0.45, - algorithm_config: Optional[Dict] = None + algorithm_config: Optional[Dict] = None, + region_polygon: Optional[List[List[int]]] = None ) -> Dict: start_time = time.time() @@ -62,6 +63,39 @@ class DetectionService: } try: + # 违停检测特殊处理:调用专门的违停检测方法 + if model_id == 'illegal_parking_detection' and hasattr(model, 'detect_illegal_parking'): + # 如果提供了禁停区域,单张图片模式下即时判定(时间阈值设为0) + parking_time = 0 if region_polygon else None + parking_result = model.detect_illegal_parking( + image, conf=confidence, illegal_parking_time=parking_time, region_polygon=region_polygon + ) + detections = [] + for vehicle in parking_result.get('illegal_parking', []): + detections.append({ + 'class': 'illegal_parking', + 'label': '违停车辆', + 'confidence': 1.0, + 'bbox': vehicle['bbox'], + 'track_id': vehicle.get('track_id'), + 'parking_duration': vehicle.get('parking_duration', 0) + }) + + processing_time = time.time() - start_time + result_data = { + 'success': parking_result['success'], + 'message': parking_result.get('message', '违停检测完成'), + 'detections': detections, + 'stats': { + **parking_result.get('stats', {}), + 'total_detections': len(detections), + 'processing_time': round(processing_time, 3), + 'model_used': model_id + } + } + result_data = self._apply_event_pipeline(result_data, model_id) + return result_data + results = model(image, conf=confidence, iou=iou, verbose=False) detections = [] @@ -202,6 +236,44 @@ class DetectionService: 'stats': None } + # 违停检测特殊处理:调用专门的违停检测方法 + if model_id == 'illegal_parking_detection' and hasattr(model, 'detect_illegal_parking'): + parking_result = model.detect_illegal_parking( + frame, conf=confidence + ) + detections = [] + for vehicle in parking_result.get('illegal_parking', []): + detections.append({ + 'class': 'illegal_parking', + 'label': '违停车辆', + 'confidence': 1.0, + 'bbox': vehicle['bbox'], + 'track_id': vehicle.get('track_id'), + 'parking_duration': vehicle.get('parking_duration', 0) + }) + + processing_time = time.time() - start_time + fps = 1.0 / processing_time if processing_time > 0 else 0 + + result_data = { + 'success': parking_result['success'], + 'message': parking_result.get('message', '违停检测完成'), + 'detections': detections, + 'stats': { + **parking_result.get('stats', {}), + 'total_detections': len(detections), + 'fps': round(fps, 2), + 'processing_time': round(processing_time, 3), + 'model_used': model_id + } + } + result_data = self._apply_event_pipeline(result_data, model_id) + + if draw: + frame = self.draw_detections(frame, detections, fps) + + return frame, result_data + results = model(frame, conf=confidence, iou=iou, verbose=False) detections = [] @@ -579,6 +651,7 @@ class DetectionService: 'helmet': (255, 255, 0), 'no_helmet': (255, 0, 255), 'cigarette': (0, 165, 255), + 'illegal_parking': (0, 0, 255), # 兼容旧模型类别 'violence': (0, 0, 255), 'fight': (0, 0, 255), diff --git a/apps/server/services/model_service.py b/apps/server/services/model_service.py index 19f0e3a..c473811 100644 --- a/apps/server/services/model_service.py +++ b/apps/server/services/model_service.py @@ -102,22 +102,22 @@ class ModelService: 'name': '徘徊检测' }, 'vehicle_detection': { - 'path': os.path.join(base_dir, 'models', 'vehicle_detection_paddle', 'mot_ppyoloe_l_36e_ppvehicle', 'model.pdmodel'), - 'type': 'paddle', - 'classes': ['vehicle'], - 'labels': {'vehicle': '车辆'}, - 'size': '181MB', - 'description': '基于PaddlePaddle PP-YOLOE-l的车辆检测和跟踪模型', - 'name': '车辆检测 (Paddle)' + 'path': os.path.join(base_dir, 'models', 'vehicle_detection_paddle', 'yolov8s.pt'), + 'type': 'yolov8', + 'classes': ['car', 'truck', 'bus', 'motorcycle'], + 'labels': {'car': '小汽车', 'truck': '卡车', 'bus': '公交车', 'motorcycle': '摩托车'}, + 'size': '23MB', + 'description': '基于YOLOv8s的园区车辆检测模型(COCO预训练)', + 'name': '车辆检测 (YOLOv8s)' }, 'illegal_parking_detection': { - 'path': os.path.join(base_dir, 'models', 'vehicle_detection_paddle', 'mot_ppyoloe_l_36e_ppvehicle', 'model.pdmodel'), - 'type': 'paddle', - 'classes': ['vehicle'], - 'labels': {'vehicle': '车辆'}, - 'size': '200MB', - 'description': '基于PaddlePaddle PP-YOLOE-l的违停检测模型,支持车牌识别', - 'name': '违停检测 (Paddle)' + 'path': os.path.join(base_dir, 'models', 'vehicle_detection_paddle', 'yolov8s.pt'), + 'type': 'yolov8', + 'classes': ['car', 'truck', 'bus', 'motorcycle'], + 'labels': {'car': '小汽车', 'truck': '卡车', 'bus': '公交车', 'motorcycle': '摩托车'}, + 'size': '23MB', + 'description': '基于YOLOv8s的园区违停检测模型,支持停车时长与区域判定', + 'name': '违停检测 (YOLOv8s)' }, 'fight_detection': { 'path': os.path.join(base_dir, 'models', 'fight_detection', 'yolov8n.pt'), @@ -157,6 +157,12 @@ class ModelService: os.path.exists(os.path.join(model_dir, f)) for f in required_files ) + elif config['type'] in ('yolov8', 'yolov10'): + # 本地路径存在,或官方预训练模型名(会自动下载) + model_exists = os.path.exists(model_path) + if not model_exists: + model_name = os.path.basename(model_path) + model_exists = model_name.startswith(('yolov8', 'yolov10', 'yolo11', 'yolo26')) else: model_exists = os.path.exists(model_path) @@ -213,10 +219,6 @@ class ModelService: from .paddle_detection_service import SmokingDetectionModel logger.info(f"正在加载 PaddlePaddle 抽烟检测服务: {model_id}") model = SmokingDetectionModel() - elif model_id in ['vehicle_detection', 'illegal_parking_detection']: - from .vehicle_detection_service import VehicleDetectionModel - logger.info(f"正在加载 PaddlePaddle 车辆检测服务: {model_id}") - model = VehicleDetectionModel() else: logger.error(f"未知的 Paddle 模型类型: {model_id}") return None @@ -230,14 +232,21 @@ class ModelService: # 处理 YOLO 模型 model_path = config['path'] - - if not os.path.exists(model_path): + model_name = os.path.basename(model_path) + is_official_model = model_name.startswith(('yolov8', 'yolov10', 'yolo11', 'yolo26')) + + if not os.path.exists(model_path) and not is_official_model: logger.warning(f"模型文件不存在: {model_path},跳过加载 {model_id}") return None try: logger.info(f"正在加载 YOLO 模型: {model_id} from {model_path}") model = YOLO(model_path) + + # 对违停检测模型,包装以支持 detect_illegal_parking 方法 + if model_id == 'illegal_parking_detection': + from .vehicle_detection_service import VehicleDetectionModel + model = VehicleDetectionModel(yolo_model=model) self.models[model_id] = model logger.info(f"YOLO 模型加载成功: {model_id}") diff --git a/apps/server/services/vehicle_detection_service.py b/apps/server/services/vehicle_detection_service.py index ed89f78..b1d042f 100644 --- a/apps/server/services/vehicle_detection_service.py +++ b/apps/server/services/vehicle_detection_service.py @@ -1,21 +1,14 @@ """ 车辆检测服务适配器 -支持车辆检测、跟踪、车牌识别和违停检测功能 +支持车辆检测、跟踪和违停检测功能(基于YOLOv8) """ -# 禁用 PIR API 以支持旧版模型格式(必须在任何导入之前设置) -import os -os.environ['FLAGS_enable_pir_api'] = '0' - import cv2 import numpy as np import logging import threading import time -import sys from typing import Dict, List, Optional, Tuple -from pathlib import Path -from collections import defaultdict from dataclasses import dataclass logger = logging.getLogger(__name__) @@ -39,117 +32,31 @@ class VehicleTrackingInfo: class VehicleDetectionService: - """车辆检测服务(本地模式)""" + """车辆检测服务(YOLOv8模式)""" - def __init__(self): + # COCO数据集中车辆相关类别映射 + VEHICLE_CLASSES = {2: 'car', 3: 'motorcycle', 5: 'bus', 7: 'truck'} + VEHICLE_LABELS = {'car': '小汽车', 'motorcycle': '摩托车', 'bus': '公交车', 'truck': '卡车'} + + def __init__(self, yolo_model=None): self.model_name = "vehicle_detection" self.threshold = 0.1 self._lock = threading.Lock() - - # 本地环境配置 - project_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(__file__)))) - self.paddle_dir = os.path.join(project_root, "third-party", "paddle-inference") - self.model_dir = os.path.join(project_root, "models", "vehicle_detection_paddle") - - # 模型路径配置 - self.mot_model_dir = os.path.join(self.model_dir, "mot_ppyoloe_l_36e_ppvehicle") - self.plate_det_model_dir = os.path.join(self.model_dir, "ch_PP-OCRv3_det_infer") - self.plate_rec_model_dir = os.path.join(self.model_dir, "ch_PP-OCRv3_rec_infer") - - # 检测器实例(延迟加载) - self._detector = None - self._detector_initialized = False + self._yolo_model = yolo_model # 车辆跟踪信息 self.vehicle_tracks: Dict[int, VehicleTrackingInfo] = {} self.track_id_counter = 0 # 违停检测配置 - self.illegal_parking_time = 5.0 # 默认5秒 - self.illegal_parking_region = None # 违停区域多边形 + self.illegal_parking_time = 5.0 + self.illegal_parking_region = None - self.available = True - logger.info(f"车辆检测服务初始化完成") - logger.info(f"车辆检测模型目录: {self.mot_model_dir}") - logger.info(f"车牌检测模型目录: {self.plate_det_model_dir}") - logger.info(f"车牌识别模型目录: {self.plate_rec_model_dir}") - - # 禁用 PIR API 以支持旧版模型格式 - os.environ['FLAGS_enable_pir_api'] = '0' - - try: - self._initialize_environment() - except Exception as e: - logger.error(f"环境初始化失败: {e}") - self.available = False - - def _initialize_environment(self): - """初始化本地 PaddlePaddle 环境""" - try: - # 添加 PaddleDetection 部署路径 - paddle_detection_path = self.paddle_dir - if paddle_detection_path not in sys.path: - sys.path.insert(0, paddle_detection_path) - logger.info(f"✅ 添加 PaddleDetection 路径: {paddle_detection_path}") - - # 检查模型目录是否存在 - required_models = { - 'MOT': self.mot_model_dir, - 'Plate Detection': self.plate_det_model_dir, - 'Plate Recognition': self.plate_rec_model_dir - } - - for model_name, model_path in required_models.items(): - if not os.path.exists(model_path): - raise Exception(f"{model_name} 模型目录不存在: {model_path}") - - required_files = ['inference.pdmodel', 'inference.pdiparams', 'inference.pdiparams.info'] - if model_name == 'MOT': - required_files = ['model.pdmodel', 'model.pdiparams', 'infer_cfg.yml'] - - for file in required_files: - file_path = os.path.join(model_path, file) - if not os.path.exists(file_path): - raise Exception(f"{model_name} 模型文件不存在: {file}") - - logger.info("✅ 环境检查通过") - - except Exception as e: - logger.error(f"环境初始化失败: {e}") - raise - - def _get_detector(self): - """获取检测器实例(单例模式)""" - if self._detector is None or not self._detector_initialized: - try: - # 设置环境变量以支持旧版模型格式 - os.environ['FLAGS_enable_pir_api'] = '0' - - # 添加 PaddleDetection 路径 - if self.paddle_dir not in sys.path: - sys.path.insert(0, self.paddle_dir) - - # 导入 PaddleDetection 模块 - from infer import Detector, PredictConfig - - # 创建检测器(使用MOT模型) - self._detector = Detector( - model_dir=self.mot_model_dir, - device='CPU', - run_mode='paddle', - batch_size=1, - output_dir='output', - threshold=self.threshold - ) - - self._detector_initialized = True - logger.info("✅ 车辆检测器初始化成功") - - except Exception as e: - logger.error(f"检测器初始化失败: {e}") - raise - - return self._detector + self.available = yolo_model is not None + if self.available: + logger.info("车辆检测服务初始化完成(YOLOv8模式)") + else: + logger.error("车辆检测服务初始化失败:未提供YOLO模型") def detect_image(self, image: np.ndarray, threshold: float = None) -> Dict: """ @@ -165,7 +72,7 @@ class VehicleDetectionService: if threshold is None: threshold = self.threshold - if not self.available: + if not self.available or self._yolo_model is None: return { 'success': False, 'message': '车辆检测服务不可用', @@ -177,35 +84,36 @@ class VehicleDetectionService: with self._lock: start_time = time.time() - # 确保检测器已初始化 - detector = self._get_detector() - - # 准备输入图片 - if not isinstance(image, np.ndarray): - raise Exception(f"不支持的图片类型: {type(image)}") - - if len(image.shape) == 2: - image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR) - elif image.shape[2] == 4: - image = cv2.cvtColor(image, cv2.COLOR_RGBA2BGR) - - # 执行推理 + # 使用YOLO模型推理 inference_start = time.time() - - results = detector.predict_image( - [image], - visual=False, - save_results=False - ) - + results = self._yolo_model(image, conf=threshold, verbose=False) inference_time = time.time() - inference_start - logger.info(f"推理耗时: {inference_time:.3f}s") - # 解析检测结果 - detections = self._parse_detection_results(results, threshold) + # 解析检测结果,只保留车辆类别 + detections = [] + result = results[0] + if result.boxes is not None: + for box in result.boxes: + cls_id = int(box.cls[0]) + if cls_id in self.VEHICLE_CLASSES: + confidence = float(box.conf[0]) + x1, y1, x2, y2 = box.xyxy[0].tolist() + + center_x = (x1 + x2) / 2 + center_y = (y1 + y2) / 2 + + class_name = self.VEHICLE_CLASSES[cls_id] + label = self.VEHICLE_LABELS[class_name] + + detections.append({ + 'class': class_name, + 'label': label, + 'confidence': round(confidence, 3), + 'bbox': [int(x1), int(y1), int(x2), int(y2)], + 'center': [round(center_x, 2), round(center_y, 2)] + }) total_time = time.time() - start_time - logger.info(f"检测总耗时: {total_time:.3f}s") return { 'success': True, @@ -213,7 +121,7 @@ class VehicleDetectionService: 'detections': detections, 'stats': { 'total_detections': len(detections), - 'model_used': 'mot_ppyoloe_l_36e_ppvehicle', + 'model_used': 'yolov8s', 'threshold': threshold, 'processing_time': round(total_time, 3), 'inference_time': round(inference_time, 3) @@ -224,9 +132,6 @@ class VehicleDetectionService: import traceback logger.error(f"检测失败: {e}") logger.error(f"错误堆栈: {traceback.format_exc()}") - - self._detector_initialized = False - return { 'success': False, 'message': f'检测失败: {e}', @@ -234,40 +139,6 @@ class VehicleDetectionService: 'stats': None } - def _parse_detection_results(self, results: Dict, threshold: float) -> List[Dict]: - """解析 PaddleDetection 返回的检测结果""" - detections = [] - - try: - if results and 'boxes' in results: - boxes = results['boxes'] - - if boxes is not None and len(boxes) > 0: - for box in boxes: - if len(box) >= 6: - class_id = int(box[0]) - confidence = float(box[1]) - x1, y1, x2, y2 = float(box[2]), float(box[3]), float(box[4]), float(box[5]) - - # 计算中心点 - center_x = (x1 + x2) / 2 - center_y = (y1 + y2) / 2 - - # 过滤低置信度检测 - if confidence >= threshold: - detections.append({ - 'class': 'vehicle', - 'label': '车辆', - 'confidence': round(confidence, 3), - 'bbox': [int(x1), int(y1), int(x2), int(y2)], - 'center': [round(center_x, 2), round(center_y, 2)] - }) - - except Exception as e: - logger.error(f"解析检测结果失败: {e}") - - return detections - def detect_illegal_parking(self, image: np.ndarray, threshold: float = None, illegal_parking_time: float = 5.0, region_polygon: List[Tuple[int, int]] = None) -> Dict: @@ -286,10 +157,20 @@ class VehicleDetectionService: if threshold is None: threshold = self.threshold + # 规范化 region_polygon 格式(支持 [{x,y},...] 和 [[x,y],...] 两种格式) + if region_polygon: + normalized = [] + for p in region_polygon: + if isinstance(p, dict): + normalized.append((int(p['x']), int(p['y']))) + elif isinstance(p, (list, tuple)) and len(p) >= 2: + normalized.append((int(p[0]), int(p[1]))) + region_polygon = normalized + # 更新违停配置 self.illegal_parking_time = illegal_parking_time self.illegal_parking_region = region_polygon - + # 基础车辆检测 detection_result = self.detect_image(image, threshold) @@ -333,17 +214,18 @@ class VehicleDetectionService: track_info.center = center track_info.last_seen = current_time track_info.trajectory.append(center) - - # 检查违停条件 - if self._check_illegal_parking(track_info, region_polygon): - track_info.is_illegal_parking = True - illegal_parking_vehicles.append({ - 'track_id': matched_track_id, - 'bbox': bbox, - 'center': center, - 'parking_duration': round(current_time - track_info.first_seen, 2), - 'plate_number': track_info.plate_number - }) + + # 检查违停条件(对新车辆和已有车辆都检查) + track_info = self.vehicle_tracks[matched_track_id] + if self._check_illegal_parking(track_info, region_polygon): + track_info.is_illegal_parking = True + illegal_parking_vehicles.append({ + 'track_id': matched_track_id, + 'bbox': bbox, + 'center': center, + 'parking_duration': round(current_time - track_info.first_seen, 2), + 'plate_number': track_info.plate_number + }) # 清理长时间未出现的车辆 self._cleanup_old_tracks(current_time) @@ -387,12 +269,12 @@ class VehicleDetectionService: if parking_duration < self.illegal_parking_time: return False - # 检查是否在违停区域内 - if region_polygon is None: - return False + # 如果指定了违停区域,检查车辆中心是否在多边形内 + if region_polygon is not None: + return self._point_in_polygon(track_info.center, region_polygon) - # 检查车辆中心是否在多边形内 - return self._point_in_polygon(track_info.center, region_polygon) + # 未指定区域时,无法判定是否违停(避免全量误报) + return False def _point_in_polygon(self, point: Tuple[float, float], polygon: List[Tuple[int, int]]) -> bool: @@ -432,12 +314,8 @@ class VehicleDetectionService: """获取性能信息""" return { 'mode': 'local', - 'environment': 'PaddlePaddle', - 'model_dir': self.model_dir, - 'mot_model_dir': self.mot_model_dir, - 'plate_det_model_dir': self.plate_det_model_dir, - 'plate_rec_model_dir': self.plate_rec_model_dir, - 'detector_loaded': self._detector_initialized, + 'environment': 'YOLOv8', + 'detector_loaded': self._yolo_model is not None, 'available': self.available, 'active_tracks': len(self.vehicle_tracks) } @@ -447,16 +325,19 @@ class VehicleDetectionService: class VehicleDetectionModel: """车辆检测模型包装器,兼容 YOLO 接口""" - def __init__(self): - self.service = VehicleDetectionService() + def __init__(self, yolo_model=None): + self._yolo_model = yolo_model + self.service = VehicleDetectionService(yolo_model=yolo_model) + # 保留 names 以兼容 YOLO 接口(实际类别由底层YOLO模型提供) self.names = {0: 'vehicle'} def __call__(self, image, conf=0.1, iou=0.45, verbose=False): """ - 模拟 YOLO 模型的调用接口 + 直接委托给底层YOLO模型,返回标准YOLO结果 """ - result = self.service.detect_image(image, threshold=conf) - return [PaddleDetectionResult(result, self.names)] + if self._yolo_model is not None: + return self._yolo_model(image, conf=conf, iou=iou, verbose=verbose) + raise RuntimeError("YOLO模型未初始化") def detect_illegal_parking(self, image, conf=0.1, illegal_parking_time=5.0, region_polygon=None): @@ -464,122 +345,3 @@ class VehicleDetectionModel: return self.service.detect_illegal_parking( image, conf, illegal_parking_time, region_polygon ) - - -class PaddleDetectionResult: - """模拟 YOLO 检测结果对象""" - - def __init__(self, detection_result: Dict, names: Dict): - self.detection_result = detection_result - self.names = names - self.boxes = self._create_boxes() - - def _create_boxes(self): - """创建模拟的 boxes 对象""" - detections = self.detection_result.get('detections', []) - - if not detections: - return MockBoxes([]) - - xyxy = [] - conf = [] - cls = [] - - for det in detections: - xyxy.append(det['bbox']) - conf.append(det['confidence']) - cls.append(0) - - return MockBoxes(xyxy, conf, cls) - - -class MockBoxes: - """模拟 YOLO boxes 对象""" - - def __init__(self, xyxy_list, conf_list=None, cls_list=None): - try: - import torch - use_torch = True - except ImportError: - use_torch = False - - if xyxy_list and len(xyxy_list) > 0: - if use_torch: - self.xyxy = torch.tensor(xyxy_list, dtype=torch.float32) - self.conf = torch.tensor(conf_list, dtype=torch.float32).reshape(-1, 1) - self.cls = torch.tensor(cls_list, dtype=torch.int64).reshape(-1, 1) - else: - self.xyxy = np.array(xyxy_list, dtype=np.float32) - self.conf = np.array(conf_list, dtype=np.float32).reshape(-1, 1) - self.cls = np.array(cls_list, dtype=np.int64).reshape(-1, 1) - else: - if use_torch: - self.xyxy = torch.empty((0, 4), dtype=torch.float32) - self.conf = torch.empty((0, 1), dtype=torch.float32) - self.cls = torch.empty((0, 1), dtype=torch.int64) - else: - self.xyxy = np.array([]).reshape(0, 4) - self.conf = np.array([]).reshape(0, 1) - self.cls = np.array([], dtype=np.int64).reshape(0, 1) - - self._use_torch = use_torch - - def __iter__(self): - for i in range(len(self.xyxy)): - yield MockBox( - self.xyxy[i], - self.conf[i][0] if len(self.conf) > i else 0.0, - self.cls[i][0] if len(self.cls) > i else 0 - ) - - def __len__(self): - return len(self.xyxy) - - def cpu(self): - return self - - def numpy(self): - if self._use_torch: - if len(self.xyxy) > 0: - return ( - self.xyxy.numpy(), - self.conf.numpy(), - self.cls.numpy() - ) - else: - return ( - np.array([]).reshape(0, 4), - np.array([]).reshape(0, 1), - np.array([], dtype=np.int64).reshape(0, 1) - ) - else: - return ( - self.xyxy, - self.conf, - self.cls - ) - - -class MockBox: - """模拟单个 YOLO box 对象""" - - def __init__(self, xyxy, conf, cls): - try: - import torch - use_torch = True - except ImportError: - use_torch = False - - if use_torch: - if isinstance(xyxy, torch.Tensor): - self.xyxy = xyxy - else: - self.xyxy = torch.tensor(xyxy, dtype=torch.float32) - else: - if isinstance(xyxy, np.ndarray): - self.xyxy = xyxy - else: - self.xyxy = np.array(xyxy, dtype=np.float32) - - self.conf = conf - self.cls = cls diff --git a/apps/web/src/api/detection.js b/apps/web/src/api/detection.js index e1715d6..0c1d592 100644 --- a/apps/web/src/api/detection.js +++ b/apps/web/src/api/detection.js @@ -25,11 +25,18 @@ export const detectionApi = { return api.get('/algorithms/config') }, - detectImage(formData, algorithmConfig = null) { - const params = {} + detectImage(formData, modelId, confidence, iou, algorithmConfig = null, regionPolygon = null) { + const params = { + model_id: modelId, + confidence, + iou + } if (algorithmConfig) { params.algorithm_config = JSON.stringify(algorithmConfig) } + if (regionPolygon) { + params.region_polygon = JSON.stringify(regionPolygon) + } return api.post('/detect/image', formData, { headers: { diff --git a/apps/web/src/components/ImageDetection.vue b/apps/web/src/components/ImageDetection.vue index caffa1d..645f635 100644 --- a/apps/web/src/components/ImageDetection.vue +++ b/apps/web/src/components/ImageDetection.vue @@ -50,10 +50,26 @@ -