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/protocol.py
ADDED
|
@@ -0,0 +1,377 @@
|
|
|
1
|
+
import time
|
|
2
|
+
import uuid
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
from typing import List, Optional, Union
|
|
5
|
+
|
|
6
|
+
from src.serving.errors import bad_request
|
|
7
|
+
|
|
8
|
+
CHAT_ROLES = ("system", "user", "assistant", "developer")
|
|
9
|
+
|
|
10
|
+
NEUTRAL_PARAMS = {
|
|
11
|
+
"n": 1,
|
|
12
|
+
"presence_penalty": 0,
|
|
13
|
+
"frequency_penalty": 0,
|
|
14
|
+
"logprobs": False,
|
|
15
|
+
"logit_bias": None,
|
|
16
|
+
"echo": False,
|
|
17
|
+
"best_of": 1,
|
|
18
|
+
"suffix": None,
|
|
19
|
+
"parallel_tool_calls": None,
|
|
20
|
+
"top_logprobs": None,
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
IGNORED_PARAMS = ("user", "metadata", "store", "service_tier", "response_format")
|
|
24
|
+
|
|
25
|
+
UNSUPPORTED_PARAMS = ("tools", "tool_choice", "functions", "function_call", "audio", "modalities", "prediction")
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass
|
|
29
|
+
class SamplingParams:
|
|
30
|
+
max_tokens: Optional[int] = None
|
|
31
|
+
temperature: Optional[float] = None
|
|
32
|
+
top_p: Optional[float] = None
|
|
33
|
+
top_k: Optional[int] = None
|
|
34
|
+
repetition_penalty: Optional[float] = None
|
|
35
|
+
seed: Optional[int] = None
|
|
36
|
+
stop: List[str] = field(default_factory=list)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass
|
|
40
|
+
class ChatRequest:
|
|
41
|
+
model: str
|
|
42
|
+
messages: List[dict]
|
|
43
|
+
sampling: SamplingParams
|
|
44
|
+
stream: bool = False
|
|
45
|
+
include_usage: bool = False
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
@dataclass
|
|
49
|
+
class CompletionRequest:
|
|
50
|
+
model: str
|
|
51
|
+
prompt: Union[str, List[int]]
|
|
52
|
+
sampling: SamplingParams
|
|
53
|
+
stream: bool = False
|
|
54
|
+
include_usage: bool = False
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _is_int(v):
|
|
58
|
+
return isinstance(v, int) and not isinstance(v, bool)
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _is_num(v):
|
|
62
|
+
return (isinstance(v, (int, float))) and not isinstance(v, bool)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _number(body, name, lo=None, hi=None, lo_open=False, integer=False):
|
|
66
|
+
if name not in body or body[name] is None:
|
|
67
|
+
return None
|
|
68
|
+
|
|
69
|
+
v = body[name]
|
|
70
|
+
|
|
71
|
+
if integer and not _is_int(v):
|
|
72
|
+
raise bad_request(f"'{name}' must be an integer", param=name)
|
|
73
|
+
|
|
74
|
+
if not integer and not _is_num(v):
|
|
75
|
+
raise bad_request(f"'{name}' must be a number", param=name)
|
|
76
|
+
|
|
77
|
+
if lo is not None and (v <= lo if lo_open else v < lo):
|
|
78
|
+
raise bad_request(f"'{name}' must be {'>' if lo_open else '>='} {lo}", param=name)
|
|
79
|
+
|
|
80
|
+
if hi is not None and v > hi:
|
|
81
|
+
raise bad_request(f"'{name}' must be <= {hi}", param=name)
|
|
82
|
+
|
|
83
|
+
return v
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _check_params(body, allowed):
|
|
87
|
+
if not isinstance(body, dict):
|
|
88
|
+
raise bad_request("request body must be a JSON object")
|
|
89
|
+
|
|
90
|
+
for name in UNSUPPORTED_PARAMS:
|
|
91
|
+
if body.get(name) not in (None, [], {}):
|
|
92
|
+
raise bad_request(f"'{name}' is not supported by this server", param=name, code="unsupported_parameter")
|
|
93
|
+
|
|
94
|
+
for name, neutral in NEUTRAL_PARAMS.items():
|
|
95
|
+
if name in body and body[name] is not None and body[name] != neutral:
|
|
96
|
+
raise bad_request(
|
|
97
|
+
f"'{name}'={body[name]!r} is not supported; only {neutral!r} is accepted",
|
|
98
|
+
param=name, code="unsupported_parameter",
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
known = set(allowed) | set(NEUTRAL_PARAMS) | set(IGNORED_PARAMS) | set(UNSUPPORTED_PARAMS)
|
|
102
|
+
unknown = sorted(k for k in body if k not in known)
|
|
103
|
+
|
|
104
|
+
if unknown:
|
|
105
|
+
raise bad_request(f"unrecognized request parameter(s): {unknown}", param=unknown[0])
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _model(body):
|
|
109
|
+
model = body.get("model")
|
|
110
|
+
|
|
111
|
+
if not isinstance(model, str) or not model:
|
|
112
|
+
raise bad_request("'model' is required and must be a string", param="model")
|
|
113
|
+
|
|
114
|
+
return model
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
def _stop(body, limits):
|
|
118
|
+
stop = body.get("stop")
|
|
119
|
+
|
|
120
|
+
if stop is None:
|
|
121
|
+
return []
|
|
122
|
+
|
|
123
|
+
if isinstance(stop, str):
|
|
124
|
+
stop = [stop]
|
|
125
|
+
|
|
126
|
+
if not isinstance(stop, list) or not all(isinstance(s, str) for s in stop):
|
|
127
|
+
raise bad_request("'stop' must be a string or a list of strings", param="stop")
|
|
128
|
+
|
|
129
|
+
stop = [s for s in stop if s]
|
|
130
|
+
|
|
131
|
+
if len(stop) > limits.max_stop_sequences:
|
|
132
|
+
raise bad_request(f"at most {limits.max_stop_sequences} stop sequences are allowed", param="stop")
|
|
133
|
+
|
|
134
|
+
if any(len(s) > limits.max_stop_sequence_chars for s in stop):
|
|
135
|
+
raise bad_request(f"stop sequences must be at most {limits.max_stop_sequence_chars} characters", param="stop")
|
|
136
|
+
|
|
137
|
+
return stop
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def _stream(body):
|
|
141
|
+
stream = body.get("stream", False)
|
|
142
|
+
|
|
143
|
+
if stream is None:
|
|
144
|
+
stream = False
|
|
145
|
+
|
|
146
|
+
if not isinstance(stream, bool):
|
|
147
|
+
raise bad_request("'stream' must be a boolean", param="stream")
|
|
148
|
+
|
|
149
|
+
opts = body.get("stream_options")
|
|
150
|
+
include_usage = False
|
|
151
|
+
|
|
152
|
+
if opts is not None:
|
|
153
|
+
if not isinstance(opts, dict) or set(opts) - {"include_usage"}:
|
|
154
|
+
raise bad_request("'stream_options' supports only 'include_usage'", param="stream_options")
|
|
155
|
+
|
|
156
|
+
if not stream:
|
|
157
|
+
raise bad_request("'stream_options' is only allowed when 'stream' is true", param="stream_options")
|
|
158
|
+
|
|
159
|
+
include_usage = bool(opts.get("include_usage", False))
|
|
160
|
+
|
|
161
|
+
return stream, include_usage
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def _sampling(body, limits, token_field_names):
|
|
165
|
+
max_tokens = None
|
|
166
|
+
source = None
|
|
167
|
+
|
|
168
|
+
for name in token_field_names:
|
|
169
|
+
v = _number(body, name, lo=1, integer=True)
|
|
170
|
+
|
|
171
|
+
if v is not None:
|
|
172
|
+
if max_tokens is not None and v != max_tokens:
|
|
173
|
+
raise bad_request(f"conflicting values for {' and '.join(token_field_names)}", param=name)
|
|
174
|
+
max_tokens = v
|
|
175
|
+
source = source or name
|
|
176
|
+
|
|
177
|
+
if max_tokens is not None and max_tokens > limits.max_generation_tokens:
|
|
178
|
+
raise bad_request(
|
|
179
|
+
f"{source}={max_tokens} exceeds this server's limit of {limits.max_generation_tokens}",
|
|
180
|
+
param=source, code="max_tokens_exceeded",
|
|
181
|
+
)
|
|
182
|
+
|
|
183
|
+
return SamplingParams(
|
|
184
|
+
max_tokens=max_tokens,
|
|
185
|
+
temperature=_number(body, "temperature", lo=0, hi=2),
|
|
186
|
+
top_p=_number(body, "top_p", lo=0, hi=1, lo_open=True),
|
|
187
|
+
top_k=_number(body, "top_k", lo=0, integer=True),
|
|
188
|
+
repetition_penalty=_number(body, "repetition_penalty", lo=0, hi=10, lo_open=True),
|
|
189
|
+
seed=_number(body, "seed", integer=True),
|
|
190
|
+
stop=_stop(body, limits),
|
|
191
|
+
)
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
def _content_text(content, idx):
|
|
195
|
+
if isinstance(content, str):
|
|
196
|
+
return content
|
|
197
|
+
|
|
198
|
+
if isinstance(content, list):
|
|
199
|
+
parts = []
|
|
200
|
+
|
|
201
|
+
for part in content:
|
|
202
|
+
if not isinstance(part, dict) or part.get("type") != "text" or not isinstance(part.get("text"), str):
|
|
203
|
+
raise bad_request(
|
|
204
|
+
f"messages[{idx}].content: only text content parts are supported", param=f"messages[{idx}].content"
|
|
205
|
+
)
|
|
206
|
+
parts.append(part["text"])
|
|
207
|
+
|
|
208
|
+
return "".join(parts)
|
|
209
|
+
|
|
210
|
+
if content is None:
|
|
211
|
+
return ""
|
|
212
|
+
|
|
213
|
+
raise bad_request(f"messages[{idx}].content must be a string or a list of text parts",
|
|
214
|
+
param=f"messages[{idx}].content")
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
CHAT_FIELDS = ("model", "messages", "stream", "stream_options", "max_tokens", "max_completion_tokens",
|
|
218
|
+
"temperature", "top_p", "top_k", "repetition_penalty", "seed", "stop")
|
|
219
|
+
|
|
220
|
+
COMPLETION_FIELDS = ("model", "prompt", "stream", "stream_options", "max_tokens", "temperature", "top_p",
|
|
221
|
+
"top_k", "repetition_penalty", "seed", "stop")
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
def parse_chat_request(body, limits):
|
|
225
|
+
_check_params(body, CHAT_FIELDS)
|
|
226
|
+
|
|
227
|
+
messages = body.get("messages")
|
|
228
|
+
|
|
229
|
+
if not isinstance(messages, list) or not messages:
|
|
230
|
+
raise bad_request("'messages' must be a non-empty list", param="messages")
|
|
231
|
+
|
|
232
|
+
if len(messages) > limits.max_messages:
|
|
233
|
+
raise bad_request(f"at most {limits.max_messages} messages are allowed", param="messages")
|
|
234
|
+
|
|
235
|
+
out = []
|
|
236
|
+
total_chars = 0
|
|
237
|
+
|
|
238
|
+
for i, m in enumerate(messages):
|
|
239
|
+
if not isinstance(m, dict):
|
|
240
|
+
raise bad_request(f"messages[{i}] must be an object", param=f"messages[{i}]")
|
|
241
|
+
|
|
242
|
+
role = m.get("role")
|
|
243
|
+
|
|
244
|
+
if role not in CHAT_ROLES:
|
|
245
|
+
raise bad_request(f"messages[{i}].role must be one of {list(CHAT_ROLES)}", param=f"messages[{i}].role")
|
|
246
|
+
|
|
247
|
+
if m.get("tool_calls") or role == "tool":
|
|
248
|
+
raise bad_request("tool messages are not supported", param=f"messages[{i}]")
|
|
249
|
+
|
|
250
|
+
text = _content_text(m.get("content"), i)
|
|
251
|
+
total_chars += len(text)
|
|
252
|
+
out.append({"role": "system" if role == "developer" else role, "content": text})
|
|
253
|
+
|
|
254
|
+
if total_chars > limits.max_prompt_chars:
|
|
255
|
+
raise bad_request(f"messages exceed the {limits.max_prompt_chars}-character limit",
|
|
256
|
+
param="messages", code="prompt_too_large")
|
|
257
|
+
|
|
258
|
+
if out[-1]["role"] == "assistant":
|
|
259
|
+
raise bad_request("the last message must not be from the assistant", param="messages")
|
|
260
|
+
|
|
261
|
+
stream, include_usage = _stream(body)
|
|
262
|
+
|
|
263
|
+
return ChatRequest(
|
|
264
|
+
model=_model(body),
|
|
265
|
+
messages=out,
|
|
266
|
+
sampling=_sampling(body, limits, ("max_tokens", "max_completion_tokens")),
|
|
267
|
+
stream=stream,
|
|
268
|
+
include_usage=include_usage,
|
|
269
|
+
)
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def parse_completion_request(body, limits):
|
|
273
|
+
_check_params(body, COMPLETION_FIELDS)
|
|
274
|
+
|
|
275
|
+
prompt = body.get("prompt")
|
|
276
|
+
|
|
277
|
+
if isinstance(prompt, list) and len(prompt) == 1 and isinstance(prompt[0], str):
|
|
278
|
+
prompt = prompt[0]
|
|
279
|
+
elif isinstance(prompt, list) and len(prompt) == 1 and isinstance(prompt[0], list):
|
|
280
|
+
prompt = prompt[0]
|
|
281
|
+
|
|
282
|
+
if isinstance(prompt, str):
|
|
283
|
+
if len(prompt) > limits.max_prompt_chars:
|
|
284
|
+
raise bad_request(f"prompt exceeds the {limits.max_prompt_chars}-character limit",
|
|
285
|
+
param="prompt", code="prompt_too_large")
|
|
286
|
+
elif isinstance(prompt, list) and prompt and all(_is_int(t) for t in prompt):
|
|
287
|
+
if len(prompt) > limits.max_prompt_chars:
|
|
288
|
+
raise bad_request("prompt is too long", param="prompt", code="prompt_too_large")
|
|
289
|
+
elif isinstance(prompt, list):
|
|
290
|
+
raise bad_request("batched prompts are not supported; send one prompt per request", param="prompt")
|
|
291
|
+
else:
|
|
292
|
+
raise bad_request("'prompt' must be a string or a list of token ids", param="prompt")
|
|
293
|
+
|
|
294
|
+
stream, include_usage = _stream(body)
|
|
295
|
+
|
|
296
|
+
return CompletionRequest(
|
|
297
|
+
model=_model(body),
|
|
298
|
+
prompt=prompt,
|
|
299
|
+
sampling=_sampling(body, limits, ("max_tokens",)),
|
|
300
|
+
stream=stream,
|
|
301
|
+
include_usage=include_usage,
|
|
302
|
+
)
|
|
303
|
+
|
|
304
|
+
|
|
305
|
+
def new_id(prefix):
|
|
306
|
+
return f"{prefix}-{uuid.uuid4().hex[:24]}"
|
|
307
|
+
|
|
308
|
+
|
|
309
|
+
def usage_dict(prompt_tokens, completion_tokens):
|
|
310
|
+
return {
|
|
311
|
+
"prompt_tokens": prompt_tokens,
|
|
312
|
+
"completion_tokens": completion_tokens,
|
|
313
|
+
"total_tokens": prompt_tokens + completion_tokens,
|
|
314
|
+
}
|
|
315
|
+
|
|
316
|
+
|
|
317
|
+
def chat_response(rid, model, fingerprint, text, finish_reason, usage):
|
|
318
|
+
return {
|
|
319
|
+
"id": rid,
|
|
320
|
+
"object": "chat.completion",
|
|
321
|
+
"created": int(time.time()),
|
|
322
|
+
"model": model,
|
|
323
|
+
"system_fingerprint": fingerprint,
|
|
324
|
+
"choices": [{
|
|
325
|
+
"index": 0,
|
|
326
|
+
"message": {"role": "assistant", "content": text},
|
|
327
|
+
"logprobs": None,
|
|
328
|
+
"finish_reason": finish_reason,
|
|
329
|
+
}],
|
|
330
|
+
"usage": usage,
|
|
331
|
+
}
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
def chat_chunk(rid, created, model, fingerprint, delta, finish_reason=None, usage=None, include_choice=True):
|
|
335
|
+
chunk = {
|
|
336
|
+
"id": rid,
|
|
337
|
+
"object": "chat.completion.chunk",
|
|
338
|
+
"created": created,
|
|
339
|
+
"model": model,
|
|
340
|
+
"system_fingerprint": fingerprint,
|
|
341
|
+
"choices": [{"index": 0, "delta": delta, "logprobs": None, "finish_reason": finish_reason}]
|
|
342
|
+
if include_choice else [],
|
|
343
|
+
}
|
|
344
|
+
|
|
345
|
+
if usage is not None:
|
|
346
|
+
chunk["usage"] = usage
|
|
347
|
+
|
|
348
|
+
return chunk
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
def completion_response(rid, model, fingerprint, text, finish_reason, usage):
|
|
352
|
+
return {
|
|
353
|
+
"id": rid,
|
|
354
|
+
"object": "text_completion",
|
|
355
|
+
"created": int(time.time()),
|
|
356
|
+
"model": model,
|
|
357
|
+
"system_fingerprint": fingerprint,
|
|
358
|
+
"choices": [{"index": 0, "text": text, "logprobs": None, "finish_reason": finish_reason}],
|
|
359
|
+
"usage": usage,
|
|
360
|
+
}
|
|
361
|
+
|
|
362
|
+
|
|
363
|
+
def completion_chunk(rid, created, model, fingerprint, text, finish_reason=None, usage=None, include_choice=True):
|
|
364
|
+
chunk = {
|
|
365
|
+
"id": rid,
|
|
366
|
+
"object": "text_completion",
|
|
367
|
+
"created": created,
|
|
368
|
+
"model": model,
|
|
369
|
+
"system_fingerprint": fingerprint,
|
|
370
|
+
"choices": [{"index": 0, "text": text, "logprobs": None, "finish_reason": finish_reason}]
|
|
371
|
+
if include_choice else [],
|
|
372
|
+
}
|
|
373
|
+
|
|
374
|
+
if usage is not None:
|
|
375
|
+
chunk["usage"] = usage
|
|
376
|
+
|
|
377
|
+
return chunk
|
src/serving/security.py
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import hmac
|
|
3
|
+
import ipaddress
|
|
4
|
+
import os
|
|
5
|
+
import time
|
|
6
|
+
from collections import OrderedDict
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
|
|
9
|
+
from src.serving import errors
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def hash_key(key):
|
|
13
|
+
return hashlib.sha256(key.encode("utf-8")).hexdigest()
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _parse_entry(entry):
|
|
17
|
+
entry = entry.strip()
|
|
18
|
+
|
|
19
|
+
if not entry or entry.startswith("#"):
|
|
20
|
+
return None
|
|
21
|
+
|
|
22
|
+
if entry.startswith("sha256:"):
|
|
23
|
+
digest = entry[len("sha256:"):].strip().lower()
|
|
24
|
+
|
|
25
|
+
if len(digest) != 64 or any(c not in "0123456789abcdef" for c in digest):
|
|
26
|
+
raise ValueError("malformed sha256 key entry")
|
|
27
|
+
|
|
28
|
+
return digest
|
|
29
|
+
|
|
30
|
+
if len(entry) < 16:
|
|
31
|
+
raise ValueError("API keys must be at least 16 characters")
|
|
32
|
+
|
|
33
|
+
return hash_key(entry)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _read_key_file(path):
|
|
37
|
+
with open(path) as f:
|
|
38
|
+
return [line for line in f.read().splitlines()]
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _collect(inline, path, env_name):
|
|
42
|
+
raw = list(inline or [])
|
|
43
|
+
|
|
44
|
+
if path:
|
|
45
|
+
raw.extend(_read_key_file(path))
|
|
46
|
+
|
|
47
|
+
if env_name and os.environ.get(env_name):
|
|
48
|
+
raw.extend(os.environ[env_name].split(","))
|
|
49
|
+
|
|
50
|
+
return [d for d in (_parse_entry(e) for e in raw) if d]
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@dataclass(frozen=True)
|
|
54
|
+
class Principal:
|
|
55
|
+
id: str
|
|
56
|
+
authenticated: bool
|
|
57
|
+
admin: bool = False
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class KeyStore:
|
|
61
|
+
|
|
62
|
+
def __init__(self, user_digests, admin_digests, allow_unauthenticated=False):
|
|
63
|
+
self._user = [bytes.fromhex(d) for d in dict.fromkeys(user_digests)]
|
|
64
|
+
self._admin = [bytes.fromhex(d) for d in dict.fromkeys(admin_digests)]
|
|
65
|
+
self.allow_unauthenticated = allow_unauthenticated
|
|
66
|
+
|
|
67
|
+
@classmethod
|
|
68
|
+
def from_config(cls, sec):
|
|
69
|
+
users = _collect(sec.api_keys, sec.api_keys_file, sec.api_keys_env)
|
|
70
|
+
admins = _collect(sec.admin_keys, sec.admin_keys_file, None)
|
|
71
|
+
return cls(users, admins, sec.allow_unauthenticated)
|
|
72
|
+
|
|
73
|
+
@property
|
|
74
|
+
def enabled(self):
|
|
75
|
+
return bool(self._user or self._admin)
|
|
76
|
+
|
|
77
|
+
@staticmethod
|
|
78
|
+
def extract(headers):
|
|
79
|
+
auth = headers.get("authorization", "")
|
|
80
|
+
|
|
81
|
+
if auth[:7].lower() == "bearer ":
|
|
82
|
+
return auth[7:].strip() or None
|
|
83
|
+
|
|
84
|
+
key = headers.get("x-api-key", "").strip()
|
|
85
|
+
|
|
86
|
+
return key or None
|
|
87
|
+
|
|
88
|
+
@staticmethod
|
|
89
|
+
def _match(digest, pool):
|
|
90
|
+
found = False
|
|
91
|
+
|
|
92
|
+
for candidate in pool:
|
|
93
|
+
found |= hmac.compare_digest(digest, candidate)
|
|
94
|
+
|
|
95
|
+
return found
|
|
96
|
+
|
|
97
|
+
def authenticate(self, headers, client_ip):
|
|
98
|
+
key = self.extract(headers)
|
|
99
|
+
|
|
100
|
+
if key is None:
|
|
101
|
+
if self.enabled and not self.allow_unauthenticated:
|
|
102
|
+
raise errors.unauthorized("missing API key; send 'Authorization: Bearer <key>'")
|
|
103
|
+
|
|
104
|
+
return Principal(f"ip:{client_ip}", authenticated=False)
|
|
105
|
+
|
|
106
|
+
digest = hashlib.sha256(key.encode("utf-8")).digest()
|
|
107
|
+
is_admin = self._match(digest, self._admin)
|
|
108
|
+
is_user = self._match(digest, self._user)
|
|
109
|
+
|
|
110
|
+
if not (is_admin or is_user):
|
|
111
|
+
raise errors.unauthorized()
|
|
112
|
+
|
|
113
|
+
return Principal(f"key:{digest.hex()[:12]}", authenticated=True, admin=is_admin)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def is_loopback(host):
|
|
117
|
+
if host in ("localhost", ""):
|
|
118
|
+
return True
|
|
119
|
+
|
|
120
|
+
try:
|
|
121
|
+
return ipaddress.ip_address(host).is_loopback
|
|
122
|
+
except ValueError:
|
|
123
|
+
return False
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
class _Bucket:
|
|
127
|
+
__slots__ = ("tokens", "updated", "in_flight")
|
|
128
|
+
|
|
129
|
+
def __init__(self, capacity, now):
|
|
130
|
+
self.tokens = float(capacity)
|
|
131
|
+
self.updated = now
|
|
132
|
+
self.in_flight = 0
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
class RateLimiter:
|
|
136
|
+
|
|
137
|
+
def __init__(self, requests_per_minute, burst, max_concurrent, max_tracked=10_000, clock=time.monotonic):
|
|
138
|
+
self.rate = requests_per_minute / 60.0
|
|
139
|
+
self.capacity = burst
|
|
140
|
+
self.max_concurrent = max_concurrent
|
|
141
|
+
self.max_tracked = max_tracked
|
|
142
|
+
self.clock = clock
|
|
143
|
+
self._buckets = OrderedDict()
|
|
144
|
+
|
|
145
|
+
@classmethod
|
|
146
|
+
def from_config(cls, rl):
|
|
147
|
+
return cls(rl.requests_per_minute, rl.burst, rl.max_concurrent_per_principal, rl.max_tracked_principals)
|
|
148
|
+
|
|
149
|
+
def _bucket(self, key, now):
|
|
150
|
+
b = self._buckets.get(key)
|
|
151
|
+
|
|
152
|
+
if b is None:
|
|
153
|
+
b = _Bucket(self.capacity, now)
|
|
154
|
+
self._buckets[key] = b
|
|
155
|
+
|
|
156
|
+
while len(self._buckets) > self.max_tracked:
|
|
157
|
+
oldest, victim = next(iter(self._buckets.items()))
|
|
158
|
+
|
|
159
|
+
if victim.in_flight:
|
|
160
|
+
self._buckets.move_to_end(oldest)
|
|
161
|
+
break
|
|
162
|
+
|
|
163
|
+
self._buckets.popitem(last=False)
|
|
164
|
+
else:
|
|
165
|
+
self._buckets.move_to_end(key)
|
|
166
|
+
|
|
167
|
+
b.tokens = min(self.capacity, b.tokens + (now - b.updated) * self.rate)
|
|
168
|
+
b.updated = now
|
|
169
|
+
|
|
170
|
+
return b
|
|
171
|
+
|
|
172
|
+
def acquire(self, key):
|
|
173
|
+
now = self.clock()
|
|
174
|
+
b = self._bucket(key, now)
|
|
175
|
+
|
|
176
|
+
if b.in_flight >= self.max_concurrent:
|
|
177
|
+
raise errors.rate_limited(
|
|
178
|
+
f"too many concurrent requests for this API key (limit {self.max_concurrent})", 1
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
if b.tokens < 1.0:
|
|
182
|
+
raise errors.rate_limited("request rate limit exceeded", (1.0 - b.tokens) / self.rate)
|
|
183
|
+
|
|
184
|
+
b.tokens -= 1.0
|
|
185
|
+
b.in_flight += 1
|
|
186
|
+
|
|
187
|
+
def release(self, key):
|
|
188
|
+
b = self._buckets.get(key)
|
|
189
|
+
|
|
190
|
+
if b is not None and b.in_flight > 0:
|
|
191
|
+
b.in_flight -= 1
|
|
192
|
+
|
|
193
|
+
def charge(self, key):
|
|
194
|
+
now = self.clock()
|
|
195
|
+
b = self._bucket(key, now)
|
|
196
|
+
|
|
197
|
+
if b.tokens < 1.0:
|
|
198
|
+
raise errors.rate_limited("request rate limit exceeded", (1.0 - b.tokens) / self.rate)
|
|
199
|
+
|
|
200
|
+
b.tokens -= 1.0
|
src/serving/server.py
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
import signal
|
|
5
|
+
import threading
|
|
6
|
+
|
|
7
|
+
from src.serving.app import APIApp
|
|
8
|
+
from src.serving.http import HTTPServer
|
|
9
|
+
from src.serving.model_server import ModelServer
|
|
10
|
+
from src.serving.security import is_loopback
|
|
11
|
+
|
|
12
|
+
log = logging.getLogger("ptf.server")
|
|
13
|
+
|
|
14
|
+
DEFAULT_UI_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))),
|
|
15
|
+
"web", "dist")
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class InsecureConfiguration(ValueError):
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def check_security(config, keys):
|
|
23
|
+
if keys.enabled:
|
|
24
|
+
return
|
|
25
|
+
|
|
26
|
+
if config.security.allow_unauthenticated:
|
|
27
|
+
if not is_loopback(config.host):
|
|
28
|
+
log.warning("serving WITHOUT authentication on non-loopback address %s", config.host)
|
|
29
|
+
return
|
|
30
|
+
|
|
31
|
+
if not is_loopback(config.host):
|
|
32
|
+
raise InsecureConfiguration(
|
|
33
|
+
f"refusing to listen on {config.host} without API keys; configure security.api_keys / "
|
|
34
|
+
f"api_keys_file / ${config.security.api_keys_env}, or set security.allow_unauthenticated"
|
|
35
|
+
)
|
|
36
|
+
|
|
37
|
+
config.security.allow_unauthenticated = True
|
|
38
|
+
keys.allow_unauthenticated = True
|
|
39
|
+
log.warning("no API keys configured; accepting unauthenticated requests on loopback %s", config.host)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class APIServer:
|
|
43
|
+
|
|
44
|
+
def __init__(self, config):
|
|
45
|
+
self.config = config.validate()
|
|
46
|
+
self.models = ModelServer(self.config)
|
|
47
|
+
ui_dir = self.config.ui.directory or DEFAULT_UI_DIR
|
|
48
|
+
self.app = APIApp(self.models, self.config, ui_directory=ui_dir)
|
|
49
|
+
check_security(self.config, self.app.keys)
|
|
50
|
+
|
|
51
|
+
cors = {"Server": "pytensorforge"}
|
|
52
|
+
self.http = HTTPServer(self.app, self.config.host, self.config.port, self.config.limits,
|
|
53
|
+
on_request_done=self.app.on_request_done, extra_headers=cors)
|
|
54
|
+
self.http.on_connection_change = lambda n: self.app.metrics.open_connections.set(n)
|
|
55
|
+
self._loop = None
|
|
56
|
+
self._thread = None
|
|
57
|
+
self._stopped = None
|
|
58
|
+
self._stop_requested = None
|
|
59
|
+
|
|
60
|
+
@property
|
|
61
|
+
def port(self):
|
|
62
|
+
return self.http.bound_port
|
|
63
|
+
|
|
64
|
+
async def _run(self, ready=None, install_signals=True):
|
|
65
|
+
self._loop = asyncio.get_running_loop()
|
|
66
|
+
self._stop_requested = asyncio.Event()
|
|
67
|
+
|
|
68
|
+
await self._loop.run_in_executor(None, self.models.preload)
|
|
69
|
+
await self.http.start()
|
|
70
|
+
log.info("listening on http://%s:%d", self.config.host, self.port)
|
|
71
|
+
|
|
72
|
+
if install_signals:
|
|
73
|
+
for sig in (signal.SIGINT, signal.SIGTERM):
|
|
74
|
+
try:
|
|
75
|
+
self._loop.add_signal_handler(sig, self._stop_requested.set)
|
|
76
|
+
except (NotImplementedError, RuntimeError):
|
|
77
|
+
pass
|
|
78
|
+
|
|
79
|
+
if ready is not None:
|
|
80
|
+
ready.set()
|
|
81
|
+
|
|
82
|
+
await self._stop_requested.wait()
|
|
83
|
+
log.info("shutting down: draining in-flight requests (up to %.0fs)", self.config.limits.shutdown_drain_s)
|
|
84
|
+
|
|
85
|
+
await self.http.shutdown(self.config.limits.shutdown_drain_s)
|
|
86
|
+
await self._loop.run_in_executor(None, self.models.shutdown, 5.0)
|
|
87
|
+
self.app.close()
|
|
88
|
+
log.info("shutdown complete")
|
|
89
|
+
|
|
90
|
+
def serve_forever(self):
|
|
91
|
+
asyncio.run(self._run())
|
|
92
|
+
|
|
93
|
+
def start_background(self, timeout=60.0):
|
|
94
|
+
ready = threading.Event()
|
|
95
|
+
failure = []
|
|
96
|
+
|
|
97
|
+
def target():
|
|
98
|
+
try:
|
|
99
|
+
asyncio.run(self._run(ready=ready, install_signals=False))
|
|
100
|
+
except BaseException as exc:
|
|
101
|
+
failure.append(exc)
|
|
102
|
+
ready.set()
|
|
103
|
+
|
|
104
|
+
self._thread = threading.Thread(target=target, name="ptf-api", daemon=True)
|
|
105
|
+
self._thread.start()
|
|
106
|
+
|
|
107
|
+
if not ready.wait(timeout):
|
|
108
|
+
raise TimeoutError("server did not start in time")
|
|
109
|
+
|
|
110
|
+
if failure:
|
|
111
|
+
raise failure[0]
|
|
112
|
+
|
|
113
|
+
return self
|
|
114
|
+
|
|
115
|
+
def stop(self, timeout=30.0):
|
|
116
|
+
if self._loop is not None and self._stop_requested is not None:
|
|
117
|
+
self._loop.call_soon_threadsafe(self._stop_requested.set)
|
|
118
|
+
|
|
119
|
+
if self._thread is not None:
|
|
120
|
+
self._thread.join(timeout)
|
|
121
|
+
self._thread = None
|
|
File without changes
|