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,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()
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)