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.
Files changed (84) hide show
  1. {tensorless-0.9.3/tensorless.egg-info → tensorless-0.9.4}/PKG-INFO +1 -1
  2. {tensorless-0.9.3 → tensorless-0.9.4}/pyproject.toml +1 -1
  3. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/backends/jax_backend.py +230 -3
  4. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/backends/mlx_backend.py +20 -1
  5. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/devices/device.py +26 -2
  6. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/trainer.py +76 -10
  7. {tensorless-0.9.3 → tensorless-0.9.4/tensorless.egg-info}/PKG-INFO +1 -1
  8. {tensorless-0.9.3 → tensorless-0.9.4}/LICENSE +0 -0
  9. {tensorless-0.9.3 → tensorless-0.9.4}/MANIFEST.in +0 -0
  10. {tensorless-0.9.3 → tensorless-0.9.4}/README.md +0 -0
  11. {tensorless-0.9.3 → tensorless-0.9.4}/docs/api_reference.md +0 -0
  12. {tensorless-0.9.3 → tensorless-0.9.4}/docs/architecture.md +0 -0
  13. {tensorless-0.9.3 → tensorless-0.9.4}/docs/automatic_mode.md +0 -0
  14. {tensorless-0.9.3 → tensorless-0.9.4}/docs/checkpointing.md +0 -0
  15. {tensorless-0.9.3 → tensorless-0.9.4}/docs/cli.md +0 -0
  16. {tensorless-0.9.3 → tensorless-0.9.4}/docs/configuration.md +0 -0
  17. {tensorless-0.9.3 → tensorless-0.9.4}/docs/contributing.md +0 -0
  18. {tensorless-0.9.3 → tensorless-0.9.4}/docs/examples.md +0 -0
  19. {tensorless-0.9.3 → tensorless-0.9.4}/docs/inference.md +0 -0
  20. {tensorless-0.9.3 → tensorless-0.9.4}/docs/installation.md +0 -0
  21. {tensorless-0.9.3 → tensorless-0.9.4}/docs/limitations.md +0 -0
  22. {tensorless-0.9.3 → tensorless-0.9.4}/docs/quickstart.md +0 -0
  23. {tensorless-0.9.3 → tensorless-0.9.4}/docs/roadmap.md +0 -0
  24. {tensorless-0.9.3 → tensorless-0.9.4}/docs/tl_format.md +0 -0
  25. {tensorless-0.9.3 → tensorless-0.9.4}/docs/training.md +0 -0
  26. {tensorless-0.9.3 → tensorless-0.9.4}/docs/troubleshooting.md +0 -0
  27. {tensorless-0.9.3 → tensorless-0.9.4}/docs/tutorial.md +0 -0
  28. {tensorless-0.9.3 → tensorless-0.9.4}/examples/tabular_classification_example.py +0 -0
  29. {tensorless-0.9.3 → tensorless-0.9.4}/examples/tabular_regression_example.py +0 -0
  30. {tensorless-0.9.3 → tensorless-0.9.4}/examples/text_classification_example.py +0 -0
  31. {tensorless-0.9.3 → tensorless-0.9.4}/examples/text_generation_example.py +0 -0
  32. {tensorless-0.9.3 → tensorless-0.9.4}/setup.cfg +0 -0
  33. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/__init__.py +0 -0
  34. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/_version.py +0 -0
  35. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/api.py +0 -0
  36. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/auto/__init__.py +0 -0
  37. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/auto/config.py +0 -0
  38. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/auto/detector.py +0 -0
  39. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/backends/__init__.py +0 -0
  40. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/checkpoint/__init__.py +0 -0
  41. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/checkpoint/manager.py +0 -0
  42. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/cli/__init__.py +0 -0
  43. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/cli/main.py +0 -0
  44. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/config.py +0 -0
  45. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/__init__.py +0 -0
  46. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/english_grammar.txt +0 -0
  47. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/fingerprint.py +0 -0
  48. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/inspector.py +0 -0
  49. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/loader.py +0 -0
  50. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/data/tabular.py +0 -0
  51. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/devices/__init__.py +0 -0
  52. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/devices/memory.py +0 -0
  53. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/engine.py +0 -0
  54. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/errors.py +0 -0
  55. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/__init__.py +0 -0
  56. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/mlp.py +0 -0
  57. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/registry.py +0 -0
  58. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/models/transformer.py +0 -0
  59. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/runtime.py +0 -0
  60. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/serialization/__init__.py +0 -0
  61. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/serialization/tl_format.py +0 -0
  62. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/tokenization/__init__.py +0 -0
  63. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/tokenization/bpe_tokenizer.py +0 -0
  64. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/tokenization/char_tokenizer.py +0 -0
  65. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/__init__.py +0 -0
  66. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/data_prep.py +0 -0
  67. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless/training/early_stopping.py +0 -0
  68. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/SOURCES.txt +0 -0
  69. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/dependency_links.txt +0 -0
  70. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/entry_points.txt +0 -0
  71. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/requires.txt +0 -0
  72. {tensorless-0.9.3 → tensorless-0.9.4}/tensorless.egg-info/top_level.txt +0 -0
  73. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_auto_detection.py +0 -0
  74. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_checkpoint_resume.py +0 -0
  75. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_cli.py +0 -0
  76. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_data_loading.py +0 -0
  77. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_end_to_end.py +0 -0
  78. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_fingerprint.py +0 -0
  79. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_jax_backend.py +0 -0
  80. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_mlx_backend.py +0 -0
  81. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_serialization.py +0 -0
  82. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_train_tabular.py +0 -0
  83. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_train_text_classification.py +0 -0
  84. {tensorless-0.9.3 → tensorless-0.9.4}/tests/test_train_text_generation.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.9.3
3
+ Version: 0.9.4
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "tensorless"
7
- version = "0.9.3"
7
+ version = "0.9.4"
8
8
  description = "ML with maximum automation and minimum setup."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.9"
@@ -39,7 +39,97 @@ def is_available() -> bool:
39
39
  return False
40
40
 
41
41
 
42
- class JaxTinyTransformer(TinyTransformer):
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 generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, eos_id: Optional[int] = None):
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, eos_id: Optional[int] = None):
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
- """Use the conservative CUDA default unless hardware is identified."""
42
- return False
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
- optimizer.zero_grad(); last_train_loss = _compute_loss(task, model_type, model, batch, prepared.meta.get("pad_id", 0));
113
- if not np.isfinite(last_train_loss):
114
- raise FloatingPointError(f"Non-finite training loss at step {global_step + 1}: {last_train_loss}")
115
- if cfg["grad_clip"]:
116
- norm = np.sqrt(sum(float(np.sum(p.grad.astype(np.float64) ** 2)) for p in model.parameters()))
117
- if not np.isfinite(norm):
118
- raise FloatingPointError(f"Non-finite gradient norm at step {global_step + 1}: {norm}")
119
- if norm > cfg["grad_clip"]:
120
- for p in model.parameters(): p.grad *= cfg["grad_clip"] / (norm + 1e-12)
121
- optimizer.step(); scheduler.step(); global_step += 1
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)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.9.3
3
+ Version: 0.9.4
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
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