File size: 1,492 Bytes
94fa27c
 
 
 
 
 
 
cd4509b
94fa27c
cd4509b
 
 
 
94fa27c
 
cd4509b
 
94fa27c
 
 
 
 
 
cd4509b
94fa27c
 
 
 
 
 
 
 
 
 
cd4509b
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
import base64, io
from typing import Any, Dict
from PIL import Image
import torch
from transformers import AutoImageProcessor, AutoModel

class EndpointHandler:
    """DINOv2 image -> 768-dim instance embedding (CLS token).

    The repo ships the full DINOv2 model, so `path` (the deployed repo dir) is a
    complete model and loads locally with no runtime download. Custom handler
    because the endpoint 'feature-extraction' task loads a TEXT pipeline and
    breaks on this vision model.
    """
    def __init__(self, path: str = ""):
        self.processor = AutoImageProcessor.from_pretrained(path)
        self.model = AutoModel.from_pretrained(path).eval()

    def _image(self, inp: Any) -> Image.Image:
        if isinstance(inp, Image.Image):
            return inp.convert("RGB")
        if isinstance(inp, (bytes, bytearray)):
            return Image.open(io.BytesIO(bytes(inp))).convert("RGB")
        if isinstance(inp, str):
            return Image.open(io.BytesIO(base64.b64decode(inp))).convert("RGB")
        raise ValueError(f"unsupported input type: {type(inp)}")

    def __call__(self, data: Dict[str, Any]):
        image = self._image(data.get("inputs"))
        inputs = self.processor(images=image, return_tensors="pt")
        with torch.no_grad():
            out = self.model(**inputs)
        pooled = getattr(out, "pooler_output", None)
        emb = pooled[0] if pooled is not None else out.last_hidden_state[:, 0].squeeze(0)
        return emb.tolist()