CoolFace
Apppublic

gbreadman13/sam2-api

sourceHugging Faceapache-2.0updated 10mo agoView on Hugging Face
0likes
app.py1094 linesDownload Raw Back to root
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