pytensorforge 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- cli.py +604 -0
- pytensorforge-0.1.0.dist-info/METADATA +103 -0
- pytensorforge-0.1.0.dist-info/RECORD +146 -0
- pytensorforge-0.1.0.dist-info/WHEEL +5 -0
- pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
- pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
- src/__init__.py +0 -0
- src/activations/Activation.py +4 -0
- src/activations/ELU.py +11 -0
- src/activations/GELU.py +6 -0
- src/activations/ReLU.py +27 -0
- src/activations/SELU.py +14 -0
- src/activations/Sigmoid.py +27 -0
- src/activations/Softmax.py +84 -0
- src/activations/Tanh.py +29 -0
- src/activations/__init__.py +17 -0
- src/config.py +120 -0
- src/core/Matrix.py +3 -0
- src/core/Scalar.py +18 -0
- src/core/Tensor.py +866 -0
- src/core/Vector.py +31 -0
- src/core/__init__.py +0 -0
- src/data/__init__.py +0 -0
- src/data/chat_dataset.py +188 -0
- src/data/corpus.py +104 -0
- src/data/document_stream.py +178 -0
- src/data/parallel_encode.py +86 -0
- src/data/prefetch.py +62 -0
- src/data/shard_builder.py +119 -0
- src/data/shard_writer.py +81 -0
- src/data/sharded_dataset.py +112 -0
- src/data/streaming_dataset.py +132 -0
- src/data/validation.py +212 -0
- src/inference/__init__.py +0 -0
- src/inference/chat_template.py +384 -0
- src/inference/config.py +48 -0
- src/inference/engine.py +241 -0
- src/inference/export.py +133 -0
- src/inference/kv_cache.py +65 -0
- src/inference/runtime.py +161 -0
- src/inference/sampling.py +42 -0
- src/inference/scheduler.py +473 -0
- src/inference/text.py +67 -0
- src/initializers/Constant.py +9 -0
- src/initializers/GlorotNormal.py +15 -0
- src/initializers/GlorotUniform.py +26 -0
- src/initializers/HeNormal.py +15 -0
- src/initializers/HeUniform.py +14 -0
- src/initializers/Initializer.py +4 -0
- src/initializers/LecunNormal.py +16 -0
- src/initializers/LecunUniform.py +14 -0
- src/initializers/Ones.py +6 -0
- src/initializers/Orthogonal.py +14 -0
- src/initializers/RandomNormal.py +14 -0
- src/initializers/RandomUniform.py +14 -0
- src/initializers/Zeros.py +8 -0
- src/initializers/__init__.py +17 -0
- src/loss/CategoricalCrossEntropy.py +9 -0
- src/loss/CrossEntropyLoss.py +34 -0
- src/loss/CrossEntropyWithLogitsLoss.py +59 -0
- src/loss/Hinge.py +5 -0
- src/loss/Huber.py +22 -0
- src/loss/Loss.py +6 -0
- src/loss/MSE.py +7 -0
- src/loss/MSELoss.py +10 -0
- src/loss/SparseCategoricalCrossEntropy.py +15 -0
- src/loss/__init__.py +18 -0
- src/loss/bce.py +34 -0
- src/loss/mae.py +16 -0
- src/math/__init__.py +0 -0
- src/math/clip.py +37 -0
- src/math/exp.py +27 -0
- src/math/log.py +25 -0
- src/math/sigmoid.py +5 -0
- src/models/__init__.py +0 -0
- src/models/embedding/Embedding.py +65 -0
- src/models/embedding/__init__.py +0 -0
- src/models/gpt/__init__.py +0 -0
- src/models/gpt/attention.py +158 -0
- src/models/gpt/block.py +74 -0
- src/models/gpt/config.py +103 -0
- src/models/gpt/context.py +44 -0
- src/models/gpt/model.py +165 -0
- src/models/gpt/recompute.py +35 -0
- src/models/gpt/rope.py +84 -0
- src/models/regression/Linear.py +51 -0
- src/models/regression/Logistic.py +36 -0
- src/models/regression/__init__.py +0 -0
- src/models/seq/Sequential.py +297 -0
- src/models/seq/__init__.py +0 -0
- src/models/svm/__init__.py +0 -0
- src/models/tokenizer/BPETokenizer.py +228 -0
- src/models/tokenizer/__init__.py +0 -0
- src/models/transformers/Dropout.py +35 -0
- src/models/transformers/LastToken.py +10 -0
- src/models/transformers/LayerNorm.py +54 -0
- src/models/transformers/Linear.py +18 -0
- src/models/transformers/MultiHeadAttention.py +130 -0
- src/models/transformers/TransformerBlock.py +79 -0
- src/models/transformers/__init__.py +0 -0
- src/neural/Dense.py +58 -0
- src/neural/LSTM.py +167 -0
- src/neural/Layer.py +72 -0
- src/neural/Parameter.py +30 -0
- src/neural/RNN.py +83 -0
- src/neural/__init__.py +0 -0
- src/ops/__init__.py +0 -0
- src/ops/stack.py +40 -0
- src/optimizers/Adagrad.py +31 -0
- src/optimizers/Adam.py +98 -0
- src/optimizers/AdamW.py +84 -0
- src/optimizers/Batch.py +11 -0
- src/optimizers/Nesterov.py +35 -0
- src/optimizers/Optimizer.py +18 -0
- src/optimizers/RMSProp.py +35 -0
- src/optimizers/SGD.py +30 -0
- src/optimizers/SGDMomentum.py +28 -0
- src/optimizers/__init__.py +9 -0
- src/scaling/StandardScaler.py +15 -0
- src/scaling/__init__.py +0 -0
- src/serialization/__init__.py +0 -0
- src/serialization/checkpoint.py +58 -0
- src/serialization/modelio.py +132 -0
- src/serving/__init__.py +0 -0
- src/serving/app.py +792 -0
- src/serving/config.py +216 -0
- src/serving/errors.py +51 -0
- src/serving/http.py +599 -0
- src/serving/metrics.py +293 -0
- src/serving/model_server.py +287 -0
- src/serving/protocol.py +377 -0
- src/serving/security.py +200 -0
- src/serving/server.py +121 -0
- src/tokenization/__init__.py +0 -0
- src/tokenization/base.py +75 -0
- src/tokenization/bpe.py +190 -0
- src/tokenization/bytebpe.py +476 -0
- src/tokenization/registry.py +28 -0
- src/training/__init__.py +0 -0
- src/training/checkpoint_manager.py +101 -0
- src/training/experiment.py +71 -0
- src/training/losses.py +42 -0
- src/training/precision.py +141 -0
- src/training/profiler.py +38 -0
- src/training/scheduler.py +50 -0
- src/training/trainer.py +594 -0
src/serving/app.py
ADDED
|
@@ -0,0 +1,792 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import json
|
|
3
|
+
import logging
|
|
4
|
+
import mimetypes
|
|
5
|
+
import os
|
|
6
|
+
import re
|
|
7
|
+
import sys
|
|
8
|
+
import time
|
|
9
|
+
import weakref
|
|
10
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
11
|
+
|
|
12
|
+
from src.inference.chat_template import ChatTemplateError, PromptTooLong
|
|
13
|
+
from src.inference.scheduler import GenerationRequest
|
|
14
|
+
from src.serving import errors, protocol
|
|
15
|
+
from src.serving.errors import APIError
|
|
16
|
+
from src.serving.http import ClientDisconnected, Response, StreamingResponse
|
|
17
|
+
from src.serving.metrics import ServerMetrics
|
|
18
|
+
from src.serving.security import KeyStore, RateLimiter
|
|
19
|
+
|
|
20
|
+
log = logging.getLogger("ptf.api")
|
|
21
|
+
|
|
22
|
+
MODEL_PATH_RE = re.compile(r"^/v1/models/([^/]+)$")
|
|
23
|
+
ADMIN_MODEL_RE = re.compile(r"^/admin/models/([^/]+)/(load|unload)$")
|
|
24
|
+
UI_TYPES = {".html", ".js", ".css", ".svg", ".png", ".ico", ".json", ".map", ".txt", ".woff2"}
|
|
25
|
+
SECURITY_HEADERS = {"X-Content-Type-Options": "nosniff", "Referrer-Policy": "no-referrer"}
|
|
26
|
+
UI_CSP = ("default-src 'self'; script-src 'self'; style-src 'self'; img-src 'self' data:; "
|
|
27
|
+
"connect-src 'self'; frame-ancestors 'none'; base-uri 'none'; form-action 'none'")
|
|
28
|
+
FINISH_TO_OPENAI = {"stop": "stop", "length": "length"}
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _json_dumps(obj):
|
|
32
|
+
return json.dumps(obj, separators=(",", ":"), ensure_ascii=False)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def sse(obj):
|
|
36
|
+
return f"data: {_json_dumps(obj)}\n\n"
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class _Bridge:
|
|
40
|
+
|
|
41
|
+
def __init__(self, handle, loop):
|
|
42
|
+
self.handle = handle
|
|
43
|
+
self._signal = asyncio.Event()
|
|
44
|
+
self._loop = loop
|
|
45
|
+
handle.set_listener(self._notify)
|
|
46
|
+
|
|
47
|
+
def _notify(self):
|
|
48
|
+
try:
|
|
49
|
+
self._loop.call_soon_threadsafe(self._signal.set)
|
|
50
|
+
except RuntimeError:
|
|
51
|
+
pass
|
|
52
|
+
|
|
53
|
+
async def events(self, ctx, tick_s):
|
|
54
|
+
disconnect_wait = asyncio.ensure_future(ctx.disconnected.wait())
|
|
55
|
+
|
|
56
|
+
try:
|
|
57
|
+
while True:
|
|
58
|
+
for ev in self.handle.drain():
|
|
59
|
+
yield ev
|
|
60
|
+
|
|
61
|
+
if ev.finished:
|
|
62
|
+
return
|
|
63
|
+
|
|
64
|
+
if ctx.disconnected.is_set():
|
|
65
|
+
raise ClientDisconnected()
|
|
66
|
+
|
|
67
|
+
self._signal.clear()
|
|
68
|
+
pending = self.handle.drain()
|
|
69
|
+
|
|
70
|
+
if pending:
|
|
71
|
+
for ev in pending:
|
|
72
|
+
yield ev
|
|
73
|
+
|
|
74
|
+
if ev.finished:
|
|
75
|
+
return
|
|
76
|
+
continue
|
|
77
|
+
|
|
78
|
+
signal_wait = asyncio.ensure_future(self._signal.wait())
|
|
79
|
+
done, _ = await asyncio.wait({signal_wait, disconnect_wait}, timeout=tick_s,
|
|
80
|
+
return_when=asyncio.FIRST_COMPLETED)
|
|
81
|
+
|
|
82
|
+
if signal_wait not in done:
|
|
83
|
+
signal_wait.cancel()
|
|
84
|
+
|
|
85
|
+
if not done:
|
|
86
|
+
yield None
|
|
87
|
+
finally:
|
|
88
|
+
disconnect_wait.cancel()
|
|
89
|
+
self.handle.set_listener(None)
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
class _LStrip:
|
|
93
|
+
|
|
94
|
+
def __init__(self, enabled):
|
|
95
|
+
self.pending = enabled
|
|
96
|
+
|
|
97
|
+
def __call__(self, text):
|
|
98
|
+
if not self.pending or not text:
|
|
99
|
+
return text
|
|
100
|
+
|
|
101
|
+
stripped = text.lstrip()
|
|
102
|
+
|
|
103
|
+
if stripped:
|
|
104
|
+
self.pending = False
|
|
105
|
+
|
|
106
|
+
return stripped
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
class APIApp:
|
|
110
|
+
|
|
111
|
+
def __init__(self, model_server, config, ui_directory=None):
|
|
112
|
+
self.models = model_server
|
|
113
|
+
self.config = config
|
|
114
|
+
self.limits = config.limits
|
|
115
|
+
self.keys = KeyStore.from_config(config.security)
|
|
116
|
+
self.limiter = RateLimiter.from_config(config.security.rate_limit)
|
|
117
|
+
self.metrics = ServerMetrics(weakref.ref(model_server))
|
|
118
|
+
self.cors = set(config.security.cors_origins)
|
|
119
|
+
self._tok_pool = ThreadPoolExecutor(max_workers=4, thread_name_prefix="ptf-tokenize")
|
|
120
|
+
self._admin_pool = ThreadPoolExecutor(max_workers=1, thread_name_prefix="ptf-admin")
|
|
121
|
+
self._access = self._open_access_log(config.logging.access_log)
|
|
122
|
+
self._ui_files = self._index_ui(ui_directory) if config.ui.enabled else {}
|
|
123
|
+
|
|
124
|
+
@staticmethod
|
|
125
|
+
def _open_access_log(target):
|
|
126
|
+
if not target:
|
|
127
|
+
return None
|
|
128
|
+
|
|
129
|
+
if target == "-":
|
|
130
|
+
return sys.stderr
|
|
131
|
+
|
|
132
|
+
return open(target, "a", buffering=1)
|
|
133
|
+
|
|
134
|
+
@staticmethod
|
|
135
|
+
def _index_ui(directory):
|
|
136
|
+
if not directory or not os.path.isdir(directory):
|
|
137
|
+
return {}
|
|
138
|
+
|
|
139
|
+
files = {}
|
|
140
|
+
root = os.path.realpath(directory)
|
|
141
|
+
|
|
142
|
+
for dirpath, _, names in os.walk(root):
|
|
143
|
+
for name in names:
|
|
144
|
+
full = os.path.realpath(os.path.join(dirpath, name))
|
|
145
|
+
|
|
146
|
+
if not full.startswith(root + os.sep) or os.path.splitext(name)[1] not in UI_TYPES:
|
|
147
|
+
continue
|
|
148
|
+
|
|
149
|
+
rel = "/" + os.path.relpath(full, root).replace(os.sep, "/")
|
|
150
|
+
files[rel] = full
|
|
151
|
+
|
|
152
|
+
return files
|
|
153
|
+
|
|
154
|
+
def close(self):
|
|
155
|
+
self._tok_pool.shutdown(wait=False, cancel_futures=True)
|
|
156
|
+
self._admin_pool.shutdown(wait=False, cancel_futures=True)
|
|
157
|
+
|
|
158
|
+
if self._access not in (None, sys.stderr):
|
|
159
|
+
self._access.close()
|
|
160
|
+
|
|
161
|
+
def _client_ip(self, request):
|
|
162
|
+
if self.config.security.trust_forwarded_for:
|
|
163
|
+
fwd = request.headers.get("x-forwarded-for")
|
|
164
|
+
|
|
165
|
+
if fwd:
|
|
166
|
+
return fwd.split(",")[0].strip()
|
|
167
|
+
|
|
168
|
+
return request.client_ip
|
|
169
|
+
|
|
170
|
+
def _cors_headers(self, request):
|
|
171
|
+
origin = request.headers.get("origin")
|
|
172
|
+
|
|
173
|
+
if not origin or not self.cors:
|
|
174
|
+
return {}
|
|
175
|
+
|
|
176
|
+
if "*" in self.cors or origin in self.cors:
|
|
177
|
+
return {
|
|
178
|
+
"Access-Control-Allow-Origin": origin,
|
|
179
|
+
"Vary": "Origin",
|
|
180
|
+
"Access-Control-Expose-Headers": "X-Request-Id, Retry-After",
|
|
181
|
+
"Access-Control-Max-Age": "600",
|
|
182
|
+
}
|
|
183
|
+
|
|
184
|
+
return {}
|
|
185
|
+
|
|
186
|
+
def _error(self, request, exc):
|
|
187
|
+
headers = dict(exc.headers)
|
|
188
|
+
headers.update(self._cors_headers(request))
|
|
189
|
+
headers.update(SECURITY_HEADERS)
|
|
190
|
+
request.info.setdefault("error_code", exc.code)
|
|
191
|
+
|
|
192
|
+
return Response(exc.status, exc.to_dict(), headers)
|
|
193
|
+
|
|
194
|
+
def _authenticate(self, request):
|
|
195
|
+
try:
|
|
196
|
+
principal = self.keys.authenticate(request.headers, self._client_ip(request))
|
|
197
|
+
except APIError:
|
|
198
|
+
self.metrics.auth_failures.inc()
|
|
199
|
+
raise
|
|
200
|
+
|
|
201
|
+
request.info["principal"] = principal.id
|
|
202
|
+
|
|
203
|
+
return principal
|
|
204
|
+
|
|
205
|
+
def _route_name(self, request):
|
|
206
|
+
path = request.path
|
|
207
|
+
|
|
208
|
+
if path in ("/v1/chat/completions", "/v1/completions", "/v1/models", "/health", "/ready", "/metrics",
|
|
209
|
+
"/admin/models"):
|
|
210
|
+
return path
|
|
211
|
+
|
|
212
|
+
if MODEL_PATH_RE.match(path):
|
|
213
|
+
return "/v1/models/{id}"
|
|
214
|
+
|
|
215
|
+
if ADMIN_MODEL_RE.match(path):
|
|
216
|
+
return "/admin/models/{name}/{action}"
|
|
217
|
+
|
|
218
|
+
if path in self._ui_files or path == "/":
|
|
219
|
+
return "ui"
|
|
220
|
+
|
|
221
|
+
return "other"
|
|
222
|
+
|
|
223
|
+
async def __call__(self, request, ctx):
|
|
224
|
+
request.info["route"] = self._route_name(request)
|
|
225
|
+
|
|
226
|
+
try:
|
|
227
|
+
response = await self._handle(request, ctx)
|
|
228
|
+
except ClientDisconnected:
|
|
229
|
+
raise
|
|
230
|
+
except APIError as exc:
|
|
231
|
+
return self._error(request, exc)
|
|
232
|
+
except Exception:
|
|
233
|
+
log.exception("request %s failed", request.request_id)
|
|
234
|
+
self.metrics.errors.inc(kind="internal")
|
|
235
|
+
return self._error(request, errors.internal(request.request_id))
|
|
236
|
+
|
|
237
|
+
response.headers.update(self._cors_headers(request))
|
|
238
|
+
|
|
239
|
+
for k, v in SECURITY_HEADERS.items():
|
|
240
|
+
response.headers.setdefault(k, v)
|
|
241
|
+
|
|
242
|
+
return response
|
|
243
|
+
|
|
244
|
+
async def _handle(self, request, ctx):
|
|
245
|
+
method, path = request.method, request.path
|
|
246
|
+
|
|
247
|
+
if method == "OPTIONS":
|
|
248
|
+
return self._preflight(request)
|
|
249
|
+
|
|
250
|
+
if path == "/health" and method in ("GET", "HEAD"):
|
|
251
|
+
return Response(200, {"status": "ok"})
|
|
252
|
+
|
|
253
|
+
if path == "/ready" and method in ("GET", "HEAD"):
|
|
254
|
+
loaded = [m.name for m in self.models.loaded_models() if m.accepting]
|
|
255
|
+
return Response(200 if loaded else 503, {"status": "ready" if loaded else "not_ready", "models": loaded})
|
|
256
|
+
|
|
257
|
+
if path == "/metrics" and method == "GET":
|
|
258
|
+
if self.config.security.metrics_require_auth and self.keys.enabled:
|
|
259
|
+
self._authenticate(request)
|
|
260
|
+
|
|
261
|
+
fmt = request.query.get("format")
|
|
262
|
+
|
|
263
|
+
if fmt == "json":
|
|
264
|
+
return Response(200, self.metrics.registry.snapshot())
|
|
265
|
+
|
|
266
|
+
return Response(200, self.metrics.registry.render_prometheus(),
|
|
267
|
+
content_type="text/plain; version=0.0.4; charset=utf-8")
|
|
268
|
+
|
|
269
|
+
if path.startswith("/v1/"):
|
|
270
|
+
return await self._handle_v1(request, ctx)
|
|
271
|
+
|
|
272
|
+
if path.startswith("/admin/"):
|
|
273
|
+
return await self._handle_admin(request)
|
|
274
|
+
|
|
275
|
+
if method in ("GET", "HEAD") and self._ui_files:
|
|
276
|
+
return self._serve_ui(request)
|
|
277
|
+
|
|
278
|
+
raise errors.not_found(f"no route for {method} {path}", code="unknown_url")
|
|
279
|
+
|
|
280
|
+
def _preflight(self, request):
|
|
281
|
+
cors = self._cors_headers(request)
|
|
282
|
+
|
|
283
|
+
if not cors:
|
|
284
|
+
raise APIError(403, "origin not allowed", "permission_error", code="cors_forbidden")
|
|
285
|
+
|
|
286
|
+
cors.update({
|
|
287
|
+
"Access-Control-Allow-Methods": "GET, POST, OPTIONS",
|
|
288
|
+
"Access-Control-Allow-Headers": "authorization, content-type, x-api-key",
|
|
289
|
+
})
|
|
290
|
+
|
|
291
|
+
return Response(204, b"", cors, content_type=None)
|
|
292
|
+
|
|
293
|
+
def _serve_ui(self, request):
|
|
294
|
+
path = "/index.html" if request.path in ("/", "/ui", "/ui/") else request.path
|
|
295
|
+
full = self._ui_files.get(path)
|
|
296
|
+
|
|
297
|
+
if full is None:
|
|
298
|
+
raise errors.not_found("not found", code="unknown_url")
|
|
299
|
+
|
|
300
|
+
with open(full, "rb") as f:
|
|
301
|
+
body = f.read()
|
|
302
|
+
|
|
303
|
+
ctype = mimetypes.guess_type(full)[0] or "application/octet-stream"
|
|
304
|
+
|
|
305
|
+
if ctype.startswith("text/") or ctype in ("application/javascript", "application/json"):
|
|
306
|
+
ctype += "; charset=utf-8"
|
|
307
|
+
|
|
308
|
+
headers = {"Cache-Control": "no-cache" if path.endswith(".html") else "public, max-age=300"}
|
|
309
|
+
|
|
310
|
+
if path.endswith(".html"):
|
|
311
|
+
headers["Content-Security-Policy"] = UI_CSP
|
|
312
|
+
headers["X-Frame-Options"] = "DENY"
|
|
313
|
+
|
|
314
|
+
return Response(200, body, headers, content_type=ctype)
|
|
315
|
+
|
|
316
|
+
def _parse_json(self, request):
|
|
317
|
+
ctype = request.headers.get("content-type", "application/json").split(";")[0].strip().lower()
|
|
318
|
+
|
|
319
|
+
if ctype != "application/json":
|
|
320
|
+
raise APIError(415, "Content-Type must be application/json", "invalid_request_error",
|
|
321
|
+
code="unsupported_media_type")
|
|
322
|
+
|
|
323
|
+
try:
|
|
324
|
+
return request.json()
|
|
325
|
+
except (ValueError, UnicodeDecodeError):
|
|
326
|
+
raise errors.bad_request("request body is not valid JSON", code="invalid_json")
|
|
327
|
+
|
|
328
|
+
async def _handle_v1(self, request, ctx):
|
|
329
|
+
method, path = request.method, request.path
|
|
330
|
+
|
|
331
|
+
if path == "/v1/chat/completions":
|
|
332
|
+
if method != "POST":
|
|
333
|
+
raise APIError(405, "use POST", "invalid_request_error", code="method_not_allowed",
|
|
334
|
+
headers={"Allow": "POST"})
|
|
335
|
+
return await self._completion(request, ctx, chat=True)
|
|
336
|
+
|
|
337
|
+
if path == "/v1/completions":
|
|
338
|
+
if method != "POST":
|
|
339
|
+
raise APIError(405, "use POST", "invalid_request_error", code="method_not_allowed",
|
|
340
|
+
headers={"Allow": "POST"})
|
|
341
|
+
return await self._completion(request, ctx, chat=False)
|
|
342
|
+
|
|
343
|
+
if method != "GET":
|
|
344
|
+
raise APIError(405, "use GET", "invalid_request_error", code="method_not_allowed",
|
|
345
|
+
headers={"Allow": "GET"})
|
|
346
|
+
|
|
347
|
+
principal = self._authenticate(request)
|
|
348
|
+
self._charge(principal)
|
|
349
|
+
|
|
350
|
+
if path == "/v1/models":
|
|
351
|
+
data = [m.describe() for m in self.models.loaded_models() if m.accepting]
|
|
352
|
+
return Response(200, {"object": "list", "data": data})
|
|
353
|
+
|
|
354
|
+
match = MODEL_PATH_RE.match(path)
|
|
355
|
+
|
|
356
|
+
if match:
|
|
357
|
+
return Response(200, self.models.get(match.group(1)).describe())
|
|
358
|
+
|
|
359
|
+
raise errors.not_found(f"no route for {method} {path}", code="unknown_url")
|
|
360
|
+
|
|
361
|
+
def _charge(self, principal):
|
|
362
|
+
try:
|
|
363
|
+
self.limiter.charge(principal.id)
|
|
364
|
+
except APIError:
|
|
365
|
+
self.metrics.rate_limited.inc(reason="rate")
|
|
366
|
+
raise
|
|
367
|
+
|
|
368
|
+
async def _handle_admin(self, request):
|
|
369
|
+
principal = self._authenticate(request)
|
|
370
|
+
|
|
371
|
+
if not principal.admin:
|
|
372
|
+
raise errors.forbidden("admin API key required")
|
|
373
|
+
|
|
374
|
+
if request.path == "/admin/models" and request.method == "GET":
|
|
375
|
+
loaded = {m.name: m for m in self.models.loaded_models()}
|
|
376
|
+
return Response(200, {"object": "list", "data": [
|
|
377
|
+
{
|
|
378
|
+
"name": e.name,
|
|
379
|
+
"loaded": e.name in loaded,
|
|
380
|
+
"accepting": loaded[e.name].accepting if e.name in loaded else False,
|
|
381
|
+
"in_flight": loaded[e.name].in_flight if e.name in loaded else 0,
|
|
382
|
+
"memory_bytes": loaded[e.name].memory_bytes if e.name in loaded else None,
|
|
383
|
+
}
|
|
384
|
+
for e in self.models.declared()
|
|
385
|
+
]})
|
|
386
|
+
|
|
387
|
+
match = ADMIN_MODEL_RE.match(request.path)
|
|
388
|
+
|
|
389
|
+
if match and request.method == "POST":
|
|
390
|
+
name, action = match.groups()
|
|
391
|
+
loop = asyncio.get_running_loop()
|
|
392
|
+
|
|
393
|
+
if action == "load":
|
|
394
|
+
model = await loop.run_in_executor(self._admin_pool, self.models.load, name)
|
|
395
|
+
return Response(200, {"status": "loaded", "model": model.describe()})
|
|
396
|
+
|
|
397
|
+
drained = await loop.run_in_executor(
|
|
398
|
+
self._admin_pool, self.models.unload, name, self.limits.shutdown_drain_s
|
|
399
|
+
)
|
|
400
|
+
return Response(200, {"status": "unloaded", "model": name, "drained": drained})
|
|
401
|
+
|
|
402
|
+
raise errors.not_found(f"no route for {request.method} {request.path}", code="unknown_url")
|
|
403
|
+
|
|
404
|
+
def _generation_config(self, model, sampling, prompt_len, chat):
|
|
405
|
+
gen = model.generator.default_config
|
|
406
|
+
limits = model.limits
|
|
407
|
+
|
|
408
|
+
max_new = sampling.max_tokens or gen.max_new_tokens
|
|
409
|
+
max_new = min(max_new, limits.max_generation_tokens, limits.context_length - prompt_len)
|
|
410
|
+
|
|
411
|
+
if max_new < 1:
|
|
412
|
+
raise errors.bad_request(
|
|
413
|
+
f"prompt of {prompt_len} tokens leaves no room to generate within the context of "
|
|
414
|
+
f"{limits.context_length} tokens",
|
|
415
|
+
param="messages" if chat else "prompt", code="context_length_exceeded",
|
|
416
|
+
)
|
|
417
|
+
|
|
418
|
+
stop_strings = list(sampling.stop)
|
|
419
|
+
stop_ids = list(gen.stop_token_ids)
|
|
420
|
+
|
|
421
|
+
if chat:
|
|
422
|
+
stop_strings = list(model.template.stop_strings) + stop_strings
|
|
423
|
+
stop_ids = stop_ids + [i for i in model.template.stop_token_ids if i not in stop_ids]
|
|
424
|
+
|
|
425
|
+
overrides = {
|
|
426
|
+
"max_new_tokens": max_new,
|
|
427
|
+
"temperature": sampling.temperature,
|
|
428
|
+
"top_p": sampling.top_p,
|
|
429
|
+
"top_k": sampling.top_k,
|
|
430
|
+
"repetition_penalty": sampling.repetition_penalty,
|
|
431
|
+
"seed": sampling.seed,
|
|
432
|
+
"stop_strings": stop_strings,
|
|
433
|
+
"stop_token_ids": stop_ids,
|
|
434
|
+
}
|
|
435
|
+
|
|
436
|
+
if sampling.temperature is not None and sampling.temperature > 0:
|
|
437
|
+
overrides["do_sample"] = True
|
|
438
|
+
|
|
439
|
+
try:
|
|
440
|
+
return gen.updated(**overrides)
|
|
441
|
+
except ValueError as exc:
|
|
442
|
+
raise errors.bad_request(str(exc))
|
|
443
|
+
|
|
444
|
+
def _tokenize_chat(self, model, messages, max_new_hint):
|
|
445
|
+
limits = model.limits
|
|
446
|
+
reserve = min(max_new_hint, limits.context_length // 2)
|
|
447
|
+
budget = min(limits.max_prompt_tokens, limits.context_length - reserve)
|
|
448
|
+
truncate = self.config.runtime.chat_truncation == "auto"
|
|
449
|
+
|
|
450
|
+
started = time.perf_counter()
|
|
451
|
+
|
|
452
|
+
try:
|
|
453
|
+
rendered = model.template.render(messages, max_prompt_tokens=budget, truncate=truncate)
|
|
454
|
+
except PromptTooLong as exc:
|
|
455
|
+
raise errors.bad_request(
|
|
456
|
+
f"the conversation needs {exc.needed} prompt tokens but at most {exc.budget} fit "
|
|
457
|
+
f"(context {limits.context_length}, {reserve} reserved for the reply)",
|
|
458
|
+
param="messages", code="context_length_exceeded",
|
|
459
|
+
)
|
|
460
|
+
except ChatTemplateError as exc:
|
|
461
|
+
raise errors.bad_request(str(exc), param="messages")
|
|
462
|
+
|
|
463
|
+
return rendered.token_ids, rendered.messages_dropped, time.perf_counter() - started
|
|
464
|
+
|
|
465
|
+
def _tokenize_completion(self, model, prompt):
|
|
466
|
+
started = time.perf_counter()
|
|
467
|
+
vocab = model.generator.config.vocab_size
|
|
468
|
+
|
|
469
|
+
if isinstance(prompt, str):
|
|
470
|
+
ids = model.generator.tokenizer.encode(prompt)
|
|
471
|
+
else:
|
|
472
|
+
ids = list(prompt)
|
|
473
|
+
|
|
474
|
+
if any(t < 0 or t >= vocab for t in ids):
|
|
475
|
+
raise errors.bad_request("prompt contains token ids outside the model vocabulary", param="prompt")
|
|
476
|
+
|
|
477
|
+
if not ids:
|
|
478
|
+
eos = model.generator.tokenizer.eos_id
|
|
479
|
+
|
|
480
|
+
if eos is None:
|
|
481
|
+
raise errors.bad_request("prompt must not be empty", param="prompt")
|
|
482
|
+
|
|
483
|
+
ids = [eos]
|
|
484
|
+
|
|
485
|
+
if len(ids) > model.limits.max_prompt_tokens:
|
|
486
|
+
raise errors.bad_request(
|
|
487
|
+
f"prompt of {len(ids)} tokens exceeds the limit of {model.limits.max_prompt_tokens}",
|
|
488
|
+
param="prompt", code="context_length_exceeded",
|
|
489
|
+
)
|
|
490
|
+
|
|
491
|
+
return ids, 0, time.perf_counter() - started
|
|
492
|
+
|
|
493
|
+
async def _completion(self, request, ctx, chat):
|
|
494
|
+
principal = self._authenticate(request)
|
|
495
|
+
body = self._parse_json(request)
|
|
496
|
+
parsed = (protocol.parse_chat_request if chat else protocol.parse_completion_request)(body, self.limits)
|
|
497
|
+
model = self.models.get(parsed.model)
|
|
498
|
+
|
|
499
|
+
request.info["model"] = model.name
|
|
500
|
+
request.info["stream"] = parsed.stream
|
|
501
|
+
|
|
502
|
+
try:
|
|
503
|
+
self.limiter.acquire(principal.id)
|
|
504
|
+
except APIError:
|
|
505
|
+
self.metrics.rate_limited.inc(reason="principal")
|
|
506
|
+
raise
|
|
507
|
+
|
|
508
|
+
try:
|
|
509
|
+
self.models.acquire_global()
|
|
510
|
+
except APIError:
|
|
511
|
+
self.limiter.release(principal.id)
|
|
512
|
+
self.metrics.rate_limited.inc(reason="global")
|
|
513
|
+
raise
|
|
514
|
+
|
|
515
|
+
try:
|
|
516
|
+
model.enter()
|
|
517
|
+
except APIError:
|
|
518
|
+
self.models.release_global()
|
|
519
|
+
self.limiter.release(principal.id)
|
|
520
|
+
raise
|
|
521
|
+
|
|
522
|
+
released = False
|
|
523
|
+
|
|
524
|
+
def release():
|
|
525
|
+
nonlocal released
|
|
526
|
+
|
|
527
|
+
if not released:
|
|
528
|
+
released = True
|
|
529
|
+
model.exit()
|
|
530
|
+
self.models.release_global()
|
|
531
|
+
self.limiter.release(principal.id)
|
|
532
|
+
self.metrics.active.dec(model=model.name)
|
|
533
|
+
|
|
534
|
+
self.metrics.active.inc(model=model.name)
|
|
535
|
+
|
|
536
|
+
try:
|
|
537
|
+
loop = asyncio.get_running_loop()
|
|
538
|
+
hint = parsed.sampling.max_tokens or model.generator.default_config.max_new_tokens
|
|
539
|
+
|
|
540
|
+
if chat:
|
|
541
|
+
ids, dropped, tok_s = await loop.run_in_executor(
|
|
542
|
+
self._tok_pool, self._tokenize_chat, model, parsed.messages, hint
|
|
543
|
+
)
|
|
544
|
+
else:
|
|
545
|
+
ids, dropped, tok_s = await loop.run_in_executor(
|
|
546
|
+
self._tok_pool, self._tokenize_completion, model, parsed.prompt
|
|
547
|
+
)
|
|
548
|
+
|
|
549
|
+
self.metrics.tokenize_s.observe(tok_s, model=model.name)
|
|
550
|
+
self.metrics.prompt_len.observe(len(ids), model=model.name)
|
|
551
|
+
request.info["messages_dropped"] = dropped
|
|
552
|
+
|
|
553
|
+
gen_cfg = self._generation_config(model, parsed.sampling, len(ids), chat)
|
|
554
|
+
timeout = self.limits.request_timeout_s
|
|
555
|
+
|
|
556
|
+
gen_request = GenerationRequest(prompt=ids, config=gen_cfg, request_id=request.request_id,
|
|
557
|
+
timeout_s=timeout, truncate_prompt=False)
|
|
558
|
+
handle = model.generator.submit(gen_request)
|
|
559
|
+
self.metrics.requests.inc(model=model.name, endpoint="chat" if chat else "completions",
|
|
560
|
+
stream=str(parsed.stream).lower())
|
|
561
|
+
|
|
562
|
+
runner = _Run(self, model, handle, request, ctx, chat, parsed, len(ids), release)
|
|
563
|
+
|
|
564
|
+
if parsed.stream:
|
|
565
|
+
return StreamingResponse(runner.stream(), on_close=runner.close)
|
|
566
|
+
|
|
567
|
+
return await runner.collect()
|
|
568
|
+
except BaseException:
|
|
569
|
+
release()
|
|
570
|
+
raise
|
|
571
|
+
|
|
572
|
+
def on_request_done(self, request, status, duration, write_s, disconnected, streamed):
|
|
573
|
+
route = request.info.get("route", "other")
|
|
574
|
+
self.metrics.http_requests.inc(route=route, status=status)
|
|
575
|
+
self.metrics.http_duration.observe(duration, route=route)
|
|
576
|
+
|
|
577
|
+
if write_s:
|
|
578
|
+
self.metrics.network_s.observe(write_s, route=route)
|
|
579
|
+
|
|
580
|
+
if self._access is None:
|
|
581
|
+
return
|
|
582
|
+
|
|
583
|
+
record = {
|
|
584
|
+
"ts": round(time.time(), 3),
|
|
585
|
+
"request_id": request.request_id,
|
|
586
|
+
"method": request.method,
|
|
587
|
+
"route": route,
|
|
588
|
+
"status": status,
|
|
589
|
+
"duration_s": round(duration, 4),
|
|
590
|
+
"network_write_s": round(write_s, 4),
|
|
591
|
+
"client_disconnected": disconnected,
|
|
592
|
+
"streamed": streamed,
|
|
593
|
+
"client_ip": self._client_ip(request),
|
|
594
|
+
}
|
|
595
|
+
|
|
596
|
+
record.update({k: v for k, v in request.info.items() if k != "route"})
|
|
597
|
+
|
|
598
|
+
try:
|
|
599
|
+
self._access.write(_json_dumps(record) + "\n")
|
|
600
|
+
except (OSError, ValueError):
|
|
601
|
+
pass
|
|
602
|
+
|
|
603
|
+
|
|
604
|
+
class _Run:
|
|
605
|
+
|
|
606
|
+
def __init__(self, app, model, handle, request, ctx, chat, parsed, prompt_len, release):
|
|
607
|
+
self.app = app
|
|
608
|
+
self.model = model
|
|
609
|
+
self.handle = handle
|
|
610
|
+
self.request = request
|
|
611
|
+
self.ctx = ctx
|
|
612
|
+
self.chat = chat
|
|
613
|
+
self.parsed = parsed
|
|
614
|
+
self.prompt_len = prompt_len
|
|
615
|
+
self.release = release
|
|
616
|
+
self.lstrip = _LStrip(chat and model.template.output_lstrip)
|
|
617
|
+
self.rid = protocol.new_id("chatcmpl" if chat else "cmpl")
|
|
618
|
+
self.created = int(time.time())
|
|
619
|
+
self._abandoned = False
|
|
620
|
+
self._done = False
|
|
621
|
+
|
|
622
|
+
def close(self):
|
|
623
|
+
self._abandon("client_disconnect")
|
|
624
|
+
self.release()
|
|
625
|
+
|
|
626
|
+
def _observe(self, ev):
|
|
627
|
+
self._done = True
|
|
628
|
+
m = self.app.metrics
|
|
629
|
+
name = self.model.name
|
|
630
|
+
h = self.handle
|
|
631
|
+
now = time.perf_counter()
|
|
632
|
+
usage = ev.usage or {}
|
|
633
|
+
|
|
634
|
+
if h.admitted_at is not None:
|
|
635
|
+
m.queue_s.observe(h.admitted_at - h.submitted_at, model=name)
|
|
636
|
+
|
|
637
|
+
m.prefill_s.observe(h.prefill_s, model=name)
|
|
638
|
+
|
|
639
|
+
ttft = usage.get("time_to_first_token_s")
|
|
640
|
+
m.ttft_s.observe(ttft, model=name)
|
|
641
|
+
m.generation_s.observe(now - h.submitted_at, model=name)
|
|
642
|
+
|
|
643
|
+
completion = usage.get("completion_tokens", 0)
|
|
644
|
+
m.prompt_tokens.inc(usage.get("prompt_tokens", 0) or 0, model=name)
|
|
645
|
+
m.completion_tokens.inc(completion, model=name)
|
|
646
|
+
|
|
647
|
+
if ttft is not None and completion > 1:
|
|
648
|
+
span = (now - h.submitted_at) - ttft
|
|
649
|
+
|
|
650
|
+
if span > 0:
|
|
651
|
+
m.decode_rate.observe((completion - 1) / span, model=name)
|
|
652
|
+
|
|
653
|
+
m.finished.inc(model=name, reason=ev.finish_reason or "unknown")
|
|
654
|
+
|
|
655
|
+
if ev.finish_reason in ("cancelled", "timeout"):
|
|
656
|
+
m.cancellations.inc(model=name, cause=ev.finish_reason)
|
|
657
|
+
|
|
658
|
+
self.request.info.update({
|
|
659
|
+
"prompt_tokens": usage.get("prompt_tokens"),
|
|
660
|
+
"completion_tokens": completion,
|
|
661
|
+
"finish_reason": ev.finish_reason,
|
|
662
|
+
"ttft_s": round(ttft, 4) if ttft is not None else None,
|
|
663
|
+
"queue_s": round(h.admitted_at - h.submitted_at, 4) if h.admitted_at else None,
|
|
664
|
+
})
|
|
665
|
+
|
|
666
|
+
def _usage(self, ev):
|
|
667
|
+
u = ev.usage or {}
|
|
668
|
+
return protocol.usage_dict(u.get("prompt_tokens", self.prompt_len), u.get("completion_tokens", 0))
|
|
669
|
+
|
|
670
|
+
def _terminal_error(self, ev):
|
|
671
|
+
if ev.finish_reason == "timeout":
|
|
672
|
+
return APIError(504, "generation timed out", "server_error", code="timeout")
|
|
673
|
+
|
|
674
|
+
if ev.finish_reason == "cancelled":
|
|
675
|
+
return errors.overloaded("generation was cancelled because the model is unloading")
|
|
676
|
+
|
|
677
|
+
log.error("generation %s failed: %s", self.request.request_id, ev.error)
|
|
678
|
+
self.app.metrics.errors.inc(kind="generation")
|
|
679
|
+
return errors.internal(self.request.request_id)
|
|
680
|
+
|
|
681
|
+
def _abandon(self, cause):
|
|
682
|
+
if not self._abandoned and not self._done and not self.handle.finished.is_set():
|
|
683
|
+
self._abandoned = True
|
|
684
|
+
self.handle.request_cancel()
|
|
685
|
+
self.app.metrics.cancellations.inc(model=self.model.name, cause=cause)
|
|
686
|
+
self.app.metrics.finished.inc(model=self.model.name, reason="client_disconnect")
|
|
687
|
+
self.request.info["finish_reason"] = "client_disconnect"
|
|
688
|
+
|
|
689
|
+
async def collect(self):
|
|
690
|
+
bridge = _Bridge(self.handle, asyncio.get_running_loop())
|
|
691
|
+
parts = []
|
|
692
|
+
last = None
|
|
693
|
+
|
|
694
|
+
try:
|
|
695
|
+
async for ev in bridge.events(self.ctx, self.app.limits.sse_keepalive_s):
|
|
696
|
+
if ev is None:
|
|
697
|
+
continue
|
|
698
|
+
|
|
699
|
+
if ev.text:
|
|
700
|
+
parts.append(self.lstrip(ev.text))
|
|
701
|
+
|
|
702
|
+
last = ev
|
|
703
|
+
except ClientDisconnected:
|
|
704
|
+
self._abandon("client_disconnect")
|
|
705
|
+
raise
|
|
706
|
+
finally:
|
|
707
|
+
self.release()
|
|
708
|
+
|
|
709
|
+
self._observe(last)
|
|
710
|
+
|
|
711
|
+
if last.finish_reason not in FINISH_TO_OPENAI:
|
|
712
|
+
raise self._terminal_error(last)
|
|
713
|
+
|
|
714
|
+
text = "".join(parts)
|
|
715
|
+
finish = FINISH_TO_OPENAI[last.finish_reason]
|
|
716
|
+
usage = self._usage(last)
|
|
717
|
+
fp = self.model.fingerprint
|
|
718
|
+
|
|
719
|
+
if self.chat:
|
|
720
|
+
body = protocol.chat_response(self.rid, self.model.name, fp, text, finish, usage)
|
|
721
|
+
else:
|
|
722
|
+
body = protocol.completion_response(self.rid, self.model.name, fp, text, finish, usage)
|
|
723
|
+
|
|
724
|
+
return Response(200, body)
|
|
725
|
+
|
|
726
|
+
def _chunk(self, text=None, finish=None, usage=None, include_choice=True, role=False):
|
|
727
|
+
fp = self.model.fingerprint
|
|
728
|
+
|
|
729
|
+
if self.chat:
|
|
730
|
+
delta = {}
|
|
731
|
+
|
|
732
|
+
if role:
|
|
733
|
+
delta["role"] = "assistant"
|
|
734
|
+
delta["content"] = ""
|
|
735
|
+
|
|
736
|
+
if text:
|
|
737
|
+
delta["content"] = text
|
|
738
|
+
|
|
739
|
+
return protocol.chat_chunk(self.rid, self.created, self.model.name, fp, delta, finish, usage,
|
|
740
|
+
include_choice)
|
|
741
|
+
|
|
742
|
+
return protocol.completion_chunk(self.rid, self.created, self.model.name, fp, text or "", finish, usage,
|
|
743
|
+
include_choice)
|
|
744
|
+
|
|
745
|
+
async def stream(self):
|
|
746
|
+
bridge = _Bridge(self.handle, asyncio.get_running_loop())
|
|
747
|
+
finished = False
|
|
748
|
+
|
|
749
|
+
try:
|
|
750
|
+
if self.chat:
|
|
751
|
+
yield sse(self._chunk(role=True))
|
|
752
|
+
|
|
753
|
+
async for ev in bridge.events(self.ctx, self.app.limits.sse_keepalive_s):
|
|
754
|
+
if ev is None:
|
|
755
|
+
yield ": keep-alive\n\n"
|
|
756
|
+
continue
|
|
757
|
+
|
|
758
|
+
text = self.lstrip(ev.text) if ev.text else ""
|
|
759
|
+
|
|
760
|
+
if not ev.finished:
|
|
761
|
+
if text:
|
|
762
|
+
yield sse(self._chunk(text=text))
|
|
763
|
+
continue
|
|
764
|
+
|
|
765
|
+
finished = True
|
|
766
|
+
self._observe(ev)
|
|
767
|
+
|
|
768
|
+
if ev.finish_reason in FINISH_TO_OPENAI:
|
|
769
|
+
if text:
|
|
770
|
+
yield sse(self._chunk(text=text))
|
|
771
|
+
|
|
772
|
+
yield sse(self._chunk(finish=FINISH_TO_OPENAI[ev.finish_reason]))
|
|
773
|
+
|
|
774
|
+
if self.parsed.include_usage:
|
|
775
|
+
yield sse(self._chunk(usage=self._usage(ev), include_choice=False))
|
|
776
|
+
else:
|
|
777
|
+
if text:
|
|
778
|
+
yield sse(self._chunk(text=text))
|
|
779
|
+
|
|
780
|
+
err = self._terminal_error(ev)
|
|
781
|
+
self.request.info["error_code"] = err.code
|
|
782
|
+
yield sse(err.to_dict())
|
|
783
|
+
|
|
784
|
+
yield "data: [DONE]\n\n"
|
|
785
|
+
return
|
|
786
|
+
except ClientDisconnected:
|
|
787
|
+
pass
|
|
788
|
+
finally:
|
|
789
|
+
if not finished:
|
|
790
|
+
self._abandon("client_disconnect")
|
|
791
|
+
|
|
792
|
+
self.release()
|