tensorless 0.9.2__tar.gz → 0.9.3__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.2/tensorless.egg-info → tensorless-0.9.3}/PKG-INFO +1 -1
  2. {tensorless-0.9.2 → tensorless-0.9.3}/pyproject.toml +1 -1
  3. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/backends/jax_backend.py +89 -36
  4. {tensorless-0.9.2 → tensorless-0.9.3/tensorless.egg-info}/PKG-INFO +1 -1
  5. {tensorless-0.9.2 → tensorless-0.9.3}/LICENSE +0 -0
  6. {tensorless-0.9.2 → tensorless-0.9.3}/MANIFEST.in +0 -0
  7. {tensorless-0.9.2 → tensorless-0.9.3}/README.md +0 -0
  8. {tensorless-0.9.2 → tensorless-0.9.3}/docs/api_reference.md +0 -0
  9. {tensorless-0.9.2 → tensorless-0.9.3}/docs/architecture.md +0 -0
  10. {tensorless-0.9.2 → tensorless-0.9.3}/docs/automatic_mode.md +0 -0
  11. {tensorless-0.9.2 → tensorless-0.9.3}/docs/checkpointing.md +0 -0
  12. {tensorless-0.9.2 → tensorless-0.9.3}/docs/cli.md +0 -0
  13. {tensorless-0.9.2 → tensorless-0.9.3}/docs/configuration.md +0 -0
  14. {tensorless-0.9.2 → tensorless-0.9.3}/docs/contributing.md +0 -0
  15. {tensorless-0.9.2 → tensorless-0.9.3}/docs/examples.md +0 -0
  16. {tensorless-0.9.2 → tensorless-0.9.3}/docs/inference.md +0 -0
  17. {tensorless-0.9.2 → tensorless-0.9.3}/docs/installation.md +0 -0
  18. {tensorless-0.9.2 → tensorless-0.9.3}/docs/limitations.md +0 -0
  19. {tensorless-0.9.2 → tensorless-0.9.3}/docs/quickstart.md +0 -0
  20. {tensorless-0.9.2 → tensorless-0.9.3}/docs/roadmap.md +0 -0
  21. {tensorless-0.9.2 → tensorless-0.9.3}/docs/tl_format.md +0 -0
  22. {tensorless-0.9.2 → tensorless-0.9.3}/docs/training.md +0 -0
  23. {tensorless-0.9.2 → tensorless-0.9.3}/docs/troubleshooting.md +0 -0
  24. {tensorless-0.9.2 → tensorless-0.9.3}/docs/tutorial.md +0 -0
  25. {tensorless-0.9.2 → tensorless-0.9.3}/examples/tabular_classification_example.py +0 -0
  26. {tensorless-0.9.2 → tensorless-0.9.3}/examples/tabular_regression_example.py +0 -0
  27. {tensorless-0.9.2 → tensorless-0.9.3}/examples/text_classification_example.py +0 -0
  28. {tensorless-0.9.2 → tensorless-0.9.3}/examples/text_generation_example.py +0 -0
  29. {tensorless-0.9.2 → tensorless-0.9.3}/setup.cfg +0 -0
  30. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/__init__.py +0 -0
  31. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/_version.py +0 -0
  32. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/api.py +0 -0
  33. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/auto/__init__.py +0 -0
  34. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/auto/config.py +0 -0
  35. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/auto/detector.py +0 -0
  36. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/backends/__init__.py +0 -0
  37. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/backends/mlx_backend.py +0 -0
  38. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/checkpoint/__init__.py +0 -0
  39. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/checkpoint/manager.py +0 -0
  40. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/cli/__init__.py +0 -0
  41. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/cli/main.py +0 -0
  42. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/config.py +0 -0
  43. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/__init__.py +0 -0
  44. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/english_grammar.txt +0 -0
  45. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/fingerprint.py +0 -0
  46. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/inspector.py +0 -0
  47. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/loader.py +0 -0
  48. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/tabular.py +0 -0
  49. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/devices/__init__.py +0 -0
  50. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/devices/device.py +0 -0
  51. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/devices/memory.py +0 -0
  52. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/engine.py +0 -0
  53. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/errors.py +0 -0
  54. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/__init__.py +0 -0
  55. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/mlp.py +0 -0
  56. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/registry.py +0 -0
  57. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/transformer.py +0 -0
  58. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/runtime.py +0 -0
  59. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/serialization/__init__.py +0 -0
  60. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/serialization/tl_format.py +0 -0
  61. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/tokenization/__init__.py +0 -0
  62. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/tokenization/bpe_tokenizer.py +0 -0
  63. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/tokenization/char_tokenizer.py +0 -0
  64. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/__init__.py +0 -0
  65. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/data_prep.py +0 -0
  66. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/early_stopping.py +0 -0
  67. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/trainer.py +0 -0
  68. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/SOURCES.txt +0 -0
  69. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/dependency_links.txt +0 -0
  70. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/entry_points.txt +0 -0
  71. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/requires.txt +0 -0
  72. {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/top_level.txt +0 -0
  73. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_auto_detection.py +0 -0
  74. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_checkpoint_resume.py +0 -0
  75. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_cli.py +0 -0
  76. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_data_loading.py +0 -0
  77. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_end_to_end.py +0 -0
  78. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_fingerprint.py +0 -0
  79. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_jax_backend.py +0 -0
  80. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_mlx_backend.py +0 -0
  81. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_serialization.py +0 -0
  82. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_train_tabular.py +0 -0
  83. {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_train_text_classification.py +0 -0
  84. {tensorless-0.9.2 → tensorless-0.9.3}/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.3
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.3"
8
8
  description = "ML with maximum automation and minimum setup."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.9"
@@ -101,33 +101,61 @@ class JaxTinyTransformer(TinyTransformer):
101
101
  return pooled @ params["head_weight"] + params["head_bias"]
102
102
 
103
103
  def forward(self, input_ids, attention_mask=None, cache=False):
104
- _, jnp = _jax()
105
- output = self._forward_jax(self._jax_params(), input_ids, attention_mask)
104
+ params = self._jax_params()
105
+ output = self._get_forward_jit()(params, input_ids, attention_mask)
106
106
  output = np.asarray(output)
107
107
  return (output, None) if cache else output
108
108
 
109
- def loss_and_backward(self, input_ids, target, attention_mask=None):
110
- jax, jnp = _jax()
111
- params = self._jax_params()
112
- target = jnp.asarray(target, dtype=jnp.int32)
113
-
114
- def loss_fn(current, current_ids, current_target, current_mask):
115
- forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
116
- logits = forward(current, current_ids, current_mask)
117
- flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
118
- flat_target = current_target.reshape(-1)
119
- log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
120
- losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
121
- if self.task == "text-generation":
122
- valid = flat_target != self.pad_id
123
- return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
124
- return jnp.mean(losses)
125
-
126
- loss_scale = 128.0 if self.precision == "fp16" else 1.0
127
- scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
128
- mask = jnp.ones_like(input_ids, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
129
- devices = jax.local_device_count()
130
- if devices > 1 and input_ids.shape[0] % devices == 0:
109
+ def _get_forward_jit(self):
110
+ # Compile once per instance and reuse. Without this, every call
111
+ # re-traces the whole forward pass in Python and dispatches ops to
112
+ # the accelerator one at a time -- which shows up as pegged CPU
113
+ # (tracing/dispatch overhead) with near-zero GPU utilization, even
114
+ # though the correct device was selected.
115
+ if getattr(self, "_forward_jit_fn", None) is None:
116
+ jax, _ = _jax()
117
+ self._forward_jit_fn = jax.jit(self._forward_jax)
118
+ return self._forward_jit_fn
119
+
120
+ def _get_loss_and_grad_jit(self):
121
+ if getattr(self, "_loss_and_grad_jit_fn", None) is None:
122
+ jax, jnp = _jax()
123
+
124
+ def loss_fn(current, current_ids, current_target, current_mask):
125
+ forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
126
+ logits = forward(current, current_ids, current_mask)
127
+ flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
128
+ flat_target = current_target.reshape(-1)
129
+ log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
130
+ losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
131
+ if self.task == "text-generation":
132
+ valid = flat_target != self.pad_id
133
+ return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
134
+ return jnp.mean(losses)
135
+
136
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
137
+ scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
138
+ self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(scaled_loss_fn))
139
+ return self._loss_and_grad_jit_fn
140
+
141
+ def _get_pmap_step(self):
142
+ if getattr(self, "_pmap_step_fn", None) is None:
143
+ jax, jnp = _jax()
144
+
145
+ def loss_fn(current, current_ids, current_target, current_mask):
146
+ forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
147
+ logits = forward(current, current_ids, current_mask)
148
+ flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
149
+ flat_target = current_target.reshape(-1)
150
+ log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
151
+ losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
152
+ if self.task == "text-generation":
153
+ valid = flat_target != self.pad_id
154
+ return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
155
+ return jnp.mean(losses)
156
+
157
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
158
+ scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
131
159
  per_device_loss = jax.value_and_grad(scaled_loss_fn)
132
160
 
133
161
  def mapped_step(current, ids, labels, current_mask):
@@ -136,14 +164,25 @@ class JaxTinyTransformer(TinyTransformer):
136
164
  lambda gradient: jax.lax.pmean(gradient, "data"), gradients
137
165
  )
138
166
 
167
+ self._pmap_step_fn = jax.pmap(mapped_step, axis_name="data", in_axes=(None, 0, 0, 0))
168
+ return self._pmap_step_fn
169
+
170
+ def loss_and_backward(self, input_ids, target, attention_mask=None):
171
+ jax, jnp = _jax()
172
+ params = self._jax_params()
173
+ target = jnp.asarray(target, dtype=jnp.int32)
174
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
175
+ mask = jnp.ones_like(input_ids, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
176
+ devices = jax.local_device_count()
177
+ if devices > 1 and input_ids.shape[0] % devices == 0:
139
178
  shard = lambda value: value.reshape((devices, value.shape[0] // devices) + value.shape[1:])
140
- loss, gradients = jax.pmap(mapped_step, axis_name="data", in_axes=(None, 0, 0, 0))(
179
+ loss, gradients = self._get_pmap_step()(
141
180
  params, shard(jnp.asarray(input_ids)), shard(target), shard(mask)
142
181
  )
143
182
  loss = loss[0]
144
183
  gradients = jax.tree_util.tree_map(lambda gradient: gradient[0] / loss_scale, gradients)
145
184
  else:
146
- loss, gradients = jax.value_and_grad(scaled_loss_fn)(params, input_ids, target, mask)
185
+ loss, gradients = self._get_loss_and_grad_jit()(params, input_ids, target, mask)
147
186
  gradients = jax.tree_util.tree_map(lambda gradient: gradient / loss_scale, gradients)
148
187
  for name, parameter in self.named_parameters():
149
188
  parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
@@ -213,9 +252,31 @@ class JaxTabularMLP(TabularMLP):
213
252
  params = self._jax_params()
214
253
  batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
215
254
  dropout_masks = self._make_dropout_masks(batch_size)
216
- out = np.asarray(self._forward_jax(params, numeric, categorical, dropout_masks))
255
+ out = np.asarray(self._get_forward_jit()(params, numeric, categorical, dropout_masks))
217
256
  return (out, None) if cache else out
218
257
 
258
+ def _get_forward_jit(self):
259
+ if getattr(self, "_forward_jit_fn", None) is None:
260
+ jax, _ = _jax()
261
+ self._forward_jit_fn = jax.jit(self._forward_jax)
262
+ return self._forward_jit_fn
263
+
264
+ def _get_loss_and_grad_jit(self):
265
+ if getattr(self, "_loss_and_grad_jit_fn", None) is None:
266
+ jax, jnp = _jax()
267
+ is_regression = self.task == "regression"
268
+
269
+ def loss_fn(current, numeric, categorical, dropout_masks, target_arr):
270
+ logits = self._forward_jax(current, numeric, categorical, dropout_masks)
271
+ if is_regression:
272
+ error = logits - target_arr
273
+ return jnp.mean(error * error)
274
+ log_probs = jax.nn.log_softmax(logits, axis=-1)
275
+ return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
276
+
277
+ self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(loss_fn))
278
+ return self._loss_and_grad_jit_fn
279
+
219
280
  def loss_and_backward(self, numeric, categorical, target):
220
281
  jax, jnp = _jax()
221
282
  params = self._jax_params()
@@ -224,15 +285,7 @@ class JaxTabularMLP(TabularMLP):
224
285
  is_regression = self.task == "regression"
225
286
  target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
226
287
 
227
- def loss_fn(current):
228
- logits = self._forward_jax(current, numeric, categorical, dropout_masks)
229
- if is_regression:
230
- error = logits - target_arr
231
- return jnp.mean(error * error)
232
- log_probs = jax.nn.log_softmax(logits, axis=-1)
233
- return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
234
-
235
- loss, gradients = jax.value_and_grad(loss_fn)(params)
288
+ loss, gradients = self._get_loss_and_grad_jit()(params, numeric, categorical, dropout_masks, target_arr)
236
289
  for name, parameter in self.named_parameters():
237
290
  parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
238
291
  return float(loss)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.9.2
3
+ Version: 0.9.3
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