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
src/training/trainer.py
ADDED
|
@@ -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
|