File size: 1,658 Bytes
cb9f77d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
37
38
39
40
41
42
43
44
from typing import Any, Dict, List
from transformers import AutoImageProcessor, AutoModelForImageClassification
import torch
import requests
from PIL import Image
from io import BytesIO

class EndpointHandler:
    def __init__(self, path=""):
        self.processor = AutoImageProcessor.from_pretrained(path)
        self.model = AutoModelForImageClassification.from_pretrained(path)
        self.model.eval()

    def __call__(self, data: Dict[str, Any]) -> List[Dict]:
        inputs_data = data.pop("inputs", data)

        # Accept a URL string, a list of URLs, or raw PIL images
        if isinstance(inputs_data, str):
            inputs_data = [inputs_data]
        if not isinstance(inputs_data, list):
            inputs_data = [inputs_data]

        images = []
        for item in inputs_data:
            if isinstance(item, str):
                resp = requests.get(item, timeout=10)
                resp.raise_for_status()
                images.append(Image.open(BytesIO(resp.content)).convert("RGB"))
            else:
                images.append(item.convert("RGB") if hasattr(item, "convert") else item)

        encoded = self.processor(images=images, return_tensors="pt")
        with torch.no_grad():
            logits = self.model(**encoded).logits
        scores = torch.softmax(logits, dim=-1)
        id2label = self.model.config.id2label
        results = []
        for row in scores:
            results.append(sorted(
                [{"label": id2label[i], "score": float(row[i])} for i in range(len(row))],
                key=lambda x: -x["score"],
            ))
        return results if len(results) > 1 else results[0]