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.
- {tensorless-0.9.1/tensorless.egg-info → tensorless-0.9.2}/PKG-INFO +6 -3
- {tensorless-0.9.1 → tensorless-0.9.2}/README.md +5 -2
- tensorless-0.9.2/docs/installation.md +49 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/pyproject.toml +1 -1
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/backends/jax_backend.py +73 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/backends/mlx_backend.py +79 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/registry.py +11 -1
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/runtime.py +4 -5
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/trainer.py +9 -5
- {tensorless-0.9.1 → tensorless-0.9.2/tensorless.egg-info}/PKG-INFO +6 -3
- tensorless-0.9.1/docs/installation.md +0 -29
- {tensorless-0.9.1 → tensorless-0.9.2}/LICENSE +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/MANIFEST.in +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/api_reference.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/architecture.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/automatic_mode.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/checkpointing.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/cli.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/configuration.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/contributing.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/examples.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/inference.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/limitations.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/quickstart.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/roadmap.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/tl_format.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/training.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/troubleshooting.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/docs/tutorial.md +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/examples/tabular_classification_example.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/examples/tabular_regression_example.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/examples/text_classification_example.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/examples/text_generation_example.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/setup.cfg +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/_version.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/api.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/auto/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/auto/config.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/auto/detector.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/backends/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/checkpoint/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/checkpoint/manager.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/cli/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/cli/main.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/config.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/english_grammar.txt +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/fingerprint.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/inspector.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/loader.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/data/tabular.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/devices/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/devices/device.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/devices/memory.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/engine.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/errors.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/mlp.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/models/transformer.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/serialization/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/serialization/tl_format.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/tokenization/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/tokenization/bpe_tokenizer.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/tokenization/char_tokenizer.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/__init__.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/data_prep.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless/training/early_stopping.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/SOURCES.txt +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/dependency_links.txt +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/entry_points.txt +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/requires.txt +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tensorless.egg-info/top_level.txt +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_auto_detection.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_checkpoint_resume.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_cli.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_data_loading.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_end_to_end.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_fingerprint.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_jax_backend.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_mlx_backend.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_serialization.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_train_tabular.py +0 -0
- {tensorless-0.9.1 → tensorless-0.9.2}/tests/test_train_text_classification.py +0 -0
- {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.
|
|
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
|
|
49
|
-
tasks
|
|
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
|
|
23
|
-
tasks
|
|
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.
|
|
@@ -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
|
-
|
|
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.
|
|
49
|
-
|
|
50
|
-
|
|
51
|
-
|
|
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
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
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.
|
|
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
|
|
49
|
-
tasks
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|