-
Notifications
You must be signed in to change notification settings - Fork 35
Expand file tree
/
Copy pathapp.py
More file actions
226 lines (192 loc) · 7.28 KB
/
Copy pathapp.py
File metadata and controls
226 lines (192 loc) · 7.28 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
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
import os
from typing import Optional, List
from logging import getLogger
from fastapi import FastAPI, Depends, Response, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from typing import Union
from config import (
TRUST_REMOTE_CODE,
USE_QUERY_PASSAGE_PREFIXES,
USE_QUERY_PROMPT,
get_allowed_tokens,
get_use_sentence_transformers_multi_process,
get_t2v_transformers_direct_tokenize,
)
from vectorizer import Vectorizer, VectorInput
from meta import Meta
import torch
logger = getLogger("uvicorn")
vec: Vectorizer
meta_config: Meta
get_bearer_token = HTTPBearer(auto_error=False)
allowed_tokens: List[str] = None
current_worker = 0
available_workers = 1
def is_authorized(auth: Optional[HTTPAuthorizationCredentials]) -> bool:
if allowed_tokens is not None and (
auth is None or auth.credentials not in allowed_tokens
):
return False
return True
def get_worker():
if available_workers == 1:
return 0
else:
global current_worker
if current_worker >= 1_000_000_000:
current_worker = 0
worker = current_worker % available_workers
current_worker += 1
return worker
async def lifespan(app: FastAPI):
global vec
global meta_config
global allowed_tokens
global available_workers
allowed_tokens = get_allowed_tokens()
model_dir = "./models/model"
def get_model_name() -> Union[str, bool]:
if os.path.exists(f"{model_dir}/model_name"):
with open(f"{model_dir}/model_name", "r") as f:
model_name = f.read()
return model_name, True
# Default model directory is ./models/model
return model_dir, False
def get_onnx_runtime() -> bool:
if os.path.exists(f"{model_dir}/onnx_runtime"):
with open(f"{model_dir}/onnx_runtime", "r") as f:
onnx_runtime = f.read()
return onnx_runtime == "true"
return False
def get_trust_remote_code() -> bool:
if os.path.exists(f"{model_dir}/trust_remote_code"):
with open(f"{model_dir}/trust_remote_code", "r") as f:
trust_remote_code = f.read()
return trust_remote_code == "true"
return TRUST_REMOTE_CODE
def get_use_query_passage_prefixes() -> bool:
if os.path.exists(f"{model_dir}/use_query_passage_prefixes"):
with open(f"{model_dir}/use_query_passage_prefixes", "r") as f:
use_query_passage_prefixes = f.read()
return use_query_passage_prefixes == "true"
return USE_QUERY_PASSAGE_PREFIXES
def get_use_query_prompt() -> bool:
if os.path.exists(f"{model_dir}/use_query_prompt"):
with open(f"{model_dir}/use_query_prompt", "r") as f:
use_query_prompt = f.read()
return use_query_prompt == "true"
return USE_QUERY_PROMPT
def log_info_about_onnx(onnx_runtime: bool):
if onnx_runtime:
onnx_quantization_info = "missing"
if os.path.exists(f"{model_dir}/onnx_quantization_info"):
with open(f"{model_dir}/onnx_quantization_info", "r") as f:
onnx_quantization_info = f.read()
logger.info(
f"Running ONNX vectorizer with quantized model for {onnx_quantization_info}"
)
model_name, use_sentence_transformers_vectorizer = get_model_name()
onnx_runtime = get_onnx_runtime()
trust_remote_code = get_trust_remote_code()
use_query_passage_prefixes = get_use_query_passage_prefixes()
use_query_prompt = get_use_query_prompt()
cuda_env = os.getenv("ENABLE_CUDA")
cuda_per_process_memory_fraction = 1.0
if "CUDA_PER_PROCESS_MEMORY_FRACTION" in os.environ:
try:
cuda_per_process_memory_fraction = float(
os.getenv("CUDA_PER_PROCESS_MEMORY_FRACTION")
)
except ValueError:
logger.error(
f"Invalid CUDA_PER_PROCESS_MEMORY_FRACTION (should be between 0.0-1.0)"
)
if 0.0 <= cuda_per_process_memory_fraction <= 1.0:
logger.info(
f"CUDA_PER_PROCESS_MEMORY_FRACTION set to {cuda_per_process_memory_fraction}"
)
cuda_support = False
cuda_core = ""
# Use all sentence transformers multi process
use_sentence_transformers_multi_process = (
get_use_sentence_transformers_multi_process()
)
if cuda_env is not None and cuda_env == "true" or cuda_env == "1":
cuda_support = True
cuda_core = os.getenv("CUDA_CORE")
if cuda_core is None or cuda_core == "":
if (
use_sentence_transformers_vectorizer
and use_sentence_transformers_multi_process
and torch.cuda.is_available()
):
available_workers = torch.cuda.device_count()
cuda_core = ",".join([f"cuda:{i}" for i in range(available_workers)])
else:
cuda_core = "cuda:0"
logger.info(f"CUDA_CORE set to {cuda_core}")
else:
logger.info("Running on CPU")
# Batch text tokenization enabled by default
direct_tokenize = get_t2v_transformers_direct_tokenize()
log_info_about_onnx(onnx_runtime)
meta_config = Meta(
model_dir,
model_name,
use_sentence_transformers_vectorizer,
trust_remote_code,
)
if cuda_support is False and meta_config.get_model_type() == "model2vec":
# in case of CPU we need to run this model explicitly on CPU device, not MPS device
cuda_core = "cpu"
vec = Vectorizer(
model_dir,
cuda_support,
cuda_core,
cuda_per_process_memory_fraction,
meta_config.get_model_type(),
meta_config.get_architecture(),
direct_tokenize,
onnx_runtime,
use_sentence_transformers_vectorizer,
use_sentence_transformers_multi_process,
use_query_passage_prefixes,
use_query_prompt,
model_name,
trust_remote_code,
available_workers,
)
yield
app = FastAPI(lifespan=lifespan)
@app.get("/.well-known/live", response_class=Response)
@app.get("/.well-known/ready", response_class=Response)
async def live_and_ready(response: Response):
response.status_code = status.HTTP_204_NO_CONTENT
@app.get("/meta")
def meta(
response: Response,
auth: Optional[HTTPAuthorizationCredentials] = Depends(get_bearer_token),
):
if is_authorized(auth):
return meta_config.get()
else:
response.status_code = status.HTTP_401_UNAUTHORIZED
return {"error": "Unauthorized"}
@app.post("/vectors")
@app.post("/vectors/")
async def vectorize(
item: VectorInput,
response: Response,
auth: Optional[HTTPAuthorizationCredentials] = Depends(get_bearer_token),
):
if is_authorized(auth):
try:
vector = await vec.vectorize(item.text, item.config, get_worker())
return {"text": item.text, "vector": vector.tolist(), "dim": len(vector)}
except Exception as e:
logger.exception("Something went wrong while vectorizing data.")
response.status_code = status.HTTP_500_INTERNAL_SERVER_ERROR
return {"error": str(e)}
else:
response.status_code = status.HTTP_401_UNAUTHORIZED
return {"error": "Unauthorized"}