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/metrics.py
ADDED
|
@@ -0,0 +1,293 @@
|
|
|
1
|
+
import bisect
|
|
2
|
+
import math
|
|
3
|
+
import os
|
|
4
|
+
import resource
|
|
5
|
+
import threading
|
|
6
|
+
import time
|
|
7
|
+
|
|
8
|
+
LATENCY_BUCKETS = (0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30, 60, 120, 300)
|
|
9
|
+
RATE_BUCKETS = (1, 2, 5, 10, 20, 50, 100, 200, 500, 1000, 2000, 5000)
|
|
10
|
+
TOKEN_BUCKETS = (16, 64, 256, 512, 1024, 2048, 4096, 8192, 16384, 32768)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _escape(value):
|
|
14
|
+
return str(value).replace("\\", "\\\\").replace("\n", "\\n").replace('"', '\\"')
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _labels(names, values):
|
|
18
|
+
if not names:
|
|
19
|
+
return ""
|
|
20
|
+
|
|
21
|
+
return "{" + ",".join(f'{n}="{_escape(v)}"' for n, v in zip(names, values)) + "}"
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _fmt(v):
|
|
25
|
+
if v == math.inf:
|
|
26
|
+
return "+Inf"
|
|
27
|
+
|
|
28
|
+
if isinstance(v, float) and v.is_integer() and abs(v) < 1e15:
|
|
29
|
+
return str(int(v))
|
|
30
|
+
|
|
31
|
+
return repr(float(v)) if isinstance(v, float) else str(v)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class _Metric:
|
|
35
|
+
kind = ""
|
|
36
|
+
|
|
37
|
+
def __init__(self, name, help_text, labels=()):
|
|
38
|
+
self.name = name
|
|
39
|
+
self.help = help_text
|
|
40
|
+
self.label_names = tuple(labels)
|
|
41
|
+
self._lock = threading.Lock()
|
|
42
|
+
|
|
43
|
+
def _key(self, labels):
|
|
44
|
+
if set(labels) != set(self.label_names):
|
|
45
|
+
raise ValueError(f"{self.name} expects labels {self.label_names}, got {tuple(labels)}")
|
|
46
|
+
|
|
47
|
+
return tuple(str(labels[n]) for n in self.label_names)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class Counter(_Metric):
|
|
51
|
+
kind = "counter"
|
|
52
|
+
|
|
53
|
+
def __init__(self, *a, **k):
|
|
54
|
+
super().__init__(*a, **k)
|
|
55
|
+
self._values = {}
|
|
56
|
+
|
|
57
|
+
def inc(self, amount=1, **labels):
|
|
58
|
+
key = self._key(labels)
|
|
59
|
+
|
|
60
|
+
with self._lock:
|
|
61
|
+
self._values[key] = self._values.get(key, 0) + amount
|
|
62
|
+
|
|
63
|
+
def get(self, **labels):
|
|
64
|
+
return self._values.get(self._key(labels), 0)
|
|
65
|
+
|
|
66
|
+
def samples(self):
|
|
67
|
+
with self._lock:
|
|
68
|
+
return [(self.name + "_total", k, v) for k, v in sorted(self._values.items())]
|
|
69
|
+
|
|
70
|
+
def snapshot(self):
|
|
71
|
+
with self._lock:
|
|
72
|
+
return {",".join(k) or "": v for k, v in self._values.items()}
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
class Gauge(_Metric):
|
|
76
|
+
kind = "gauge"
|
|
77
|
+
|
|
78
|
+
def __init__(self, *a, fn=None, **k):
|
|
79
|
+
super().__init__(*a, **k)
|
|
80
|
+
self._values = {}
|
|
81
|
+
self._fn = fn
|
|
82
|
+
|
|
83
|
+
def set(self, value, **labels):
|
|
84
|
+
key = self._key(labels)
|
|
85
|
+
|
|
86
|
+
with self._lock:
|
|
87
|
+
self._values[key] = value
|
|
88
|
+
|
|
89
|
+
def inc(self, amount=1, **labels):
|
|
90
|
+
key = self._key(labels)
|
|
91
|
+
|
|
92
|
+
with self._lock:
|
|
93
|
+
self._values[key] = self._values.get(key, 0) + amount
|
|
94
|
+
|
|
95
|
+
def dec(self, amount=1, **labels):
|
|
96
|
+
self.inc(-amount, **labels)
|
|
97
|
+
|
|
98
|
+
def remove(self, **labels):
|
|
99
|
+
with self._lock:
|
|
100
|
+
self._values.pop(self._key(labels), None)
|
|
101
|
+
|
|
102
|
+
def _current(self):
|
|
103
|
+
if self._fn is not None:
|
|
104
|
+
return {self._key(lbl): v for lbl, v in self._fn()}
|
|
105
|
+
|
|
106
|
+
with self._lock:
|
|
107
|
+
return dict(self._values)
|
|
108
|
+
|
|
109
|
+
def samples(self):
|
|
110
|
+
return [(self.name, k, v) for k, v in sorted(self._current().items())]
|
|
111
|
+
|
|
112
|
+
def snapshot(self):
|
|
113
|
+
return {",".join(k) or "": v for k, v in self._current().items()}
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class Histogram(_Metric):
|
|
117
|
+
kind = "histogram"
|
|
118
|
+
|
|
119
|
+
def __init__(self, *a, buckets=LATENCY_BUCKETS, **k):
|
|
120
|
+
super().__init__(*a, **k)
|
|
121
|
+
self.buckets = tuple(sorted(buckets))
|
|
122
|
+
self._data = {}
|
|
123
|
+
|
|
124
|
+
def observe(self, value, **labels):
|
|
125
|
+
if value is None or not math.isfinite(value):
|
|
126
|
+
return
|
|
127
|
+
|
|
128
|
+
key = self._key(labels)
|
|
129
|
+
idx = bisect.bisect_left(self.buckets, value)
|
|
130
|
+
|
|
131
|
+
with self._lock:
|
|
132
|
+
d = self._data.get(key)
|
|
133
|
+
|
|
134
|
+
if d is None:
|
|
135
|
+
d = self._data[key] = [[0] * (len(self.buckets) + 1), 0.0, 0]
|
|
136
|
+
|
|
137
|
+
d[0][idx] += 1
|
|
138
|
+
d[1] += value
|
|
139
|
+
d[2] += 1
|
|
140
|
+
|
|
141
|
+
def samples(self):
|
|
142
|
+
out = []
|
|
143
|
+
|
|
144
|
+
with self._lock:
|
|
145
|
+
items = sorted((k, (list(v[0]), v[1], v[2])) for k, v in self._data.items())
|
|
146
|
+
|
|
147
|
+
for key, (counts, total, n) in items:
|
|
148
|
+
running = 0
|
|
149
|
+
|
|
150
|
+
for bound, c in zip(self.buckets + (math.inf,), counts):
|
|
151
|
+
running += c
|
|
152
|
+
out.append((self.name + "_bucket", key + (_fmt(bound),), running))
|
|
153
|
+
|
|
154
|
+
out.append((self.name + "_sum", key, total))
|
|
155
|
+
out.append((self.name + "_count", key, n))
|
|
156
|
+
|
|
157
|
+
return out
|
|
158
|
+
|
|
159
|
+
def label_names_for(self, sample_name):
|
|
160
|
+
if sample_name.endswith("_bucket"):
|
|
161
|
+
return self.label_names + ("le",)
|
|
162
|
+
|
|
163
|
+
return self.label_names
|
|
164
|
+
|
|
165
|
+
def snapshot(self):
|
|
166
|
+
with self._lock:
|
|
167
|
+
return {
|
|
168
|
+
",".join(k) or "": {"count": v[2], "sum": v[1], "mean": (v[1] / v[2]) if v[2] else None}
|
|
169
|
+
for k, v in self._data.items()
|
|
170
|
+
}
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
class Registry:
|
|
174
|
+
|
|
175
|
+
def __init__(self, namespace="ptf"):
|
|
176
|
+
self.namespace = namespace
|
|
177
|
+
self._metrics = []
|
|
178
|
+
|
|
179
|
+
def _add(self, metric):
|
|
180
|
+
metric.name = f"{self.namespace}_{metric.name}"
|
|
181
|
+
self._metrics.append(metric)
|
|
182
|
+
return metric
|
|
183
|
+
|
|
184
|
+
def counter(self, name, help_text, labels=()):
|
|
185
|
+
return self._add(Counter(name, help_text, labels))
|
|
186
|
+
|
|
187
|
+
def gauge(self, name, help_text, labels=(), fn=None):
|
|
188
|
+
return self._add(Gauge(name, help_text, labels, fn=fn))
|
|
189
|
+
|
|
190
|
+
def histogram(self, name, help_text, labels=(), buckets=LATENCY_BUCKETS):
|
|
191
|
+
return self._add(Histogram(name, help_text, labels, buckets=buckets))
|
|
192
|
+
|
|
193
|
+
def render_prometheus(self):
|
|
194
|
+
lines = []
|
|
195
|
+
|
|
196
|
+
for m in self._metrics:
|
|
197
|
+
lines.append(f"# HELP {m.name} {m.help}")
|
|
198
|
+
lines.append(f"# TYPE {m.name} {m.kind}")
|
|
199
|
+
|
|
200
|
+
for sample_name, key, value in m.samples():
|
|
201
|
+
names = m.label_names_for(sample_name) if isinstance(m, Histogram) else m.label_names
|
|
202
|
+
lines.append(f"{sample_name}{_labels(names, key)} {_fmt(value)}")
|
|
203
|
+
|
|
204
|
+
return "\n".join(lines) + "\n"
|
|
205
|
+
|
|
206
|
+
def snapshot(self):
|
|
207
|
+
return {m.name: m.snapshot() for m in self._metrics}
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
def _rss_bytes():
|
|
211
|
+
try:
|
|
212
|
+
with open("/proc/self/statm") as f:
|
|
213
|
+
return int(f.read().split()[1]) * os.sysconf("SC_PAGE_SIZE")
|
|
214
|
+
except (OSError, ValueError, IndexError):
|
|
215
|
+
return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * 1024
|
|
216
|
+
|
|
217
|
+
|
|
218
|
+
class ServerMetrics:
|
|
219
|
+
|
|
220
|
+
def __init__(self, model_server_ref):
|
|
221
|
+
r = self.registry = Registry("ptf")
|
|
222
|
+
self._server = model_server_ref
|
|
223
|
+
self.started = time.time()
|
|
224
|
+
|
|
225
|
+
self.http_requests = r.counter("http_requests", "HTTP requests by route and status", ("route", "status"))
|
|
226
|
+
self.http_duration = r.histogram("http_request_duration_seconds", "Full HTTP request duration", ("route",))
|
|
227
|
+
self.open_connections = r.gauge("open_connections", "Open client TCP connections")
|
|
228
|
+
|
|
229
|
+
self.requests = r.counter("completion_requests", "Completion requests accepted", ("model", "endpoint", "stream"))
|
|
230
|
+
self.finished = r.counter("completion_finished", "Completions by finish reason", ("model", "reason"))
|
|
231
|
+
self.cancellations = r.counter("cancellations", "Generations cancelled before finishing", ("model", "cause"))
|
|
232
|
+
self.errors = r.counter("errors", "Errors by kind", ("kind",))
|
|
233
|
+
self.auth_failures = r.counter("auth_failures", "Rejected authentication attempts")
|
|
234
|
+
self.rate_limited = r.counter("rate_limited", "Requests rejected by rate limits", ("reason",))
|
|
235
|
+
|
|
236
|
+
self.prompt_tokens = r.counter("prompt_tokens", "Prompt tokens processed", ("model",))
|
|
237
|
+
self.completion_tokens = r.counter("completion_tokens", "Completion tokens generated", ("model",))
|
|
238
|
+
self.active = r.gauge("active_requests", "Requests in flight (queued or generating)", ("model",))
|
|
239
|
+
|
|
240
|
+
self.tokenize_s = r.histogram("tokenize_seconds", "Chat templating + prompt tokenization time", ("model",))
|
|
241
|
+
self.queue_s = r.histogram("queue_wait_seconds", "Time from submit to admission into a batch", ("model",))
|
|
242
|
+
self.prefill_s = r.histogram("prefill_seconds", "Prompt prefill compute time", ("model",))
|
|
243
|
+
self.ttft_s = r.histogram("time_to_first_token_seconds", "Submit to first generated token", ("model",))
|
|
244
|
+
self.generation_s = r.histogram("generation_duration_seconds", "Submit to final token", ("model",))
|
|
245
|
+
self.decode_rate = r.histogram("decode_tokens_per_second", "Per-request generation rate after first token",
|
|
246
|
+
("model",), buckets=RATE_BUCKETS)
|
|
247
|
+
self.network_s = r.histogram("network_write_seconds", "Time blocked writing response bytes to the client",
|
|
248
|
+
("route",))
|
|
249
|
+
self.prompt_len = r.histogram("prompt_length_tokens", "Prompt length", ("model",), buckets=TOKEN_BUCKETS)
|
|
250
|
+
|
|
251
|
+
r.gauge("process_resident_memory_bytes", "Resident set size", fn=lambda: [({}, _rss_bytes())])
|
|
252
|
+
r.gauge("process_cpu_seconds", "CPU seconds used by the server process",
|
|
253
|
+
fn=lambda: [({}, time.process_time())])
|
|
254
|
+
r.gauge("uptime_seconds", "Seconds since server start", fn=lambda: [({}, time.time() - self.started)])
|
|
255
|
+
|
|
256
|
+
self._scheduler_gauges(r)
|
|
257
|
+
|
|
258
|
+
def _each_model(self):
|
|
259
|
+
server = self._server()
|
|
260
|
+
|
|
261
|
+
if server is None:
|
|
262
|
+
return []
|
|
263
|
+
|
|
264
|
+
return server.loaded_models()
|
|
265
|
+
|
|
266
|
+
def _scheduler_gauges(self, r):
|
|
267
|
+
def per_model(fn):
|
|
268
|
+
return lambda: [({"model": m.name}, fn(m)) for m in self._each_model()]
|
|
269
|
+
|
|
270
|
+
def sched(key):
|
|
271
|
+
return per_model(lambda m: m.generator.scheduler.stats[key])
|
|
272
|
+
|
|
273
|
+
r.gauge("scheduler_waiting", "Requests waiting for a batch slot or cache", ("model",),
|
|
274
|
+
fn=per_model(lambda m: m.generator.scheduler.num_waiting))
|
|
275
|
+
r.gauge("scheduler_running", "Sequences in the running batch", ("model",),
|
|
276
|
+
fn=per_model(lambda m: m.generator.scheduler.num_active))
|
|
277
|
+
r.gauge("kv_cache_used_bytes", "KV cache bytes allocated", ("model",),
|
|
278
|
+
fn=per_model(lambda m: m.generator.scheduler.cache_manager.used_bytes))
|
|
279
|
+
r.gauge("kv_cache_budget_bytes", "KV cache byte budget", ("model",),
|
|
280
|
+
fn=per_model(lambda m: m.generator.scheduler.cache_manager.max_bytes or 0))
|
|
281
|
+
r.gauge("model_weights_bytes", "Weights resident in memory", ("model",),
|
|
282
|
+
fn=per_model(lambda m: m.generator.weights_nbytes))
|
|
283
|
+
r.gauge("scheduler_decode_seconds", "Cumulative batched decode compute", ("model",),
|
|
284
|
+
fn=sched("decode_seconds"))
|
|
285
|
+
r.gauge("scheduler_prefill_seconds", "Cumulative prefill compute", ("model",),
|
|
286
|
+
fn=sched("prefill_seconds"))
|
|
287
|
+
r.gauge("scheduler_decode_steps", "Batched decode steps executed", ("model",),
|
|
288
|
+
fn=sched("decode_steps"))
|
|
289
|
+
r.gauge("scheduler_mean_batch_size", "Generated tokens per decode step", ("model",),
|
|
290
|
+
fn=per_model(lambda m: (m.generator.scheduler.stats["generated_tokens"]
|
|
291
|
+
/ max(1, m.generator.scheduler.stats["decode_steps"]))))
|
|
292
|
+
r.gauge("scheduler_max_batch_seen", "Largest decode batch so far", ("model",),
|
|
293
|
+
fn=sched("max_batch_seen"))
|
|
@@ -0,0 +1,287 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import threading
|
|
3
|
+
import time
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
|
|
6
|
+
from src.inference.runtime import load_model
|
|
7
|
+
from src.serving import errors
|
|
8
|
+
from src.models.gpt.context import describe_context
|
|
9
|
+
|
|
10
|
+
log = logging.getLogger("ptf.server")
|
|
11
|
+
|
|
12
|
+
SUPPORTED_DEVICES = ("cpu",)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def resolve_device(requested):
|
|
16
|
+
requested = (requested or "auto").lower()
|
|
17
|
+
|
|
18
|
+
if requested in ("auto", "cpu"):
|
|
19
|
+
return "cpu"
|
|
20
|
+
|
|
21
|
+
raise ValueError(
|
|
22
|
+
f"device '{requested}' is not available: this build's inference engine runs on {SUPPORTED_DEVICES} only"
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass
|
|
27
|
+
class ModelLimits:
|
|
28
|
+
context_length: int
|
|
29
|
+
max_generation_tokens: int
|
|
30
|
+
max_prompt_tokens: int
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class ServedModel:
|
|
34
|
+
|
|
35
|
+
def __init__(self, entry, generator, template, template_source, limits, device):
|
|
36
|
+
self.entry = entry
|
|
37
|
+
self.name = entry.name
|
|
38
|
+
self.generator = generator
|
|
39
|
+
self.template = template
|
|
40
|
+
self.template_source = template_source
|
|
41
|
+
self.limits = limits
|
|
42
|
+
self.device = device
|
|
43
|
+
self.loaded_at = int(time.time())
|
|
44
|
+
self.in_flight = 0
|
|
45
|
+
self.accepting = True
|
|
46
|
+
self._idle = threading.Condition()
|
|
47
|
+
|
|
48
|
+
@property
|
|
49
|
+
def fingerprint(self):
|
|
50
|
+
return "ptf-" + self.generator.metadata.get("weights_sha256", "")[:12]
|
|
51
|
+
|
|
52
|
+
@property
|
|
53
|
+
def cache_budget_bytes(self):
|
|
54
|
+
return self.generator.scheduler.cache_manager.max_bytes
|
|
55
|
+
|
|
56
|
+
@property
|
|
57
|
+
def memory_bytes(self):
|
|
58
|
+
return self.generator.weights_nbytes + (self.cache_budget_bytes or 0)
|
|
59
|
+
|
|
60
|
+
def enter(self):
|
|
61
|
+
with self._idle:
|
|
62
|
+
if not self.accepting:
|
|
63
|
+
raise errors.overloaded(f"model '{self.name}' is being unloaded")
|
|
64
|
+
|
|
65
|
+
self.in_flight += 1
|
|
66
|
+
|
|
67
|
+
def exit(self):
|
|
68
|
+
with self._idle:
|
|
69
|
+
self.in_flight -= 1
|
|
70
|
+
|
|
71
|
+
if self.in_flight <= 0:
|
|
72
|
+
self._idle.notify_all()
|
|
73
|
+
|
|
74
|
+
def wait_idle(self, timeout):
|
|
75
|
+
deadline = time.monotonic() + timeout
|
|
76
|
+
|
|
77
|
+
with self._idle:
|
|
78
|
+
while self.in_flight > 0:
|
|
79
|
+
left = deadline - time.monotonic()
|
|
80
|
+
|
|
81
|
+
if left <= 0:
|
|
82
|
+
return False
|
|
83
|
+
|
|
84
|
+
self._idle.wait(left)
|
|
85
|
+
|
|
86
|
+
return True
|
|
87
|
+
|
|
88
|
+
def describe(self):
|
|
89
|
+
cfg = self.generator.config
|
|
90
|
+
|
|
91
|
+
return {
|
|
92
|
+
"id": self.name,
|
|
93
|
+
"object": "model",
|
|
94
|
+
"created": self.loaded_at,
|
|
95
|
+
"owned_by": "pytensorforge",
|
|
96
|
+
"context_length": self.limits.context_length,
|
|
97
|
+
"context": describe_context(cfg),
|
|
98
|
+
"max_generation_tokens": self.limits.max_generation_tokens,
|
|
99
|
+
"chat_template": self.template.name,
|
|
100
|
+
"chat_template_source": self.template_source,
|
|
101
|
+
"architecture": cfg.architecture,
|
|
102
|
+
"vocab_size": cfg.vocab_size,
|
|
103
|
+
"n_layers": cfg.n_layers,
|
|
104
|
+
"d_model": cfg.d_model,
|
|
105
|
+
"n_heads": cfg.n_heads,
|
|
106
|
+
"device": self.device,
|
|
107
|
+
"fingerprint": self.fingerprint,
|
|
108
|
+
}
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
class ModelServer:
|
|
112
|
+
|
|
113
|
+
def __init__(self, config):
|
|
114
|
+
self.config = config
|
|
115
|
+
self.device = resolve_device(config.runtime.device)
|
|
116
|
+
self._entries = {m.name: m for m in config.models}
|
|
117
|
+
self._models = {}
|
|
118
|
+
self._lock = threading.RLock()
|
|
119
|
+
self._loading = set()
|
|
120
|
+
self._global_in_flight = 0
|
|
121
|
+
self._global_lock = threading.Lock()
|
|
122
|
+
|
|
123
|
+
def declared(self):
|
|
124
|
+
return list(self._entries.values())
|
|
125
|
+
|
|
126
|
+
def loaded_models(self):
|
|
127
|
+
with self._lock:
|
|
128
|
+
return list(self._models.values())
|
|
129
|
+
|
|
130
|
+
def get(self, name):
|
|
131
|
+
with self._lock:
|
|
132
|
+
model = self._models.get(name)
|
|
133
|
+
|
|
134
|
+
if model is None or not model.accepting:
|
|
135
|
+
if name in self._entries:
|
|
136
|
+
raise errors.APIError(503, f"model '{name}' is not loaded", "server_error", code="model_not_loaded")
|
|
137
|
+
|
|
138
|
+
raise errors.not_found(f"the model '{name}' does not exist", code="model_not_found")
|
|
139
|
+
|
|
140
|
+
return model
|
|
141
|
+
|
|
142
|
+
def _default_budget(self, generator, entry):
|
|
143
|
+
ctx = generator.config.context_length
|
|
144
|
+
|
|
145
|
+
if entry.max_context_tokens:
|
|
146
|
+
ctx = min(ctx, entry.max_context_tokens)
|
|
147
|
+
|
|
148
|
+
return generator.scheduler.cache_manager.bytes_for(ctx) * entry.max_batch_size
|
|
149
|
+
|
|
150
|
+
def _committed_bytes(self, exclude=None):
|
|
151
|
+
return sum(m.memory_bytes for n, m in self._models.items() if n != exclude)
|
|
152
|
+
|
|
153
|
+
def load(self, name):
|
|
154
|
+
entry = self._entries.get(name)
|
|
155
|
+
|
|
156
|
+
if entry is None:
|
|
157
|
+
raise errors.not_found(f"model '{name}' is not declared in the server configuration", code="model_not_found")
|
|
158
|
+
|
|
159
|
+
with self._lock:
|
|
160
|
+
if name in self._models:
|
|
161
|
+
return self._models[name]
|
|
162
|
+
|
|
163
|
+
if name in self._loading:
|
|
164
|
+
raise errors.APIError(409, f"model '{name}' is already loading", "invalid_request_error", code="conflict")
|
|
165
|
+
|
|
166
|
+
self._loading.add(name)
|
|
167
|
+
|
|
168
|
+
try:
|
|
169
|
+
return self._load(entry)
|
|
170
|
+
finally:
|
|
171
|
+
with self._lock:
|
|
172
|
+
self._loading.discard(name)
|
|
173
|
+
|
|
174
|
+
def _load(self, entry):
|
|
175
|
+
started = time.perf_counter()
|
|
176
|
+
budget = int(entry.kv_cache_budget_mb * (1 << 20)) if entry.kv_cache_budget_mb else None
|
|
177
|
+
|
|
178
|
+
generator = load_model(
|
|
179
|
+
entry.path,
|
|
180
|
+
max_batch_size=entry.max_batch_size,
|
|
181
|
+
cache_budget_bytes=budget,
|
|
182
|
+
prefill_chunk=entry.prefill_chunk,
|
|
183
|
+
chat_template=entry.chat_template,
|
|
184
|
+
context_length=entry.extend_context_to,
|
|
185
|
+
context_extension=entry.context_extension,
|
|
186
|
+
)
|
|
187
|
+
|
|
188
|
+
if budget is None:
|
|
189
|
+
generator.scheduler.cache_manager.max_bytes = self._default_budget(generator, entry)
|
|
190
|
+
|
|
191
|
+
if generator.chat_template is not None:
|
|
192
|
+
source = "config" if entry.chat_template else "model"
|
|
193
|
+
template = generator.bound_chat_template()
|
|
194
|
+
else:
|
|
195
|
+
source = "fallback"
|
|
196
|
+
template = generator.bound_chat_template(self.config.runtime.fallback_chat_template)
|
|
197
|
+
log.warning("model '%s' has no chat template; using fallback '%s'", entry.name, template.name)
|
|
198
|
+
|
|
199
|
+
ctx = generator.config.context_length
|
|
200
|
+
|
|
201
|
+
if entry.max_context_tokens:
|
|
202
|
+
ctx = min(ctx, entry.max_context_tokens)
|
|
203
|
+
|
|
204
|
+
max_gen = min(entry.max_generation_tokens or self.config.limits.max_generation_tokens, ctx - 1)
|
|
205
|
+
max_prompt = min(self.config.limits.max_prompt_tokens or ctx, ctx - 1)
|
|
206
|
+
limits = ModelLimits(context_length=ctx, max_generation_tokens=max_gen, max_prompt_tokens=max_prompt)
|
|
207
|
+
|
|
208
|
+
model = ServedModel(entry, generator, template, source, limits, self.device)
|
|
209
|
+
|
|
210
|
+
with self._lock:
|
|
211
|
+
limit_mb = self.config.runtime.memory_limit_mb
|
|
212
|
+
|
|
213
|
+
if limit_mb is not None:
|
|
214
|
+
needed = self._committed_bytes() + model.memory_bytes
|
|
215
|
+
|
|
216
|
+
if needed > limit_mb * (1 << 20):
|
|
217
|
+
raise errors.APIError(
|
|
218
|
+
507,
|
|
219
|
+
f"loading '{entry.name}' needs {model.memory_bytes / (1 << 20):.1f} MiB "
|
|
220
|
+
f"(weights + KV cache budget); memory limit of {limit_mb:.0f} MiB would be exceeded",
|
|
221
|
+
"server_error",
|
|
222
|
+
code="insufficient_memory",
|
|
223
|
+
)
|
|
224
|
+
|
|
225
|
+
generator.start()
|
|
226
|
+
self._models[entry.name] = model
|
|
227
|
+
|
|
228
|
+
log.info("loaded model '%s' in %.2fs (%.1f MiB weights, %.1f MiB kv budget, template %s/%s)",
|
|
229
|
+
entry.name, time.perf_counter() - started, generator.weights_nbytes / (1 << 20),
|
|
230
|
+
(model.cache_budget_bytes or 0) / (1 << 20), template.name, source)
|
|
231
|
+
|
|
232
|
+
return model
|
|
233
|
+
|
|
234
|
+
def unload(self, name, drain_timeout_s=30.0):
|
|
235
|
+
with self._lock:
|
|
236
|
+
model = self._models.get(name)
|
|
237
|
+
|
|
238
|
+
if model is None:
|
|
239
|
+
raise errors.not_found(f"model '{name}' is not loaded", code="model_not_loaded")
|
|
240
|
+
|
|
241
|
+
model.accepting = False
|
|
242
|
+
|
|
243
|
+
drained = model.wait_idle(drain_timeout_s)
|
|
244
|
+
|
|
245
|
+
if not drained:
|
|
246
|
+
sched = model.generator.scheduler
|
|
247
|
+
|
|
248
|
+
for rid in list(sched._active.keys()) + [h.request_id for h in list(sched._waiting)]:
|
|
249
|
+
sched.request_cancel(rid)
|
|
250
|
+
|
|
251
|
+
model.wait_idle(5.0)
|
|
252
|
+
|
|
253
|
+
model.generator.stop()
|
|
254
|
+
|
|
255
|
+
with self._lock:
|
|
256
|
+
self._models.pop(name, None)
|
|
257
|
+
|
|
258
|
+
log.info("unloaded model '%s' (%s)", name, "drained" if drained else "cancelled in-flight requests")
|
|
259
|
+
|
|
260
|
+
return drained
|
|
261
|
+
|
|
262
|
+
def preload(self):
|
|
263
|
+
for entry in self._entries.values():
|
|
264
|
+
if entry.preload:
|
|
265
|
+
self.load(entry.name)
|
|
266
|
+
|
|
267
|
+
def acquire_global(self):
|
|
268
|
+
with self._global_lock:
|
|
269
|
+
if self._global_in_flight >= self.config.limits.max_concurrent_requests:
|
|
270
|
+
raise errors.overloaded()
|
|
271
|
+
|
|
272
|
+
self._global_in_flight += 1
|
|
273
|
+
|
|
274
|
+
def release_global(self):
|
|
275
|
+
with self._global_lock:
|
|
276
|
+
self._global_in_flight -= 1
|
|
277
|
+
|
|
278
|
+
@property
|
|
279
|
+
def global_in_flight(self):
|
|
280
|
+
return self._global_in_flight
|
|
281
|
+
|
|
282
|
+
def shutdown(self, drain_timeout_s):
|
|
283
|
+
for model in self.loaded_models():
|
|
284
|
+
model.accepting = False
|
|
285
|
+
|
|
286
|
+
for model in self.loaded_models():
|
|
287
|
+
self.unload(model.name, drain_timeout_s)
|