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/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)