CoolFace
Modelpublic

poolside-laguna-hackathon/laguna-vision

sourceHugging Faceotherupdated 4mo agoView on Hugging Face
2likes
handler.py160 linesDownload Raw Back to root
1from __future__ import annotations2 3import asyncio4import base645import io6import json7import os8import tempfile9from pathlib import Path10from typing import Any11from urllib.parse import urlparse12 13import requests14from PIL import Image15 16from lagunavision.visual_pipeline import LagunaVisionImagePipeline, VisualProjectorSpec17 18 19class EndpointHandler:20    """Hugging Face Inference Endpoint handler for Laguna Vision checkpoints."""21 22    def __init__(self, path: str = "") -> None:23        self.model_dir = Path(path or ".").resolve()24        self.checkpoint_dir = self._resolve_checkpoint_dir()25        self.pipeline: LagunaVisionImagePipeline | None = None26 27    def __call__(self, data: dict[str, Any]) -> dict[str, Any]:28        payload = data.get("inputs", data)29        if not isinstance(payload, dict):30            raise ValueError("inputs must be an object with image and question fields")31 32        payload = _normalize_payload(payload)33        question = str(payload.get("question") or "").strip()34        if not question:35            raise ValueError("question is required")36 37        image_value = payload.get("image")38        if image_value is None:39            raise ValueError(40                "image is required as base64, a data URI, an HTTPS URL, a local path, "41                "or OpenAI-style messages content with image_url"42            )43 44        context = str(payload.get("context") or "")45        max_new_tokens = int(payload.get("max_new_tokens") or os.environ.get("LAGUNA_MAX_NEW_TOKENS", "128"))46 47        image = _load_image(image_value)48        with tempfile.NamedTemporaryFile(suffix=".png") as tmp:49            image.save(tmp.name)50            answer = asyncio.run(self._answer(Path(tmp.name), question, context, max_new_tokens))51        return {"answer": answer, "checkpoint": str(self.checkpoint_dir.relative_to(self.model_dir))}52 53    async def _answer(self, image: Path, question: str, context: str, max_new_tokens: int) -> str:54        if self.pipeline is None:55            self.pipeline = await self._load_pipeline()56        return await self.pipeline.answer_image(57            image=image,58            question=question,59            context=context,60            max_new_tokens=max_new_tokens,61        )62 63    async def _load_pipeline(self) -> LagunaVisionImagePipeline:64        spec_path = self.checkpoint_dir / "projector_spec.json"65        if not spec_path.exists():66            raise FileNotFoundError(f"missing projector_spec.json in {self.checkpoint_dir}")67        spec_row = json.loads(spec_path.read_text(encoding="utf-8"))68        spec = VisualProjectorSpec(69            input_dim=int(spec_row["input_dim"]),70            embedding_dim=int(spec_row["embedding_dim"]),71            hidden_dim=int(spec_row["hidden_dim"]),72            projector=spec_row.get("projector", "mlp"),73            visual_tokens=int(spec_row.get("visual_tokens", 64)),74            encoder=spec_row.get("encoder", "hf"),75            encoder_id=spec_row.get("encoder_id", ""),76            max_tiles=int(spec_row.get("max_tiles", 4)),77            patch_px=int(spec_row.get("patch_px", 32)),78        )79        lora_dir = self.checkpoint_dir / "lora"80        return await LagunaVisionImagePipeline.from_checkpoint(81            checkpoint=self.checkpoint_dir / "projector.pt",82            spec=spec,83            backbone_name=os.environ.get("LAGUNA_BACKBONE") or spec_row.get("backbone", "laguna"),84            model_id=os.environ.get("LAGUNA_MODEL_ID") or spec_row["model_id"],85            device=os.environ.get("LAGUNA_DEVICE", "auto"),86            vision_device=os.environ.get("LAGUNA_VISION_DEVICE", "auto"),87            lora_dir=lora_dir if lora_dir.exists() else None,88        )89 90    def _resolve_checkpoint_dir(self) -> Path:91        requested = os.environ.get("LAGUNA_CHECKPOINT_PATH", "latest").strip("/")92        candidates = [self.model_dir / requested, self.model_dir]93        for candidate in candidates:94            if (candidate / "projector.pt").exists() and (candidate / "projector_spec.json").exists():95                return candidate96        matches = sorted(self.model_dir.glob("**/projector.pt"), key=lambda path: path.stat().st_mtime, reverse=True)97        if matches:98            return matches[0].parent99        raise FileNotFoundError(100            f"no Laguna Vision checkpoint found under {self.model_dir}; expected latest/projector.pt"101        )102 103 104def _load_image(value: Any) -> Image.Image:105    if isinstance(value, bytes):106        return Image.open(io.BytesIO(value)).convert("RGB")107    if not isinstance(value, str):108        raise ValueError("image must be bytes or a string")109 110    if value.startswith("data:image/"):111        _, encoded = value.split(",", 1)112        return Image.open(io.BytesIO(base64.b64decode(encoded))).convert("RGB")113 114    parsed = urlparse(value)115    if parsed.scheme in {"http", "https"}:116        response = requests.get(value, timeout=20)117        response.raise_for_status()118        return Image.open(io.BytesIO(response.content)).convert("RGB")119 120    path = Path(value)121    if path.exists():122        return Image.open(path).convert("RGB")123 124    return Image.open(io.BytesIO(base64.b64decode(value))).convert("RGB")125 126 127def _normalize_payload(payload: dict[str, Any]) -> dict[str, Any]:128    if payload.get("messages") is None:129        return payload130 131    question_parts: list[str] = []132    image_value: Any = payload.get("image")133    for message in payload.get("messages") or []:134        if not isinstance(message, dict):135            continue136        content = message.get("content")137        if isinstance(content, str):138            question_parts.append(content)139            continue140        if not isinstance(content, list):141            continue142        for item in content:143            if not isinstance(item, dict):144                continue145            item_type = item.get("type")146            if item_type in {"text", "input_text"} and item.get("text"):147                question_parts.append(str(item["text"]))148            elif item_type in {"image_url", "input_image"}:149                image_url = item.get("image_url")150                if isinstance(image_url, dict):151                    image_value = image_url.get("url") or image_url.get("image_url") or image_value152                else:153                    image_value = image_url or item.get("url") or image_value154 155    normalized = dict(payload)156    normalized.setdefault("question", "\n".join(part.strip() for part in question_parts if part.strip()))157    if image_value is not None:158        normalized["image"] = image_value159    return normalized160