File size: 1,728 Bytes
1dfe66b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
45
"""Custom HF Inference Endpoint handler for pplx-embed-context-v1-4b.

TEI cannot serve this custom `bidirectional_pplx_qwen3` architecture, so we run
it under the Default (transformers) container via this handler.

Accepts {"inputs": ...} where inputs is either:
  - a string                       -> {"embeddings": [vec]}
  - a list of strings              -> one vector per string (each = 1-chunk doc)
  - a list of list-of-strings      -> contextual: one vector per chunk per doc
Returns {"embeddings": [...]} (lists of floats), cosine-comparable.
"""
from typing import Any, Dict, List

import torch
from transformers import AutoModel


class EndpointHandler:
    def __init__(self, path: str = ""):
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        self.model = AutoModel.from_pretrained(
            path, trust_remote_code=True, torch_dtype=torch.float16
        ).to(self.device)
        self.model.eval()

    @staticmethod
    def _tolist(x):
        return x.tolist() if hasattr(x, "tolist") else x

    def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
        inputs = data.get("inputs", data)
        if isinstance(inputs, str):
            inputs = [inputs]
        contextual = bool(inputs) and isinstance(inputs[0], list)
        docs: List[List[str]] = inputs if contextual else [[t] for t in inputs]

        with torch.no_grad():
            embs = self.model.encode(docs)  # list, one (n_chunks, dim) array per doc

        if contextual:
            out = [self._tolist(e) for e in embs]            # per-chunk vectors per doc
        else:
            out = [self._tolist(e[0]) for e in embs]         # single vector per input text
        return {"embeddings": out}