tensorless 0.9.1__tar.gz → 0.9.2__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.1/tensorless.egg-info → tensorless-0.9.2}/PKG-INFO +6 -3
  2. {tensorless-0.9.1 → tensorless-0.9.2}/README.md +5 -2
  3. tensorless-0.9.2/docs/installation.md +49 -0
  4. {tensorless-0.9.1 → tensorless-0.9.2}/pyproject.toml +1 -1
  5. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/backends/jax_backend.py +73 -0
  6. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/backends/mlx_backend.py +79 -0
  7. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/registry.py +11 -1
  8. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/runtime.py +4 -5
  9. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/trainer.py +9 -5
  10. {tensorless-0.9.1 → tensorless-0.9.2/tensorless.egg-info}/PKG-INFO +6 -3
  11. tensorless-0.9.1/docs/installation.md +0 -29
  12. {tensorless-0.9.1 → tensorless-0.9.2}/LICENSE +0 -0
  13. {tensorless-0.9.1 → tensorless-0.9.2}/MANIFEST.in +0 -0
  14. {tensorless-0.9.1 → tensorless-0.9.2}/docs/api_reference.md +0 -0
  15. {tensorless-0.9.1 → tensorless-0.9.2}/docs/architecture.md +0 -0
  16. {tensorless-0.9.1 → tensorless-0.9.2}/docs/automatic_mode.md +0 -0
  17. {tensorless-0.9.1 → tensorless-0.9.2}/docs/checkpointing.md +0 -0
  18. {tensorless-0.9.1 → tensorless-0.9.2}/docs/cli.md +0 -0
  19. {tensorless-0.9.1 → tensorless-0.9.2}/docs/configuration.md +0 -0
  20. {tensorless-0.9.1 → tensorless-0.9.2}/docs/contributing.md +0 -0
  21. {tensorless-0.9.1 → tensorless-0.9.2}/docs/examples.md +0 -0
  22. {tensorless-0.9.1 → tensorless-0.9.2}/docs/inference.md +0 -0
  23. {tensorless-0.9.1 → tensorless-0.9.2}/docs/limitations.md +0 -0
  24. {tensorless-0.9.1 → tensorless-0.9.2}/docs/quickstart.md +0 -0
  25. {tensorless-0.9.1 → tensorless-0.9.2}/docs/roadmap.md +0 -0
  26. {tensorless-0.9.1 → tensorless-0.9.2}/docs/tl_format.md +0 -0
  27. {tensorless-0.9.1 → tensorless-0.9.2}/docs/training.md +0 -0
  28. {tensorless-0.9.1 → tensorless-0.9.2}/docs/troubleshooting.md +0 -0
  29. {tensorless-0.9.1 → tensorless-0.9.2}/docs/tutorial.md +0 -0
  30. {tensorless-0.9.1 → tensorless-0.9.2}/examples/tabular_classification_example.py +0 -0
  31. {tensorless-0.9.1 → tensorless-0.9.2}/examples/tabular_regression_example.py +0 -0
  32. {tensorless-0.9.1 → tensorless-0.9.2}/examples/text_classification_example.py +0 -0
  33. {tensorless-0.9.1 → tensorless-0.9.2}/examples/text_generation_example.py +0 -0
  34. {tensorless-0.9.1 → tensorless-0.9.2}/setup.cfg +0 -0
  35. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/__init__.py +0 -0
  36. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/_version.py +0 -0
  37. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/api.py +0 -0
  38. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/auto/__init__.py +0 -0
  39. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/auto/config.py +0 -0
  40. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/auto/detector.py +0 -0
  41. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/backends/__init__.py +0 -0
  42. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/checkpoint/__init__.py +0 -0
  43. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/checkpoint/manager.py +0 -0
  44. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/cli/__init__.py +0 -0
  45. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/cli/main.py +0 -0
  46. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/config.py +0 -0
  47. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/__init__.py +0 -0
  48. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/english_grammar.txt +0 -0
  49. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/fingerprint.py +0 -0
  50. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/inspector.py +0 -0
  51. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/loader.py +0 -0
  52. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/tabular.py +0 -0
  53. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/devices/__init__.py +0 -0
  54. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/devices/device.py +0 -0
  55. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/devices/memory.py +0 -0
  56. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/engine.py +0 -0
  57. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/errors.py +0 -0
  58. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/__init__.py +0 -0
  59. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/mlp.py +0 -0
  60. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/transformer.py +0 -0
  61. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/serialization/__init__.py +0 -0
  62. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/serialization/tl_format.py +0 -0
  63. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/tokenization/__init__.py +0 -0
  64. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/tokenization/bpe_tokenizer.py +0 -0
  65. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/tokenization/char_tokenizer.py +0 -0
  66. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/__init__.py +0 -0
  67. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/data_prep.py +0 -0
  68. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/early_stopping.py +0 -0
  69. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/SOURCES.txt +0 -0
  70. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/dependency_links.txt +0 -0
  71. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/entry_points.txt +0 -0
  72. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/requires.txt +0 -0
  73. {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/top_level.txt +0 -0
  74. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_auto_detection.py +0 -0
  75. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_checkpoint_resume.py +0 -0
  76. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_cli.py +0 -0
  77. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_data_loading.py +0 -0
  78. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_end_to_end.py +0 -0
  79. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_fingerprint.py +0 -0
  80. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_jax_backend.py +0 -0
  81. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_mlx_backend.py +0 -0
  82. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_serialization.py +0 -0
  83. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_train_tabular.py +0 -0
  84. {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_train_text_classification.py +0 -0
  85. {tensorless-0.9.1 → tensorless-0.9.2}/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.1
3
+ Version: 0.9.2
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
@@ -45,8 +45,11 @@ pip install -e '.[tpu]' # JAX TPU
45
45
  pip install -e '.[mps]' # Apple Silicon MLX
46
46
  ```
47
47
 
48
- CUDA and TPU backends currently accelerate transformer text tasks. Tabular
49
- tasks and unsupported platforms use the native CPU engine.
48
+ CUDA, TPU, and MPS backends accelerate both transformer text tasks and
49
+ tabular MLP tasks. Whatever device is auto-detected (or passed via
50
+ `device=...`) is what training and inference actually run on; only
51
+ unsupported platforms (no CUDA/TPU/MPS available) fall back to the native
52
+ CPU engine.
50
53
 
51
54
  ## Train on your data
52
55
 
@@ -19,8 +19,11 @@ pip install -e '.[tpu]' # JAX TPU
19
19
  pip install -e '.[mps]' # Apple Silicon MLX
20
20
  ```
21
21
 
22
- CUDA and TPU backends currently accelerate transformer text tasks. Tabular
23
- tasks and unsupported platforms use the native CPU engine.
22
+ CUDA, TPU, and MPS backends accelerate both transformer text tasks and
23
+ tabular MLP tasks. Whatever device is auto-detected (or passed via
24
+ `device=...`) is what training and inference actually run on; only
25
+ unsupported platforms (no CUDA/TPU/MPS available) fall back to the native
26
+ CPU engine.
24
27
 
25
28
  ## Train on your data
26
29
 
@@ -0,0 +1,49 @@
1
+ # Installation
2
+
3
+ ## Requirements
4
+
5
+ - Python 3.9 or later
6
+ - NumPy 1.24 or later (installed automatically as a dependency)
7
+ - Tensorless always runs its native vectorized engine, on CPU by default.
8
+ Installing the optional accelerator extras below lets it detect and use a
9
+ GPU/TPU automatically, for both transformer and tabular MLP models.
10
+
11
+ ## GPU / TPU / Apple Silicon support
12
+
13
+ By default `pip install tensorless` only pulls in NumPy, so even on a
14
+ machine with a GPU, device auto-detection will correctly report that no
15
+ accelerator backend is *installed* and fall back to CPU. To actually train
16
+ and run on your hardware, install the matching extra:
17
+
18
+ ```bash
19
+ pip install 'tensorless[cuda]' # NVIDIA GPUs, via JAX
20
+ pip install 'tensorless[tpu]' # Google TPUs, via JAX
21
+ pip install 'tensorless[mps]' # Apple Silicon, via MLX
22
+ ```
23
+
24
+ Once installed, `tl.train(...)` auto-detects the best available device
25
+ (`tpu` > `cuda` > `mps` > `cpu`) and trains on it without any extra
26
+ configuration; you can also force a specific device with
27
+ `tl.train(..., device="cuda")`.
28
+
29
+ ## Verify your installation
30
+
31
+ ```bash
32
+ python -c "import tensorless as tl; print(tl.__version__)"
33
+ tensorless --help
34
+ ```
35
+
36
+ You should see a version string printed and the CLI's help text.
37
+
38
+ ## Runtime support
39
+
40
+ Training and inference resolve the device automatically: the best available
41
+ accelerator (`tpu` > `cuda` > `mps`) is used if its extra is installed and
42
+ the hardware is detected, otherwise Tensorless falls back to the NumPy CPU
43
+ engine. You can always force a specific device with `device="cpu"`,
44
+ `device="cuda"`, etc. in `tl.train(...)`.
45
+
46
+ ## Troubleshooting installation
47
+
48
+ See [troubleshooting.md](troubleshooting.md#installation-issues) for
49
+ common installation problems.
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "tensorless"
7
- version = "0.9.1"
7
+ version = "0.9.2"
8
8
  description = "ML with maximum automation and minimum setup."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.9"
@@ -12,6 +12,7 @@ import os
12
12
  import numpy as np
13
13
 
14
14
  from ..models.transformer import TinyTransformer
15
+ from ..models.mlp import TabularMLP
15
16
 
16
17
 
17
18
  def _jax():
@@ -163,3 +164,75 @@ class JaxTinyTransformer(TinyTransformer):
163
164
  if eos_id is not None and np.all(next_ids == eos_id):
164
165
  break
165
166
  return ids
167
+
168
+
169
+ class JaxTabularMLP(TabularMLP):
170
+ """The native tabular-MLP parameter layout with JAX math and gradients.
171
+
172
+ Used automatically instead of the plain NumPy `TabularMLP` whenever the
173
+ resolved device is `cuda`/`tpu`, so tabular training actually runs on
174
+ the detected accelerator instead of silently staying on CPU.
175
+ """
176
+
177
+ def __init__(self, *args, **kwargs):
178
+ self.precision = kwargs.pop("precision", "fp32")
179
+ super().__init__(*args, **kwargs)
180
+
181
+ def _jax_params(self):
182
+ _, jnp = _jax()
183
+ dtype = {"fp16": jnp.float16, "bf16": jnp.bfloat16}.get(self.precision, jnp.float32)
184
+ return {name: jnp.asarray(value.data, dtype=dtype) for name, value in self.named_parameters()}
185
+
186
+ def _forward_jax(self, params, numeric, categorical, dropout_masks=None):
187
+ jax, jnp = _jax()
188
+ parts = [jnp.asarray(numeric, dtype=jnp.float32)] if self.n_numeric else []
189
+ categorical = jnp.asarray(categorical, dtype=jnp.int32)
190
+ for i in range(len(self.embeddings)):
191
+ parts.append(params[f"embeddings.{i}"][categorical[:, i]])
192
+ x = jnp.concatenate(parts, axis=1) if parts else jnp.asarray(numeric, dtype=jnp.float32)
193
+ for i in range(len(self.weights)):
194
+ preactivation = x @ params[f"weights.{i}"] + params[f"biases.{i}"]
195
+ x = jax.nn.gelu(preactivation, approximate=True)
196
+ if dropout_masks is not None and dropout_masks[i] is not None:
197
+ x = x * dropout_masks[i] / (1.0 - self.dropout)
198
+ out = x @ params["head_weight"] + params["head_bias"]
199
+ if self.task == "regression":
200
+ out = out[:, 0]
201
+ return out
202
+
203
+ def _make_dropout_masks(self, batch_size):
204
+ if not self.training or self.dropout <= 0:
205
+ return None
206
+ _, jnp = _jax()
207
+ return [
208
+ jnp.asarray((np.random.random((batch_size, weight.data.shape[1])) >= self.dropout).astype(np.float32))
209
+ for weight in self.weights
210
+ ]
211
+
212
+ def forward(self, numeric, categorical, cache=False):
213
+ params = self._jax_params()
214
+ batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
215
+ dropout_masks = self._make_dropout_masks(batch_size)
216
+ out = np.asarray(self._forward_jax(params, numeric, categorical, dropout_masks))
217
+ return (out, None) if cache else out
218
+
219
+ def loss_and_backward(self, numeric, categorical, target):
220
+ jax, jnp = _jax()
221
+ params = self._jax_params()
222
+ batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
223
+ dropout_masks = self._make_dropout_masks(batch_size)
224
+ is_regression = self.task == "regression"
225
+ target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
226
+
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)
236
+ for name, parameter in self.named_parameters():
237
+ parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
238
+ return float(loss)
@@ -11,6 +11,7 @@ from typing import Optional
11
11
  import numpy as np
12
12
 
13
13
  from ..models.transformer import TinyTransformer
14
+ from ..models.mlp import TabularMLP
14
15
 
15
16
 
16
17
  def _mlx():
@@ -131,3 +132,81 @@ class MlxTinyTransformer(TinyTransformer):
131
132
  if eos_id is not None and np.all(next_ids == eos_id):
132
133
  break
133
134
  return ids
135
+
136
+
137
+ def _mlx_gelu(x, mx):
138
+ scale = float(np.sqrt(2.0 / np.pi))
139
+ scaled = scale * (x + 0.044715 * x ** 3)
140
+ return 0.5 * x * (1.0 + mx.tanh(scaled))
141
+
142
+
143
+ class MlxTabularMLP(TabularMLP):
144
+ """The native tabular-MLP parameter layout with MLX math and gradients.
145
+
146
+ Used automatically instead of the plain NumPy `TabularMLP` whenever the
147
+ resolved device is `mps`, so tabular training actually runs on Apple
148
+ Silicon's GPU instead of silently staying on CPU.
149
+ """
150
+
151
+ def __init__(self, *args, **kwargs):
152
+ self.precision = kwargs.pop("precision", "fp32")
153
+ super().__init__(*args, **kwargs)
154
+
155
+ def _mlx_params(self):
156
+ mx = _mlx()
157
+ dtype = {"fp16": mx.float16, "bf16": mx.bfloat16}.get(self.precision, mx.float32)
158
+ return {name: mx.array(parameter.data, dtype=dtype) for name, parameter in self.named_parameters()}
159
+
160
+ def _forward_mlx(self, params, numeric, categorical, dropout_masks=None):
161
+ mx = _mlx()
162
+ parts = [mx.array(numeric, dtype=mx.float32)] if self.n_numeric else []
163
+ categorical = mx.array(categorical, dtype=mx.int32)
164
+ for i in range(len(self.embeddings)):
165
+ parts.append(params[f"embeddings.{i}"][categorical[:, i]])
166
+ x = mx.concatenate(parts, axis=1) if parts else mx.array(numeric, dtype=mx.float32)
167
+ for i in range(len(self.weights)):
168
+ preactivation = x @ params[f"weights.{i}"] + params[f"biases.{i}"]
169
+ x = _mlx_gelu(preactivation, mx)
170
+ if dropout_masks is not None and dropout_masks[i] is not None:
171
+ x = x * dropout_masks[i] / (1.0 - self.dropout)
172
+ out = x @ params["head_weight"] + params["head_bias"]
173
+ if self.task == "regression":
174
+ out = out[:, 0]
175
+ return out
176
+
177
+ def _make_dropout_masks(self, batch_size):
178
+ if not self.training or self.dropout <= 0:
179
+ return None
180
+ mx = _mlx()
181
+ return [
182
+ mx.array((np.random.random((batch_size, weight.data.shape[1])) >= self.dropout).astype(np.float32))
183
+ for weight in self.weights
184
+ ]
185
+
186
+ def forward(self, numeric, categorical, cache=False):
187
+ params = self._mlx_params()
188
+ batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
189
+ dropout_masks = self._make_dropout_masks(batch_size)
190
+ out = np.asarray(self._forward_mlx(params, numeric, categorical, dropout_masks))
191
+ return (out, None) if cache else out
192
+
193
+ def loss_and_backward(self, numeric, categorical, target):
194
+ mx = _mlx()
195
+ params = self._mlx_params()
196
+ batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
197
+ dropout_masks = self._make_dropout_masks(batch_size)
198
+ is_regression = self.task == "regression"
199
+ target_arr = mx.array(target, dtype=mx.float32 if is_regression else mx.int32)
200
+
201
+ def loss_fn(current):
202
+ logits = self._forward_mlx(current, numeric, categorical, dropout_masks)
203
+ if is_regression:
204
+ error = logits - target_arr
205
+ return mx.mean(error * error)
206
+ log_probs = logits - mx.logsumexp(logits, axis=-1, keepdims=True)
207
+ return -mx.mean(mx.take_along_axis(log_probs, target_arr[:, None], axis=1))
208
+
209
+ loss, gradients = mx.value_and_grad(loss_fn)(params)
210
+ for name, parameter in self.named_parameters():
211
+ parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
212
+ return float(np.asarray(loss))
@@ -53,7 +53,14 @@ def build_model(task: str, model_type: str, cfg: Dict[str, Any], meta: Dict[str,
53
53
  model_kwargs["gradient_checkpointing"] = cfg.get("gradient_checkpointing", False)
54
54
  return model_class(**model_kwargs)
55
55
  elif model_type == "mlp":
56
- return TabularMLP(
56
+ model_class = TabularMLP
57
+ if backend == "jax":
58
+ from ..backends.jax_backend import JaxTabularMLP
59
+ model_class = JaxTabularMLP
60
+ elif backend == "mlx":
61
+ from ..backends.mlx_backend import MlxTabularMLP
62
+ model_class = MlxTabularMLP
63
+ model_kwargs = dict(
57
64
  n_numeric=meta["n_numeric"],
58
65
  categorical_vocab_sizes=meta["categorical_vocab_sizes"],
59
66
  d_model=cfg["d_model"],
@@ -62,5 +69,8 @@ def build_model(task: str, model_type: str, cfg: Dict[str, Any], meta: Dict[str,
62
69
  task=task,
63
70
  n_classes=meta.get("n_classes", 0),
64
71
  )
72
+ if backend in ("jax", "mlx"):
73
+ model_kwargs["precision"] = cfg.get("precision", "fp32")
74
+ return model_class(**model_kwargs)
65
75
  else:
66
76
  raise ModelError(f"Unknown model_type '{model_type}'.")
@@ -45,11 +45,10 @@ class LoadedModel:
45
45
  self.device = get_device(device_name)
46
46
 
47
47
  backend = "numpy"
48
- if self.model_type == "transformer" and self.task in ("text-generation", "text-classification"):
49
- if self.device in ("cuda", "tpu"):
50
- backend = "jax"
51
- elif self.device == "mps":
52
- backend = "mlx"
48
+ if self.device in ("cuda", "tpu"):
49
+ backend = "jax"
50
+ elif self.device == "mps":
51
+ backend = "mlx"
53
52
  self.model = build_model(self.task, self.model_type, self.config, self.meta, backend=backend)
54
53
  self.model.load_state_dict(payload["model_state_dict"])
55
54
  self.model.eval()
@@ -66,12 +66,16 @@ def run_training(ds: Dataset, cfg: Dict[str, Any], checkpoint_mgr: CheckpointMan
66
66
  elif task == "text-classification": prepared = dp.prepare_text_classification(ds, cfg, tokenizer=tokenizer, classes=resume_state["meta"]["classes"] if resume_state else None)
67
67
  elif task in ("classification", "regression"): prepared = dp.prepare_tabular(ds, cfg, task, preprocessor)
68
68
  else: raise ValueError(f"Unsupported task '{task}'")
69
+ # Pick a backend that actually runs on the resolved device. Both model
70
+ # types (transformer and mlp) have jax/mlx implementations, so whatever
71
+ # device was auto-detected (or requested) is honored for training --
72
+ # we don't silently fall back to the NumPy/CPU engine just because the
73
+ # model happens to be an MLP.
69
74
  backend = "numpy"
70
- if model_type == "transformer" and task in ("text-generation", "text-classification"):
71
- if device in ("cuda", "tpu"):
72
- backend = "jax"
73
- elif device == "mps":
74
- backend = "mlx"
75
+ if device in ("cuda", "tpu"):
76
+ backend = "jax"
77
+ elif device == "mps":
78
+ backend = "mlx"
75
79
  model = build_model(task, model_type, cfg, prepared.meta, backend=backend)
76
80
  if pretrained_state is not None:
77
81
  try:
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.9.1
3
+ Version: 0.9.2
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
@@ -45,8 +45,11 @@ pip install -e '.[tpu]' # JAX TPU
45
45
  pip install -e '.[mps]' # Apple Silicon MLX
46
46
  ```
47
47
 
48
- CUDA and TPU backends currently accelerate transformer text tasks. Tabular
49
- tasks and unsupported platforms use the native CPU engine.
48
+ CUDA, TPU, and MPS backends accelerate both transformer text tasks and
49
+ tabular MLP tasks. Whatever device is auto-detected (or passed via
50
+ `device=...`) is what training and inference actually run on; only
51
+ unsupported platforms (no CUDA/TPU/MPS available) fall back to the native
52
+ CPU engine.
50
53
 
51
54
  ## Train on your data
52
55
 
@@ -1,29 +0,0 @@
1
- # Installation
2
-
3
- ## Requirements
4
-
5
- - Python 3.9 or later
6
- - NumPy 1.24 or later (installed automatically as a dependency)
7
- - Tensorless currently runs its native vectorized engine on CPU; accelerator
8
- backends are planned without changing the public API.
9
-
10
- ## Verify your installation
11
-
12
- ```bash
13
- python -c "import tensorless as tl; print(tl.__version__)"
14
- tensorless --help
15
- ```
16
-
17
- You should see a version string printed and the CLI's help text.
18
-
19
- ## Runtime support
20
-
21
- The native engine uses NumPy vectorized CPU kernels and automatically resolves
22
- the device to CPU. You can still specify `device="cpu"` explicitly in
23
- `tl.train(...)`; the device option remains forward-compatible with future
24
- accelerator backends.
25
-
26
- ## Troubleshooting installation
27
-
28
- See [troubleshooting.md](troubleshooting.md#installation-issues) for
29
- common installation problems.
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