tensorless 0.9.2__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 (85) hide show
  1. {tensorless-0.9.2/tensorless.egg-info → tensorless-0.9.4}/PKG-INFO +1 -1
  2. {tensorless-0.9.2 → tensorless-0.9.4}/pyproject.toml +1 -1
  3. tensorless-0.9.4/tensorless/backends/jax_backend.py +518 -0
  4. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/backends/mlx_backend.py +20 -1
  5. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/devices/device.py +26 -2
  6. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/trainer.py +76 -10
  7. {tensorless-0.9.2 → tensorless-0.9.4/tensorless.egg-info}/PKG-INFO +1 -1
  8. tensorless-0.9.2/tensorless/backends/jax_backend.py +0 -238
  9. {tensorless-0.9.2 → tensorless-0.9.4}/LICENSE +0 -0
  10. {tensorless-0.9.2 → tensorless-0.9.4}/MANIFEST.in +0 -0
  11. {tensorless-0.9.2 → tensorless-0.9.4}/README.md +0 -0
  12. {tensorless-0.9.2 → tensorless-0.9.4}/docs/api_reference.md +0 -0
  13. {tensorless-0.9.2 → tensorless-0.9.4}/docs/architecture.md +0 -0
  14. {tensorless-0.9.2 → tensorless-0.9.4}/docs/automatic_mode.md +0 -0
  15. {tensorless-0.9.2 → tensorless-0.9.4}/docs/checkpointing.md +0 -0
  16. {tensorless-0.9.2 → tensorless-0.9.4}/docs/cli.md +0 -0
  17. {tensorless-0.9.2 → tensorless-0.9.4}/docs/configuration.md +0 -0
  18. {tensorless-0.9.2 → tensorless-0.9.4}/docs/contributing.md +0 -0
  19. {tensorless-0.9.2 → tensorless-0.9.4}/docs/examples.md +0 -0
  20. {tensorless-0.9.2 → tensorless-0.9.4}/docs/inference.md +0 -0
  21. {tensorless-0.9.2 → tensorless-0.9.4}/docs/installation.md +0 -0
  22. {tensorless-0.9.2 → tensorless-0.9.4}/docs/limitations.md +0 -0
  23. {tensorless-0.9.2 → tensorless-0.9.4}/docs/quickstart.md +0 -0
  24. {tensorless-0.9.2 → tensorless-0.9.4}/docs/roadmap.md +0 -0
  25. {tensorless-0.9.2 → tensorless-0.9.4}/docs/tl_format.md +0 -0
  26. {tensorless-0.9.2 → tensorless-0.9.4}/docs/training.md +0 -0
  27. {tensorless-0.9.2 → tensorless-0.9.4}/docs/troubleshooting.md +0 -0
  28. {tensorless-0.9.2 → tensorless-0.9.4}/docs/tutorial.md +0 -0
  29. {tensorless-0.9.2 → tensorless-0.9.4}/examples/tabular_classification_example.py +0 -0
  30. {tensorless-0.9.2 → tensorless-0.9.4}/examples/tabular_regression_example.py +0 -0
  31. {tensorless-0.9.2 → tensorless-0.9.4}/examples/text_classification_example.py +0 -0
  32. {tensorless-0.9.2 → tensorless-0.9.4}/examples/text_generation_example.py +0 -0
  33. {tensorless-0.9.2 → tensorless-0.9.4}/setup.cfg +0 -0
  34. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/__init__.py +0 -0
  35. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/_version.py +0 -0
  36. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/api.py +0 -0
  37. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/auto/__init__.py +0 -0
  38. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/auto/config.py +0 -0
  39. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/auto/detector.py +0 -0
  40. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/backends/__init__.py +0 -0
  41. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/checkpoint/__init__.py +0 -0
  42. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/checkpoint/manager.py +0 -0
  43. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/cli/__init__.py +0 -0
  44. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/cli/main.py +0 -0
  45. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/config.py +0 -0
  46. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/__init__.py +0 -0
  47. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/english_grammar.txt +0 -0
  48. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/fingerprint.py +0 -0
  49. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/inspector.py +0 -0
  50. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/loader.py +0 -0
  51. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/tabular.py +0 -0
  52. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/devices/__init__.py +0 -0
  53. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/devices/memory.py +0 -0
  54. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/engine.py +0 -0
  55. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/errors.py +0 -0
  56. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/__init__.py +0 -0
  57. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/mlp.py +0 -0
  58. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/registry.py +0 -0
  59. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/transformer.py +0 -0
  60. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/runtime.py +0 -0
  61. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/serialization/__init__.py +0 -0
  62. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/serialization/tl_format.py +0 -0
  63. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/tokenization/__init__.py +0 -0
  64. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/tokenization/bpe_tokenizer.py +0 -0
  65. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/tokenization/char_tokenizer.py +0 -0
  66. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/__init__.py +0 -0
  67. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/data_prep.py +0 -0
  68. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/early_stopping.py +0 -0
  69. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/SOURCES.txt +0 -0
  70. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/dependency_links.txt +0 -0
  71. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/entry_points.txt +0 -0
  72. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/requires.txt +0 -0
  73. {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/top_level.txt +0 -0
  74. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_auto_detection.py +0 -0
  75. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_checkpoint_resume.py +0 -0
  76. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_cli.py +0 -0
  77. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_data_loading.py +0 -0
  78. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_end_to_end.py +0 -0
  79. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_fingerprint.py +0 -0
  80. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_jax_backend.py +0 -0
  81. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_mlx_backend.py +0 -0
  82. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_serialization.py +0 -0
  83. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_train_tabular.py +0 -0
  84. {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_train_text_classification.py +0 -0
  85. {tensorless-0.9.2 → 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.2
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.2"
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"
@@ -0,0 +1,518 @@
1
+ """JAX implementation of the text transformer.
2
+
3
+ JAX is imported lazily so importing Tensorless remains CPU/NumPy-only when the
4
+ optional dependency is not installed.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from typing import Optional
10
+ import os
11
+
12
+ import numpy as np
13
+
14
+ from ..models.transformer import TinyTransformer
15
+ from ..models.mlp import TabularMLP
16
+
17
+
18
+ def _jax():
19
+ try:
20
+ import jax
21
+ import jax.numpy as jnp
22
+ distributed_vars = ("JAX_COORDINATOR_ADDRESS", "JAX_PROCESS_COUNT", "JAX_PROCESS_ID")
23
+ if all(name in os.environ for name in distributed_vars) and not jax.distributed.is_initialized():
24
+ jax.distributed.initialize(
25
+ coordinator_address=os.environ["JAX_COORDINATOR_ADDRESS"],
26
+ num_processes=int(os.environ["JAX_PROCESS_COUNT"]),
27
+ process_id=int(os.environ["JAX_PROCESS_ID"]),
28
+ )
29
+ except ImportError as exc:
30
+ raise ImportError("JAX is required for the cuda/tpu backend; install tensorless[jax].") from exc
31
+ return jax, jnp
32
+
33
+
34
+ def is_available() -> bool:
35
+ try:
36
+ _jax()
37
+ return True
38
+ except ImportError:
39
+ return False
40
+
41
+
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):
133
+ """The native transformer parameter layout with JAX math and gradients."""
134
+
135
+ def __init__(self, *args, **kwargs):
136
+ self.precision = kwargs.pop("precision", "fp32")
137
+ super().__init__(*args, **kwargs)
138
+ d_model = kwargs.get("d_model", args[1] if len(args) > 1 else None)
139
+ self.heads = kwargs.get("heads", args[3] if len(args) > 3 else None)
140
+ if d_model is None or self.heads is None:
141
+ raise ValueError("JAX transformer requires d_model and heads")
142
+ self.d_model = d_model
143
+ self.head_dim = d_model // self.heads
144
+
145
+ def _jax_params(self):
146
+ _, jnp = _jax()
147
+ dtype = {"fp16": jnp.float16, "bf16": jnp.bfloat16}.get(self.precision, jnp.float32)
148
+ return {name: jnp.asarray(value.data, dtype=dtype) for name, value in self.named_parameters()}
149
+
150
+ def _forward_jax(self, params, input_ids, attention_mask=None):
151
+ _, jnp = _jax()
152
+
153
+ def layer_norm(x, gain, bias, eps=1e-5):
154
+ mean = jnp.mean(x, axis=-1, keepdims=True)
155
+ var = jnp.var(x, axis=-1, keepdims=True)
156
+ return (x - mean) / jnp.sqrt(var + eps) * gain + bias
157
+
158
+ input_ids = jnp.asarray(input_ids, dtype=jnp.int32)
159
+ batch, length = input_ids.shape
160
+ hidden = params["tok_emb"][input_ids] + params["pos_emb"][jnp.arange(length)]
161
+ causal = jnp.tril(jnp.ones((length, length), dtype=bool))
162
+ for index in range(len(self.blocks)):
163
+ prefix = f"blocks.{index}"
164
+ normed1 = layer_norm(hidden, params[f"{prefix}.ln1_gain"], params[f"{prefix}.ln1_bias"])
165
+ qkv = normed1 @ params[f"{prefix}.qkv_weight"] + params[f"{prefix}.qkv_bias"]
166
+ q, k, v = jnp.split(qkv, 3, axis=-1)
167
+ q = q.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
168
+ k = k.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
169
+ v = v.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
170
+ scores = q @ k.transpose(0, 1, 3, 2) / np.sqrt(self.head_dim)
171
+ scores = jnp.where(causal[None, None, :, :], scores, -1e9)
172
+ if attention_mask is not None:
173
+ mask = jnp.asarray(attention_mask) > 0
174
+ scores = jnp.where(mask[:, None, None, :], scores, -1e9)
175
+ probabilities = jnp.exp(scores - scores.max(axis=-1, keepdims=True))
176
+ probabilities /= probabilities.sum(axis=-1, keepdims=True)
177
+ context = probabilities @ v
178
+ context = context.transpose(0, 2, 1, 3).reshape(batch, length, self.d_model)
179
+ attention = context @ params[f"{prefix}.out_weight"] + params[f"{prefix}.out_bias"]
180
+ residual = hidden + attention
181
+ normed2 = layer_norm(residual, params[f"{prefix}.ln2_gain"], params[f"{prefix}.ln2_bias"])
182
+ ff_pre = normed2 @ params[f"{prefix}.ff1_weight"] + params[f"{prefix}.ff1_bias"]
183
+ ff_hidden = jnp.maximum(ff_pre, 0)
184
+ ff = ff_hidden @ params[f"{prefix}.ff2_weight"] + params[f"{prefix}.ff2_bias"]
185
+ hidden = residual + ff
186
+ hidden = layer_norm(hidden, params["ln_f_gain"], params["ln_f_bias"])
187
+ if self.task == "text-generation":
188
+ return hidden @ params["tok_emb"].T + params["head_bias"]
189
+ mask = jnp.ones((batch, length), dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
190
+ pooled = (hidden * mask[:, :, None]).sum(1) / jnp.maximum(mask.sum(1, keepdims=True), 1)
191
+ return pooled @ params["head_weight"] + params["head_bias"]
192
+
193
+ def forward(self, input_ids, attention_mask=None, cache=False):
194
+ params = self._jax_params()
195
+ output = self._get_forward_jit()(params, input_ids, attention_mask)
196
+ output = np.asarray(output)
197
+ return (output, None) if cache else output
198
+
199
+ def _get_forward_jit(self):
200
+ # Compile once per instance and reuse. Without this, every call
201
+ # re-traces the whole forward pass in Python and dispatches ops to
202
+ # the accelerator one at a time -- which shows up as pegged CPU
203
+ # (tracing/dispatch overhead) with near-zero GPU utilization, even
204
+ # though the correct device was selected.
205
+ if getattr(self, "_forward_jit_fn", None) is None:
206
+ jax, _ = _jax()
207
+ self._forward_jit_fn = jax.jit(self._forward_jax)
208
+ return self._forward_jit_fn
209
+
210
+ def _get_loss_and_grad_jit(self):
211
+ if getattr(self, "_loss_and_grad_jit_fn", None) is None:
212
+ jax, jnp = _jax()
213
+
214
+ def loss_fn(current, current_ids, current_target, current_mask):
215
+ forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
216
+ logits = forward(current, current_ids, current_mask)
217
+ flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
218
+ flat_target = current_target.reshape(-1)
219
+ log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
220
+ losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
221
+ if self.task == "text-generation":
222
+ valid = flat_target != self.pad_id
223
+ return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
224
+ return jnp.mean(losses)
225
+
226
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
227
+ scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
228
+ self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(scaled_loss_fn))
229
+ return self._loss_and_grad_jit_fn
230
+
231
+ def _get_pmap_step(self):
232
+ if getattr(self, "_pmap_step_fn", None) is None:
233
+ jax, jnp = _jax()
234
+
235
+ def loss_fn(current, current_ids, current_target, current_mask):
236
+ forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
237
+ logits = forward(current, current_ids, current_mask)
238
+ flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
239
+ flat_target = current_target.reshape(-1)
240
+ log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
241
+ losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
242
+ if self.task == "text-generation":
243
+ valid = flat_target != self.pad_id
244
+ return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
245
+ return jnp.mean(losses)
246
+
247
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
248
+ scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
249
+ per_device_loss = jax.value_and_grad(scaled_loss_fn)
250
+
251
+ def mapped_step(current, ids, labels, current_mask):
252
+ value, gradients = per_device_loss(current, ids, labels, current_mask)
253
+ return jax.lax.pmean(value, "data"), jax.tree_util.tree_map(
254
+ lambda gradient: jax.lax.pmean(gradient, "data"), gradients
255
+ )
256
+
257
+ self._pmap_step_fn = jax.pmap(mapped_step, axis_name="data", in_axes=(None, 0, 0, 0))
258
+ return self._pmap_step_fn
259
+
260
+ def loss_and_backward(self, input_ids, target, attention_mask=None):
261
+ jax, jnp = _jax()
262
+ params = self._jax_params()
263
+ target = jnp.asarray(target, dtype=jnp.int32)
264
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
265
+ mask = jnp.ones_like(input_ids, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
266
+ devices = jax.local_device_count()
267
+ if devices > 1 and input_ids.shape[0] % devices == 0:
268
+ shard = lambda value: value.reshape((devices, value.shape[0] // devices) + value.shape[1:])
269
+ loss, gradients = self._get_pmap_step()(
270
+ params, shard(jnp.asarray(input_ids)), shard(target), shard(mask)
271
+ )
272
+ loss = loss[0]
273
+ gradients = jax.tree_util.tree_map(lambda gradient: gradient[0] / loss_scale, gradients)
274
+ else:
275
+ loss, gradients = self._get_loss_and_grad_jit()(params, input_ids, target, mask)
276
+ gradients = jax.tree_util.tree_map(lambda gradient: gradient / loss_scale, gradients)
277
+ for name, parameter in self.named_parameters():
278
+ parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
279
+ return float(loss) / loss_scale
280
+
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.
349
+ ids = np.asarray(input_ids, dtype=np.int64).copy()
350
+ for _ in range(max_new_tokens):
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
+ )
360
+ if top_k is not None:
361
+ k = min(top_k, logits.shape[-1])
362
+ excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
363
+ logits[np.arange(len(ids))[:, None], excluded] = -np.inf
364
+ probs = np.exp(logits - logits.max(1, keepdims=True))
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)
373
+ next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
374
+ ids = np.concatenate([ids, next_ids], axis=1)
375
+ if eos_id is not None and np.all(next_ids == eos_id):
376
+ break
377
+ return ids
378
+
379
+
380
+ class JaxTabularMLP(_JaxFusedAdamMixin, TabularMLP):
381
+ """The native tabular-MLP parameter layout with JAX math and gradients.
382
+
383
+ Used automatically instead of the plain NumPy `TabularMLP` whenever the
384
+ resolved device is `cuda`/`tpu`, so tabular training actually runs on
385
+ the detected accelerator instead of silently staying on CPU.
386
+ """
387
+
388
+ def __init__(self, *args, **kwargs):
389
+ self.precision = kwargs.pop("precision", "fp32")
390
+ super().__init__(*args, **kwargs)
391
+
392
+ def _jax_params(self):
393
+ _, jnp = _jax()
394
+ dtype = {"fp16": jnp.float16, "bf16": jnp.bfloat16}.get(self.precision, jnp.float32)
395
+ return {name: jnp.asarray(value.data, dtype=dtype) for name, value in self.named_parameters()}
396
+
397
+ def _forward_jax(self, params, numeric, categorical, dropout_masks=None):
398
+ jax, jnp = _jax()
399
+ parts = [jnp.asarray(numeric, dtype=jnp.float32)] if self.n_numeric else []
400
+ categorical = jnp.asarray(categorical, dtype=jnp.int32)
401
+ for i in range(len(self.embeddings)):
402
+ parts.append(params[f"embeddings.{i}"][categorical[:, i]])
403
+ x = jnp.concatenate(parts, axis=1) if parts else jnp.asarray(numeric, dtype=jnp.float32)
404
+ for i in range(len(self.weights)):
405
+ preactivation = x @ params[f"weights.{i}"] + params[f"biases.{i}"]
406
+ x = jax.nn.gelu(preactivation, approximate=True)
407
+ if dropout_masks is not None and dropout_masks[i] is not None:
408
+ x = x * dropout_masks[i] / (1.0 - self.dropout)
409
+ out = x @ params["head_weight"] + params["head_bias"]
410
+ if self.task == "regression":
411
+ out = out[:, 0]
412
+ return out
413
+
414
+ def _make_dropout_masks(self, batch_size):
415
+ if not self.training or self.dropout <= 0:
416
+ return None
417
+ _, jnp = _jax()
418
+ return [
419
+ jnp.asarray((np.random.random((batch_size, weight.data.shape[1])) >= self.dropout).astype(np.float32))
420
+ for weight in self.weights
421
+ ]
422
+
423
+ def forward(self, numeric, categorical, cache=False):
424
+ params = self._jax_params()
425
+ batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
426
+ dropout_masks = self._make_dropout_masks(batch_size)
427
+ out = np.asarray(self._get_forward_jit()(params, numeric, categorical, dropout_masks))
428
+ return (out, None) if cache else out
429
+
430
+ def _get_forward_jit(self):
431
+ if getattr(self, "_forward_jit_fn", None) is None:
432
+ jax, _ = _jax()
433
+ self._forward_jit_fn = jax.jit(self._forward_jax)
434
+ return self._forward_jit_fn
435
+
436
+ def _get_loss_and_grad_jit(self):
437
+ if getattr(self, "_loss_and_grad_jit_fn", None) is None:
438
+ jax, jnp = _jax()
439
+ is_regression = self.task == "regression"
440
+
441
+ def loss_fn(current, numeric, categorical, dropout_masks, target_arr):
442
+ logits = self._forward_jax(current, numeric, categorical, dropout_masks)
443
+ if is_regression:
444
+ error = logits - target_arr
445
+ return jnp.mean(error * error)
446
+ log_probs = jax.nn.log_softmax(logits, axis=-1)
447
+ return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
448
+
449
+ self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(loss_fn))
450
+ return self._loss_and_grad_jit_fn
451
+
452
+ def loss_and_backward(self, numeric, categorical, target):
453
+ jax, jnp = _jax()
454
+ params = self._jax_params()
455
+ batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
456
+ dropout_masks = self._make_dropout_masks(batch_size)
457
+ is_regression = self.task == "regression"
458
+ target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
459
+
460
+ loss, gradients = self._get_loss_and_grad_jit()(params, numeric, categorical, dropout_masks, target_arr)
461
+ for name, parameter in self.named_parameters():
462
+ parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
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]: