poolside-laguna-hackathon/laguna-vision
2
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 