Repository navigation
Expand file tree
/
Copy pathencode.py
More file actions
95 lines (76 loc) · 2.7 KB
/
Copy pathencode.py
File metadata and controls
95 lines (76 loc) · 2.7 KB
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
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
from __future__ import annotations
import os
from typing import Callable, Optional, Protocol
DEFAULT_MODEL = "naver/splade-cocondenser-ensembledistil"
SparseEncoder = Callable[[str], dict[str, float]]
_encode: Optional[SparseEncoder] = None
_model_name = os.getenv("SEARCH_SPLADE_MODEL", DEFAULT_MODEL)
_load_error: Optional[str] = None
class SpladeEncoder(Protocol):
def encode(self, text: str) -> dict[str, float]:
...
def set_encoder(fn: Optional[SparseEncoder]) -> None:
global _encode, _load_error
_encode = fn
_load_error = None
def model_name() -> str:
return _model_name
def status() -> dict:
return {
"plugin": "search-splade",
"channel": "plugin:splade",
"model": _model_name,
"loaded": _encode is not None,
"error": _load_error,
}
def encode_text(text: str) -> dict[str, float]:
encoder = _ensure_encoder()
return encoder(text or "")
def _ensure_encoder() -> SparseEncoder:
global _encode, _load_error
if _encode is not None:
return _encode
try:
import torch
from transformers import AutoModelForMaskedLM, AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(_model_name)
model = AutoModelForMaskedLM.from_pretrained(_model_name)
model.eval()
skip = {
tokenizer.cls_token,
tokenizer.sep_token,
tokenizer.pad_token,
tokenizer.unk_token,
}
def _run(text: str) -> dict[str, float]:
tokens = tokenizer(
text,
return_tensors="pt",
truncation=True,
max_length=256,
)
with torch.no_grad():
logits = model(**tokens).logits
weights, _ = torch.max(torch.log1p(torch.relu(logits)), dim=1)
vector = weights.squeeze(0)
sparse: dict[str, float] = {}
nz = torch.nonzero(vector > 0, as_tuple=False).squeeze(-1)
if nz.ndim == 0:
nz = nz.unsqueeze(0)
ids = nz.tolist()
values = vector[nz].tolist() if ids else []
token_ids = [int(i) for i in (ids if isinstance(ids, list) else [ids])]
for token_id, value in zip(token_ids, values):
token = tokenizer.convert_ids_to_tokens(token_id)
if not token or token in skip:
continue
sparse[token] = float(value)
return sparse
_encode = _run
_load_error = None
return _encode
except Exception as exc:
_load_error = str(exc)
raise RuntimeError(
f"Failed to load SPLADE model {_model_name!r}: {exc}"
) from exc