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
@@ -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,9 @@
1
+ import numpy as np
2
+
3
+ class Constant:
4
+
5
+ def __init__(self, value):
6
+ self.value = value
7
+
8
+ def __call__(self, shape):
9
+ return np.full(shape, self.value, dtype=np.float32)
@@ -0,0 +1,15 @@
1
+ import numpy as np
2
+
3
+ class GlorotNormal:
4
+
5
+ def __call__(self, shape):
6
+
7
+ fan_in, fan_out = shape
8
+
9
+ std = np.sqrt(2 / (fan_in + fan_out))
10
+
11
+ return np.random.normal(
12
+ 0,
13
+ std,
14
+ shape,
15
+ ).astype(np.float32)
@@ -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,15 @@
1
+ import numpy as np
2
+
3
+ class HeNormal:
4
+
5
+ def __call__(self, shape):
6
+
7
+ fan_in = shape[0]
8
+
9
+ std = np.sqrt(2 / fan_in)
10
+
11
+ return np.random.normal(
12
+ 0,
13
+ std,
14
+ shape,
15
+ ).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,4 @@
1
+
2
+ class Initializer:
3
+ def __call__(self, shape):
4
+ raise NotImplementedError
@@ -0,0 +1,16 @@
1
+ import numpy as np
2
+
3
+
4
+ class LecunNormal:
5
+
6
+ def __call__(self, shape):
7
+
8
+ fan_in = shape[0]
9
+
10
+ std = np.sqrt(1 / fan_in)
11
+
12
+ return np.random.normal(
13
+ 0,
14
+ std,
15
+ shape,
16
+ ).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)
@@ -0,0 +1,6 @@
1
+ import numpy as np
2
+
3
+ class Ones:
4
+
5
+ def __call__(self, shape):
6
+ return np.ones(shape, dtype=np.float32)
@@ -0,0 +1,14 @@
1
+ import numpy as np
2
+
3
+
4
+ class Orthogonal:
5
+
6
+ def __call__(self, shape):
7
+
8
+ rows, cols = shape
9
+
10
+ a = np.random.randn(rows, cols)
11
+
12
+ q, r = np.linalg.qr(a)
13
+
14
+ return q.astype(np.float32)
@@ -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)