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,384 @@
1
+ import copy
2
+ import json
3
+ import os
4
+ from dataclasses import dataclass, field
5
+ from typing import List
6
+
7
+ ROLES = ("system", "user", "assistant")
8
+
9
+
10
+ class ChatTemplateError(ValueError):
11
+ pass
12
+
13
+
14
+ class PromptTooLong(ValueError):
15
+
16
+ def __init__(self, needed, budget):
17
+ super().__init__(f"prompt needs {needed} tokens but only {budget} fit in the context window")
18
+ self.needed = needed
19
+ self.budget = budget
20
+
21
+
22
+ BUILTIN_TEMPLATES = {
23
+ "plain": {
24
+ "name": "plain",
25
+ "add_bos": False,
26
+ "roles": {
27
+ "system": {"prefix": ["System: "], "suffix": ["\n"]},
28
+ "user": {"prefix": ["User: "], "suffix": ["\n"]},
29
+ "assistant": {"prefix": ["Assistant: "], "suffix": ["\n"]},
30
+ },
31
+ "generation_prompt": ["Assistant:"],
32
+ "stop_tokens": [],
33
+ "stop_strings": ["\nUser:", "\nSystem:"],
34
+ "output_lstrip": True,
35
+ },
36
+ "ptf-chat": {
37
+ "name": "ptf-chat",
38
+ "add_bos": False,
39
+ "roles": {
40
+ "system": {"prefix": [{"special": "<|system|>"}, "\n"], "suffix": [{"special": "<|end|>"}, "\n"]},
41
+ "user": {"prefix": [{"special": "<|user|>"}, "\n"], "suffix": [{"special": "<|end|>"}, "\n"]},
42
+ "assistant": {"prefix": [{"special": "<|assistant|>"}, "\n"], "suffix": [{"special": "<|end|>"}, "\n"]},
43
+ },
44
+ "generation_prompt": [{"special": "<|assistant|>"}, "\n"],
45
+ "stop_tokens": ["<|end|>"],
46
+ "stop_strings": [],
47
+ "output_lstrip": False,
48
+ },
49
+ }
50
+
51
+
52
+ def _normalize_segments(segments, where):
53
+ if segments is None:
54
+ return []
55
+
56
+ if isinstance(segments, (str, dict)):
57
+ segments = [segments]
58
+
59
+ out = []
60
+
61
+ for seg in segments:
62
+ if isinstance(seg, str):
63
+ out.append(("text", seg))
64
+ elif isinstance(seg, dict) and set(seg) == {"text"} and isinstance(seg["text"], str):
65
+ out.append(("text", seg["text"]))
66
+ elif isinstance(seg, dict) and set(seg) == {"special"} and isinstance(seg["special"], str) and seg["special"]:
67
+ out.append(("special", seg["special"]))
68
+ else:
69
+ raise ChatTemplateError(f"invalid segment in {where}: {seg!r}")
70
+
71
+ return out
72
+
73
+
74
+ def _segments_to_json(segments):
75
+ return [text if kind == "text" else {"special": text} for kind, text in segments]
76
+
77
+
78
+ @dataclass
79
+ class ChatTemplate:
80
+ name: str
81
+ roles: dict
82
+ generation_prompt: list
83
+ add_bos: bool = False
84
+ stop_tokens: List[str] = field(default_factory=list)
85
+ stop_strings: List[str] = field(default_factory=list)
86
+ output_lstrip: bool = False
87
+ default_system: str = None
88
+
89
+ @classmethod
90
+ def from_dict(cls, data):
91
+ if not isinstance(data, dict):
92
+ raise ChatTemplateError("chat template must be a JSON object")
93
+
94
+ name = data.get("name") or "custom"
95
+ roles_in = data.get("roles")
96
+
97
+ if not isinstance(roles_in, dict) or not roles_in:
98
+ raise ChatTemplateError("chat template needs a 'roles' object")
99
+
100
+ roles = {}
101
+
102
+ for role, spec in roles_in.items():
103
+ if role not in ROLES:
104
+ raise ChatTemplateError(f"unsupported role '{role}' in chat template")
105
+
106
+ if not isinstance(spec, dict):
107
+ raise ChatTemplateError(f"role '{role}' must be an object with prefix/suffix")
108
+
109
+ roles[role] = {
110
+ "prefix": _normalize_segments(spec.get("prefix"), f"{role}.prefix"),
111
+ "suffix": _normalize_segments(spec.get("suffix"), f"{role}.suffix"),
112
+ }
113
+
114
+ for required in ("user", "assistant"):
115
+ if required not in roles:
116
+ raise ChatTemplateError(f"chat template must define the '{required}' role")
117
+
118
+ stop_tokens = data.get("stop_tokens", [])
119
+ stop_strings = data.get("stop_strings", [])
120
+
121
+ if not all(isinstance(s, str) and s for s in stop_tokens):
122
+ raise ChatTemplateError("stop_tokens must be non-empty strings")
123
+
124
+ if not all(isinstance(s, str) and s for s in stop_strings):
125
+ raise ChatTemplateError("stop_strings must be non-empty strings")
126
+
127
+ default_system = data.get("default_system")
128
+
129
+ if default_system is not None and not isinstance(default_system, str):
130
+ raise ChatTemplateError("default_system must be a string")
131
+
132
+ return cls(
133
+ name=name,
134
+ roles=roles,
135
+ generation_prompt=_normalize_segments(data.get("generation_prompt"), "generation_prompt"),
136
+ add_bos=bool(data.get("add_bos", False)),
137
+ stop_tokens=list(stop_tokens),
138
+ stop_strings=list(stop_strings),
139
+ output_lstrip=bool(data.get("output_lstrip", False)),
140
+ default_system=default_system,
141
+ )
142
+
143
+ def to_dict(self):
144
+ return {
145
+ "name": self.name,
146
+ "add_bos": self.add_bos,
147
+ "roles": {
148
+ role: {"prefix": _segments_to_json(s["prefix"]), "suffix": _segments_to_json(s["suffix"])}
149
+ for role, s in self.roles.items()
150
+ },
151
+ "generation_prompt": _segments_to_json(self.generation_prompt),
152
+ "stop_tokens": list(self.stop_tokens),
153
+ "stop_strings": list(self.stop_strings),
154
+ "output_lstrip": self.output_lstrip,
155
+ "default_system": self.default_system,
156
+ }
157
+
158
+ @classmethod
159
+ def builtin(cls, name):
160
+ if name not in BUILTIN_TEMPLATES:
161
+ raise ChatTemplateError(f"unknown built-in chat template '{name}'; choose one of {sorted(BUILTIN_TEMPLATES)}")
162
+
163
+ return cls.from_dict(copy.deepcopy(BUILTIN_TEMPLATES[name]))
164
+
165
+ @classmethod
166
+ def resolve(cls, spec):
167
+ if spec is None:
168
+ return None
169
+
170
+ if isinstance(spec, ChatTemplate):
171
+ return spec
172
+
173
+ if isinstance(spec, dict):
174
+ return cls.from_dict(spec)
175
+
176
+ if isinstance(spec, str):
177
+ if spec in BUILTIN_TEMPLATES:
178
+ return cls.builtin(spec)
179
+
180
+ if os.path.isfile(spec):
181
+ with open(spec) as f:
182
+ return cls.from_dict(json.load(f))
183
+
184
+ raise ChatTemplateError(f"chat template '{spec}' is neither a built-in name nor a readable JSON file")
185
+
186
+ def special_names(self):
187
+ names = set(self.stop_tokens)
188
+
189
+ for spec in self.roles.values():
190
+ for kind, text in spec["prefix"] + spec["suffix"]:
191
+ if kind == "special":
192
+ names.add(text)
193
+
194
+ for kind, text in self.generation_prompt:
195
+ if kind == "special":
196
+ names.add(text)
197
+
198
+ return names
199
+
200
+ def bind(self, tokenizer):
201
+ return BoundChatTemplate(self, tokenizer)
202
+
203
+
204
+ @dataclass
205
+ class RenderedPrompt:
206
+ token_ids: List[int]
207
+ messages_used: int
208
+ messages_dropped: int
209
+
210
+
211
+ class BoundChatTemplate:
212
+
213
+ def __init__(self, template, tokenizer):
214
+ self.template = template
215
+ self.tokenizer = tokenizer
216
+ specials = tokenizer.special_tokens or {}
217
+
218
+ missing = sorted(n for n in template.special_names() if n not in specials)
219
+
220
+ if missing:
221
+ raise ChatTemplateError(
222
+ f"chat template '{template.name}' uses special tokens the tokenizer does not have: {missing}. "
223
+ "Train the tokenizer with them (tokenize --special-token ...) or choose another template."
224
+ )
225
+
226
+ if template.add_bos and tokenizer.bos_id is None:
227
+ raise ChatTemplateError("chat template requires a BOS token but the tokenizer has none")
228
+
229
+ self._specials = dict(specials)
230
+ self._control_ids = frozenset(
231
+ i for name, i in specials.items()
232
+ if name != getattr(tokenizer, "unk_token", None)
233
+ )
234
+
235
+ self.stop_token_ids = [specials[n] for n in template.stop_tokens]
236
+ self.stop_strings = list(template.stop_strings)
237
+ self.output_lstrip = template.output_lstrip
238
+ self._gen_prompt_ids = self._encode_pieces(template.generation_prompt, None)
239
+
240
+ @property
241
+ def name(self):
242
+ return self.template.name
243
+
244
+ def _encode_text(self, text):
245
+ if not text:
246
+ return []
247
+
248
+ ids = self.tokenizer.encode(text)
249
+
250
+ return [i for i in ids if i not in self._control_ids]
251
+
252
+ def _encode_pieces(self, segments, content):
253
+ ids = []
254
+ run = []
255
+
256
+ def flush():
257
+ if run:
258
+ ids.extend(self._encode_text("".join(run)))
259
+ run.clear()
260
+
261
+ for kind, text in segments:
262
+ if kind == "special":
263
+ flush()
264
+ ids.append(self._specials[text])
265
+ elif kind == "content":
266
+ run.append(content)
267
+ else:
268
+ run.append(text)
269
+
270
+ flush()
271
+
272
+ return ids
273
+
274
+ def encode_message(self, role, content):
275
+ spec = self.template.roles.get(role)
276
+
277
+ if spec is None:
278
+ raise ChatTemplateError(f"this model's chat template does not support the '{role}' role")
279
+
280
+ return self._encode_pieces(spec["prefix"] + [("content", None)] + spec["suffix"], content)
281
+
282
+ def _continuation_segments(self):
283
+ prefix = list(self.template.roles.get("assistant", {}).get("prefix", []))
284
+ gen = list(self.template.generation_prompt)
285
+
286
+ if "assistant" not in self.template.roles:
287
+ raise ChatTemplateError("chat template has no assistant role to train on")
288
+
289
+ rest = list(prefix)
290
+
291
+ for i, (kind, text) in enumerate(gen):
292
+ if not rest:
293
+ raise ChatTemplateError("generation prompt is longer than the assistant prefix")
294
+
295
+ head_kind, head_text = rest[0]
296
+
297
+ if kind == "special" or head_kind == "special":
298
+ if (kind, text) != (head_kind, head_text):
299
+ raise ChatTemplateError("generation prompt must be a prefix of the assistant prefix to train")
300
+ rest.pop(0)
301
+ elif head_text == text:
302
+ rest.pop(0)
303
+ elif i == len(gen) - 1 and head_text.startswith(text):
304
+ rest[0] = (head_kind, head_text[len(text):])
305
+ else:
306
+ raise ChatTemplateError("generation prompt must be a prefix of the assistant prefix to train")
307
+
308
+ return rest
309
+
310
+ def training_tokens(self, messages, eos_after_reply=True):
311
+ messages = list(messages)
312
+
313
+ if self.template.default_system and not any(m["role"] == "system" for m in messages):
314
+ messages.insert(0, {"role": "system", "content": self.template.default_system})
315
+
316
+ continuation = self._continuation_segments()
317
+ suffix = list(self.template.roles["assistant"]["suffix"])
318
+ ids = [self.tokenizer.bos_id] if self.template.add_bos else []
319
+ learn = [0] * len(ids)
320
+
321
+ for m in messages:
322
+ if m["role"] == "assistant":
323
+ ids.extend(self._gen_prompt_ids)
324
+ learn.extend([0] * len(self._gen_prompt_ids))
325
+ reply = self._encode_pieces(continuation + [("content", None)] + suffix, m["content"])
326
+ ids.extend(reply)
327
+ learn.extend([1] * len(reply))
328
+ else:
329
+ encoded = self.encode_message(m["role"], m["content"])
330
+ ids.extend(encoded)
331
+ learn.extend([0] * len(encoded))
332
+
333
+ if eos_after_reply and messages and messages[-1]["role"] == "assistant" and self.tokenizer.eos_id is not None:
334
+ ids.append(self.tokenizer.eos_id)
335
+ learn.append(1)
336
+
337
+ return ids, learn
338
+
339
+ def render(self, messages, max_prompt_tokens=None, truncate=True):
340
+ messages = list(messages)
341
+
342
+ if self.template.default_system and not any(m["role"] == "system" for m in messages):
343
+ messages.insert(0, {"role": "system", "content": self.template.default_system})
344
+
345
+ if not messages:
346
+ raise ChatTemplateError("messages must not be empty")
347
+
348
+ encoded = [self.encode_message(m["role"], m["content"]) for m in messages]
349
+ head = [self.tokenizer.bos_id] if self.template.add_bos else []
350
+ tail = self._gen_prompt_ids
351
+
352
+ pinned = 0
353
+
354
+ while pinned < len(messages) and messages[pinned]["role"] == "system":
355
+ pinned += 1
356
+
357
+ fixed = len(head) + len(tail) + sum(len(e) for e in encoded[:pinned])
358
+ body = encoded[pinned:]
359
+ total = fixed + sum(len(e) for e in body)
360
+ dropped = 0
361
+
362
+ if max_prompt_tokens is not None and total > max_prompt_tokens:
363
+ if not truncate:
364
+ raise PromptTooLong(total, max_prompt_tokens)
365
+
366
+ while body and total > max_prompt_tokens and len(body) > 1:
367
+ total -= len(body[0])
368
+ body = body[1:]
369
+ dropped += 1
370
+
371
+ if total > max_prompt_tokens:
372
+ raise PromptTooLong(total, max_prompt_tokens)
373
+
374
+ ids = list(head)
375
+
376
+ for e in encoded[:pinned]:
377
+ ids.extend(e)
378
+
379
+ for e in body:
380
+ ids.extend(e)
381
+
382
+ ids.extend(tail)
383
+
384
+ return RenderedPrompt(ids, len(messages) - dropped, dropped)
@@ -0,0 +1,48 @@
1
+ from dataclasses import asdict, dataclass, field, replace
2
+ from typing import List, Optional
3
+
4
+
5
+ @dataclass
6
+ class GenerationConfig:
7
+ max_new_tokens: int = 128
8
+ temperature: float = 1.0
9
+ top_k: int = 0
10
+ top_p: float = 1.0
11
+ repetition_penalty: float = 1.0
12
+ do_sample: bool = True
13
+ seed: Optional[int] = None
14
+ stop_token_ids: List[int] = field(default_factory=list)
15
+ stop_strings: List[str] = field(default_factory=list)
16
+ stop_on_eos: bool = True
17
+
18
+ def __post_init__(self):
19
+ if self.max_new_tokens < 1:
20
+ raise ValueError("max_new_tokens must be at least 1")
21
+
22
+ if self.temperature < 0:
23
+ raise ValueError("temperature must be >= 0")
24
+
25
+ if self.top_k < 0:
26
+ raise ValueError("top_k must be >= 0")
27
+
28
+ if not 0 < self.top_p <= 1:
29
+ raise ValueError("top_p must be in (0, 1]")
30
+
31
+ if self.repetition_penalty <= 0:
32
+ raise ValueError("repetition_penalty must be > 0")
33
+
34
+ @property
35
+ def greedy(self):
36
+ return (not self.do_sample) or self.temperature == 0
37
+
38
+ def updated(self, **overrides):
39
+ overrides = {k: v for k, v in overrides.items() if v is not None}
40
+ return replace(self, **overrides)
41
+
42
+ def to_dict(self):
43
+ return asdict(self)
44
+
45
+ @classmethod
46
+ def from_dict(cls, data):
47
+ known = set(cls.__dataclass_fields__.keys())
48
+ return cls(**{k: v for k, v in data.items() if k in known})
@@ -0,0 +1,241 @@
1
+ import numpy as np
2
+
3
+ from src.inference.kv_cache import KVCache
4
+
5
+
6
+ def _layer_norm(x, gamma, beta, eps):
7
+ mean = x.mean(axis=-1, keepdims=True)
8
+ var = ((x - mean) ** 2).mean(axis=-1, keepdims=True)
9
+ return gamma * ((x - mean) / np.sqrt(var + eps)) + beta
10
+
11
+
12
+ def _softmax(x):
13
+ x = x - x.max(axis=-1, keepdims=True)
14
+ np.exp(x, out=x)
15
+ x /= x.sum(axis=-1, keepdims=True)
16
+ return x
17
+
18
+
19
+ def _gelu(x):
20
+ return 0.5 * x * (1.0 + np.tanh(0.7978845608028654 * (x + 0.044715 * (x * x * x))))
21
+
22
+
23
+ def _sigmoid(x):
24
+ return 1.0 / (1.0 + np.exp(-x))
25
+
26
+
27
+ ACTIVATIONS = {
28
+ "gelu": _gelu,
29
+ "relu": lambda x: np.maximum(0, x),
30
+ "tanh": np.tanh,
31
+ "sigmoid": _sigmoid,
32
+ }
33
+
34
+
35
+ class _Layer:
36
+
37
+ def __init__(self, w, prefix):
38
+ g = lambda name: np.ascontiguousarray(w[f"{prefix}.{name}"], dtype=np.float32)
39
+
40
+ self.ln1_g = g("norm1.gamma")
41
+ self.ln1_b = g("norm1.beta")
42
+ self.wqkv = np.concatenate([g("attn.Wq"), g("attn.Wk"), g("attn.Wv")], axis=1)
43
+ self.wo = g("attn.Wo")
44
+ self.ln2_g = g("norm2.gamma")
45
+ self.ln2_b = g("norm2.beta")
46
+ self.fc1_k = g("fc1.kernel")
47
+ self.fc1_b = g("fc1.bias")
48
+ self.fc2_k = g("fc2.kernel")
49
+ self.fc2_b = g("fc2.bias")
50
+
51
+
52
+ class InferenceModel:
53
+
54
+ def __init__(self, config, weights):
55
+ if config.activation not in ACTIVATIONS:
56
+ raise ValueError(f"unsupported activation '{config.activation}'")
57
+
58
+ self.config = config
59
+ self.n_layers = config.n_layers
60
+ self.n_heads = config.n_heads
61
+ self.d_model = config.d_model
62
+ self.head_dim = config.d_model // config.n_heads
63
+ self.context_length = config.context_length
64
+ self.eps = config.norm_eps
65
+ self.scale = np.float32(1.0 / np.sqrt(self.head_dim))
66
+ self.act = ACTIVATIONS[config.activation]
67
+
68
+ f32 = lambda name: np.ascontiguousarray(weights[name], dtype=np.float32)
69
+
70
+ self.tok_emb = f32("token_emb.embedding")
71
+ self.rope = None
72
+ self.pos_emb = None
73
+
74
+ if getattr(config, "position_encoding", "learned") == "rope":
75
+ self.rope = config.rotary_tables()
76
+
77
+ if "pos_emb.embedding" in weights:
78
+ raise ValueError("rope model must not carry learned positional embeddings")
79
+ else:
80
+ self.pos_emb = f32("pos_emb.embedding")
81
+ self.ln_f_g = f32("final_norm.gamma")
82
+ self.ln_f_b = f32("final_norm.beta")
83
+ self.layers = [_Layer(weights, f"blocks.{i}") for i in range(config.n_layers)]
84
+
85
+ if config.tie_weights:
86
+ self.head_w = self.tok_emb.T
87
+ self.head_b = None
88
+ else:
89
+ self.head_w = f32("lm_head.kernel")
90
+ self.head_b = f32("lm_head.bias")
91
+
92
+ if self.tok_emb.shape != (config.vocab_size, config.d_model):
93
+ raise ValueError("token embedding shape does not match config")
94
+
95
+ if self.pos_emb is not None and self.pos_emb.shape != (config.context_length, config.d_model):
96
+ raise ValueError("positional embedding shape does not match config")
97
+
98
+ @property
99
+ def nbytes(self):
100
+ total = self.tok_emb.nbytes + self.ln_f_g.nbytes + self.ln_f_b.nbytes
101
+
102
+ if self.pos_emb is not None:
103
+ total += self.pos_emb.nbytes
104
+
105
+ if self.head_b is not None:
106
+ total += self.head_w.nbytes + self.head_b.nbytes
107
+
108
+ for layer in self.layers:
109
+ total += sum(v.nbytes for v in vars(layer).values())
110
+
111
+ return total
112
+
113
+ def new_cache(self, capacity):
114
+ return KVCache(self.n_layers, self.n_heads, self.head_dim, min(capacity, self.context_length))
115
+
116
+ def _logits(self, x):
117
+ out = _layer_norm(x, self.ln_f_g, self.ln_f_b, self.eps) @ self.head_w
118
+
119
+ if self.head_b is not None:
120
+ out = out + self.head_b
121
+
122
+ return out
123
+
124
+ def _ffn(self, layer, x):
125
+ h = _layer_norm(x, layer.ln2_g, layer.ln2_b, self.eps)
126
+ return self.act(h @ layer.fc1_k + layer.fc1_b) @ layer.fc2_k + layer.fc2_b
127
+
128
+ def forward_chunk(self, token_ids, cache, want_logits=True):
129
+ ids = np.asarray(token_ids, dtype=np.int64)
130
+ T = ids.size
131
+ p0 = cache.length
132
+
133
+ if T == 0:
134
+ raise ValueError("empty chunk")
135
+
136
+ if p0 + T > cache.capacity:
137
+ raise ValueError("chunk exceeds kv cache capacity")
138
+
139
+ x = self.tok_emb[ids]
140
+ L = p0 + T
141
+
142
+ if self.pos_emb is not None:
143
+ x = x + self.pos_emb[p0:L]
144
+ else:
145
+ cos, sin = self.rope.tables(np.arange(p0, L))
146
+
147
+ mask = None
148
+ if T > 1:
149
+ mask = np.arange(L)[None, :] > (p0 + np.arange(T))[:, None]
150
+
151
+ H, hd, d = self.n_heads, self.head_dim, self.d_model
152
+
153
+ for li, layer in enumerate(self.layers):
154
+ h = _layer_norm(x, layer.ln1_g, layer.ln1_b, self.eps)
155
+ qkv = h @ layer.wqkv
156
+
157
+ q = qkv[:, :d].reshape(T, H, hd).transpose(1, 0, 2)
158
+ k = qkv[:, d:2 * d].reshape(T, H, hd).transpose(1, 0, 2)
159
+
160
+ if self.rope is not None:
161
+ q = self.rope.rotate(q, cos, sin)
162
+ k = self.rope.rotate(k, cos, sin)
163
+
164
+ cache.k[li, :, p0:L] = k
165
+ cache.v[li, :, p0:L] = qkv[:, 2 * d:].reshape(T, H, hd).transpose(1, 0, 2)
166
+
167
+ scores = (q @ cache.k[li, :, :L].transpose(0, 2, 1)) * self.scale
168
+
169
+ if mask is not None:
170
+ scores[:, mask] = -1e9
171
+
172
+ ctx = _softmax(scores) @ cache.v[li, :, :L]
173
+ x = x + ctx.transpose(1, 0, 2).reshape(T, d) @ layer.wo
174
+ x = x + self._ffn(layer, x)
175
+
176
+ cache.length = L
177
+
178
+ if not want_logits:
179
+ return None
180
+
181
+ return self._logits(x[-1:])[0]
182
+
183
+ def prefill(self, token_ids, cache, chunk_size=256):
184
+ ids = np.asarray(token_ids, dtype=np.int64)
185
+ logits = None
186
+
187
+ for start in range(0, ids.size, chunk_size):
188
+ end = min(start + chunk_size, ids.size)
189
+ logits = self.forward_chunk(ids[start:end], cache, want_logits=(end == ids.size))
190
+
191
+ return logits
192
+
193
+ def decode_batch(self, token_ids, caches):
194
+ B = len(caches)
195
+ ids = np.asarray(token_ids, dtype=np.int64)
196
+ pos = np.array([c.length for c in caches], dtype=np.int64)
197
+
198
+ for c in caches:
199
+ if c.length + 1 > c.capacity:
200
+ raise ValueError("kv cache capacity exceeded")
201
+
202
+ x = self.tok_emb[ids]
203
+
204
+ if self.pos_emb is not None:
205
+ x = x + self.pos_emb[pos]
206
+ else:
207
+ cos, sin = self.rope.tables(pos)
208
+ q_cos, q_sin = cos[:, None, None, :], sin[:, None, None, :]
209
+ k_cos, k_sin = cos[:, None, :], sin[:, None, :]
210
+
211
+ H, hd, d = self.n_heads, self.head_dim, self.d_model
212
+
213
+ for li, layer in enumerate(self.layers):
214
+ h = _layer_norm(x, layer.ln1_g, layer.ln1_b, self.eps)
215
+ qkv = h @ layer.wqkv
216
+
217
+ q = qkv[:, :d].reshape(B, H, 1, hd)
218
+ k = qkv[:, d:2 * d].reshape(B, H, hd)
219
+ v = qkv[:, 2 * d:].reshape(B, H, hd)
220
+
221
+ if self.rope is not None:
222
+ q = self.rope.rotate(q, q_cos, q_sin)
223
+ k = self.rope.rotate(k, k_cos, k_sin)
224
+
225
+ ctx = np.empty((B, d), dtype=np.float32)
226
+
227
+ for b, c in enumerate(caches):
228
+ n = int(pos[b])
229
+ c.k[li, :, n] = k[b]
230
+ c.v[li, :, n] = v[b]
231
+
232
+ scores = (q[b] @ c.k[li, :, :n + 1].transpose(0, 2, 1)) * self.scale
233
+ ctx[b] = (_softmax(scores) @ c.v[li, :, :n + 1]).reshape(d)
234
+
235
+ x = x + ctx @ layer.wo
236
+ x = x + self._ffn(layer, x)
237
+
238
+ for c in caches:
239
+ c.length += 1
240
+
241
+ return self._logits(x)