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,594 @@
1
+ import contextlib
2
+ import json
3
+ import os
4
+ import random
5
+ import resource
6
+ import signal
7
+ import time
8
+ from dataclasses import asdict
9
+
10
+ import numpy as np
11
+
12
+ from src.core.Tensor import Tensor, no_grad
13
+ from src.data.prefetch import PrefetchLoader
14
+ from src.loss.CrossEntropyWithLogitsLoss import CrossEntropyWithLogitsLoss
15
+ from src.optimizers.AdamW import AdamW
16
+ from src.tokenization.base import identity_matches
17
+ from src.training.checkpoint_manager import CheckpointManager
18
+ from src.training.experiment import Experiment
19
+ from src.training.profiler import Profiler
20
+ from src.training.losses import masked_cross_entropy
21
+ from src.training.precision import PRECISIONS, GradScaler, autocast, resolve_loss_scaling
22
+ from src.training.scheduler import LRScheduler
23
+
24
+ FRAMEWORK_VERSION = "0.4.0-phase7"
25
+
26
+ ARCHITECTURE_FIELDS = ("vocab_size", "d_model", "n_layers", "n_heads", "ff_dim", "activation", "norm_eps",
27
+ "tie_weights", "position_encoding", "rope_theta")
28
+
29
+
30
+ def describe_incompatibility(source, target):
31
+ problems = [f"{name}: {getattr(source, name)} vs {getattr(target, name)}"
32
+ for name in ARCHITECTURE_FIELDS if getattr(source, name) != getattr(target, name)]
33
+
34
+ if source.context_length != target.context_length and not target.uses_rope:
35
+ problems.append(
36
+ f"context_length: {source.context_length} vs {target.context_length} "
37
+ "(learned position tables cannot change size; only rope models can change context)"
38
+ )
39
+
40
+ return "; ".join(problems)
41
+
42
+
43
+ class Trainer:
44
+
45
+ def __init__(
46
+ self,
47
+ model,
48
+ train_dataset,
49
+ config,
50
+ tokenizer,
51
+ val_dataset=None,
52
+ save_on_exit=True,
53
+ install_signal_handlers=True,
54
+ ):
55
+ if config.training.precision not in PRECISIONS:
56
+ raise ValueError(f"unknown precision '{config.training.precision}'; choose one of {PRECISIONS}")
57
+
58
+ if config.runtime.device != "cpu":
59
+ raise ValueError("only the cpu device is available in this backend")
60
+
61
+ self.model = model
62
+ self.train_dataset = train_dataset
63
+ self.val_dataset = val_dataset
64
+ self.config = config
65
+ self.tokenizer = tokenizer
66
+ self.save_on_exit = save_on_exit
67
+
68
+ self.loss_fn = CrossEntropyWithLogitsLoss()
69
+
70
+ self.optimizer = AdamW(
71
+ lr=config.training.learning_rate,
72
+ weight_decay=config.training.weight_decay,
73
+ )
74
+
75
+ self.scheduler = LRScheduler(
76
+ base_lr=config.training.learning_rate,
77
+ warmup_steps=config.training.warmup_steps,
78
+ decay=config.training.decay,
79
+ total_steps=config.training.total_steps,
80
+ min_lr=config.training.min_lr,
81
+ )
82
+
83
+ self.checkpoints = CheckpointManager(
84
+ config.checkpoint.directory,
85
+ keep_last=config.checkpoint.keep_last,
86
+ keep_every=config.checkpoint.keep_every,
87
+ )
88
+
89
+ self.profiler = Profiler()
90
+
91
+ self.precision = config.training.precision
92
+ self.scaler = GradScaler(
93
+ enabled=resolve_loss_scaling(self.precision, config.training.loss_scaling),
94
+ init_scale=config.training.initial_loss_scale,
95
+ growth_interval=config.training.loss_scale_growth_interval,
96
+ )
97
+ self.model.activation_checkpointing = bool(config.training.activation_checkpointing)
98
+
99
+ if self.precision != "fp32":
100
+ print(
101
+ f"precision {self.precision}: matmul operands, outputs and gradients are rounded to "
102
+ f"{self.precision} with fp32 master weights; on the NumPy CPU backend this reproduces "
103
+ f"reduced-precision numerics, and it makes training slower and uses more memory, not less"
104
+ )
105
+
106
+ self.global_step = 0
107
+ self.optimizer_steps = 0
108
+ self.epoch = 0
109
+ self.tokens_processed = 0
110
+ self.target_tokens = 0
111
+ self.init_source = None
112
+ self.examples_processed = 0
113
+ self.last_loss = None
114
+ self.last_eval = None
115
+
116
+ self._consumed_state = None
117
+ self._tokens_at_last_eval = 0
118
+ self._tokens_at_last_ckpt = 0
119
+ self._stop_requested = False
120
+ self._last_saved_step = None
121
+
122
+ os.makedirs(config.checkpoint.directory, exist_ok=True)
123
+ self.log_path = os.path.join(config.checkpoint.directory, "train_log.jsonl")
124
+
125
+ self.experiment = Experiment.load_or_create(
126
+ config.checkpoint.directory,
127
+ asdict(config),
128
+ tokenizer.identity,
129
+ {"train_files": len(train_dataset.corpus), "sequence_length": config.data.sequence_length},
130
+ )
131
+
132
+ self._seed_everything(config.training.seed)
133
+
134
+ if install_signal_handlers:
135
+ self._install_signal_handlers()
136
+
137
+ def _seed_everything(self, seed):
138
+ random.seed(seed)
139
+ np.random.seed(seed)
140
+
141
+ def _install_signal_handlers(self):
142
+ def handler(signum, frame):
143
+ print(f"received signal {signum}; finishing the current step, then checkpointing")
144
+ self._stop_requested = True
145
+
146
+ signal.signal(signal.SIGINT, handler)
147
+ signal.signal(signal.SIGTERM, handler)
148
+
149
+ def request_stop(self):
150
+ self._stop_requested = True
151
+
152
+ def pause(self):
153
+ self._stop_requested = True
154
+ return self.save_checkpoint()
155
+
156
+ def _clip_grad_norm(self, params, max_norm):
157
+ total_sq = 0.0
158
+
159
+ for p in params:
160
+ if p.requires_grad:
161
+ total_sq += float(np.sum(p.grad ** 2))
162
+
163
+ total_norm = total_sq ** 0.5
164
+
165
+ if max_norm and total_norm > max_norm:
166
+ scale = max_norm / (total_norm + 1e-6)
167
+ for p in params:
168
+ if p.requires_grad:
169
+ p.grad *= scale
170
+
171
+ return total_norm
172
+
173
+ def _loss(self, logits, targets, mask):
174
+ batch, seq, vocab = logits.shape
175
+
176
+ if mask is None:
177
+ return self.loss_fn(logits.reshape(batch * seq, vocab), targets.reshape(batch * seq))
178
+
179
+ return masked_cross_entropy(logits.reshape(batch * seq, vocab), targets.data, mask)
180
+
181
+ def _forward_backward(self, input_ids, targets, accum_steps, mask=None):
182
+ with self.profiler.section("host_prep"):
183
+ x = Tensor(input_ids, requires_grad=False)
184
+ y = Tensor(targets, requires_grad=False)
185
+
186
+ numerics = np.errstate(over="ignore", invalid="ignore") if self.scaler.enabled else contextlib.nullcontext()
187
+
188
+ with numerics:
189
+ with self.profiler.section("forward"), autocast(self.precision):
190
+ logits = self.model(x)
191
+ loss = self._loss(logits, y, mask)
192
+ scaled = loss * (self.scaler.loss_multiplier() / accum_steps)
193
+
194
+ with self.profiler.section("backward"):
195
+ scaled.backward(release=True)
196
+
197
+ return float(loss.data)
198
+
199
+ def _batch_source(self):
200
+ micro_bs = self.config.training.micro_batch_size
201
+ dataset = self.train_dataset
202
+
203
+ def factory():
204
+ for batch in dataset.batches(micro_bs):
205
+ yield batch, dataset.state_dict()
206
+
207
+ if self.config.data.prefetch and self.config.data.prefetch > 0:
208
+ return iter(PrefetchLoader(factory, prefetch_size=self.config.data.prefetch))
209
+
210
+ return factory()
211
+
212
+ def _limits_reached(self):
213
+ t = self.config.training
214
+
215
+ if t.max_tokens is not None and self.tokens_processed >= t.max_tokens:
216
+ return True
217
+
218
+ if t.total_steps and self.optimizer_steps >= t.total_steps:
219
+ return True
220
+
221
+ return False
222
+
223
+ def _epoch_limit_reached(self):
224
+ t = self.config.training
225
+
226
+ if t.max_epochs is not None:
227
+ return self.epoch >= t.max_epochs
228
+
229
+ if t.max_tokens is None and not t.total_steps:
230
+ return self.epoch >= 1
231
+
232
+ return False
233
+
234
+ def _next_batch(self, source):
235
+ with self.profiler.section("data_wait"):
236
+ try:
237
+ return next(source), source
238
+ except StopIteration:
239
+ pass
240
+
241
+ self.epoch += 1
242
+
243
+ if self._epoch_limit_reached() or self._limits_reached():
244
+ return None, source
245
+
246
+ self.train_dataset.start_new_epoch()
247
+ self._consumed_state = self.train_dataset.state_dict()
248
+
249
+ source = self._batch_source()
250
+
251
+ with self.profiler.section("data_wait"):
252
+ try:
253
+ return next(source), source
254
+ except StopIteration:
255
+ return None, source
256
+
257
+ def train(self):
258
+ params = self.model.parameters()
259
+ t = self.config.training
260
+ accum_steps = t.gradient_accumulation_steps
261
+
262
+ source = self._batch_source()
263
+
264
+ started = time.time()
265
+ tokens_at_start = self.tokens_processed
266
+
267
+ try:
268
+ while not self._stop_requested and not self._limits_reached():
269
+ self.optimizer.zero_grad(params)
270
+
271
+ step_loss = 0.0
272
+ micro_done = 0
273
+ exhausted = False
274
+
275
+ for _ in range(accum_steps):
276
+ item, source = self._next_batch(source)
277
+
278
+ if item is None:
279
+ exhausted = True
280
+ break
281
+
282
+ batch, state = item
283
+ input_ids, targets = batch[0], batch[1]
284
+ mask = batch[2] if len(batch) > 2 else None
285
+ self._consumed_state = state
286
+
287
+ step_loss += self._forward_backward(input_ids, targets, accum_steps, mask)
288
+ micro_done += 1
289
+
290
+ self.tokens_processed += int(input_ids.size)
291
+ self.target_tokens += int(input_ids.size) if mask is None else int(mask.sum())
292
+ self.examples_processed += int(input_ids.shape[0])
293
+
294
+ if micro_done == 0:
295
+ break
296
+
297
+ if micro_done < accum_steps:
298
+ for p in params:
299
+ if p.requires_grad:
300
+ p.grad *= accum_steps / micro_done
301
+
302
+ step_loss /= micro_done
303
+ self.global_step += micro_done
304
+
305
+ with self.profiler.section("optimizer"):
306
+ finite = self.scaler.unscale_and_check(params)
307
+ self.scaler.update(not finite)
308
+
309
+ if not finite:
310
+ self._log_skipped(step_loss)
311
+
312
+ if exhausted:
313
+ break
314
+
315
+ continue
316
+
317
+ grad_norm = self._clip_grad_norm(params, t.gradient_clip)
318
+ lr = self.scheduler.step()
319
+ self.optimizer.lr = lr
320
+ self.optimizer.step(params)
321
+
322
+ self.optimizer_steps += 1
323
+ self.last_loss = step_loss
324
+
325
+ elapsed = time.time() - started
326
+ tps = (self.tokens_processed - tokens_at_start) / max(elapsed, 1e-6)
327
+
328
+ self._log(step_loss, lr, grad_norm, tps, elapsed)
329
+
330
+ if self.global_step and self.optimizer_steps % self.config.checkpoint.interval_steps == 0:
331
+ self.save_checkpoint()
332
+
333
+ self._maybe_evaluate()
334
+
335
+ if exhausted:
336
+ break
337
+ finally:
338
+ if self.save_on_exit and self._last_saved_step != self.optimizer_steps:
339
+ self.save_checkpoint()
340
+
341
+ self.experiment.update_metrics({
342
+ "optimizer_steps": self.optimizer_steps,
343
+ "tokens_processed": self.tokens_processed,
344
+ "epoch": self.epoch,
345
+ "last_loss": self.last_loss,
346
+ "last_eval": self.last_eval,
347
+ "profile": self.profiler.summary(),
348
+ })
349
+
350
+ def _maybe_evaluate(self):
351
+ if self.val_dataset is None:
352
+ return
353
+
354
+ e = self.config.evaluation
355
+ due = False
356
+
357
+ if e.interval_steps and self.optimizer_steps % e.interval_steps == 0:
358
+ due = True
359
+
360
+ if e.interval_tokens and self.tokens_processed - self._tokens_at_last_eval >= e.interval_tokens:
361
+ due = True
362
+
363
+ if due:
364
+ self._tokens_at_last_eval = self.tokens_processed
365
+ self.evaluate()
366
+
367
+ def _log(self, loss, lr, grad_norm, tps, elapsed):
368
+ t = self.config.training
369
+ remaining = None
370
+
371
+ if t.max_tokens:
372
+ remaining = max(t.max_tokens - self.tokens_processed, 0)
373
+ eta = remaining / tps if remaining is not None and tps > 0 else None
374
+
375
+ record = {
376
+ "step": self.optimizer_steps,
377
+ "micro_steps": self.global_step,
378
+ "epoch": self.epoch,
379
+ "tokens": self.tokens_processed,
380
+ "target_tokens": self.target_tokens,
381
+ "tokens_remaining": remaining,
382
+ "examples": self.examples_processed,
383
+ "loss": loss,
384
+ "lr": lr,
385
+ "grad_norm": grad_norm,
386
+ "tokens_per_sec": tps,
387
+ "samples_per_sec": self.examples_processed / max(elapsed, 1e-6),
388
+ "elapsed_s": elapsed,
389
+ "eta_s": eta,
390
+ "max_rss_mb": resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024,
391
+ "profile_s": self.profiler.take_window(),
392
+ "precision": self.precision,
393
+ "loss_scale": self.scaler.scale if self.scaler.enabled else None,
394
+ "skipped_steps": self.scaler.skipped_steps,
395
+ "time": time.time(),
396
+ }
397
+
398
+ with open(self.log_path, "a") as f:
399
+ f.write(json.dumps(record) + "\n")
400
+
401
+ eta_txt = f"{eta:.0f}s" if eta is not None else "n/a"
402
+
403
+ print(
404
+ f"step {self.optimizer_steps} | loss {loss:.4f} | lr {lr:.2e} "
405
+ f"| gnorm {grad_norm:.3f} | tok/s {tps:.0f} | tokens {self.tokens_processed} "
406
+ f"| eta {eta_txt}"
407
+ )
408
+
409
+ def _log_skipped(self, loss):
410
+ record = {
411
+ "step": self.optimizer_steps,
412
+ "micro_steps": self.global_step,
413
+ "tokens": self.tokens_processed,
414
+ "loss": loss,
415
+ "skipped": "non_finite_gradients",
416
+ "loss_scale": self.scaler.scale,
417
+ "skipped_steps": self.scaler.skipped_steps,
418
+ "time": time.time(),
419
+ }
420
+
421
+ with open(self.log_path, "a") as f:
422
+ f.write(json.dumps(record) + "\n")
423
+
424
+ print(f"step {self.optimizer_steps} skipped: non-finite gradients; loss scale now {self.scaler.scale:g}")
425
+
426
+ def initialize_from(self, path):
427
+ from src.training.checkpoint_manager import read_checkpoint
428
+
429
+ payload = read_checkpoint(path)
430
+
431
+ if not identity_matches(payload["tokenizer_identity"], self.tokenizer.identity):
432
+ raise ValueError("init_from checkpoint was trained with a different tokenizer")
433
+
434
+ config = self.model.config
435
+ source = type(config).from_dict(payload["model_config"])
436
+ problems = describe_incompatibility(source, config)
437
+
438
+ if problems:
439
+ raise ValueError(f"init_from checkpoint is incompatible with this model: {problems}")
440
+
441
+ self.model.load_state_dict(payload["model_state"])
442
+ self.init_source = {
443
+ "path": os.path.abspath(path),
444
+ "run_id": payload.get("run_id"),
445
+ "optimizer_steps": payload.get("optimizer_steps"),
446
+ "tokens_processed": payload.get("tokens_processed"),
447
+ "context_length": source.context_length,
448
+ }
449
+ print(f"initialized weights from {path} (step {payload.get('optimizer_steps')}); "
450
+ f"optimizer, schedule and data start fresh")
451
+ return self.init_source
452
+
453
+ def evaluate(self):
454
+ if self.val_dataset is None:
455
+ return None
456
+
457
+ self.val_dataset.start_new_epoch()
458
+
459
+ params = self.model.parameters()
460
+ flags = [p.requires_grad for p in params]
461
+
462
+ for p in params:
463
+ p.requires_grad = False
464
+
465
+ was_training = self.model.set_training(False)
466
+
467
+ total_loss = 0.0
468
+ total_tokens = 0
469
+ batches = 0
470
+
471
+ try:
472
+ for val_batch in self.val_dataset.batches(self.config.training.micro_batch_size):
473
+ input_ids, targets = val_batch[0], val_batch[1]
474
+ mask = val_batch[2] if len(val_batch) > 2 else None
475
+
476
+ with no_grad(), autocast(self.precision):
477
+ logits = self.model(Tensor(input_ids, requires_grad=False))
478
+ loss = self._loss(logits, Tensor(targets, requires_grad=False), mask)
479
+
480
+ count = int(input_ids.size) if mask is None else int(mask.sum())
481
+ total_loss += float(loss.data) * count
482
+ total_tokens += count
483
+ batches += 1
484
+
485
+ if batches >= self.config.evaluation.max_batches:
486
+ break
487
+ finally:
488
+ self.model.set_training(was_training)
489
+
490
+ for p, flag in zip(params, flags):
491
+ p.requires_grad = flag
492
+
493
+ if batches == 0 or total_tokens == 0:
494
+ return None
495
+
496
+ avg = total_loss / total_tokens
497
+ ppl = float(np.exp(min(avg, 20.0)))
498
+
499
+ self.last_eval = {"loss": avg, "perplexity": ppl, "tokens": total_tokens}
500
+
501
+ print(f"eval | step {self.optimizer_steps} | loss {avg:.4f} | ppl {ppl:.2f}")
502
+
503
+ with open(self.log_path, "a") as f:
504
+ f.write(json.dumps({
505
+ "step": self.optimizer_steps,
506
+ "eval_loss": avg,
507
+ "eval_perplexity": ppl,
508
+ "eval_tokens": total_tokens,
509
+ "time": time.time(),
510
+ }) + "\n")
511
+
512
+ return avg, ppl
513
+
514
+ def save_checkpoint(self):
515
+ with self.profiler.section("checkpoint"):
516
+ dataset_state = self._consumed_state or self.train_dataset.state_dict()
517
+
518
+ payload = {
519
+ "framework_version": FRAMEWORK_VERSION,
520
+ "run_id": self.experiment.run_id,
521
+ "global_step": self.global_step,
522
+ "optimizer_steps": self.optimizer_steps,
523
+ "epoch": self.epoch,
524
+ "tokens_processed": self.tokens_processed,
525
+ "target_tokens": self.target_tokens,
526
+ "init_from": self.init_source,
527
+ "examples_processed": self.examples_processed,
528
+ "last_loss": self.last_loss,
529
+ "model_config": self.model.config.to_dict(),
530
+ "model_state": self.model.state_dict(),
531
+ "optimizer_state": self.optimizer.state_dict(),
532
+ "scheduler_state": self.scheduler.state_dict(),
533
+ "grad_scaler_state": self.scaler.state_dict(),
534
+ "dataset_state": dataset_state,
535
+ "corpus_state": self.train_dataset.corpus.state_dict(),
536
+ "tokenizer_identity": self.tokenizer.identity,
537
+ "training_config": asdict(self.config.training),
538
+ "data_config": asdict(self.config.data),
539
+ "python_rng_state": random.getstate(),
540
+ "numpy_rng_state": np.random.get_state(),
541
+ }
542
+
543
+ path = self.checkpoints.save(self.optimizer_steps, payload)
544
+
545
+ self._tokens_at_last_ckpt = self.tokens_processed
546
+ self._last_saved_step = self.optimizer_steps
547
+
548
+ print(f"checkpoint saved: {path}")
549
+
550
+ return path
551
+
552
+ def load_checkpoint(self, path=None):
553
+ payload = self.checkpoints.load(path)
554
+
555
+ if payload is None:
556
+ return False
557
+
558
+ if not identity_matches(payload["tokenizer_identity"], self.tokenizer.identity):
559
+ raise ValueError(
560
+ "checkpoint tokenizer identity does not match the tokenizer "
561
+ "passed to this run; resuming would corrupt training"
562
+ )
563
+
564
+ if type(self.model.config).normalized(payload["model_config"]) != self.model.config.to_dict():
565
+ raise ValueError("checkpoint model architecture does not match the configured model")
566
+
567
+ self.model.build()
568
+ self.model.load_state_dict(payload["model_state"])
569
+
570
+ self.optimizer.load_state_dict(payload["optimizer_state"])
571
+ self.scheduler.load_state_dict(payload["scheduler_state"])
572
+
573
+ if "grad_scaler_state" in payload:
574
+ self.scaler.load_state_dict(payload["grad_scaler_state"])
575
+
576
+ self.train_dataset.corpus.load_state_dict(payload["corpus_state"])
577
+ self.train_dataset.load_state_dict(payload["dataset_state"])
578
+ self._consumed_state = payload["dataset_state"]
579
+
580
+ self.global_step = payload["global_step"]
581
+ self.optimizer_steps = payload["optimizer_steps"]
582
+ self.epoch = payload["epoch"]
583
+ self.tokens_processed = payload["tokens_processed"]
584
+ self.target_tokens = payload.get("target_tokens", self.tokens_processed)
585
+ self.init_source = payload.get("init_from")
586
+ self.examples_processed = payload["examples_processed"]
587
+ self.last_loss = payload["last_loss"]
588
+
589
+ self._tokens_at_last_eval = self.tokens_processed
590
+
591
+ random.setstate(payload["python_rng_state"])
592
+ np.random.set_state(payload["numpy_rng_state"])
593
+
594
+ return True