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,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)
|
src/inference/config.py
ADDED
|
@@ -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})
|
src/inference/engine.py
ADDED
|
@@ -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)
|