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,476 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import heapq
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
import re
|
|
6
|
+
from collections import Counter
|
|
7
|
+
|
|
8
|
+
from src.tokenization.base import Tokenizer
|
|
9
|
+
|
|
10
|
+
FORMAT_VERSION = 1
|
|
11
|
+
BASE_VOCAB = 256
|
|
12
|
+
HEAP_PATH_THRESHOLD = 128
|
|
13
|
+
CACHE_LIMIT = 1_000_000
|
|
14
|
+
|
|
15
|
+
PATTERNS = {
|
|
16
|
+
"gpt2": re.compile(
|
|
17
|
+
r"'(?:[sdmt]|ll|ve|re)| ?[^\W\d_]+| ?\d+| ?(?:[^\s\w]|_)+|\s+(?!\S)|\s+"
|
|
18
|
+
),
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _bpe_simple(ids, ranks):
|
|
23
|
+
while len(ids) > 1:
|
|
24
|
+
best_rank = None
|
|
25
|
+
|
|
26
|
+
for i in range(len(ids) - 1):
|
|
27
|
+
rank = ranks.get((ids[i], ids[i + 1]))
|
|
28
|
+
|
|
29
|
+
if rank is not None and (best_rank is None or rank < best_rank):
|
|
30
|
+
best_rank = rank
|
|
31
|
+
best = (ids[i], ids[i + 1])
|
|
32
|
+
|
|
33
|
+
if best_rank is None:
|
|
34
|
+
break
|
|
35
|
+
|
|
36
|
+
new_id = BASE_VOCAB + best_rank
|
|
37
|
+
merged = []
|
|
38
|
+
i = 0
|
|
39
|
+
n = len(ids)
|
|
40
|
+
|
|
41
|
+
while i < n:
|
|
42
|
+
if i < n - 1 and ids[i] == best[0] and ids[i + 1] == best[1]:
|
|
43
|
+
merged.append(new_id)
|
|
44
|
+
i += 2
|
|
45
|
+
else:
|
|
46
|
+
merged.append(ids[i])
|
|
47
|
+
i += 1
|
|
48
|
+
|
|
49
|
+
ids = merged
|
|
50
|
+
|
|
51
|
+
return ids
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _bpe_heap(ids, ranks):
|
|
55
|
+
n = len(ids)
|
|
56
|
+
ids = list(ids)
|
|
57
|
+
nxt = list(range(1, n + 1))
|
|
58
|
+
prv = list(range(-1, n - 1))
|
|
59
|
+
alive = [True] * n
|
|
60
|
+
heap = []
|
|
61
|
+
|
|
62
|
+
for i in range(n - 1):
|
|
63
|
+
rank = ranks.get((ids[i], ids[i + 1]))
|
|
64
|
+
if rank is not None:
|
|
65
|
+
heap.append((rank, i))
|
|
66
|
+
|
|
67
|
+
heapq.heapify(heap)
|
|
68
|
+
|
|
69
|
+
while heap:
|
|
70
|
+
rank, i = heapq.heappop(heap)
|
|
71
|
+
|
|
72
|
+
if not alive[i]:
|
|
73
|
+
continue
|
|
74
|
+
|
|
75
|
+
j = nxt[i]
|
|
76
|
+
|
|
77
|
+
if j >= n or not alive[j]:
|
|
78
|
+
continue
|
|
79
|
+
|
|
80
|
+
if ranks.get((ids[i], ids[j])) != rank:
|
|
81
|
+
continue
|
|
82
|
+
|
|
83
|
+
ids[i] = BASE_VOCAB + rank
|
|
84
|
+
alive[j] = False
|
|
85
|
+
after = nxt[j]
|
|
86
|
+
nxt[i] = after
|
|
87
|
+
|
|
88
|
+
if after < n:
|
|
89
|
+
prv[after] = i
|
|
90
|
+
|
|
91
|
+
before = prv[i]
|
|
92
|
+
|
|
93
|
+
if before >= 0:
|
|
94
|
+
r = ranks.get((ids[before], ids[i]))
|
|
95
|
+
if r is not None:
|
|
96
|
+
heapq.heappush(heap, (r, before))
|
|
97
|
+
|
|
98
|
+
if after < n:
|
|
99
|
+
r = ranks.get((ids[i], ids[after]))
|
|
100
|
+
if r is not None:
|
|
101
|
+
heapq.heappush(heap, (r, i))
|
|
102
|
+
|
|
103
|
+
return [ids[i] for i in range(n) if alive[i]]
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
class ByteLevelBPETokenizer(Tokenizer):
|
|
107
|
+
|
|
108
|
+
def __init__(
|
|
109
|
+
self,
|
|
110
|
+
merges=(),
|
|
111
|
+
special_tokens=None,
|
|
112
|
+
bos_token="<bos>",
|
|
113
|
+
eos_token="<eos>",
|
|
114
|
+
pad_token="<pad>",
|
|
115
|
+
pattern="gpt2",
|
|
116
|
+
):
|
|
117
|
+
if pattern not in PATTERNS:
|
|
118
|
+
raise ValueError(f"unknown pretokenizer pattern '{pattern}'")
|
|
119
|
+
|
|
120
|
+
self.pattern_name = pattern
|
|
121
|
+
self._pat = PATTERNS[pattern]
|
|
122
|
+
|
|
123
|
+
self.bos_token = bos_token
|
|
124
|
+
self.eos_token = eos_token
|
|
125
|
+
self.pad_token = pad_token
|
|
126
|
+
|
|
127
|
+
names = [bos_token, eos_token, pad_token]
|
|
128
|
+
|
|
129
|
+
for token in special_tokens or []:
|
|
130
|
+
if token not in names:
|
|
131
|
+
names.append(token)
|
|
132
|
+
|
|
133
|
+
if len(set(names)) != len(names) or any(not n for n in names):
|
|
134
|
+
raise ValueError("special tokens must be unique and non-empty")
|
|
135
|
+
|
|
136
|
+
self._special_names = names
|
|
137
|
+
self._set_merges([tuple(m) for m in merges])
|
|
138
|
+
|
|
139
|
+
def _set_merges(self, merges):
|
|
140
|
+
ranks = {}
|
|
141
|
+
token_bytes = [bytes([i]) for i in range(BASE_VOCAB)]
|
|
142
|
+
|
|
143
|
+
for rank, (a, b) in enumerate(merges):
|
|
144
|
+
if not (0 <= a < len(token_bytes) and 0 <= b < len(token_bytes)):
|
|
145
|
+
raise ValueError(f"merge {rank} references an id that does not exist yet")
|
|
146
|
+
|
|
147
|
+
if (a, b) in ranks:
|
|
148
|
+
raise ValueError(f"duplicate merge {(a, b)}")
|
|
149
|
+
|
|
150
|
+
ranks[(a, b)] = rank
|
|
151
|
+
token_bytes.append(token_bytes[a] + token_bytes[b])
|
|
152
|
+
|
|
153
|
+
self.merges = merges
|
|
154
|
+
self._ranks = ranks
|
|
155
|
+
self._token_bytes = token_bytes
|
|
156
|
+
self._cache = {}
|
|
157
|
+
|
|
158
|
+
base = len(token_bytes)
|
|
159
|
+
self._special_ids = {name: base + i for i, name in enumerate(self._special_names)}
|
|
160
|
+
self._id_to_special = {v: k for k, v in self._special_ids.items()}
|
|
161
|
+
|
|
162
|
+
ordered = sorted(self._special_names, key=len, reverse=True)
|
|
163
|
+
self._special_re = re.compile("|".join(re.escape(n) for n in ordered))
|
|
164
|
+
self._fingerprint_value = None
|
|
165
|
+
|
|
166
|
+
@property
|
|
167
|
+
def vocab_size(self):
|
|
168
|
+
return len(self._token_bytes) + len(self._special_names)
|
|
169
|
+
|
|
170
|
+
@property
|
|
171
|
+
def bos_id(self):
|
|
172
|
+
return self._special_ids[self.bos_token]
|
|
173
|
+
|
|
174
|
+
@property
|
|
175
|
+
def eos_id(self):
|
|
176
|
+
return self._special_ids[self.eos_token]
|
|
177
|
+
|
|
178
|
+
@property
|
|
179
|
+
def pad_id(self):
|
|
180
|
+
return self._special_ids[self.pad_token]
|
|
181
|
+
|
|
182
|
+
@property
|
|
183
|
+
def special_tokens(self):
|
|
184
|
+
return dict(self._special_ids)
|
|
185
|
+
|
|
186
|
+
@property
|
|
187
|
+
def fingerprint(self):
|
|
188
|
+
if self._fingerprint_value is None:
|
|
189
|
+
digest = hashlib.sha256()
|
|
190
|
+
digest.update(json.dumps(self.merges).encode())
|
|
191
|
+
digest.update(json.dumps(self._special_names).encode())
|
|
192
|
+
digest.update(self.pattern_name.encode())
|
|
193
|
+
self._fingerprint_value = digest.hexdigest()[:16]
|
|
194
|
+
|
|
195
|
+
return self._fingerprint_value
|
|
196
|
+
|
|
197
|
+
@property
|
|
198
|
+
def identity(self):
|
|
199
|
+
return {
|
|
200
|
+
"type": "bytebpe",
|
|
201
|
+
"format_version": FORMAT_VERSION,
|
|
202
|
+
"vocab_size": self.vocab_size,
|
|
203
|
+
"fingerprint": self.fingerprint,
|
|
204
|
+
}
|
|
205
|
+
|
|
206
|
+
def token_bytes(self, token_id):
|
|
207
|
+
if 0 <= token_id < len(self._token_bytes):
|
|
208
|
+
return self._token_bytes[token_id]
|
|
209
|
+
|
|
210
|
+
if token_id in self._id_to_special:
|
|
211
|
+
return b""
|
|
212
|
+
|
|
213
|
+
raise ValueError(f"token id {token_id} is outside the vocabulary of {self.vocab_size}")
|
|
214
|
+
|
|
215
|
+
def _encode_piece(self, piece):
|
|
216
|
+
cached = self._cache.get(piece)
|
|
217
|
+
|
|
218
|
+
if cached is not None:
|
|
219
|
+
return cached
|
|
220
|
+
|
|
221
|
+
ids = list(piece.encode("utf-8", errors="replace"))
|
|
222
|
+
|
|
223
|
+
if len(ids) > 1:
|
|
224
|
+
ids = _bpe_heap(ids, self._ranks) if len(ids) > HEAP_PATH_THRESHOLD else _bpe_simple(ids, self._ranks)
|
|
225
|
+
|
|
226
|
+
result = tuple(ids)
|
|
227
|
+
|
|
228
|
+
if len(self._cache) >= CACHE_LIMIT:
|
|
229
|
+
self._cache.clear()
|
|
230
|
+
|
|
231
|
+
self._cache[piece] = result
|
|
232
|
+
|
|
233
|
+
return result
|
|
234
|
+
|
|
235
|
+
def encode_ordinary(self, text):
|
|
236
|
+
out = []
|
|
237
|
+
|
|
238
|
+
for piece in self._pat.findall(text):
|
|
239
|
+
out.extend(self._encode_piece(piece))
|
|
240
|
+
|
|
241
|
+
return out
|
|
242
|
+
|
|
243
|
+
def encode(self, text, add_bos=False, add_eos=False, allow_special=False):
|
|
244
|
+
out = [self.bos_id] if add_bos else []
|
|
245
|
+
|
|
246
|
+
if allow_special:
|
|
247
|
+
pos = 0
|
|
248
|
+
|
|
249
|
+
for match in self._special_re.finditer(text):
|
|
250
|
+
if match.start() > pos:
|
|
251
|
+
out.extend(self.encode_ordinary(text[pos:match.start()]))
|
|
252
|
+
|
|
253
|
+
out.append(self._special_ids[match.group()])
|
|
254
|
+
pos = match.end()
|
|
255
|
+
|
|
256
|
+
if pos < len(text):
|
|
257
|
+
out.extend(self.encode_ordinary(text[pos:]))
|
|
258
|
+
else:
|
|
259
|
+
out.extend(self.encode_ordinary(text))
|
|
260
|
+
|
|
261
|
+
if add_eos:
|
|
262
|
+
out.append(self.eos_id)
|
|
263
|
+
|
|
264
|
+
return out
|
|
265
|
+
|
|
266
|
+
def decode_bytes(self, ids, skip_special=True):
|
|
267
|
+
parts = []
|
|
268
|
+
|
|
269
|
+
for i in ids:
|
|
270
|
+
i = int(i)
|
|
271
|
+
|
|
272
|
+
if i in self._id_to_special:
|
|
273
|
+
if not skip_special:
|
|
274
|
+
parts.append(self._id_to_special[i].encode("utf-8"))
|
|
275
|
+
else:
|
|
276
|
+
parts.append(self.token_bytes(i))
|
|
277
|
+
|
|
278
|
+
return b"".join(parts)
|
|
279
|
+
|
|
280
|
+
def decode(self, ids, skip_special=True):
|
|
281
|
+
return self.decode_bytes(ids, skip_special).decode("utf-8", errors="replace")
|
|
282
|
+
|
|
283
|
+
def save(self, path):
|
|
284
|
+
state = {
|
|
285
|
+
"type": "bytebpe",
|
|
286
|
+
"format_version": FORMAT_VERSION,
|
|
287
|
+
"pattern": self.pattern_name,
|
|
288
|
+
"bos_token": self.bos_token,
|
|
289
|
+
"eos_token": self.eos_token,
|
|
290
|
+
"pad_token": self.pad_token,
|
|
291
|
+
"special_tokens": self._special_names,
|
|
292
|
+
"merges": [list(m) for m in self.merges],
|
|
293
|
+
"vocab_size": self.vocab_size,
|
|
294
|
+
"fingerprint": self.fingerprint,
|
|
295
|
+
}
|
|
296
|
+
|
|
297
|
+
tmp = path + ".tmp"
|
|
298
|
+
|
|
299
|
+
with open(tmp, "w") as f:
|
|
300
|
+
json.dump(state, f)
|
|
301
|
+
|
|
302
|
+
os.replace(tmp, path)
|
|
303
|
+
|
|
304
|
+
@classmethod
|
|
305
|
+
def load(cls, path):
|
|
306
|
+
with open(path) as f:
|
|
307
|
+
state = json.load(f)
|
|
308
|
+
|
|
309
|
+
if state.get("format_version") != FORMAT_VERSION:
|
|
310
|
+
raise ValueError(f"unsupported byte-level tokenizer format {state.get('format_version')}")
|
|
311
|
+
|
|
312
|
+
tok = cls(
|
|
313
|
+
merges=state["merges"],
|
|
314
|
+
special_tokens=state["special_tokens"],
|
|
315
|
+
bos_token=state["bos_token"],
|
|
316
|
+
eos_token=state["eos_token"],
|
|
317
|
+
pad_token=state["pad_token"],
|
|
318
|
+
pattern=state["pattern"],
|
|
319
|
+
)
|
|
320
|
+
|
|
321
|
+
if state.get("fingerprint") and state["fingerprint"] != tok.fingerprint:
|
|
322
|
+
raise ValueError("tokenizer file failed its integrity check")
|
|
323
|
+
|
|
324
|
+
return tok
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
def count_pretokens(texts, pattern="gpt2", max_unique=1_000_000):
|
|
328
|
+
pat = PATTERNS[pattern]
|
|
329
|
+
counts = Counter()
|
|
330
|
+
threshold = 1
|
|
331
|
+
|
|
332
|
+
for text in texts:
|
|
333
|
+
counts.update(pat.findall(text))
|
|
334
|
+
|
|
335
|
+
if len(counts) > max_unique:
|
|
336
|
+
target = int(max_unique * 0.75)
|
|
337
|
+
|
|
338
|
+
while len(counts) > target:
|
|
339
|
+
counts = Counter({k: v for k, v in counts.items() if v > threshold})
|
|
340
|
+
threshold += 1
|
|
341
|
+
|
|
342
|
+
return counts
|
|
343
|
+
|
|
344
|
+
|
|
345
|
+
def learn_merges(counts, num_merges, min_frequency=2, progress=None):
|
|
346
|
+
merged = {}
|
|
347
|
+
|
|
348
|
+
for piece, count in counts.items():
|
|
349
|
+
key = piece.encode("utf-8", errors="replace")
|
|
350
|
+
merged[key] = merged.get(key, 0) + count
|
|
351
|
+
|
|
352
|
+
words = [list(k) for k in merged]
|
|
353
|
+
freqs = list(merged.values())
|
|
354
|
+
|
|
355
|
+
pair_counts = {}
|
|
356
|
+
pair_index = {}
|
|
357
|
+
|
|
358
|
+
for wi, word in enumerate(words):
|
|
359
|
+
f = freqs[wi]
|
|
360
|
+
|
|
361
|
+
for pair in zip(word, word[1:]):
|
|
362
|
+
pair_counts[pair] = pair_counts.get(pair, 0) + f
|
|
363
|
+
pair_index.setdefault(pair, set()).add(wi)
|
|
364
|
+
|
|
365
|
+
heap = [(-c, p) for p, c in pair_counts.items()]
|
|
366
|
+
heapq.heapify(heap)
|
|
367
|
+
|
|
368
|
+
merges = []
|
|
369
|
+
|
|
370
|
+
for step in range(num_merges):
|
|
371
|
+
best = None
|
|
372
|
+
|
|
373
|
+
while heap:
|
|
374
|
+
neg, pair = heapq.heappop(heap)
|
|
375
|
+
current = pair_counts.get(pair, 0)
|
|
376
|
+
|
|
377
|
+
if current <= 0:
|
|
378
|
+
continue
|
|
379
|
+
|
|
380
|
+
if current != -neg:
|
|
381
|
+
heapq.heappush(heap, (-current, pair))
|
|
382
|
+
continue
|
|
383
|
+
|
|
384
|
+
best = pair
|
|
385
|
+
break
|
|
386
|
+
|
|
387
|
+
if best is None or pair_counts[best] < min_frequency:
|
|
388
|
+
break
|
|
389
|
+
|
|
390
|
+
a, b = best
|
|
391
|
+
new_id = BASE_VOCAB + step
|
|
392
|
+
merges.append(best)
|
|
393
|
+
touched = set()
|
|
394
|
+
|
|
395
|
+
for wi in list(pair_index.pop(best, ())):
|
|
396
|
+
word = words[wi]
|
|
397
|
+
|
|
398
|
+
if not any(word[i] == a and word[i + 1] == b for i in range(len(word) - 1)):
|
|
399
|
+
continue
|
|
400
|
+
|
|
401
|
+
f = freqs[wi]
|
|
402
|
+
|
|
403
|
+
for pair in zip(word, word[1:]):
|
|
404
|
+
pair_counts[pair] -= f
|
|
405
|
+
|
|
406
|
+
new_word = []
|
|
407
|
+
i = 0
|
|
408
|
+
n = len(word)
|
|
409
|
+
|
|
410
|
+
while i < n:
|
|
411
|
+
if i < n - 1 and word[i] == a and word[i + 1] == b:
|
|
412
|
+
new_word.append(new_id)
|
|
413
|
+
i += 2
|
|
414
|
+
else:
|
|
415
|
+
new_word.append(word[i])
|
|
416
|
+
i += 1
|
|
417
|
+
|
|
418
|
+
for pair in zip(new_word, new_word[1:]):
|
|
419
|
+
pair_counts[pair] = pair_counts.get(pair, 0) + f
|
|
420
|
+
pair_index.setdefault(pair, set()).add(wi)
|
|
421
|
+
touched.add(pair)
|
|
422
|
+
|
|
423
|
+
words[wi] = new_word
|
|
424
|
+
|
|
425
|
+
pair_counts.pop(best, None)
|
|
426
|
+
|
|
427
|
+
for pair in touched:
|
|
428
|
+
c = pair_counts.get(pair, 0)
|
|
429
|
+
if c > 0:
|
|
430
|
+
heapq.heappush(heap, (-c, pair))
|
|
431
|
+
|
|
432
|
+
if progress is not None and (step + 1) % 500 == 0:
|
|
433
|
+
progress(step + 1, num_merges)
|
|
434
|
+
|
|
435
|
+
return merges
|
|
436
|
+
|
|
437
|
+
|
|
438
|
+
def train_byte_bpe(
|
|
439
|
+
texts,
|
|
440
|
+
vocab_size,
|
|
441
|
+
special_tokens=None,
|
|
442
|
+
min_frequency=2,
|
|
443
|
+
max_unique_words=1_000_000,
|
|
444
|
+
pattern="gpt2",
|
|
445
|
+
bos_token="<bos>",
|
|
446
|
+
eos_token="<eos>",
|
|
447
|
+
pad_token="<pad>",
|
|
448
|
+
progress=None,
|
|
449
|
+
):
|
|
450
|
+
shell = ByteLevelBPETokenizer(
|
|
451
|
+
special_tokens=special_tokens,
|
|
452
|
+
bos_token=bos_token,
|
|
453
|
+
eos_token=eos_token,
|
|
454
|
+
pad_token=pad_token,
|
|
455
|
+
pattern=pattern,
|
|
456
|
+
)
|
|
457
|
+
|
|
458
|
+
num_merges = vocab_size - BASE_VOCAB - len(shell.special_tokens)
|
|
459
|
+
|
|
460
|
+
if num_merges < 0:
|
|
461
|
+
raise ValueError(
|
|
462
|
+
f"vocab_size {vocab_size} is too small: {BASE_VOCAB} byte tokens plus "
|
|
463
|
+
f"{len(shell.special_tokens)} special tokens are always present"
|
|
464
|
+
)
|
|
465
|
+
|
|
466
|
+
counts = count_pretokens(texts, pattern, max_unique_words)
|
|
467
|
+
merges = learn_merges(counts, num_merges, min_frequency, progress)
|
|
468
|
+
|
|
469
|
+
return ByteLevelBPETokenizer(
|
|
470
|
+
merges=merges,
|
|
471
|
+
special_tokens=special_tokens,
|
|
472
|
+
bos_token=bos_token,
|
|
473
|
+
eos_token=eos_token,
|
|
474
|
+
pad_token=pad_token,
|
|
475
|
+
pattern=pattern,
|
|
476
|
+
)
|
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
import json
|
|
2
|
+
|
|
3
|
+
_REGISTRY = {}
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def register_tokenizer(kind, cls):
|
|
7
|
+
_REGISTRY[kind] = cls
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def load_tokenizer(path):
|
|
11
|
+
with open(path) as f:
|
|
12
|
+
kind = json.load(f).get("type", "bpe")
|
|
13
|
+
|
|
14
|
+
if kind not in _REGISTRY:
|
|
15
|
+
raise ValueError(f"no tokenizer registered for type '{kind}'")
|
|
16
|
+
|
|
17
|
+
return _REGISTRY[kind].load(path)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _register_builtin():
|
|
21
|
+
from src.tokenization.bpe import PTFBPETokenizer
|
|
22
|
+
from src.tokenization.bytebpe import ByteLevelBPETokenizer
|
|
23
|
+
|
|
24
|
+
register_tokenizer("bpe", PTFBPETokenizer)
|
|
25
|
+
register_tokenizer("bytebpe", ByteLevelBPETokenizer)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
_register_builtin()
|
src/training/__init__.py
ADDED
|
File without changes
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
import glob
|
|
2
|
+
import os
|
|
3
|
+
import pickle
|
|
4
|
+
import re
|
|
5
|
+
import tempfile
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def read_checkpoint(path):
|
|
9
|
+
target = os.path.join(path, "checkpoint_latest.ptf") if os.path.isdir(path) else path
|
|
10
|
+
|
|
11
|
+
if not os.path.isfile(target):
|
|
12
|
+
raise FileNotFoundError(f"no checkpoint at {path}")
|
|
13
|
+
|
|
14
|
+
with open(target, "rb") as f:
|
|
15
|
+
return pickle.load(f)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class CheckpointManager:
|
|
19
|
+
|
|
20
|
+
def __init__(self, directory, keep_last=3, keep_every=None):
|
|
21
|
+
self.directory = directory
|
|
22
|
+
self.keep_last = keep_last
|
|
23
|
+
self.keep_every = keep_every
|
|
24
|
+
|
|
25
|
+
os.makedirs(directory, exist_ok=True)
|
|
26
|
+
|
|
27
|
+
def _atomic_write(self, path, payload):
|
|
28
|
+
fd, tmp_path = tempfile.mkstemp(dir=self.directory, prefix=".tmp_ckpt_")
|
|
29
|
+
|
|
30
|
+
try:
|
|
31
|
+
with os.fdopen(fd, "wb") as f:
|
|
32
|
+
pickle.dump(payload, f, protocol=pickle.HIGHEST_PROTOCOL)
|
|
33
|
+
f.flush()
|
|
34
|
+
os.fsync(f.fileno())
|
|
35
|
+
|
|
36
|
+
os.replace(tmp_path, path)
|
|
37
|
+
except BaseException:
|
|
38
|
+
if os.path.exists(tmp_path):
|
|
39
|
+
os.remove(tmp_path)
|
|
40
|
+
raise
|
|
41
|
+
|
|
42
|
+
def save(self, step, payload):
|
|
43
|
+
step_path = os.path.join(self.directory, f"checkpoint_step_{step}.ptf")
|
|
44
|
+
self._atomic_write(step_path, payload)
|
|
45
|
+
|
|
46
|
+
latest_path = os.path.join(self.directory, "checkpoint_latest.ptf")
|
|
47
|
+
self._atomic_write(latest_path, payload)
|
|
48
|
+
|
|
49
|
+
self._rotate()
|
|
50
|
+
|
|
51
|
+
return step_path
|
|
52
|
+
|
|
53
|
+
def _step_checkpoints(self):
|
|
54
|
+
pattern = os.path.join(self.directory, "checkpoint_step_*.ptf")
|
|
55
|
+
found = []
|
|
56
|
+
|
|
57
|
+
for path in glob.glob(pattern):
|
|
58
|
+
match = re.search(r"checkpoint_step_(\d+)\.ptf$", path)
|
|
59
|
+
if match:
|
|
60
|
+
found.append((int(match.group(1)), path))
|
|
61
|
+
|
|
62
|
+
found.sort(key=lambda item: item[0])
|
|
63
|
+
|
|
64
|
+
return found
|
|
65
|
+
|
|
66
|
+
def _rotate(self):
|
|
67
|
+
checkpoints = self._step_checkpoints()
|
|
68
|
+
|
|
69
|
+
if not checkpoints:
|
|
70
|
+
return
|
|
71
|
+
|
|
72
|
+
if self.keep_last is None and self.keep_every is None:
|
|
73
|
+
return
|
|
74
|
+
|
|
75
|
+
keep_steps = {checkpoints[-1][0]}
|
|
76
|
+
|
|
77
|
+
if self.keep_last:
|
|
78
|
+
keep_steps.update(step for step, _ in checkpoints[-self.keep_last:])
|
|
79
|
+
|
|
80
|
+
if self.keep_every:
|
|
81
|
+
keep_steps.update(
|
|
82
|
+
step for step, _ in checkpoints if step % self.keep_every == 0
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
for step, path in checkpoints:
|
|
86
|
+
if step not in keep_steps:
|
|
87
|
+
os.remove(path)
|
|
88
|
+
|
|
89
|
+
def load(self, path=None):
|
|
90
|
+
if path is None:
|
|
91
|
+
path = os.path.join(self.directory, "checkpoint_latest.ptf")
|
|
92
|
+
|
|
93
|
+
if not os.path.exists(path):
|
|
94
|
+
return None
|
|
95
|
+
|
|
96
|
+
with open(path, "rb") as f:
|
|
97
|
+
return pickle.load(f)
|
|
98
|
+
|
|
99
|
+
def latest_path(self):
|
|
100
|
+
path = os.path.join(self.directory, "checkpoint_latest.ptf")
|
|
101
|
+
return path if os.path.exists(path) else None
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
import platform
|
|
4
|
+
import sys
|
|
5
|
+
import time
|
|
6
|
+
import uuid
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def hardware_info():
|
|
12
|
+
return {
|
|
13
|
+
"platform": platform.platform(),
|
|
14
|
+
"processor": platform.processor(),
|
|
15
|
+
"cpu_count": os.cpu_count(),
|
|
16
|
+
"python": sys.version.split()[0],
|
|
17
|
+
"numpy": np.__version__,
|
|
18
|
+
"device": "cpu",
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class Experiment:
|
|
23
|
+
|
|
24
|
+
def __init__(self, directory, config_dict, tokenizer_identity, dataset_info, run_id=None):
|
|
25
|
+
self.directory = directory
|
|
26
|
+
self.path = os.path.join(directory, "experiment.json")
|
|
27
|
+
|
|
28
|
+
self.record = {
|
|
29
|
+
"run_id": run_id or f"{time.strftime('%Y%m%d-%H%M%S')}-{uuid.uuid4().hex[:6]}",
|
|
30
|
+
"start_time": time.time(),
|
|
31
|
+
"config": config_dict,
|
|
32
|
+
"tokenizer": tokenizer_identity,
|
|
33
|
+
"dataset": dataset_info,
|
|
34
|
+
"hardware": hardware_info(),
|
|
35
|
+
"seed": config_dict.get("training", {}).get("seed"),
|
|
36
|
+
"resumes": [],
|
|
37
|
+
"final_metrics": None,
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
@classmethod
|
|
41
|
+
def load_or_create(cls, directory, config_dict, tokenizer_identity, dataset_info):
|
|
42
|
+
path = os.path.join(directory, "experiment.json")
|
|
43
|
+
|
|
44
|
+
exp = cls(directory, config_dict, tokenizer_identity, dataset_info)
|
|
45
|
+
|
|
46
|
+
if os.path.isfile(path):
|
|
47
|
+
with open(path) as f:
|
|
48
|
+
exp.record = json.load(f)
|
|
49
|
+
|
|
50
|
+
exp.record["resumes"].append({"time": time.time(), "hardware": hardware_info()})
|
|
51
|
+
|
|
52
|
+
exp.save()
|
|
53
|
+
|
|
54
|
+
return exp
|
|
55
|
+
|
|
56
|
+
@property
|
|
57
|
+
def run_id(self):
|
|
58
|
+
return self.record["run_id"]
|
|
59
|
+
|
|
60
|
+
def update_metrics(self, metrics):
|
|
61
|
+
self.record["final_metrics"] = metrics
|
|
62
|
+
self.save()
|
|
63
|
+
|
|
64
|
+
def save(self):
|
|
65
|
+
os.makedirs(self.directory, exist_ok=True)
|
|
66
|
+
tmp = self.path + ".tmp"
|
|
67
|
+
|
|
68
|
+
with open(tmp, "w") as f:
|
|
69
|
+
json.dump(self.record, f, indent=2, default=str)
|
|
70
|
+
|
|
71
|
+
os.replace(tmp, self.path)
|