tensorless 0.9.3__tar.gz → 0.9.4__tar.gz
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.
- {tensorless-0.9.3/tensorless.egg-info → tensorless-0.9.4}/PKG-INFO +1 -1
- {tensorless-0.9.3 → tensorless-0.9.4}/pyproject.toml +1 -1
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/backends/jax_backend.py +230 -3
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/backends/mlx_backend.py +20 -1
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/devices/device.py +26 -2
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/trainer.py +76 -10
- {tensorless-0.9.3 → tensorless-0.9.4/tensorless.egg-info}/PKG-INFO +1 -1
- {tensorless-0.9.3 → tensorless-0.9.4}/LICENSE +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/MANIFEST.in +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/README.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/api_reference.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/architecture.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/automatic_mode.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/checkpointing.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/cli.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/configuration.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/contributing.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/examples.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/inference.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/installation.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/limitations.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/quickstart.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/roadmap.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/tl_format.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/training.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/troubleshooting.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/docs/tutorial.md +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/examples/tabular_classification_example.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/examples/tabular_regression_example.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/examples/text_classification_example.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/examples/text_generation_example.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/setup.cfg +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/_version.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/api.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/auto/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/auto/config.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/auto/detector.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/backends/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/checkpoint/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/checkpoint/manager.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/cli/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/cli/main.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/config.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/english_grammar.txt +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/fingerprint.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/inspector.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/loader.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/tabular.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/devices/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/devices/memory.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/engine.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/errors.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/mlp.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/registry.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/transformer.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/runtime.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/serialization/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/serialization/tl_format.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/tokenization/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/tokenization/bpe_tokenizer.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/tokenization/char_tokenizer.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/__init__.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/data_prep.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/early_stopping.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/SOURCES.txt +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/dependency_links.txt +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/entry_points.txt +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/requires.txt +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/top_level.txt +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_auto_detection.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_checkpoint_resume.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_cli.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_data_loading.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_end_to_end.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_fingerprint.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_jax_backend.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_mlx_backend.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_serialization.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_train_tabular.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_train_text_classification.py +0 -0
- {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_train_text_generation.py +0 -0
|
@@ -39,7 +39,97 @@ def is_available() -> bool:
|
|
|
39
39
|
return False
|
|
40
40
|
|
|
41
41
|
|
|
42
|
-
|
|
42
|
+
def supports_fused_training(model) -> bool:
|
|
43
|
+
"""Whether `model` can use the device-resident fused AdamW/Adam train
|
|
44
|
+
step (see `_JaxFusedAdamMixin`) instead of round-tripping full
|
|
45
|
+
parameter/gradient dicts through host NumPy every step."""
|
|
46
|
+
return isinstance(model, _JaxFusedAdamMixin)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _adam_update_tree(jnp, params, m, v, grads, step_count, lr, weight_decay, decoupled):
|
|
50
|
+
"""Fused AdamW/Adam parameter update over a pytree dict of device arrays.
|
|
51
|
+
|
|
52
|
+
Runs entirely inside `jax.jit` as part of the fused train step below, so
|
|
53
|
+
params/moments never have to leave the accelerator to be updated.
|
|
54
|
+
"""
|
|
55
|
+
beta1_correction = 1.0 - 0.9 ** step_count
|
|
56
|
+
beta2_correction = 1.0 - 0.999 ** step_count
|
|
57
|
+
new_params, new_m, new_v = {}, {}, {}
|
|
58
|
+
for name, p in params.items():
|
|
59
|
+
g = grads[name]
|
|
60
|
+
if decoupled:
|
|
61
|
+
p = p * (1.0 - lr * weight_decay)
|
|
62
|
+
else:
|
|
63
|
+
g = g + weight_decay * p
|
|
64
|
+
m_i = 0.9 * m[name] + 0.1 * g
|
|
65
|
+
v_i = 0.999 * v[name] + 0.001 * g * g
|
|
66
|
+
m_hat = m_i / beta1_correction
|
|
67
|
+
v_hat = v_i / beta2_correction
|
|
68
|
+
new_params[name] = p - lr * m_hat / (jnp.sqrt(v_hat) + 1e-8)
|
|
69
|
+
new_m[name] = m_i
|
|
70
|
+
new_v[name] = v_i
|
|
71
|
+
return new_params, new_m, new_v
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class _JaxFusedAdamMixin:
|
|
75
|
+
"""Keeps parameters + Adam moment estimates resident on the GPU/TPU
|
|
76
|
+
across many training steps, instead of re-uploading the full parameter
|
|
77
|
+
dict and downloading the full gradient dict through host NumPy on
|
|
78
|
+
*every single step*.
|
|
79
|
+
|
|
80
|
+
Compiling the forward pass with `jax.jit` (see `_get_forward_jit`) fixes
|
|
81
|
+
re-tracing overhead, but each `loss_and_backward()` call still does two
|
|
82
|
+
full host<->device transfers of the whole model (params up, grads down),
|
|
83
|
+
and the NumPy `Adam`/`AdamW` optimizer in `engine.py` then runs the
|
|
84
|
+
update on the CPU. For small/medium models that transfer, not compute,
|
|
85
|
+
dominates the step time and is why GPU utilization can stay low even
|
|
86
|
+
after JIT is fixed. This mixin fuses forward+backward+optimizer-update
|
|
87
|
+
into one jitted function so state only crosses the host/device boundary
|
|
88
|
+
at explicit sync points (checkpointing, validation, end of training) --
|
|
89
|
+
see `training/trainer.py`'s use of `init_fused_state`/`sync_fused_to_host`.
|
|
90
|
+
"""
|
|
91
|
+
|
|
92
|
+
_fused_params = None
|
|
93
|
+
_fused_m = None
|
|
94
|
+
_fused_v = None
|
|
95
|
+
_fused_step = 0
|
|
96
|
+
_fused_train_step_key = None
|
|
97
|
+
|
|
98
|
+
def fused_training_supported(self) -> bool:
|
|
99
|
+
return True
|
|
100
|
+
|
|
101
|
+
def init_fused_state(self, weight_decay: float, decoupled: bool, step: int = 0):
|
|
102
|
+
_, jnp = _jax()
|
|
103
|
+
params = self._jax_params()
|
|
104
|
+
self._fused_params = params
|
|
105
|
+
self._fused_m = {name: jnp.zeros_like(value) for name, value in params.items()}
|
|
106
|
+
self._fused_v = {name: jnp.zeros_like(value) for name, value in params.items()}
|
|
107
|
+
self._fused_step = int(step)
|
|
108
|
+
self._fused_train_step_key = None
|
|
109
|
+
|
|
110
|
+
def load_fused_optimizer_moments(self, m_by_name, v_by_name):
|
|
111
|
+
"""Restore Adam moment estimates (as NumPy, from a resumed checkpoint)."""
|
|
112
|
+
_, jnp = _jax()
|
|
113
|
+
self._fused_m = {name: jnp.asarray(value, dtype=jnp.float32) for name, value in m_by_name.items()}
|
|
114
|
+
self._fused_v = {name: jnp.asarray(value, dtype=jnp.float32) for name, value in v_by_name.items()}
|
|
115
|
+
|
|
116
|
+
def sync_fused_to_host(self):
|
|
117
|
+
"""Write device-resident params/moments back to host NumPy.
|
|
118
|
+
|
|
119
|
+
Called only at checkpoint/validation/end-of-training boundaries, so
|
|
120
|
+
this cost is paid a handful of times per epoch rather than every step.
|
|
121
|
+
"""
|
|
122
|
+
if self._fused_params is None:
|
|
123
|
+
return {}, {}
|
|
124
|
+
for name, parameter in self.named_parameters():
|
|
125
|
+
parameter.data[...] = np.asarray(self._fused_params[name], dtype=np.float32)
|
|
126
|
+
m_by_name = {name: np.asarray(value, dtype=np.float32) for name, value in self._fused_m.items()}
|
|
127
|
+
v_by_name = {name: np.asarray(value, dtype=np.float32) for name, value in self._fused_v.items()}
|
|
128
|
+
return m_by_name, v_by_name
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
class JaxTinyTransformer(_JaxFusedAdamMixin, TinyTransformer):
|
|
43
133
|
"""The native transformer parameter layout with JAX math and gradients."""
|
|
44
134
|
|
|
45
135
|
def __init__(self, *args, **kwargs):
|
|
@@ -188,16 +278,98 @@ class JaxTinyTransformer(TinyTransformer):
|
|
|
188
278
|
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
|
|
189
279
|
return float(loss) / loss_scale
|
|
190
280
|
|
|
191
|
-
def
|
|
281
|
+
def _get_fused_train_step(self, weight_decay: float, decoupled: bool):
|
|
282
|
+
cache_key = (weight_decay, decoupled)
|
|
283
|
+
if self._fused_train_step_key != cache_key:
|
|
284
|
+
jax, jnp = _jax()
|
|
285
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
286
|
+
|
|
287
|
+
def loss_fn(params, ids, target, mask):
|
|
288
|
+
forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
|
|
289
|
+
logits = forward(params, ids, mask)
|
|
290
|
+
flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
|
|
291
|
+
flat_target = target.reshape(-1)
|
|
292
|
+
log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
|
|
293
|
+
losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
|
|
294
|
+
if self.task == "text-generation":
|
|
295
|
+
valid = flat_target != self.pad_id
|
|
296
|
+
loss = jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
|
|
297
|
+
else:
|
|
298
|
+
loss = jnp.mean(losses)
|
|
299
|
+
return loss * loss_scale
|
|
300
|
+
|
|
301
|
+
grad_fn = jax.value_and_grad(loss_fn)
|
|
302
|
+
|
|
303
|
+
def step(params, m, v, step_count, ids, target, mask, lr, grad_clip):
|
|
304
|
+
loss, grads = grad_fn(params, ids, target, mask)
|
|
305
|
+
grads = {name: g / loss_scale for name, g in grads.items()}
|
|
306
|
+
grad_norm = jnp.sqrt(sum(jnp.sum(g.astype(jnp.float32) ** 2) for g in grads.values()))
|
|
307
|
+
clip_scale = jnp.where(grad_clip > 0, jnp.minimum(1.0, grad_clip / (grad_norm + 1e-12)), 1.0)
|
|
308
|
+
grads = {name: g * clip_scale for name, g in grads.items()}
|
|
309
|
+
new_params, new_m, new_v = _adam_update_tree(
|
|
310
|
+
jnp, params, m, v, grads, step_count, lr, weight_decay, decoupled
|
|
311
|
+
)
|
|
312
|
+
return new_params, new_m, new_v, loss / loss_scale, grad_norm
|
|
313
|
+
|
|
314
|
+
self._fused_train_step_fn = jax.jit(step)
|
|
315
|
+
self._fused_train_step_key = cache_key
|
|
316
|
+
return self._fused_train_step_fn
|
|
317
|
+
|
|
318
|
+
def fused_train_step(self, input_ids, target, attention_mask, lr, weight_decay, grad_clip, decoupled):
|
|
319
|
+
"""One fused forward+backward+AdamW step, entirely on-device.
|
|
320
|
+
|
|
321
|
+
Call `init_fused_state()` once before the first step and
|
|
322
|
+
`sync_fused_to_host()` before reading `state_dict()`/`forward()` on
|
|
323
|
+
the host (e.g. for validation or checkpointing).
|
|
324
|
+
"""
|
|
325
|
+
if self._fused_params is None:
|
|
326
|
+
raise RuntimeError("init_fused_state() must be called before fused_train_step().")
|
|
327
|
+
jax, jnp = _jax()
|
|
328
|
+
step_fn = self._get_fused_train_step(weight_decay, decoupled)
|
|
329
|
+
self._fused_step += 1
|
|
330
|
+
ids_arr = jnp.asarray(input_ids, dtype=jnp.int32)
|
|
331
|
+
target_arr = jnp.asarray(target, dtype=jnp.int32)
|
|
332
|
+
mask_arr = jnp.ones_like(ids_arr, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask, dtype=jnp.float32)
|
|
333
|
+
new_params, new_m, new_v, loss, grad_norm = step_fn(
|
|
334
|
+
self._fused_params, self._fused_m, self._fused_v,
|
|
335
|
+
jnp.asarray(self._fused_step, dtype=jnp.float32),
|
|
336
|
+
ids_arr, target_arr, mask_arr,
|
|
337
|
+
jnp.asarray(lr, dtype=jnp.float32), jnp.asarray(grad_clip, dtype=jnp.float32),
|
|
338
|
+
)
|
|
339
|
+
self._fused_params, self._fused_m, self._fused_v = new_params, new_m, new_v
|
|
340
|
+
return float(loss), float(grad_norm)
|
|
341
|
+
|
|
342
|
+
def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, top_p=.9,
|
|
343
|
+
repetition_penalty=1.1, eos_id: Optional[int] = None):
|
|
344
|
+
# Mirrors `TinyTransformer.generate` (the NumPy path) exactly, so
|
|
345
|
+
# generation quality doesn't quietly regress just because a model
|
|
346
|
+
# happens to run on cuda/tpu. Sampling itself stays in NumPy since
|
|
347
|
+
# it's inherently sequential and tiny relative to the forward pass;
|
|
348
|
+
# only the forward pass benefits from JIT.
|
|
192
349
|
ids = np.asarray(input_ids, dtype=np.int64).copy()
|
|
193
350
|
for _ in range(max_new_tokens):
|
|
194
351
|
logits = self.forward(ids[:, -self.max_seq_len:])[:, -1, :] / max(temperature, 1e-5)
|
|
352
|
+
if repetition_penalty and repetition_penalty > 1.0:
|
|
353
|
+
for row, sequence in enumerate(ids):
|
|
354
|
+
seen = np.unique(sequence)
|
|
355
|
+
logits[row, seen] = np.where(
|
|
356
|
+
logits[row, seen] < 0,
|
|
357
|
+
logits[row, seen] * repetition_penalty,
|
|
358
|
+
logits[row, seen] / repetition_penalty,
|
|
359
|
+
)
|
|
195
360
|
if top_k is not None:
|
|
196
361
|
k = min(top_k, logits.shape[-1])
|
|
197
362
|
excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
|
|
198
363
|
logits[np.arange(len(ids))[:, None], excluded] = -np.inf
|
|
199
364
|
probs = np.exp(logits - logits.max(1, keepdims=True))
|
|
200
365
|
probs /= probs.sum(1, keepdims=True)
|
|
366
|
+
if top_p is not None and 0 < top_p < 1:
|
|
367
|
+
order = np.argsort(-probs, axis=1)
|
|
368
|
+
sorted_probs = np.take_along_axis(probs, order, axis=1)
|
|
369
|
+
cutoff = np.cumsum(sorted_probs, axis=1) > top_p
|
|
370
|
+
cutoff[:, 0] = False
|
|
371
|
+
probs[np.arange(len(ids))[:, None], order] = np.where(cutoff, 0.0, sorted_probs)
|
|
372
|
+
probs /= probs.sum(1, keepdims=True)
|
|
201
373
|
next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
|
|
202
374
|
ids = np.concatenate([ids, next_ids], axis=1)
|
|
203
375
|
if eos_id is not None and np.all(next_ids == eos_id):
|
|
@@ -205,7 +377,7 @@ class JaxTinyTransformer(TinyTransformer):
|
|
|
205
377
|
return ids
|
|
206
378
|
|
|
207
379
|
|
|
208
|
-
class JaxTabularMLP(TabularMLP):
|
|
380
|
+
class JaxTabularMLP(_JaxFusedAdamMixin, TabularMLP):
|
|
209
381
|
"""The native tabular-MLP parameter layout with JAX math and gradients.
|
|
210
382
|
|
|
211
383
|
Used automatically instead of the plain NumPy `TabularMLP` whenever the
|
|
@@ -289,3 +461,58 @@ class JaxTabularMLP(TabularMLP):
|
|
|
289
461
|
for name, parameter in self.named_parameters():
|
|
290
462
|
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
|
|
291
463
|
return float(loss)
|
|
464
|
+
|
|
465
|
+
def _get_fused_train_step(self, weight_decay: float, decoupled: bool):
|
|
466
|
+
cache_key = (weight_decay, decoupled)
|
|
467
|
+
if self._fused_train_step_key != cache_key:
|
|
468
|
+
jax, jnp = _jax()
|
|
469
|
+
is_regression = self.task == "regression"
|
|
470
|
+
|
|
471
|
+
def loss_fn(params, numeric, categorical, dropout_masks, target_arr):
|
|
472
|
+
logits = self._forward_jax(params, numeric, categorical, dropout_masks)
|
|
473
|
+
if is_regression:
|
|
474
|
+
error = logits - target_arr
|
|
475
|
+
return jnp.mean(error * error)
|
|
476
|
+
log_probs = jax.nn.log_softmax(logits, axis=-1)
|
|
477
|
+
return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
|
|
478
|
+
|
|
479
|
+
grad_fn = jax.value_and_grad(loss_fn)
|
|
480
|
+
|
|
481
|
+
def step(params, m, v, step_count, numeric, categorical, dropout_masks, target_arr, lr, grad_clip):
|
|
482
|
+
loss, grads = grad_fn(params, numeric, categorical, dropout_masks, target_arr)
|
|
483
|
+
grad_norm = jnp.sqrt(sum(jnp.sum(g.astype(jnp.float32) ** 2) for g in grads.values()))
|
|
484
|
+
clip_scale = jnp.where(grad_clip > 0, jnp.minimum(1.0, grad_clip / (grad_norm + 1e-12)), 1.0)
|
|
485
|
+
grads = {name: g * clip_scale for name, g in grads.items()}
|
|
486
|
+
new_params, new_m, new_v = _adam_update_tree(
|
|
487
|
+
jnp, params, m, v, grads, step_count, lr, weight_decay, decoupled
|
|
488
|
+
)
|
|
489
|
+
return new_params, new_m, new_v, loss, grad_norm
|
|
490
|
+
|
|
491
|
+
self._fused_train_step_fn = jax.jit(step)
|
|
492
|
+
self._fused_train_step_key = cache_key
|
|
493
|
+
return self._fused_train_step_fn
|
|
494
|
+
|
|
495
|
+
def fused_train_step(self, numeric, categorical, target, lr, weight_decay, grad_clip, decoupled):
|
|
496
|
+
"""One fused forward+backward+AdamW step, entirely on-device.
|
|
497
|
+
|
|
498
|
+
Call `init_fused_state()` once before the first step and
|
|
499
|
+
`sync_fused_to_host()` before reading `state_dict()`/`forward()` on
|
|
500
|
+
the host (e.g. for validation or checkpointing).
|
|
501
|
+
"""
|
|
502
|
+
if self._fused_params is None:
|
|
503
|
+
raise RuntimeError("init_fused_state() must be called before fused_train_step().")
|
|
504
|
+
jax, jnp = _jax()
|
|
505
|
+
step_fn = self._get_fused_train_step(weight_decay, decoupled)
|
|
506
|
+
self._fused_step += 1
|
|
507
|
+
batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
|
|
508
|
+
dropout_masks = self._make_dropout_masks(batch_size)
|
|
509
|
+
is_regression = self.task == "regression"
|
|
510
|
+
target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
|
|
511
|
+
new_params, new_m, new_v, loss, grad_norm = step_fn(
|
|
512
|
+
self._fused_params, self._fused_m, self._fused_v,
|
|
513
|
+
jnp.asarray(self._fused_step, dtype=jnp.float32),
|
|
514
|
+
numeric, categorical, dropout_masks, target_arr,
|
|
515
|
+
jnp.asarray(lr, dtype=jnp.float32), jnp.asarray(grad_clip, dtype=jnp.float32),
|
|
516
|
+
)
|
|
517
|
+
self._fused_params, self._fused_m, self._fused_v = new_params, new_m, new_v
|
|
518
|
+
return float(loss), float(grad_norm)
|
|
@@ -117,16 +117,35 @@ class MlxTinyTransformer(TinyTransformer):
|
|
|
117
117
|
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32) / loss_scale
|
|
118
118
|
return float(np.asarray(loss)) / loss_scale
|
|
119
119
|
|
|
120
|
-
def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40,
|
|
120
|
+
def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, top_p=.9,
|
|
121
|
+
repetition_penalty=1.1, eos_id: Optional[int] = None):
|
|
122
|
+
# Mirrors `TinyTransformer.generate` (the NumPy path) exactly, so
|
|
123
|
+
# generation quality doesn't quietly regress just because a model
|
|
124
|
+
# happens to run on mps.
|
|
121
125
|
ids = np.asarray(input_ids, dtype=np.int64).copy()
|
|
122
126
|
for _ in range(max_new_tokens):
|
|
123
127
|
logits = self.forward(ids[:, -self.max_seq_len:])[:, -1, :] / max(temperature, 1e-5)
|
|
128
|
+
if repetition_penalty and repetition_penalty > 1.0:
|
|
129
|
+
for row, sequence in enumerate(ids):
|
|
130
|
+
seen = np.unique(sequence)
|
|
131
|
+
logits[row, seen] = np.where(
|
|
132
|
+
logits[row, seen] < 0,
|
|
133
|
+
logits[row, seen] * repetition_penalty,
|
|
134
|
+
logits[row, seen] / repetition_penalty,
|
|
135
|
+
)
|
|
124
136
|
if top_k is not None:
|
|
125
137
|
k = min(top_k, logits.shape[-1])
|
|
126
138
|
excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
|
|
127
139
|
logits[np.arange(len(ids))[:, None], excluded] = -np.inf
|
|
128
140
|
probs = np.exp(logits - logits.max(1, keepdims=True))
|
|
129
141
|
probs /= probs.sum(1, keepdims=True)
|
|
142
|
+
if top_p is not None and 0 < top_p < 1:
|
|
143
|
+
order = np.argsort(-probs, axis=1)
|
|
144
|
+
sorted_probs = np.take_along_axis(probs, order, axis=1)
|
|
145
|
+
cutoff = np.cumsum(sorted_probs, axis=1) > top_p
|
|
146
|
+
cutoff[:, 0] = False
|
|
147
|
+
probs[np.arange(len(ids))[:, None], order] = np.where(cutoff, 0.0, sorted_probs)
|
|
148
|
+
probs /= probs.sum(1, keepdims=True)
|
|
130
149
|
next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
|
|
131
150
|
ids = np.concatenate([ids, next_ids], axis=1)
|
|
132
151
|
if eos_id is not None and np.all(next_ids == eos_id):
|
|
@@ -37,9 +37,33 @@ def _mps_available() -> bool:
|
|
|
37
37
|
return False
|
|
38
38
|
|
|
39
39
|
|
|
40
|
+
_BF16_CAPABLE_GPU_HINTS = (
|
|
41
|
+
"a100", "a10", "a30", "a40", "l4", "l40", "l20",
|
|
42
|
+
"h100", "h200", "h800", "b100", "b200", "gh200",
|
|
43
|
+
"rtx 30", "rtx 40", "rtx 50", "rtx a", "geforce rtx 30", "geforce rtx 40", "geforce rtx 50",
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
|
|
40
47
|
def _cuda_supports_bf16() -> bool:
|
|
41
|
-
"""
|
|
42
|
-
|
|
48
|
+
"""Best-effort detection of real (fast) bf16 tensor-core support.
|
|
49
|
+
|
|
50
|
+
Older GPUs (Turing/Volta/Pascal -- T4, V100, P100, K80, RTX 20xx) either
|
|
51
|
+
lack bf16 tensor cores or emulate bf16 slowly, so fp16 (with loss
|
|
52
|
+
scaling) is the safer/faster default there. Ampere-or-newer GPUs (A100,
|
|
53
|
+
L4, H100, RTX 30xx+, ...) run bf16 at full tensor-core speed and don't
|
|
54
|
+
need loss scaling, so we opt in automatically when the detected GPU name
|
|
55
|
+
matches a known Ampere-or-newer architecture. Anything unrecognized
|
|
56
|
+
conservatively falls back to fp16, matching the previous behavior.
|
|
57
|
+
"""
|
|
58
|
+
try:
|
|
59
|
+
import jax
|
|
60
|
+
gpus = [d for d in jax.devices() if d.platform == "gpu"]
|
|
61
|
+
if not gpus:
|
|
62
|
+
return False
|
|
63
|
+
name = getattr(gpus[0], "device_kind", "").lower()
|
|
64
|
+
return any(hint in name for hint in _BF16_CAPABLE_GPU_HINTS)
|
|
65
|
+
except (ImportError, RuntimeError):
|
|
66
|
+
return False
|
|
43
67
|
|
|
44
68
|
|
|
45
69
|
def auto_select_device(user_device: Optional[str], user_precision: Optional[str]) -> Tuple[str, str]:
|
|
@@ -86,6 +86,26 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
|
|
|
86
86
|
"Keep the task, tokenizer, and architecture dimensions compatible."
|
|
87
87
|
) from exc
|
|
88
88
|
optimizer = _build_optimizer(model, cfg)
|
|
89
|
+
|
|
90
|
+
# Fused device-resident training: for jax models trained with Adam/AdamW
|
|
91
|
+
# on a single accelerator, keep params + optimizer moments on the
|
|
92
|
+
# GPU/TPU for the whole run instead of re-uploading the full parameter
|
|
93
|
+
# dict and downloading the full gradient dict through host NumPy on
|
|
94
|
+
# every step (what the plain `Adam`/`AdamW` classes in `engine.py` force
|
|
95
|
+
# `loss_and_backward` to do). This is what actually keeps the
|
|
96
|
+
# accelerator busy once JIT tracing overhead is no longer the
|
|
97
|
+
# bottleneck. Multi-device (pmap) runs and non-Adam optimizers keep
|
|
98
|
+
# using the existing host-side path.
|
|
99
|
+
use_fused = False
|
|
100
|
+
fused_decoupled = cfg["optimizer"].lower() == "adamw"
|
|
101
|
+
if backend == "jax" and cfg["optimizer"].lower() in ("adam", "adamw"):
|
|
102
|
+
from ..backends.jax_backend import supports_fused_training
|
|
103
|
+
if supports_fused_training(model):
|
|
104
|
+
try:
|
|
105
|
+
import jax
|
|
106
|
+
use_fused = jax.local_device_count() == 1
|
|
107
|
+
except ImportError:
|
|
108
|
+
use_fused = False
|
|
89
109
|
total_steps = cfg.get("max_steps") or max(1, len(prepared.train_loader)) * cfg["epochs"]
|
|
90
110
|
scheduler = LambdaScheduler(optimizer, cfg["warmup_steps"], total_steps)
|
|
91
111
|
early_stopper = EarlyStopping(patience=cfg["patience"], min_delta=cfg["min_delta"])
|
|
@@ -98,7 +118,30 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
|
|
|
98
118
|
early_stopper.best = resume_state.get("early_stopping_best", float("inf")); early_stopper.num_bad_checks = resume_state.get("early_stopping_bad_checks", 0)
|
|
99
119
|
if resume_state.get("best_model_state_dict") is not None:
|
|
100
120
|
best_model_state = resume_state["best_model_state_dict"]
|
|
121
|
+
|
|
122
|
+
if use_fused:
|
|
123
|
+
param_names = [name for name, _ in model.named_parameters()]
|
|
124
|
+
model.init_fused_state(cfg["weight_decay"], fused_decoupled, step=optimizer.step_count)
|
|
125
|
+
if resume_state:
|
|
126
|
+
model.load_fused_optimizer_moments(
|
|
127
|
+
dict(zip(param_names, optimizer.m)), dict(zip(param_names, optimizer.v))
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
def sync_fused():
|
|
131
|
+
"""Bring device-resident params/moments back to host NumPy so
|
|
132
|
+
`model.state_dict()`/`model.forward()`/`optimizer.state_dict()`
|
|
133
|
+
reflect the latest step. Only called at checkpoint/validation/
|
|
134
|
+
end-of-training boundaries -- not every step."""
|
|
135
|
+
if not use_fused:
|
|
136
|
+
return
|
|
137
|
+
param_names = [name for name, _ in model.named_parameters()]
|
|
138
|
+
m_by_name, v_by_name = model.sync_fused_to_host()
|
|
139
|
+
optimizer.m = [m_by_name[name] for name in param_names]
|
|
140
|
+
optimizer.v = [v_by_name[name] for name in param_names]
|
|
141
|
+
optimizer.step_count = model._fused_step
|
|
142
|
+
|
|
101
143
|
def checkpoint(epoch, complete):
|
|
144
|
+
sync_fused()
|
|
102
145
|
checkpoint_mgr.save({"epoch": epoch, "global_step": global_step, "train_loader_epoch": prepared.train_loader.epoch, "model_state_dict": model.state_dict(), "best_model_state_dict": best_model_state, "optimizer_state_dict": optimizer.state_dict(), "scheduler_state_dict": scheduler.state_dict(), "early_stopping_best": early_stopper.best, "early_stopping_bad_checks": early_stopper.num_bad_checks, "config": cfg, "meta": prepared.meta, "tokenizer_state": prepared.tokenizer.state_dict() if prepared.tokenizer else None, "preprocessor_state": prepared.preprocessor.state_dict() if prepared.preprocessor else None, "dataset_fingerprint": dataset_fingerprint, "training_complete": complete})
|
|
103
146
|
last_train_loss = last_val_loss = None
|
|
104
147
|
t0 = time.time(); stop = False
|
|
@@ -109,16 +152,37 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
|
|
|
109
152
|
total_epoch_steps = len(prepared.train_loader)
|
|
110
153
|
for batch in prepared.train_loader:
|
|
111
154
|
epoch_step += 1
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
155
|
+
if use_fused:
|
|
156
|
+
grad_clip = cfg["grad_clip"] or 0.0
|
|
157
|
+
if model_type == "transformer" and task == "text-generation":
|
|
158
|
+
last_train_loss, grad_norm = model.fused_train_step(
|
|
159
|
+
batch[0], batch[1], None, optimizer.lr, cfg["weight_decay"], grad_clip, fused_decoupled
|
|
160
|
+
)
|
|
161
|
+
elif model_type == "transformer":
|
|
162
|
+
last_train_loss, grad_norm = model.fused_train_step(
|
|
163
|
+
batch[0], batch[2], batch[1], optimizer.lr, cfg["weight_decay"], grad_clip, fused_decoupled
|
|
164
|
+
)
|
|
165
|
+
else:
|
|
166
|
+
numeric, categorical, target = batch
|
|
167
|
+
last_train_loss, grad_norm = model.fused_train_step(
|
|
168
|
+
numeric, categorical, target, optimizer.lr, cfg["weight_decay"], grad_clip, fused_decoupled
|
|
169
|
+
)
|
|
170
|
+
if not np.isfinite(last_train_loss):
|
|
171
|
+
raise FloatingPointError(f"Non-finite training loss at step {global_step + 1}: {last_train_loss}")
|
|
172
|
+
if grad_clip and not np.isfinite(grad_norm):
|
|
173
|
+
raise FloatingPointError(f"Non-finite gradient norm at step {global_step + 1}: {grad_norm}")
|
|
174
|
+
else:
|
|
175
|
+
optimizer.zero_grad(); last_train_loss = _compute_loss(task, model_type, model, batch, prepared.meta.get("pad_id", 0));
|
|
176
|
+
if not np.isfinite(last_train_loss):
|
|
177
|
+
raise FloatingPointError(f"Non-finite training loss at step {global_step + 1}: {last_train_loss}")
|
|
178
|
+
if cfg["grad_clip"]:
|
|
179
|
+
norm = np.sqrt(sum(float(np.sum(p.grad.astype(np.float64) ** 2)) for p in model.parameters()))
|
|
180
|
+
if not np.isfinite(norm):
|
|
181
|
+
raise FloatingPointError(f"Non-finite gradient norm at step {global_step + 1}: {norm}")
|
|
182
|
+
if norm > cfg["grad_clip"]:
|
|
183
|
+
for p in model.parameters(): p.grad *= cfg["grad_clip"] / (norm + 1e-12)
|
|
184
|
+
optimizer.step()
|
|
185
|
+
scheduler.step(); global_step += 1
|
|
122
186
|
if progress_enabled:
|
|
123
187
|
_show_progress(epoch + 1, cfg["epochs"], epoch_step, total_epoch_steps, last_train_loss)
|
|
124
188
|
if global_step % cfg["checkpoint_every"] == 0: checkpoint(epoch, False)
|
|
@@ -128,6 +192,7 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
|
|
|
128
192
|
if stop:
|
|
129
193
|
break
|
|
130
194
|
if prepared.val_loader:
|
|
195
|
+
sync_fused()
|
|
131
196
|
model.eval(); losses = [_compute_loss(task, model_type, model, batch, prepared.meta.get("pad_id", 0), False) for batch in prepared.val_loader]; last_val_loss = sum(losses) / max(1, len(losses))
|
|
132
197
|
if not np.isfinite(last_val_loss):
|
|
133
198
|
raise FloatingPointError(f"Non-finite validation loss after epoch {epoch + 1}: {last_val_loss}")
|
|
@@ -146,6 +211,7 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
|
|
|
146
211
|
f"train_loss={last_train_loss:.4f}"
|
|
147
212
|
)
|
|
148
213
|
checkpoint(epoch + 1, epoch + 1 >= cfg["epochs"])
|
|
214
|
+
sync_fused()
|
|
149
215
|
elapsed = time.time() - t0
|
|
150
216
|
if best_model_state is not None:
|
|
151
217
|
model.load_state_dict(best_model_state)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|