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.
Files changed (146) hide show
  1. cli.py +604 -0
  2. pytensorforge-0.1.0.dist-info/METADATA +103 -0
  3. pytensorforge-0.1.0.dist-info/RECORD +146 -0
  4. pytensorforge-0.1.0.dist-info/WHEEL +5 -0
  5. pytensorforge-0.1.0.dist-info/entry_points.txt +2 -0
  6. pytensorforge-0.1.0.dist-info/top_level.txt +2 -0
  7. src/__init__.py +0 -0
  8. src/activations/Activation.py +4 -0
  9. src/activations/ELU.py +11 -0
  10. src/activations/GELU.py +6 -0
  11. src/activations/ReLU.py +27 -0
  12. src/activations/SELU.py +14 -0
  13. src/activations/Sigmoid.py +27 -0
  14. src/activations/Softmax.py +84 -0
  15. src/activations/Tanh.py +29 -0
  16. src/activations/__init__.py +17 -0
  17. src/config.py +120 -0
  18. src/core/Matrix.py +3 -0
  19. src/core/Scalar.py +18 -0
  20. src/core/Tensor.py +866 -0
  21. src/core/Vector.py +31 -0
  22. src/core/__init__.py +0 -0
  23. src/data/__init__.py +0 -0
  24. src/data/chat_dataset.py +188 -0
  25. src/data/corpus.py +104 -0
  26. src/data/document_stream.py +178 -0
  27. src/data/parallel_encode.py +86 -0
  28. src/data/prefetch.py +62 -0
  29. src/data/shard_builder.py +119 -0
  30. src/data/shard_writer.py +81 -0
  31. src/data/sharded_dataset.py +112 -0
  32. src/data/streaming_dataset.py +132 -0
  33. src/data/validation.py +212 -0
  34. src/inference/__init__.py +0 -0
  35. src/inference/chat_template.py +384 -0
  36. src/inference/config.py +48 -0
  37. src/inference/engine.py +241 -0
  38. src/inference/export.py +133 -0
  39. src/inference/kv_cache.py +65 -0
  40. src/inference/runtime.py +161 -0
  41. src/inference/sampling.py +42 -0
  42. src/inference/scheduler.py +473 -0
  43. src/inference/text.py +67 -0
  44. src/initializers/Constant.py +9 -0
  45. src/initializers/GlorotNormal.py +15 -0
  46. src/initializers/GlorotUniform.py +26 -0
  47. src/initializers/HeNormal.py +15 -0
  48. src/initializers/HeUniform.py +14 -0
  49. src/initializers/Initializer.py +4 -0
  50. src/initializers/LecunNormal.py +16 -0
  51. src/initializers/LecunUniform.py +14 -0
  52. src/initializers/Ones.py +6 -0
  53. src/initializers/Orthogonal.py +14 -0
  54. src/initializers/RandomNormal.py +14 -0
  55. src/initializers/RandomUniform.py +14 -0
  56. src/initializers/Zeros.py +8 -0
  57. src/initializers/__init__.py +17 -0
  58. src/loss/CategoricalCrossEntropy.py +9 -0
  59. src/loss/CrossEntropyLoss.py +34 -0
  60. src/loss/CrossEntropyWithLogitsLoss.py +59 -0
  61. src/loss/Hinge.py +5 -0
  62. src/loss/Huber.py +22 -0
  63. src/loss/Loss.py +6 -0
  64. src/loss/MSE.py +7 -0
  65. src/loss/MSELoss.py +10 -0
  66. src/loss/SparseCategoricalCrossEntropy.py +15 -0
  67. src/loss/__init__.py +18 -0
  68. src/loss/bce.py +34 -0
  69. src/loss/mae.py +16 -0
  70. src/math/__init__.py +0 -0
  71. src/math/clip.py +37 -0
  72. src/math/exp.py +27 -0
  73. src/math/log.py +25 -0
  74. src/math/sigmoid.py +5 -0
  75. src/models/__init__.py +0 -0
  76. src/models/embedding/Embedding.py +65 -0
  77. src/models/embedding/__init__.py +0 -0
  78. src/models/gpt/__init__.py +0 -0
  79. src/models/gpt/attention.py +158 -0
  80. src/models/gpt/block.py +74 -0
  81. src/models/gpt/config.py +103 -0
  82. src/models/gpt/context.py +44 -0
  83. src/models/gpt/model.py +165 -0
  84. src/models/gpt/recompute.py +35 -0
  85. src/models/gpt/rope.py +84 -0
  86. src/models/regression/Linear.py +51 -0
  87. src/models/regression/Logistic.py +36 -0
  88. src/models/regression/__init__.py +0 -0
  89. src/models/seq/Sequential.py +297 -0
  90. src/models/seq/__init__.py +0 -0
  91. src/models/svm/__init__.py +0 -0
  92. src/models/tokenizer/BPETokenizer.py +228 -0
  93. src/models/tokenizer/__init__.py +0 -0
  94. src/models/transformers/Dropout.py +35 -0
  95. src/models/transformers/LastToken.py +10 -0
  96. src/models/transformers/LayerNorm.py +54 -0
  97. src/models/transformers/Linear.py +18 -0
  98. src/models/transformers/MultiHeadAttention.py +130 -0
  99. src/models/transformers/TransformerBlock.py +79 -0
  100. src/models/transformers/__init__.py +0 -0
  101. src/neural/Dense.py +58 -0
  102. src/neural/LSTM.py +167 -0
  103. src/neural/Layer.py +72 -0
  104. src/neural/Parameter.py +30 -0
  105. src/neural/RNN.py +83 -0
  106. src/neural/__init__.py +0 -0
  107. src/ops/__init__.py +0 -0
  108. src/ops/stack.py +40 -0
  109. src/optimizers/Adagrad.py +31 -0
  110. src/optimizers/Adam.py +98 -0
  111. src/optimizers/AdamW.py +84 -0
  112. src/optimizers/Batch.py +11 -0
  113. src/optimizers/Nesterov.py +35 -0
  114. src/optimizers/Optimizer.py +18 -0
  115. src/optimizers/RMSProp.py +35 -0
  116. src/optimizers/SGD.py +30 -0
  117. src/optimizers/SGDMomentum.py +28 -0
  118. src/optimizers/__init__.py +9 -0
  119. src/scaling/StandardScaler.py +15 -0
  120. src/scaling/__init__.py +0 -0
  121. src/serialization/__init__.py +0 -0
  122. src/serialization/checkpoint.py +58 -0
  123. src/serialization/modelio.py +132 -0
  124. src/serving/__init__.py +0 -0
  125. src/serving/app.py +792 -0
  126. src/serving/config.py +216 -0
  127. src/serving/errors.py +51 -0
  128. src/serving/http.py +599 -0
  129. src/serving/metrics.py +293 -0
  130. src/serving/model_server.py +287 -0
  131. src/serving/protocol.py +377 -0
  132. src/serving/security.py +200 -0
  133. src/serving/server.py +121 -0
  134. src/tokenization/__init__.py +0 -0
  135. src/tokenization/base.py +75 -0
  136. src/tokenization/bpe.py +190 -0
  137. src/tokenization/bytebpe.py +476 -0
  138. src/tokenization/registry.py +28 -0
  139. src/training/__init__.py +0 -0
  140. src/training/checkpoint_manager.py +101 -0
  141. src/training/experiment.py +71 -0
  142. src/training/losses.py +42 -0
  143. src/training/precision.py +141 -0
  144. src/training/profiler.py +38 -0
  145. src/training/scheduler.py +50 -0
  146. 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()