gbreadman13/sam2-api
0
1"""2REST API сервер для сегментации изображений через SAM2.3Уставший сеньор кодит это в 3 часа ночи, поэтому код местами будет грязный.4"""5 6from contextlib import asynccontextmanager7from fastapi import FastAPI, File, UploadFile, HTTPException, Query, Body8from fastapi.responses import JSONResponse, HTMLResponse, FileResponse9from fastapi.middleware.cors import CORSMiddleware10from pydantic import BaseModel, Field11from PIL import Image12import numpy as np13import torch14import io15import os16import base6417import cv218from typing import List, Dict, Any, Optional, Literal19import logging20from datetime import datetime21import json22 23# Настройка логирования, потому что дебажить это говно иначе невозможно24logging.basicConfig(level=logging.INFO)25logger = logging.getLogger(__name__)26 27# Глобальные переменные для модели (лень каждый раз загружать)28predictor = None29device = None30 31# ===== Pydantic модели для батчинг API =====32 33class BBoxModel(BaseModel):34 """Bounding box в нормализованных координатах (0.0 - 1.0) или пиксельных"""35 x_min: float = Field(..., description="X координата левого верхнего угла")36 y_min: float = Field(..., description="Y координата левого верхнего угла")37 x_max: float = Field(..., description="X координата правого нижнего угла")38 y_max: float = Field(..., description="Y координата правого нижнего угла")39 40class PromptModel(BaseModel):41 """Промпт для сегментации одного объекта"""42 id: int = Field(..., description="Уникальный ID объекта")43 type: Literal["mask", "box", "points"] = Field(..., description="Тип промпта")44 data: str = Field(..., description="Данные промпта (base64 для mask, JSON для points)")45 bbox: Optional[BBoxModel] = Field(None, description="Опциональный bounding box")46 label: Optional[str] = Field(None, description="Метка объекта (person, car, etc)")47 selected: bool = Field(True, description="Обрабатывать ли этот промпт")48 49class SegmentOptionsModel(BaseModel):50 """Опции сегментации"""51 extract_objects: bool = Field(True, description="Вернуть вырезанные объекты")52 include_masks: bool = Field(False, description="Включить контуры масок")53 clean_masks: bool = Field(True, description="Очистить маски от артефактов")54 55class BatchSegmentRequest(BaseModel):56 """Запрос на батчинг сегментацию"""57 image: str = Field(..., description="Изображение в base64 (с data URL или без)")58 prompts: List[PromptModel] = Field(..., description="Массив промптов")59 options: Optional[SegmentOptionsModel] = Field(default_factory=SegmentOptionsModel)60 61class SegmentResultModel(BaseModel):62 """Результат сегментации одного объекта"""63 id: int64 label: Optional[str] = None65 bbox: Dict[str, Any]66 area: int67 center: Dict[str, int]68 confidence: float69 extracted_image: Optional[str] = None70 contours: Optional[List[Dict[str, Any]]] = None71 mask_rle: Optional[Dict[str, Any]] = None72 73class BatchSegmentResponse(BaseModel):74 """Ответ батчинг сегментации"""75 success: bool76 image_size: Dict[str, int]77 results: List[SegmentResultModel]78 79def save_batch_request_log(request_data: dict, response_data: dict, image_width: int, image_height: int):80 """81 Сохраняет запрос батчинга для аудита и дебага.82 Создает папку с timestamp и сохраняет только метаданные:83 1. Лог запроса (request.json) - параметры без base6484 2. Лог ответа (response.json) - результаты без base6485 3. Краткую сводку (summary.json)86 87 ⚠️ Изображения и маски НЕ сохраняются для безопасности!88 """89 try:90 # Создаем корневую папку для логов91 logs_dir = "batch_logs"92 os.makedirs(logs_dir, exist_ok=True)93 94 # Создаем папку с timestamp95 timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f")[:-3] # Миллисекунды96 request_dir = os.path.join(logs_dir, timestamp)97 os.makedirs(request_dir, exist_ok=True)98 99 logger.info(f"📁 Сохраняю лог запроса в: {request_dir}")100 101 # Сохраняем запрос (без base64 для безопасности)102 request_log = {103 "timestamp": timestamp,104 "image_size": {105 "width": image_width,106 "height": image_height107 },108 "prompts": [109 {110 "id": p.get("id"),111 "type": p.get("type"),112 "label": p.get("label"),113 "bbox": p.get("bbox"),114 "selected": p.get("selected"),115 "data_length": len(p.get("data", "")) # Длина вместо самих данных116 }117 for p in request_data.get("prompts", [])118 ],119 "options": request_data.get("options", {})120 }121 122 request_path = os.path.join(request_dir, "request.json")123 with open(request_path, "w", encoding="utf-8") as f:124 json.dump(request_log, f, indent=2, ensure_ascii=False)125 logger.info(f" ✓ Сохранен лог запроса: {request_path}")126 127 # 4. Сохраняем ответ (без base64 объектов)128 response_log = {129 "timestamp": timestamp,130 "success": response_data.get("success"),131 "image_size": response_data.get("image_size"),132 "results": [133 {134 "id": r.get("id"),135 "label": r.get("label"),136 "bbox": r.get("bbox"),137 "area": r.get("area"),138 "center": r.get("center"),139 "confidence": r.get("confidence"),140 "has_extracted_image": "extracted_image" in r,141 "has_contours": "contours" in r142 }143 for r in response_data.get("results", [])144 ]145 }146 147 response_path = os.path.join(request_dir, "response.json")148 with open(response_path, "w", encoding="utf-8") as f:149 json.dump(response_log, f, indent=2, ensure_ascii=False)150 logger.info(f" ✓ Сохранен лог ответа: {response_path}")151 152 # 3. Создаем summary файл153 summary = {154 "timestamp": timestamp,155 "processed_prompts": len(response_data.get("results", [])),156 "total_prompts": len(request_data.get("prompts", [])),157 "selected_prompts": len([p for p in request_data.get("prompts", []) if p.get("selected", True)]),158 "image_size": f"{image_width}x{image_height}",159 "prompt_types": [p.get("type") for p in request_data.get("prompts", [])],160 "files": {161 "request": "request.json",162 "response": "response.json"163 }164 }165 166 summary_path = os.path.join(request_dir, "summary.json")167 with open(summary_path, "w", encoding="utf-8") as f:168 json.dump(summary, f, indent=2, ensure_ascii=False)169 170 logger.info(f"✅ Лог запроса сохранен: {request_dir}")171 172 except Exception as e:173 logger.error(f"❌ Ошибка при сохранении лога: {e}")174 # Не прерываем обработку запроса если не удалось сохранить лог175 176def load_model(checkpoint_path: str = "checkpoints/sam2.1_hiera_tiny.pt"):177 """178 Загружает модель SAM2. 179 Вызывается один раз при старте сервера.180 """181 global predictor, device182 183 try:184 from sam2.build_sam import build_sam2185 from sam2.sam2_image_predictor import SAM2ImagePredictor186 187 # Проверяем CUDA188 device = "cuda" if torch.cuda.is_available() else "cpu"189 logger.info(f"Используем устройство: {device}")190 191 if device == "cpu":192 logger.warning("CUDA недоступна, работаем на CPU (будет медленно как черепаха)")193 194 # Определяем конфиг по имени файла чекпоинта195 # Указываем путь относительно configs/ директории в пакете sam2196 checkpoint_name = os.path.basename(checkpoint_path)197 if "tiny" in checkpoint_name:198 config = "configs/sam2.1/sam2.1_hiera_t.yaml"199 elif "small" in checkpoint_name:200 config = "configs/sam2.1/sam2.1_hiera_s.yaml"201 elif "base_plus" in checkpoint_name:202 config = "configs/sam2.1/sam2.1_hiera_b+.yaml"203 elif "large" in checkpoint_name:204 config = "configs/sam2.1/sam2.1_hiera_l.yaml"205 else:206 logger.warning(f"Неизвестный тип модели, пробую tiny конфиг")207 config = "configs/sam2.1/sam2.1_hiera_t.yaml"208 209 logger.info(f"Загружаю модель из {checkpoint_path}")210 logger.info(f"Конфиг: {config}")211 212 sam2_model = build_sam2(config, checkpoint_path, device=device)213 predictor = SAM2ImagePredictor(sam2_model)214 215 logger.info("✓ Модель загружена успешно")216 217 except Exception as e:218 logger.error(f"Не удалось загрузить модель: {e}")219 logger.error("Убедись что SAM2 установлен (./install_sam2.sh)")220 raise221 222@asynccontextmanager223async def lifespan(app: FastAPI):224 """Загружаем модель при старте, выгружаем при остановке"""225 # Startup226 checkpoint_dir = "checkpoints"227 if os.path.exists(checkpoint_dir):228 checkpoints = [f for f in os.listdir(checkpoint_dir) if f.endswith(".pt")]229 if checkpoints:230 checkpoint_path = os.path.join(checkpoint_dir, checkpoints[0])231 load_model(checkpoint_path)232 else:233 logger.error("Нет чекпоинтов в директории checkpoints/")234 logger.error("Запусти: python download_model.py")235 else:236 logger.error("Директория checkpoints/ не найдена")237 238 yield # Сервер работает239 240 # Shutdown (если нужна очистка)241 242# Создаем FastAPI приложение с lifespan243app = FastAPI(244 title="SAM2 Segmentation API",245 description="API для автоматической сегментации объектов на изображениях",246 version="1.0.0",247 lifespan=lifespan248)249 250# Добавляем CORS для работы с веб-интерфейсом251app.add_middleware(252 CORSMiddleware,253 allow_origins=["*"], # В продакшене указать конкретные домены254 allow_credentials=True,255 allow_methods=["*"],256 allow_headers=["*"],257)258 259@app.get("/")260async def root():261 """Главная страница - информация об API"""262 return {263 "message": "SAM2 Segmentation API работает",264 "version": "2.0.0",265 "web_ui": {266 "simple": "/web - Box промпты",267 "advanced": "/web/advanced - Box + Brush промпты (рисование)"268 },269 "docs": "/docs",270 "endpoints": {271 "POST /segment": "Сегментация изображения (поддерживает points, box, mask via query params)",272 "POST /segment/batch": "🔥 Батчинг сегментация (JSON API для множественных объектов)",273 "POST /segment/auto": "Автоматическая сегментация всех объектов",274 "GET /health": "Проверка здоровья сервиса"275 }276 }277 278@app.get("/web", response_class=HTMLResponse)279async def web_interface():280 """Веб-интерфейс для тестирования Box Prompts (простой)"""281 web_demo_path = os.path.join(os.path.dirname(__file__), "web_demo.html")282 if os.path.exists(web_demo_path):283 with open(web_demo_path, "r", encoding="utf-8") as f:284 return f.read()285 else:286 return "<h1>Веб-интерфейс не найден</h1><p>Файл web_demo.html отсутствует</p>"287 288@app.get("/web/advanced", response_class=HTMLResponse)289async def web_interface_advanced():290 """Продвинутый веб-интерфейс с Box + Brush промптами"""291 web_demo_path = os.path.join(os.path.dirname(__file__), "web_demo_advanced.html")292 if os.path.exists(web_demo_path):293 with open(web_demo_path, "r", encoding="utf-8") as f:294 return f.read()295 else:296 return "<h1>Продвинутый интерфейс не найден</h1><p>Файл web_demo_advanced.html отсутствует</p>"297 298@app.get("/health")299async def health():300 """Проверка что всё ок"""301 return {302 "status": "healthy" if predictor is not None else "model not loaded",303 "device": str(device) if device else "unknown"304 }305 306def process_image(image_bytes: bytes) -> np.ndarray:307 """Конвертирует байты в numpy array"""308 image = Image.open(io.BytesIO(image_bytes))309 if image.mode != "RGB":310 image = image.convert("RGB")311 return np.array(image)312 313def masks_to_coords(masks: np.ndarray, include_contours: bool = False) -> List[Dict[str, Any]]:314 """315 Конвертирует маски в координаты bounding box и контуров.316 masks: (N, H, W) - N масок317 include_contours: если True, добавляет контуры масок318 """319 results = []320 321 for i, mask in enumerate(masks):322 # Находим координаты пикселей маски323 y_coords, x_coords = np.where(mask > 0)324 325 if len(x_coords) == 0:326 continue327 328 # Bounding box329 x_min, x_max = int(x_coords.min()), int(x_coords.max())330 y_min, y_max = int(y_coords.min()), int(y_coords.max())331 332 # Площадь сегмента333 area = int(mask.sum())334 335 segment_data = {336 "segment_id": i,337 "bbox": {338 "x_min": x_min,339 "y_min": y_min,340 "x_max": x_max,341 "y_max": y_max,342 "width": x_max - x_min,343 "height": y_max - y_min344 },345 "area": area,346 "center": {347 "x": int(x_coords.mean()),348 "y": int(y_coords.mean())349 }350 }351 352 # Добавляем контуры если нужно353 if include_contours:354 try:355 # Конвертируем маску в uint8 (защита от булевых масок)356 if mask.dtype == bool:357 mask_uint8 = mask.astype(np.uint8) * 255358 else:359 mask_uint8 = (mask * 255).astype(np.uint8)360 361 # Находим контуры с иерархией для поддержки "дыр"362 # RETR_CCOMP: находит внешние контуры И внутренние дыры (holes)363 # CHAIN_APPROX_NONE: сохраняет ВСЕ точки для pixel-perfect результата364 contours, hierarchy = cv2.findContours(mask_uint8, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_NONE)365 except Exception as e:366 logger.warning(f"Ошибка при извлечении контуров: {e}, использую fallback")367 # Fallback на простое извлечение без иерархии368 contours, hierarchy = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)369 hierarchy = None370 371 # Конвертируем контуры в список точек с учетом иерархии372 contour_data = []373 374 if hierarchy is not None and len(contours) > 0:375 hierarchy = hierarchy[0] # OpenCV возвращает hierarchy в странном формате376 377 for i, contour in enumerate(contours):378 try:379 # Небольшое упрощение только для очень больших контуров380 if len(contour) > 1000:381 arc_length = cv2.arcLength(contour, True)382 if arc_length > 0: # Защита от деления на 0383 epsilon = 0.0005 * arc_length384 approx = cv2.approxPolyDP(contour, epsilon, True)385 else:386 approx = contour387 else:388 approx = contour389 390 # Конвертируем в список [x, y]391 points = [[int(point[0][0]), int(point[0][1])] for point in approx]392 393 if len(points) > 2:394 # hierarchy[i] = [Next, Previous, First_Child, Parent]395 # Если Parent == -1, это внешний контур396 # Если Parent >= 0, это дыра (hole) внутри родительского контура397 is_hole = hierarchy[i][3] != -1398 399 contour_data.append({400 "points": points,401 "is_hole": is_hole402 })403 except Exception as e:404 logger.warning(f"Ошибка при обработке контура {i}: {e}")405 continue406 else:407 # Fallback если hierarchy не вернулась408 for contour in contours:409 try:410 if len(contour) > 1000:411 arc_length = cv2.arcLength(contour, True)412 if arc_length > 0:413 epsilon = 0.0005 * arc_length414 approx = cv2.approxPolyDP(contour, epsilon, True)415 else:416 approx = contour417 else:418 approx = contour419 420 points = [[int(point[0][0]), int(point[0][1])] for point in approx]421 if len(points) > 2:422 contour_data.append({423 "points": points,424 "is_hole": False425 })426 except Exception as e:427 logger.warning(f"Ошибка при обработке контура: {e}")428 continue429 430 segment_data["contours"] = contour_data if len(contour_data) > 0 else []431 432 # Также добавляем RLE (Run-Length Encoding) для компактного представления433 # Это полезно если нужно восстановить точную маску434 segment_data["mask_rle"] = mask_to_rle(mask)435 436 results.append(segment_data)437 438 return results439 440def mask_to_rle(mask: np.ndarray) -> Dict[str, Any]:441 """442 Конвертирует бинарную маску в RLE (Run-Length Encoding)443 Компактное представление маски444 """445 # Конвертируем в int если это bool446 if mask.dtype == bool:447 pixels = mask.astype(np.uint8).flatten()448 else:449 pixels = mask.flatten()450 451 pixels = np.concatenate([[0], pixels, [0]])452 runs = np.where(pixels[1:] != pixels[:-1])[0] + 1453 runs[1::2] -= runs[::2]454 455 return {456 "counts": [int(x) for x in runs], # Конвертируем numpy int в Python int457 "size": [int(x) for x in mask.shape] # Конвертируем в Python int458 }459 460def convert_to_native_types(obj):461 """462 Рекурсивно конвертирует numpy типы в нативные Python типы463 Нужно для сериализации в JSON через FastAPI464 """465 if isinstance(obj, np.integer):466 return int(obj)467 elif isinstance(obj, np.floating):468 return float(obj)469 elif isinstance(obj, np.ndarray):470 return obj.tolist()471 elif isinstance(obj, np.bool_):472 return bool(obj)473 elif isinstance(obj, dict):474 return {key: convert_to_native_types(value) for key, value in obj.items()}475 elif isinstance(obj, list):476 return [convert_to_native_types(item) for item in obj]477 return obj478 479def clean_mask(mask: np.ndarray, min_area: int = 100) -> np.ndarray:480 """481 Очищает маску от мелких артефактов и дыр.482 Более мягкий вариант - не убивает тонкие детали типа лямок.483 484 mask: бинарная маска (H, W)485 min_area: минимальная площадь компонента в пикселях486 487 Returns: очищенная маска488 """489 # Конвертируем в uint8 если нужно490 if mask.dtype == bool:491 mask_uint8 = mask.astype(np.uint8) * 255492 else:493 mask_uint8 = (mask * 255).astype(np.uint8)494 495 # Только легкое закрытие для удаления мелких дыр внутри объекта496 # Используем маленький kernel чтобы не убить тонкие детали (лямки, пальцы и т.д.)497 kernel = np.ones((2, 2), np.uint8)498 mask_uint8 = cv2.morphologyEx(mask_uint8, cv2.MORPH_CLOSE, kernel, iterations=1)499 500 # УБРАЛ MORPH_OPEN - он убивал тонкие элементы типа лямок портфеля501 # УБРАЛ фильтрацию по площади компонентов - она тоже могла вырезать лямки502 503 return (mask_uint8 > 127).astype(bool)504 505def extract_object_image(image: np.ndarray, mask: np.ndarray, clean: bool = True) -> str:506 """507 Вырезает объект из изображения по маске и возвращает base64 PNG с прозрачностью.508 509 image: RGB изображение (H, W, 3)510 mask: бинарная маска (H, W)511 clean: применить постобработку для удаления артефактов512 513 Returns: base64 строка PNG изображения с альфа-каналом514 """515 # Конвертируем маску в bool если нужно516 if mask.dtype != bool:517 mask = mask > 0.5518 519 # Очищаем маску от артефактов520 if clean:521 mask = clean_mask(mask, min_area=100)522 523 # Создаем RGBA изображение524 h, w = image.shape[:2]525 rgba = np.zeros((h, w, 4), dtype=np.uint8)526 rgba[:, :, :3] = image # RGB каналы527 rgba[:, :, 3] = (mask * 255).astype(np.uint8) # Alpha канал из маски528 529 # Конвертируем в PIL Image530 pil_image = Image.fromarray(rgba, 'RGBA')531 532 # Конвертируем в base64533 buffer = io.BytesIO()534 pil_image.save(buffer, format='PNG')535 buffer.seek(0)536 img_base64 = base64.b64encode(buffer.read()).decode('utf-8')537 538 return f"data:image/png;base64,{img_base64}"539 540@app.post("/segment")541async def segment_image(542 file: UploadFile = File(...),543 point_x: List[float] = Query(None, description="X координаты точек промпта"),544 point_y: List[float] = Query(None, description="Y координаты точек промпта"),545 point_labels: List[int] = Query(None, description="Лейблы точек (1=foreground, 0=background)"),546 box_x1: float = Query(None, description="X координата левого верхнего угла бокса"),547 box_y1: float = Query(None, description="Y координата левого верхнего угла бокса"),548 box_x2: float = Query(None, description="X координата правого нижнего угла бокса"),549 box_y2: float = Query(None, description="Y координата правого нижнего угла бокса"),550 mask_data: str = Query(None, description="Base64 закодированная маска (PNG с альфа-каналом)"),551 include_masks: bool = Query(True, description="Включить контуры масок в ответ"),552 extract_objects: bool = Query(False, description="Вернуть вырезанные объекты как base64 PNG"),553):554 """555 Сегментирует изображение по промпту (точкам, боксу, маске или их комбинации).556 557 Поддерживаемые промпты:558 - Точки (point_x, point_y, point_labels) - клики пользователя559 - Бокс (box_x1, box_y1, box_x2, box_y2) - прямоугольное выделение560 - Маска (mask_data) - нарисованная кистью маска (зеленый=foreground, красный=background)561 - Комбинация промптов - для максимальной точности562 563 Если промпты не указаны, сегментирует центральный объект.564 Если include_masks=True, возвращает контуры масок для точной отрисовки.565 Если extract_objects=True, возвращает готовые вырезанные объекты как base64 PNG.566 """567 if predictor is None:568 raise HTTPException(status_code=503, detail="Модель не загружена, перезапусти сервер")569 570 try:571 # Читаем изображение572 image_bytes = await file.read()573 image = process_image(image_bytes)574 575 logger.info(f"Обрабатываю изображение: {image.shape}")576 logger.info(f"Параметры: include_masks={include_masks}, extract_objects={extract_objects}")577 578 # Устанавливаем изображение в предиктор579 predictor.set_image(image)580 581 # Подготавливаем промпты582 points = None583 labels = None584 box = None585 586 # Проверяем наличие точек587 if point_x and point_y:588 if len(point_x) != len(point_y):589 raise HTTPException(status_code=400, detail="Количество X и Y координат должно совпадать")590 points = np.array([[x, y] for x, y in zip(point_x, point_y)])591 labels = np.array(point_labels) if point_labels else np.ones(len(points))592 logger.info(f"Промпт: {len(points)} точек")593 594 # Проверяем наличие бокса595 if all(v is not None for v in [box_x1, box_y1, box_x2, box_y2]):596 box = np.array([box_x1, box_y1, box_x2, box_y2])597 logger.info(f"Промпт: бокс [{box_x1:.1f}, {box_y1:.1f}, {box_x2:.1f}, {box_y2:.1f}]")598 599 # Валидация бокса600 if box_x2 <= box_x1 or box_y2 <= box_y1:601 raise HTTPException(602 status_code=400, 603 detail="Некорректный бокс: x2 должен быть больше x1, y2 больше y1"604 )605 606 # Проверяем наличие нарисованной маски607 if mask_data:608 logger.info("Обрабатываю нарисованную маску...")609 try:610 # Декодируем base64611 if ',' in mask_data:612 mask_data = mask_data.split(',')[1] # Убираем data:image/png;base64,613 614 mask_bytes = base64.b64decode(mask_data)615 mask_image = Image.open(io.BytesIO(mask_bytes)).convert('RGBA')616 mask_array = np.array(mask_image)617 618 # Извлекаем foreground и background пиксели619 # Поддерживаем несколько форматов:620 # 1. Зеленый (R<100, G>150, B<100) - классический foreground621 # 2. Белый/светлый (R>200, G>200, B>200) - часто используется фронтами622 # 3. Красный (R>150, G<100, B<100) - background623 624 green_mask = (mask_array[:, :, 0] < 100) & (mask_array[:, :, 1] > 150) & (mask_array[:, :, 2] < 100) & (mask_array[:, :, 3] > 0)625 white_mask = (mask_array[:, :, 0] > 200) & (mask_array[:, :, 1] > 200) & (mask_array[:, :, 2] > 200) & (mask_array[:, :, 3] > 0)626 red_mask = (mask_array[:, :, 0] > 150) & (mask_array[:, :, 1] < 100) & (mask_array[:, :, 2] < 100) & (mask_array[:, :, 3] > 0)627 628 # Объединяем зеленые и белые как foreground629 foreground_mask = green_mask | white_mask630 631 # Сэмплируем точки из закрашенных областей632 mask_points = []633 mask_labels = []634 635 # Foreground точки (зеленые + белые)636 foreground_coords = np.argwhere(foreground_mask)637 if len(foreground_coords) > 0:638 # Масштабируем к размеру исходного изображения639 scale_y = image.shape[0] / mask_array.shape[0]640 scale_x = image.shape[1] / mask_array.shape[1]641 642 # Сэмплируем до 20 точек равномерно (меньше = стабильнее)643 step = max(1, len(foreground_coords) // 20)644 sampled = foreground_coords[::step][:20] # Максимум 20 точек645 646 for y, x in sampled:647 mask_points.append([x * scale_x, y * scale_y])648 mask_labels.append(1) # foreground649 650 # Background точки (красные)651 red_coords = np.argwhere(red_mask)652 if len(red_coords) > 0:653 scale_y = image.shape[0] / mask_array.shape[0]654 scale_x = image.shape[1] / mask_array.shape[1]655 656 step = max(1, len(red_coords) // 20)657 sampled = red_coords[::step][:20] # Максимум 20 точек658 659 for y, x in sampled:660 mask_points.append([x * scale_x, y * scale_y])661 mask_labels.append(0) # background662 663 if mask_points:664 # Объединяем с существующими точками665 if points is not None:666 points = np.vstack([points, np.array(mask_points)])667 labels = np.concatenate([labels, np.array(mask_labels)])668 else:669 points = np.array(mask_points)670 labels = np.array(mask_labels)671 672 logger.info(f"Промпт из маски: {len(mask_points)} точек ({np.sum(np.array(mask_labels) == 1)} foreground, {np.sum(np.array(mask_labels) == 0)} background)")673 else:674 logger.warning("Маска пустая или не содержит foreground (зеленых/белых) или background (красных) пикселей")675 676 except Exception as e:677 logger.error(f"Ошибка обработки маски: {e}")678 raise HTTPException(status_code=400, detail=f"Некорректная маска: {str(e)}")679 680 # Делаем предсказание с промптами681 if points is not None or box is not None:682 logger.info(f"Используем промпты: points={points is not None}, box={box is not None}")683 684 # Если много точек (>10), используем single mask для стабильности685 # Если мало точек или только box, используем multimask для вариативности686 use_multimask = True687 if points is not None and len(points) > 10:688 use_multimask = False689 logger.info("Много точек, используем single mask mode для стабильности")690 691 masks, scores, logits = predictor.predict(692 point_coords=points,693 point_labels=labels,694 box=box,695 multimask_output=use_multimask,696 )697 698 # Если multimask, выбираем лучшую по score699 if use_multimask and len(masks) > 1:700 best_idx = np.argmax(scores)701 masks = masks[best_idx:best_idx+1]702 scores = scores[best_idx:best_idx+1]703 logger.info(f"Выбрана маска {best_idx} с confidence {scores[0]:.3f}")704 else:705 # Автоматическая сегментация - берем центральную точку706 logger.info("Промпты не указаны, сегментирую центральный объект")707 h, w = image.shape[:2]708 point = np.array([[w // 2, h // 2]])709 label = np.array([1])710 711 masks, scores, logits = predictor.predict(712 point_coords=point,713 point_labels=label,714 multimask_output=True,715 )716 717 # Конвертируем маски в координаты (с контурами если нужно)718 segments = masks_to_coords(masks, include_contours=include_masks)719 720 logger.info(f"Найдено сегментов: {len(segments)}, масок: {len(masks)}")721 logger.info(f"extract_objects = {extract_objects}")722 723 # Добавляем confidence scores724 for i, seg in enumerate(segments):725 seg["confidence"] = float(scores[i]) if i < len(scores) else 0.0726 727 # Если нужно - вырезаем объект и добавляем base64728 logger.info(f"Обрабатываю сегмент {i}: extract_objects={extract_objects}, i < len(masks) = {i < len(masks)}")729 if extract_objects and i < len(masks):730 logger.info(f"Вырезаю объект {i}...")731 seg["extracted_image"] = extract_object_image(image, masks[i])732 logger.info(f"✓ Вырезан объект {i}, размер маски: {masks[i].sum()} пикселей")733 else:734 logger.warning(f"❌ Пропускаю объект {i}: extract_objects={extract_objects}")735 736 result = {737 "success": True,738 "image_size": {739 "width": int(image.shape[1]),740 "height": int(image.shape[0])741 },742 "segments_count": len(segments),743 "segments": segments744 }745 746 # Конвертируем все numpy типы в нативные Python типы747 return convert_to_native_types(result)748 749 except Exception as e:750 logger.error(f"Ошибка при сегментации: {e}")751 raise HTTPException(status_code=500, detail=f"Ошибка обработки: {str(e)}")752 753@app.post("/segment/auto")754async def segment_auto(755 file: UploadFile = File(...),756 points_per_side: int = Query(32, description="Количество точек на сторону для автосегментации"),757 include_masks: bool = Query(True, description="Включить контуры масок в ответ"),758):759 """760 Автоматическая сегментация всех объектов на изображении.761 Использует grid of points для поиска всех возможных объектов.762 Если include_masks=True, возвращает контуры масок для точной отрисовки.763 """764 if predictor is None:765 raise HTTPException(status_code=503, detail="Модель не загружена")766 767 try:768 image_bytes = await file.read()769 image = process_image(image_bytes)770 771 logger.info(f"Автосегментация изображения: {image.shape}")772 773 predictor.set_image(image)774 775 # Создаем сетку точек776 h, w = image.shape[:2]777 x_coords = np.linspace(0, w, points_per_side)778 y_coords = np.linspace(0, h, points_per_side)779 780 all_segments = []781 segment_id = 0782 783 # Для каждой точки в сетке пытаемся найти объект784 for y in y_coords:785 for x in x_coords:786 point = np.array([[x, y]])787 label = np.array([1])788 789 masks, scores, _ = predictor.predict(790 point_coords=point,791 point_labels=label,792 multimask_output=False,793 )794 795 if masks.shape[0] > 0 and scores[0] > 0.5: # Порог confidence796 segments = masks_to_coords(masks, include_contours=include_masks)797 for seg in segments:798 seg["segment_id"] = segment_id799 seg["confidence"] = float(scores[0])800 all_segments.append(seg)801 segment_id += 1802 803 # Убираем дубликаты (примерно)804 # Два сегмента считаем дубликатами если их центры близко805 unique_segments = []806 for seg in all_segments:807 is_duplicate = False808 for unique_seg in unique_segments:809 dx = seg["center"]["x"] - unique_seg["center"]["x"]810 dy = seg["center"]["y"] - unique_seg["center"]["y"]811 dist = (dx**2 + dy**2) ** 0.5812 813 if dist < 50: # Порог расстояния между центрами814 is_duplicate = True815 break816 817 if not is_duplicate:818 unique_segments.append(seg)819 820 result = {821 "success": True,822 "image_size": {823 "width": int(image.shape[1]),824 "height": int(image.shape[0])825 },826 "segments_count": len(unique_segments),827 "segments": unique_segments828 }829 830 # Конвертируем все numpy типы в нативные Python типы831 return convert_to_native_types(result)832 833 except Exception as e:834 logger.error(f"Ошибка при автосегментации: {e}")835 raise HTTPException(status_code=500, detail=f"Ошибка обработки: {str(e)}")836 837@app.post("/segment/batch", response_model=BatchSegmentResponse)838async def segment_batch(request: BatchSegmentRequest = Body(...)):839 """840 Батчинг сегментация нескольких объектов.841 842 Принимает изображение и массив промптов (mask/box/points).843 Обрабатывает каждый selected промпт отдельно.844 Возвращает массив результатов с метаданными.845 846 Идеально для:847 - Множественных объектов848 - Мобильных приложений849 - Когда фронт уже разделил объекты850 """851 if predictor is None:852 raise HTTPException(status_code=503, detail="Модель не загружена, перезапусти сервер")853 854 try:855 # Декодируем изображение из base64856 image_data = request.image857 if ',' in image_data:858 image_data = image_data.split(',')[1] # Убираем data:image/...;base64,859 860 image_bytes = base64.b64decode(image_data)861 image = process_image(image_bytes)862 863 logger.info(f"Батчинг сегментация: {image.shape}, промптов: {len(request.prompts)}")864 865 # Устанавливаем изображение один раз866 predictor.set_image(image)867 868 results = []869 870 # Фильтруем только selected промпты871 selected_prompts = [p for p in request.prompts if p.selected]872 logger.info(f"Обрабатываем {len(selected_prompts)} из {len(request.prompts)} промптов")873 874 # Обрабатываем каждый промпт отдельно875 for prompt in selected_prompts:876 logger.info(f"Обрабатываю промпт #{prompt.id}, тип: {prompt.type}, label: {prompt.label}")877 878 try:879 # Подготавливаем промпт в зависимости от типа880 points = None881 labels = None882 box = None883 884 if prompt.type == "mask":885 # Декодируем маску и извлекаем точки886 mask_data = prompt.data887 if ',' in mask_data:888 mask_data = mask_data.split(',')[1]889 890 mask_bytes = base64.b64decode(mask_data)891 mask_image = Image.open(io.BytesIO(mask_bytes)).convert('RGBA')892 mask_array = np.array(mask_image)893 894 # Извлекаем foreground и background пиксели895 # Поддерживаем несколько форматов:896 # 1. Зеленый (R<100, G>150, B<100) - классический foreground897 # 2. Белый/светлый (R>200, G>200, B>200) - часто используется фронтами898 # 3. Красный (R>150, G<100, B<100) - background899 900 green_mask = (mask_array[:, :, 0] < 100) & (mask_array[:, :, 1] > 150) & (mask_array[:, :, 2] < 100) & (mask_array[:, :, 3] > 0)901 white_mask = (mask_array[:, :, 0] > 200) & (mask_array[:, :, 1] > 200) & (mask_array[:, :, 2] > 200) & (mask_array[:, :, 3] > 0)902 red_mask = (mask_array[:, :, 0] > 150) & (mask_array[:, :, 1] < 100) & (mask_array[:, :, 2] < 100) & (mask_array[:, :, 3] > 0)903 904 # Объединяем зеленые и белые как foreground905 foreground_mask = green_mask | white_mask906 907 mask_points = []908 mask_labels = []909 910 # Foreground точки (зеленые + белые)911 foreground_coords = np.argwhere(foreground_mask)912 if len(foreground_coords) > 0:913 scale_y = image.shape[0] / mask_array.shape[0]914 scale_x = image.shape[1] / mask_array.shape[1]915 step = max(1, len(foreground_coords) // 20)916 sampled = foreground_coords[::step][:20]917 918 for y, x in sampled:919 mask_points.append([x * scale_x, y * scale_y])920 mask_labels.append(1)921 922 # Background точки923 red_coords = np.argwhere(red_mask)924 if len(red_coords) > 0:925 scale_y = image.shape[0] / mask_array.shape[0]926 scale_x = image.shape[1] / mask_array.shape[1]927 step = max(1, len(red_coords) // 20)928 sampled = red_coords[::step][:20]929 930 for y, x in sampled:931 mask_points.append([x * scale_x, y * scale_y])932 mask_labels.append(0)933 934 if mask_points:935 points = np.array(mask_points)936 labels = np.array(mask_labels)937 938 elif prompt.type == "box":939 # Парсим bbox - может быть нормализованный (0-1) или пиксельный940 bbox_data = prompt.bbox if prompt.bbox else None941 942 if bbox_data:943 x1 = bbox_data.x_min944 y1 = bbox_data.y_min945 x2 = bbox_data.x_max946 y2 = bbox_data.y_max947 948 # Если нормализованные координаты (0-1), конвертируем в пиксели949 if x2 <= 1.0 and y2 <= 1.0:950 x1 *= image.shape[1]951 x2 *= image.shape[1]952 y1 *= image.shape[0]953 y2 *= image.shape[0]954 955 box = np.array([x1, y1, x2, y2])956 957 elif prompt.type == "points":958 # Ожидаем JSON в формате [[x, y, label], ...]959 import json960 points_data = json.loads(prompt.data)961 962 points_list = []963 labels_list = []964 965 for point in points_data:966 x, y = point[0], point[1]967 label = point[2] if len(point) > 2 else 1968 969 # Если нормализованные, конвертируем970 if x <= 1.0 and y <= 1.0:971 x *= image.shape[1]972 y *= image.shape[0]973 974 points_list.append([x, y])975 labels_list.append(label)976 977 points = np.array(points_list)978 labels = np.array(labels_list)979 980 # Делаем предсказание981 if points is not None or box is not None:982 # Решаем использовать ли multimask983 use_multimask = True984 if points is not None and len(points) > 10:985 use_multimask = False986 987 masks, scores, logits = predictor.predict(988 point_coords=points,989 point_labels=labels,990 box=box,991 multimask_output=use_multimask,992 )993 994 # Если multimask, выбираем лучшую995 if use_multimask and len(masks) > 1:996 best_idx = np.argmax(scores)997 masks = masks[best_idx:best_idx+1]998 scores = scores[best_idx:best_idx+1]999 1000 # Берем первую маску1001 mask = masks[0]1002 score = float(scores[0])1003 1004 # Очищаем маску если нужно1005 if request.options.clean_masks:1006 mask = clean_mask(mask, min_area=100)1007 1008 # Вычисляем метрики1009 y_coords, x_coords = np.where(mask > 0)1010 1011 if len(x_coords) > 0:1012 x_min, x_max = int(x_coords.min()), int(x_coords.max())1013 y_min, y_max = int(y_coords.min()), int(y_coords.max())1014 area = int(mask.sum())1015 center_x = int(x_coords.mean())1016 center_y = int(y_coords.mean())1017 1018 # Формируем результат1019 result = {1020 "id": prompt.id,1021 "label": prompt.label,1022 "bbox": {1023 "x_min": x_min,1024 "y_min": y_min,1025 "x_max": x_max,1026 "y_max": y_max,1027 "width": x_max - x_min,1028 "height": y_max - y_min1029 },1030 "area": area,1031 "center": {1032 "x": center_x,1033 "y": center_y1034 },1035 "confidence": score1036 }1037 1038 # Добавляем вырезанный объект если нужно1039 if request.options.extract_objects:1040 result["extracted_image"] = extract_object_image(1041 image, mask, clean=request.options.clean_masks1042 )1043 1044 # Добавляем контуры если нужно1045 if request.options.include_masks:1046 segments = masks_to_coords(masks, include_contours=True)1047 if segments:1048 result["contours"] = segments[0].get("contours", [])1049 result["mask_rle"] = segments[0].get("mask_rle", {})1050 1051 results.append(result)1052 logger.info(f"✓ Промпт #{prompt.id} обработан, confidence: {score:.3f}")1053 else:1054 logger.warning(f"✗ Промпт #{prompt.id} не дал результата")1055 else:1056 logger.warning(f"✗ Промпт #{prompt.id}: нет данных для сегментации")1057 1058 except Exception as e:1059 logger.error(f"✗ Ошибка обработки промпта #{prompt.id}: {e}")1060 # Продолжаем обработку остальных промптов1061 continue1062 1063 response = {1064 "success": True,1065 "image_size": {1066 "width": int(image.shape[1]),1067 "height": int(image.shape[0])1068 },1069 "results": results1070 }1071 1072 logger.info(f"Батчинг завершен: обработано {len(results)} объектов")1073 1074 # Сохраняем лог запроса для аудита (только метаданные, без изображений)1075 try:1076 request_dict = request.dict()1077 save_batch_request_log(request_dict, response, image.shape[1], image.shape[0])1078 except Exception as e:1079 logger.warning(f"Не удалось сохранить лог запроса: {e}")1080 1081 return convert_to_native_types(response)1082 1083 except Exception as e:1084 logger.error(f"Ошибка при батчинг сегментации: {e}")1085 raise HTTPException(status_code=500, detail=f"Ошибка обработки: {str(e)}")1086 1087if __name__ == "__main__":1088 import uvicorn1089 import os1090 1091 # Порт из переменной окружения (для HF Spaces) или 8000 по умолчанию1092 port = int(os.getenv("PORT", 8000))1093 uvicorn.run(app, host="0.0.0.0", port=port)1094 