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
|
@@ -0,0 +1,473 @@
|
|
|
1
|
+
import queue
|
|
2
|
+
import threading
|
|
3
|
+
import time
|
|
4
|
+
import uuid
|
|
5
|
+
from collections import deque
|
|
6
|
+
from dataclasses import dataclass, field
|
|
7
|
+
from typing import List, Optional, Union
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
from src.inference.config import GenerationConfig
|
|
12
|
+
from src.inference.kv_cache import KVCacheManager
|
|
13
|
+
from src.inference.sampling import sample_token
|
|
14
|
+
from src.inference.text import StreamDecoder
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@dataclass
|
|
18
|
+
class GenerationRequest:
|
|
19
|
+
prompt: Union[str, List[int]]
|
|
20
|
+
config: GenerationConfig = field(default_factory=GenerationConfig)
|
|
21
|
+
request_id: str = field(default_factory=lambda: uuid.uuid4().hex)
|
|
22
|
+
timeout_s: Optional[float] = None
|
|
23
|
+
truncate_prompt: bool = False
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass
|
|
27
|
+
class StreamEvent:
|
|
28
|
+
request_id: str
|
|
29
|
+
token_id: Optional[int]
|
|
30
|
+
text: str
|
|
31
|
+
finished: bool = False
|
|
32
|
+
finish_reason: Optional[str] = None
|
|
33
|
+
usage: Optional[dict] = None
|
|
34
|
+
error: Optional[str] = None
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass
|
|
38
|
+
class GenerationResult:
|
|
39
|
+
request_id: str
|
|
40
|
+
text: str
|
|
41
|
+
token_ids: List[int]
|
|
42
|
+
finish_reason: str
|
|
43
|
+
usage: dict
|
|
44
|
+
error: Optional[str] = None
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class RequestHandle:
|
|
48
|
+
|
|
49
|
+
def __init__(self, request, scheduler):
|
|
50
|
+
self.request = request
|
|
51
|
+
self.request_id = request.request_id
|
|
52
|
+
self._scheduler = scheduler
|
|
53
|
+
self._events = queue.Queue()
|
|
54
|
+
self.finished = threading.Event()
|
|
55
|
+
self.submitted_at = time.perf_counter()
|
|
56
|
+
self.admitted_at = None
|
|
57
|
+
self.prefill_s = None
|
|
58
|
+
self.prompt_tokens = None
|
|
59
|
+
self._listener = None
|
|
60
|
+
|
|
61
|
+
def cancel(self):
|
|
62
|
+
self._scheduler.cancel(self.request_id)
|
|
63
|
+
|
|
64
|
+
def request_cancel(self):
|
|
65
|
+
self._scheduler.request_cancel(self.request_id)
|
|
66
|
+
|
|
67
|
+
def set_listener(self, fn):
|
|
68
|
+
self._listener = fn
|
|
69
|
+
|
|
70
|
+
if fn is not None and (not self._events.empty() or self.finished.is_set()):
|
|
71
|
+
fn()
|
|
72
|
+
|
|
73
|
+
def _emit(self, ev):
|
|
74
|
+
self._events.put(ev)
|
|
75
|
+
fn = self._listener
|
|
76
|
+
|
|
77
|
+
if fn is not None:
|
|
78
|
+
try:
|
|
79
|
+
fn()
|
|
80
|
+
except Exception:
|
|
81
|
+
pass
|
|
82
|
+
|
|
83
|
+
def drain(self):
|
|
84
|
+
out = []
|
|
85
|
+
|
|
86
|
+
while True:
|
|
87
|
+
try:
|
|
88
|
+
out.append(self._events.get_nowait())
|
|
89
|
+
except queue.Empty:
|
|
90
|
+
return out
|
|
91
|
+
|
|
92
|
+
def events(self, drive=True, poll_s=0.02):
|
|
93
|
+
try:
|
|
94
|
+
while True:
|
|
95
|
+
try:
|
|
96
|
+
ev = self._events.get_nowait()
|
|
97
|
+
except queue.Empty:
|
|
98
|
+
if self.finished.is_set():
|
|
99
|
+
return
|
|
100
|
+
|
|
101
|
+
if drive:
|
|
102
|
+
progressed = self._scheduler.step()
|
|
103
|
+
|
|
104
|
+
if not progressed and self._events.empty():
|
|
105
|
+
self.finished.wait(poll_s)
|
|
106
|
+
|
|
107
|
+
continue
|
|
108
|
+
|
|
109
|
+
try:
|
|
110
|
+
ev = self._events.get(timeout=poll_s)
|
|
111
|
+
except queue.Empty:
|
|
112
|
+
continue
|
|
113
|
+
|
|
114
|
+
yield ev
|
|
115
|
+
|
|
116
|
+
if ev.finished:
|
|
117
|
+
return
|
|
118
|
+
finally:
|
|
119
|
+
if not self.finished.is_set():
|
|
120
|
+
self.cancel()
|
|
121
|
+
|
|
122
|
+
def result(self, drive=True):
|
|
123
|
+
parts = []
|
|
124
|
+
ids = []
|
|
125
|
+
last = None
|
|
126
|
+
|
|
127
|
+
for ev in self.events(drive=drive):
|
|
128
|
+
parts.append(ev.text)
|
|
129
|
+
if ev.token_id is not None:
|
|
130
|
+
ids.append(ev.token_id)
|
|
131
|
+
last = ev
|
|
132
|
+
|
|
133
|
+
return GenerationResult(
|
|
134
|
+
request_id=self.request_id,
|
|
135
|
+
text="".join(parts),
|
|
136
|
+
token_ids=ids,
|
|
137
|
+
finish_reason=last.finish_reason,
|
|
138
|
+
usage=last.usage,
|
|
139
|
+
error=last.error,
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
class _Sequence:
|
|
144
|
+
|
|
145
|
+
def __init__(self, handle, prompt_ids, cache, decoder, rng, deadline):
|
|
146
|
+
self.handle = handle
|
|
147
|
+
self.request = handle.request
|
|
148
|
+
self.config = handle.request.config
|
|
149
|
+
self.prompt_ids = prompt_ids
|
|
150
|
+
self.cache = cache
|
|
151
|
+
self.decoder = decoder
|
|
152
|
+
self.rng = rng
|
|
153
|
+
self.deadline = deadline
|
|
154
|
+
self.generated = []
|
|
155
|
+
self.seen = set(prompt_ids)
|
|
156
|
+
self.last_token = None
|
|
157
|
+
self.submitted = time.perf_counter()
|
|
158
|
+
self.first_token_at = None
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
class BatchScheduler:
|
|
162
|
+
|
|
163
|
+
def __init__(
|
|
164
|
+
self,
|
|
165
|
+
engine,
|
|
166
|
+
tokenizer,
|
|
167
|
+
max_batch_size=8,
|
|
168
|
+
cache_budget_bytes=None,
|
|
169
|
+
prefill_chunk=256,
|
|
170
|
+
):
|
|
171
|
+
self.engine = engine
|
|
172
|
+
self.tokenizer = tokenizer
|
|
173
|
+
self.max_batch_size = max_batch_size
|
|
174
|
+
self.prefill_chunk = prefill_chunk
|
|
175
|
+
|
|
176
|
+
self.cache_manager = KVCacheManager(
|
|
177
|
+
engine.n_layers, engine.n_heads, engine.head_dim, max_bytes=cache_budget_bytes
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
self._waiting = deque()
|
|
181
|
+
self._waiting_lock = threading.Lock()
|
|
182
|
+
self._active = {}
|
|
183
|
+
self._step_lock = threading.RLock()
|
|
184
|
+
self._runner = None
|
|
185
|
+
self._stop = threading.Event()
|
|
186
|
+
self._cancel_lock = threading.Lock()
|
|
187
|
+
self._cancel_requests = set()
|
|
188
|
+
|
|
189
|
+
self.stats = {
|
|
190
|
+
"requests_submitted": 0,
|
|
191
|
+
"requests_completed": 0,
|
|
192
|
+
"requests_cancelled": 0,
|
|
193
|
+
"requests_errored": 0,
|
|
194
|
+
"requests_timed_out": 0,
|
|
195
|
+
"prompt_tokens": 0,
|
|
196
|
+
"generated_tokens": 0,
|
|
197
|
+
"prefill_seconds": 0.0,
|
|
198
|
+
"decode_seconds": 0.0,
|
|
199
|
+
"decode_steps": 0,
|
|
200
|
+
"max_batch_seen": 0,
|
|
201
|
+
"ttft_seconds_total": 0.0,
|
|
202
|
+
"ttft_count": 0,
|
|
203
|
+
}
|
|
204
|
+
|
|
205
|
+
def submit(self, request):
|
|
206
|
+
handle = RequestHandle(request, self)
|
|
207
|
+
|
|
208
|
+
with self._waiting_lock:
|
|
209
|
+
self._waiting.append(handle)
|
|
210
|
+
self.stats["requests_submitted"] += 1
|
|
211
|
+
|
|
212
|
+
return handle
|
|
213
|
+
|
|
214
|
+
@property
|
|
215
|
+
def num_active(self):
|
|
216
|
+
return len(self._active)
|
|
217
|
+
|
|
218
|
+
@property
|
|
219
|
+
def num_waiting(self):
|
|
220
|
+
return len(self._waiting)
|
|
221
|
+
|
|
222
|
+
def _usage(self, seq_or_none, handle, prompt_len=0, completion=0, ttft=None):
|
|
223
|
+
return {
|
|
224
|
+
"prompt_tokens": prompt_len,
|
|
225
|
+
"completion_tokens": completion,
|
|
226
|
+
"total_tokens": prompt_len + completion,
|
|
227
|
+
"time_to_first_token_s": ttft,
|
|
228
|
+
}
|
|
229
|
+
|
|
230
|
+
def _finalize_unstarted(self, handle, reason, error=None):
|
|
231
|
+
handle._emit(StreamEvent(
|
|
232
|
+
handle.request_id, None, "", True, reason, self._usage(None, handle), error
|
|
233
|
+
))
|
|
234
|
+
handle.finished.set()
|
|
235
|
+
|
|
236
|
+
key = {"cancelled": "requests_cancelled", "timeout": "requests_timed_out"}.get(reason)
|
|
237
|
+
self.stats[key or "requests_errored"] += 1
|
|
238
|
+
|
|
239
|
+
def _finish(self, seq, reason, error=None, token_id=None, text=""):
|
|
240
|
+
tail = seq.decoder.flush()
|
|
241
|
+
ttft = None if seq.first_token_at is None else seq.first_token_at - seq.submitted
|
|
242
|
+
|
|
243
|
+
seq.handle._emit(StreamEvent(
|
|
244
|
+
seq.request.request_id,
|
|
245
|
+
token_id,
|
|
246
|
+
text + tail,
|
|
247
|
+
True,
|
|
248
|
+
reason,
|
|
249
|
+
self._usage(seq, seq.handle, len(seq.prompt_ids), len(seq.generated), ttft),
|
|
250
|
+
error,
|
|
251
|
+
))
|
|
252
|
+
seq.handle.finished.set()
|
|
253
|
+
|
|
254
|
+
self.cache_manager.release(seq.request.request_id)
|
|
255
|
+
self._active.pop(seq.request.request_id, None)
|
|
256
|
+
|
|
257
|
+
key = {"cancelled": "requests_cancelled", "timeout": "requests_timed_out", "error": "requests_errored"}
|
|
258
|
+
self.stats[key.get(reason, "requests_completed")] += 1
|
|
259
|
+
|
|
260
|
+
def cancel(self, request_id):
|
|
261
|
+
with self._step_lock:
|
|
262
|
+
seq = self._active.get(request_id)
|
|
263
|
+
|
|
264
|
+
if seq is not None:
|
|
265
|
+
self._finish(seq, "cancelled")
|
|
266
|
+
return
|
|
267
|
+
|
|
268
|
+
with self._waiting_lock:
|
|
269
|
+
for handle in list(self._waiting):
|
|
270
|
+
if handle.request_id == request_id:
|
|
271
|
+
self._waiting.remove(handle)
|
|
272
|
+
self._finalize_unstarted(handle, "cancelled")
|
|
273
|
+
return
|
|
274
|
+
|
|
275
|
+
def request_cancel(self, request_id):
|
|
276
|
+
with self._cancel_lock:
|
|
277
|
+
self._cancel_requests.add(request_id)
|
|
278
|
+
|
|
279
|
+
def _apply_cancel_requests(self):
|
|
280
|
+
with self._cancel_lock:
|
|
281
|
+
if not self._cancel_requests:
|
|
282
|
+
return
|
|
283
|
+
pending = self._cancel_requests
|
|
284
|
+
self._cancel_requests = set()
|
|
285
|
+
|
|
286
|
+
for request_id in pending:
|
|
287
|
+
self.cancel(request_id)
|
|
288
|
+
|
|
289
|
+
def _prepare_prompt(self, request):
|
|
290
|
+
ctx = self.engine.context_length
|
|
291
|
+
cfg = request.config
|
|
292
|
+
|
|
293
|
+
if isinstance(request.prompt, str):
|
|
294
|
+
ids = self.tokenizer.encode(request.prompt)
|
|
295
|
+
else:
|
|
296
|
+
ids = [int(t) for t in request.prompt]
|
|
297
|
+
|
|
298
|
+
if not ids:
|
|
299
|
+
if self.tokenizer.eos_id is None:
|
|
300
|
+
raise ValueError("empty prompt and tokenizer has no EOS token to start from")
|
|
301
|
+
ids = [self.tokenizer.eos_id]
|
|
302
|
+
|
|
303
|
+
if max(ids) >= self.engine.config.vocab_size or min(ids) < 0:
|
|
304
|
+
raise ValueError("prompt contains token ids outside the model vocabulary")
|
|
305
|
+
|
|
306
|
+
if len(ids) >= ctx:
|
|
307
|
+
if not request.truncate_prompt:
|
|
308
|
+
raise ValueError(f"prompt of {len(ids)} tokens does not fit the context length of {ctx}")
|
|
309
|
+
|
|
310
|
+
keep = ctx - min(cfg.max_new_tokens, ctx // 2)
|
|
311
|
+
ids = ids[-keep:]
|
|
312
|
+
|
|
313
|
+
return ids
|
|
314
|
+
|
|
315
|
+
def _admit(self):
|
|
316
|
+
while len(self._active) < self.max_batch_size:
|
|
317
|
+
with self._waiting_lock:
|
|
318
|
+
if not self._waiting:
|
|
319
|
+
return
|
|
320
|
+
handle = self._waiting[0]
|
|
321
|
+
|
|
322
|
+
req = handle.request
|
|
323
|
+
|
|
324
|
+
try:
|
|
325
|
+
ids = self._prepare_prompt(req)
|
|
326
|
+
except Exception as exc:
|
|
327
|
+
with self._waiting_lock:
|
|
328
|
+
self._waiting.popleft()
|
|
329
|
+
self._finalize_unstarted(handle, "error", str(exc))
|
|
330
|
+
continue
|
|
331
|
+
|
|
332
|
+
capacity = min(len(ids) + req.config.max_new_tokens, self.engine.context_length)
|
|
333
|
+
|
|
334
|
+
if not self.cache_manager.fits_at_all(capacity):
|
|
335
|
+
with self._waiting_lock:
|
|
336
|
+
self._waiting.popleft()
|
|
337
|
+
self._finalize_unstarted(handle, "error", "request needs more kv cache than the configured budget")
|
|
338
|
+
continue
|
|
339
|
+
|
|
340
|
+
if not self.cache_manager.can_allocate(capacity):
|
|
341
|
+
return
|
|
342
|
+
|
|
343
|
+
with self._waiting_lock:
|
|
344
|
+
self._waiting.popleft()
|
|
345
|
+
|
|
346
|
+
cache = self.cache_manager.allocate(req.request_id, capacity)
|
|
347
|
+
handle.admitted_at = time.perf_counter()
|
|
348
|
+
handle.prompt_tokens = len(ids)
|
|
349
|
+
deadline = handle.submitted_at + req.timeout_s if req.timeout_s else None
|
|
350
|
+
seq = _Sequence(
|
|
351
|
+
handle,
|
|
352
|
+
ids,
|
|
353
|
+
cache,
|
|
354
|
+
StreamDecoder(self.tokenizer, req.config.stop_strings),
|
|
355
|
+
np.random.default_rng(req.config.seed),
|
|
356
|
+
deadline,
|
|
357
|
+
)
|
|
358
|
+
self._active[req.request_id] = seq
|
|
359
|
+
self.stats["prompt_tokens"] += len(ids)
|
|
360
|
+
|
|
361
|
+
started = time.perf_counter()
|
|
362
|
+
|
|
363
|
+
try:
|
|
364
|
+
logits = self.engine.prefill(ids, cache, self.prefill_chunk)
|
|
365
|
+
except Exception as exc:
|
|
366
|
+
self._finish(seq, "error", str(exc))
|
|
367
|
+
continue
|
|
368
|
+
|
|
369
|
+
handle.prefill_s = time.perf_counter() - started
|
|
370
|
+
self.stats["prefill_seconds"] += handle.prefill_s
|
|
371
|
+
self._consume(seq, logits)
|
|
372
|
+
|
|
373
|
+
def _consume(self, seq, logits):
|
|
374
|
+
cfg = seq.config
|
|
375
|
+
token = sample_token(logits, cfg, seq.rng, seq.seen)
|
|
376
|
+
|
|
377
|
+
if seq.first_token_at is None:
|
|
378
|
+
seq.first_token_at = time.perf_counter()
|
|
379
|
+
self.stats["ttft_seconds_total"] += seq.first_token_at - seq.submitted
|
|
380
|
+
self.stats["ttft_count"] += 1
|
|
381
|
+
|
|
382
|
+
seq.generated.append(token)
|
|
383
|
+
seq.seen.add(token)
|
|
384
|
+
seq.last_token = token
|
|
385
|
+
self.stats["generated_tokens"] += 1
|
|
386
|
+
|
|
387
|
+
is_stop_token = token in cfg.stop_token_ids or (cfg.stop_on_eos and token == self.tokenizer.eos_id)
|
|
388
|
+
|
|
389
|
+
delta = "" if is_stop_token else seq.decoder.push(token)
|
|
390
|
+
|
|
391
|
+
if is_stop_token or seq.decoder.stopped:
|
|
392
|
+
self._finish(seq, "stop", token_id=token, text=delta)
|
|
393
|
+
return
|
|
394
|
+
|
|
395
|
+
if len(seq.generated) >= cfg.max_new_tokens or seq.cache.length + 1 >= seq.cache.capacity:
|
|
396
|
+
self._finish(seq, "length", token_id=token, text=delta)
|
|
397
|
+
return
|
|
398
|
+
|
|
399
|
+
seq.handle._emit(StreamEvent(seq.request.request_id, token, delta))
|
|
400
|
+
|
|
401
|
+
def step(self):
|
|
402
|
+
with self._step_lock:
|
|
403
|
+
self._apply_cancel_requests()
|
|
404
|
+
now = time.perf_counter()
|
|
405
|
+
|
|
406
|
+
with self._waiting_lock:
|
|
407
|
+
expired = [
|
|
408
|
+
h for h in self._waiting
|
|
409
|
+
if h.request.timeout_s and now - h.submitted_at > h.request.timeout_s
|
|
410
|
+
]
|
|
411
|
+
for h in expired:
|
|
412
|
+
self._waiting.remove(h)
|
|
413
|
+
|
|
414
|
+
for h in expired:
|
|
415
|
+
self._finalize_unstarted(h, "timeout")
|
|
416
|
+
|
|
417
|
+
for seq in list(self._active.values()):
|
|
418
|
+
if seq.deadline is not None and now > seq.deadline:
|
|
419
|
+
self._finish(seq, "timeout")
|
|
420
|
+
|
|
421
|
+
self._admit()
|
|
422
|
+
|
|
423
|
+
seqs = list(self._active.values())
|
|
424
|
+
|
|
425
|
+
if not seqs:
|
|
426
|
+
return bool(self._waiting)
|
|
427
|
+
|
|
428
|
+
self.stats["max_batch_seen"] = max(self.stats["max_batch_seen"], len(seqs))
|
|
429
|
+
started = time.perf_counter()
|
|
430
|
+
|
|
431
|
+
try:
|
|
432
|
+
logits = self.engine.decode_batch([s.last_token for s in seqs], [s.cache for s in seqs])
|
|
433
|
+
except Exception as exc:
|
|
434
|
+
for s in seqs:
|
|
435
|
+
self._finish(s, "error", str(exc))
|
|
436
|
+
return True
|
|
437
|
+
|
|
438
|
+
self.stats["decode_seconds"] += time.perf_counter() - started
|
|
439
|
+
self.stats["decode_steps"] += 1
|
|
440
|
+
|
|
441
|
+
for s, row in zip(seqs, logits):
|
|
442
|
+
self._consume(s, row)
|
|
443
|
+
|
|
444
|
+
return True
|
|
445
|
+
|
|
446
|
+
def run_until_idle(self):
|
|
447
|
+
while self.step():
|
|
448
|
+
pass
|
|
449
|
+
|
|
450
|
+
def start(self, idle_sleep_s=0.002):
|
|
451
|
+
if self._runner is not None:
|
|
452
|
+
return
|
|
453
|
+
|
|
454
|
+
self._stop.clear()
|
|
455
|
+
|
|
456
|
+
def loop():
|
|
457
|
+
while not self._stop.is_set():
|
|
458
|
+
if not self.step():
|
|
459
|
+
time.sleep(idle_sleep_s)
|
|
460
|
+
|
|
461
|
+
self._runner = threading.Thread(target=loop, daemon=True)
|
|
462
|
+
self._runner.start()
|
|
463
|
+
|
|
464
|
+
def stop(self):
|
|
465
|
+
self._stop.set()
|
|
466
|
+
|
|
467
|
+
if self._runner is not None:
|
|
468
|
+
self._runner.join(timeout=2.0)
|
|
469
|
+
self._runner = None
|
|
470
|
+
|
|
471
|
+
@property
|
|
472
|
+
def background(self):
|
|
473
|
+
return self._runner is not None
|
src/inference/text.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
import codecs
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class StreamDecoder:
|
|
5
|
+
|
|
6
|
+
def __init__(self, tokenizer, stop_strings=None):
|
|
7
|
+
self.tokenizer = tokenizer
|
|
8
|
+
self.stop_strings = [s for s in (stop_strings or []) if s]
|
|
9
|
+
self.holdback = max((len(s) for s in self.stop_strings), default=1) - 1
|
|
10
|
+
self.ids = []
|
|
11
|
+
self.emitted = 0
|
|
12
|
+
self.text = ""
|
|
13
|
+
self.stopped = False
|
|
14
|
+
|
|
15
|
+
self._incremental = hasattr(tokenizer, "token_bytes")
|
|
16
|
+
self._utf8 = codecs.getincrementaldecoder("utf-8")(errors="replace") if self._incremental else None
|
|
17
|
+
|
|
18
|
+
def _advance(self, token_id):
|
|
19
|
+
if self._incremental:
|
|
20
|
+
return self.text + self._utf8.decode(self.tokenizer.token_bytes(token_id))
|
|
21
|
+
|
|
22
|
+
self.ids.append(token_id)
|
|
23
|
+
return self.tokenizer.decode(self.ids)
|
|
24
|
+
|
|
25
|
+
def push(self, token_id):
|
|
26
|
+
if self.stopped:
|
|
27
|
+
return ""
|
|
28
|
+
|
|
29
|
+
full = self._advance(token_id)
|
|
30
|
+
|
|
31
|
+
if not full.startswith(self.text):
|
|
32
|
+
common = 0
|
|
33
|
+
|
|
34
|
+
for a, b in zip(full, self.text):
|
|
35
|
+
if a != b:
|
|
36
|
+
break
|
|
37
|
+
common += 1
|
|
38
|
+
|
|
39
|
+
self.emitted = min(self.emitted, common)
|
|
40
|
+
|
|
41
|
+
self.text = full
|
|
42
|
+
|
|
43
|
+
for s in self.stop_strings:
|
|
44
|
+
at = full.find(s, max(self.emitted - len(s) + 1, 0))
|
|
45
|
+
|
|
46
|
+
if at != -1:
|
|
47
|
+
self.text = full[:at]
|
|
48
|
+
self.stopped = True
|
|
49
|
+
break
|
|
50
|
+
|
|
51
|
+
safe = len(self.text) if self.stopped else max(len(self.text) - self.holdback, self.emitted)
|
|
52
|
+
delta = self.text[self.emitted:safe]
|
|
53
|
+
self.emitted = safe
|
|
54
|
+
|
|
55
|
+
return delta
|
|
56
|
+
|
|
57
|
+
def flush(self):
|
|
58
|
+
if self._incremental and not self.stopped:
|
|
59
|
+
tail = self._utf8.decode(b"", final=True)
|
|
60
|
+
|
|
61
|
+
if tail:
|
|
62
|
+
self.text += tail
|
|
63
|
+
|
|
64
|
+
delta = self.text[self.emitted:]
|
|
65
|
+
self.emitted = len(self.text)
|
|
66
|
+
|
|
67
|
+
return delta
|
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.initializers.Initializer import Initializer
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
# class GlorotUniform:
|
|
7
|
+
# def __call__(self, shape):
|
|
8
|
+
# fan_in, fan_out = shape
|
|
9
|
+
# limit = np.sqrt(6 / (fan_in + fan_out))
|
|
10
|
+
# return np.random.uniform(-limit, limit, shape)
|
|
11
|
+
|
|
12
|
+
# class GlorotUniform(Initializer):
|
|
13
|
+
# def __call__(self, shape):
|
|
14
|
+
# fan_in, fan_out = shape
|
|
15
|
+
# limit = np.sqrt(6.0 / (fan_in + fan_out))
|
|
16
|
+
# return np.random.uniform(
|
|
17
|
+
# -limit,
|
|
18
|
+
# limit,
|
|
19
|
+
# shape,
|
|
20
|
+
# ).astype(np.float32)
|
|
21
|
+
|
|
22
|
+
class GlorotUniform:
|
|
23
|
+
def __call__(self, shape):
|
|
24
|
+
fan_in, fan_out = shape
|
|
25
|
+
limit = np.sqrt(6.0 / (fan_in + fan_out))
|
|
26
|
+
return np.random.uniform(-limit, limit, shape).astype(np.float32)
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.initializers.Initializer import Initializer
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class HeUniform(Initializer):
|
|
7
|
+
def __call__(self, shape):
|
|
8
|
+
fan_in = shape[0]
|
|
9
|
+
limit = np.sqrt(6.0 / fan_in)
|
|
10
|
+
return np.random.uniform(
|
|
11
|
+
-limit,
|
|
12
|
+
limit,
|
|
13
|
+
shape,
|
|
14
|
+
).astype(np.float32)
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
from src.initializers.Initializer import Initializer
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class LecunUniform(Initializer):
|
|
7
|
+
def __call__(self, shape):
|
|
8
|
+
fan_in = shape[0]
|
|
9
|
+
limit = np.sqrt(3.0 / fan_in)
|
|
10
|
+
return np.random.uniform(
|
|
11
|
+
-limit,
|
|
12
|
+
limit,
|
|
13
|
+
shape,
|
|
14
|
+
).astype(np.float32)
|
src/initializers/Ones.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
class RandomNormal:
|
|
4
|
+
|
|
5
|
+
def __init__(self, mean=0.0, stddev=0.05):
|
|
6
|
+
self.mean = mean
|
|
7
|
+
self.stddev = stddev
|
|
8
|
+
|
|
9
|
+
def __call__(self, shape):
|
|
10
|
+
return np.random.normal(
|
|
11
|
+
self.mean,
|
|
12
|
+
self.stddev,
|
|
13
|
+
shape,
|
|
14
|
+
).astype(np.float32)
|