tensorless 0.4.0__tar.gz → 0.6.0__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 (90) hide show
  1. {tensorless-0.4.0 → tensorless-0.6.0}/PKG-INFO +33 -2
  2. tensorless-0.4.0/tensorless.egg-info/PKG-INFO → tensorless-0.6.0/README.md +20 -15
  3. {tensorless-0.4.0 → tensorless-0.6.0}/docs/architecture.md +19 -3
  4. {tensorless-0.4.0 → tensorless-0.6.0}/docs/checkpointing.md +15 -0
  5. {tensorless-0.4.0 → tensorless-0.6.0}/docs/configuration.md +9 -2
  6. {tensorless-0.4.0 → tensorless-0.6.0}/docs/roadmap.md +11 -8
  7. {tensorless-0.4.0 → tensorless-0.6.0}/pyproject.toml +6 -1
  8. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/api.py +2 -1
  9. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/auto/config.py +2 -0
  10. tensorless-0.6.0/tensorless/backends/__init__.py +1 -0
  11. tensorless-0.6.0/tensorless/backends/jax_backend.py +165 -0
  12. tensorless-0.6.0/tensorless/backends/mlx_backend.py +133 -0
  13. tensorless-0.6.0/tensorless/checkpoint/manager.py +129 -0
  14. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/config.py +4 -0
  15. tensorless-0.6.0/tensorless/devices/__init__.py +4 -0
  16. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/devices/device.py +28 -4
  17. tensorless-0.6.0/tensorless/devices/memory.py +36 -0
  18. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/engine.py +73 -17
  19. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/models/mlp.py +3 -4
  20. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/models/registry.py +17 -2
  21. tensorless-0.6.0/tensorless/models/transformer.py +222 -0
  22. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/runtime.py +7 -1
  23. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/training/data_prep.py +8 -5
  24. tensorless-0.6.0/tensorless/training/trainer.py +138 -0
  25. tensorless-0.6.0/tensorless.egg-info/PKG-INFO +108 -0
  26. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/SOURCES.txt +6 -0
  27. tensorless-0.6.0/tensorless.egg-info/requires.txt +21 -0
  28. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_checkpoint_resume.py +16 -0
  29. tensorless-0.6.0/tests/test_jax_backend.py +42 -0
  30. tensorless-0.6.0/tests/test_mlx_backend.py +44 -0
  31. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_train_tabular.py +28 -0
  32. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_train_text_generation.py +8 -0
  33. tensorless-0.4.0/README.md +0 -63
  34. tensorless-0.4.0/tensorless/checkpoint/manager.py +0 -68
  35. tensorless-0.4.0/tensorless/devices/__init__.py +0 -3
  36. tensorless-0.4.0/tensorless/models/transformer.py +0 -255
  37. tensorless-0.4.0/tensorless/training/trainer.py +0 -317
  38. tensorless-0.4.0/tensorless.egg-info/requires.txt +0 -4
  39. {tensorless-0.4.0 → tensorless-0.6.0}/LICENSE +0 -0
  40. {tensorless-0.4.0 → tensorless-0.6.0}/MANIFEST.in +0 -0
  41. {tensorless-0.4.0 → tensorless-0.6.0}/docs/api_reference.md +0 -0
  42. {tensorless-0.4.0 → tensorless-0.6.0}/docs/automatic_mode.md +0 -0
  43. {tensorless-0.4.0 → tensorless-0.6.0}/docs/cli.md +0 -0
  44. {tensorless-0.4.0 → tensorless-0.6.0}/docs/contributing.md +0 -0
  45. {tensorless-0.4.0 → tensorless-0.6.0}/docs/examples.md +0 -0
  46. {tensorless-0.4.0 → tensorless-0.6.0}/docs/inference.md +0 -0
  47. {tensorless-0.4.0 → tensorless-0.6.0}/docs/installation.md +0 -0
  48. {tensorless-0.4.0 → tensorless-0.6.0}/docs/limitations.md +0 -0
  49. {tensorless-0.4.0 → tensorless-0.6.0}/docs/quickstart.md +0 -0
  50. {tensorless-0.4.0 → tensorless-0.6.0}/docs/tl_format.md +0 -0
  51. {tensorless-0.4.0 → tensorless-0.6.0}/docs/training.md +0 -0
  52. {tensorless-0.4.0 → tensorless-0.6.0}/docs/troubleshooting.md +0 -0
  53. {tensorless-0.4.0 → tensorless-0.6.0}/docs/tutorial.md +0 -0
  54. {tensorless-0.4.0 → tensorless-0.6.0}/examples/tabular_classification_example.py +0 -0
  55. {tensorless-0.4.0 → tensorless-0.6.0}/examples/tabular_regression_example.py +0 -0
  56. {tensorless-0.4.0 → tensorless-0.6.0}/examples/text_classification_example.py +0 -0
  57. {tensorless-0.4.0 → tensorless-0.6.0}/examples/text_generation_example.py +0 -0
  58. {tensorless-0.4.0 → tensorless-0.6.0}/setup.cfg +0 -0
  59. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/__init__.py +0 -0
  60. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/_version.py +0 -0
  61. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/auto/__init__.py +0 -0
  62. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/auto/detector.py +0 -0
  63. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/checkpoint/__init__.py +0 -0
  64. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/cli/__init__.py +0 -0
  65. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/cli/main.py +0 -0
  66. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/__init__.py +0 -0
  67. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/english_grammar.txt +0 -0
  68. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/fingerprint.py +0 -0
  69. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/inspector.py +0 -0
  70. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/loader.py +0 -0
  71. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/tabular.py +0 -0
  72. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/errors.py +0 -0
  73. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/models/__init__.py +0 -0
  74. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/serialization/__init__.py +0 -0
  75. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/serialization/tl_format.py +0 -0
  76. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/tokenization/__init__.py +0 -0
  77. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/tokenization/bpe_tokenizer.py +0 -0
  78. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/tokenization/char_tokenizer.py +0 -0
  79. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/training/__init__.py +0 -0
  80. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/training/early_stopping.py +0 -0
  81. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/dependency_links.txt +0 -0
  82. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/entry_points.txt +0 -0
  83. {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/top_level.txt +0 -0
  84. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_auto_detection.py +0 -0
  85. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_cli.py +0 -0
  86. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_data_loading.py +0 -0
  87. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_end_to_end.py +0 -0
  88. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_fingerprint.py +0 -0
  89. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_serialization.py +0 -0
  90. {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_train_text_classification.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tensorless
3
- Version: 0.4.0
3
+ Version: 0.6.0
4
4
  Summary: ML with maximum automation and minimum setup.
5
5
  Author: Tensorless Contributors
6
6
  License: MIT
@@ -10,11 +10,25 @@ License-File: LICENSE
10
10
  Requires-Dist: numpy>=1.24
11
11
  Provides-Extra: dev
12
12
  Requires-Dist: pytest>=7.0; extra == "dev"
13
+ Provides-Extra: jax
14
+ Requires-Dist: jax>=0.4.30; extra == "jax"
15
+ Provides-Extra: cuda
16
+ Requires-Dist: jax[cuda12]>=0.4.30; extra == "cuda"
17
+ Provides-Extra: tpu
18
+ Requires-Dist: jax[tpu]>=0.4.30; extra == "tpu"
19
+ Provides-Extra: mps
20
+ Requires-Dist: mlx>=0.18; extra == "mps"
21
+ Provides-Extra: accelerators
22
+ Requires-Dist: jax[cuda12]>=0.4.30; extra == "accelerators"
23
+ Requires-Dist: jax[tpu]>=0.4.30; extra == "accelerators"
24
+ Requires-Dist: mlx>=0.18; extra == "accelerators"
13
25
  Dynamic: license-file
14
26
 
15
27
  # Tensorless
16
28
 
17
- Tensorless trains small native NumPy models with sensible defaults. It supports
29
+ Tensorless trains small custom models with sensible defaults. It uses a native
30
+ NumPy engine on CPU and optional JAX or MLX backends for accelerator execution.
31
+ It supports
18
32
  text generation, text classification, tabular classification, and regression.
19
33
 
20
34
  ## Install
@@ -23,6 +37,17 @@ text generation, text classification, tabular classification, and regression.
23
37
  pip install -e .
24
38
  ```
25
39
 
40
+ Optional accelerator backends:
41
+
42
+ ```bash
43
+ pip install -e '.[cuda]' # JAX CUDA
44
+ pip install -e '.[tpu]' # JAX TPU
45
+ pip install -e '.[mps]' # Apple Silicon MLX
46
+ ```
47
+
48
+ CUDA and TPU backends currently accelerate transformer text tasks. Tabular
49
+ tasks and unsupported platforms use the native CPU engine.
50
+
26
51
  ## Train on your data
27
52
 
28
53
  ```python
@@ -73,5 +98,11 @@ model = tl.load("model.tl")
73
98
  print(model.info())
74
99
  ```
75
100
 
101
+ The native extension API is in `tensorless.engine`: `Module`, `Parameter`,
102
+ `Adam`, and `SGD` provide model parameters, gradients, and optimization
103
+ without a PyTorch dependency. Accelerator cache helpers are available as
104
+ `tensorless.devices.clear_memory()` and `tensorless.devices.memory_stats()`.
105
+
76
106
  See the [documentation](docs/quickstart.md) for data formats, configuration,
77
107
  checkpointing, and the command-line interface.
108
+ pypi-AgEIcHlwaS5vcmcCJDJmYzZiYmQ1LTI1YTAtNGNlZi05OWE2LTliMjg3ZjY5MThiZQACElsxLFsidGVuc29ybGVzcyJdXQACLFsyLFsiNjcxZGZmNDQtZmVmMC00MWNiLWIzYzYtZWQzMTI3OWU4NWM5Il1dAAAGIDZAp3AAlp0idrTOMPJ227qF_7W0LDsyUpi6pfqf0aV0
@@ -1,20 +1,8 @@
1
- Metadata-Version: 2.4
2
- Name: tensorless
3
- Version: 0.4.0
4
- Summary: ML with maximum automation and minimum setup.
5
- Author: Tensorless Contributors
6
- License: MIT
7
- Requires-Python: >=3.9
8
- Description-Content-Type: text/markdown
9
- License-File: LICENSE
10
- Requires-Dist: numpy>=1.24
11
- Provides-Extra: dev
12
- Requires-Dist: pytest>=7.0; extra == "dev"
13
- Dynamic: license-file
14
-
15
1
  # Tensorless
16
2
 
17
- Tensorless trains small native NumPy models with sensible defaults. It supports
3
+ Tensorless trains small custom models with sensible defaults. It uses a native
4
+ NumPy engine on CPU and optional JAX or MLX backends for accelerator execution.
5
+ It supports
18
6
  text generation, text classification, tabular classification, and regression.
19
7
 
20
8
  ## Install
@@ -23,6 +11,17 @@ text generation, text classification, tabular classification, and regression.
23
11
  pip install -e .
24
12
  ```
25
13
 
14
+ Optional accelerator backends:
15
+
16
+ ```bash
17
+ pip install -e '.[cuda]' # JAX CUDA
18
+ pip install -e '.[tpu]' # JAX TPU
19
+ pip install -e '.[mps]' # Apple Silicon MLX
20
+ ```
21
+
22
+ CUDA and TPU backends currently accelerate transformer text tasks. Tabular
23
+ tasks and unsupported platforms use the native CPU engine.
24
+
26
25
  ## Train on your data
27
26
 
28
27
  ```python
@@ -73,5 +72,11 @@ model = tl.load("model.tl")
73
72
  print(model.info())
74
73
  ```
75
74
 
75
+ The native extension API is in `tensorless.engine`: `Module`, `Parameter`,
76
+ `Adam`, and `SGD` provide model parameters, gradients, and optimization
77
+ without a PyTorch dependency. Accelerator cache helpers are available as
78
+ `tensorless.devices.clear_memory()` and `tensorless.devices.memory_stats()`.
79
+
76
80
  See the [documentation](docs/quickstart.md) for data formats, configuration,
77
81
  checkpointing, and the command-line interface.
82
+ pypi-AgEIcHlwaS5vcmcCJDJmYzZiYmQ1LTI1YTAtNGNlZi05OWE2LTliMjg3ZjY5MThiZQACElsxLFsidGVuc29ybGVzcyJdXQACLFsyLFsiNjcxZGZmNDQtZmVmMC00MWNiLWIzYzYtZWQzMTI3OWU4NWM5Il1dAAAGIDZAp3AAlp0idrTOMPJ227qF_7W0LDsyUpi6pfqf0aV0
@@ -36,8 +36,12 @@ tensorless/
36
36
  │ └── manager.py CheckpointManager: atomic save/load/clear
37
37
  ├── serialization/
38
38
  │ └── tl_format.py save_tl/load_tl: the .tl file format
39
+ ├── backends/
40
+ │ ├── jax_backend.py optional CUDA/TPU transformer backend and local data parallelism
41
+ │ └── mlx_backend.py optional Apple Silicon MPS backend
39
42
  ├── devices/
40
- │ └── device.py native CPU device resolution
43
+ │ ├── device.py accelerator detection and resolution
44
+ │ └── memory.py backend-neutral cache and memory controls
41
45
  └── cli/
42
46
  └── main.py argparse-based CLI
43
47
  ```
@@ -60,7 +64,7 @@ auto.config.resolve_config(ds, TrainConfig) -> ResolvedConfig (every field concr
60
64
  training.data_prep.prepare_*(ds, cfg) -> PreparedData (train/val DataLoaders,
61
65
  │ meta, tokenizer/preprocessor)
62
66
  ▼
63
- models.registry.build_model(task, model_type, cfg, meta) -> nn.Module
67
+ models.registry.build_model(task, model_type, cfg, meta) -> native/backend model
64
68
  │
65
69
  ▼
66
70
  training.trainer.run_training(...) -> trains, checkpoints periodically,
@@ -97,6 +101,9 @@ matching the dataset's fingerprint.
97
101
  - **A `.tl` file is the unit of portability.** Nothing about inference
98
102
  should require the original dataset, training script, or checkpoint
99
103
  directory to still exist.
104
+ - **The public model API is backend-neutral.** `engine.Module` and
105
+ `engine.Parameter` expose parameters and gradients consistently while JAX
106
+ and MLX provide accelerator execution behind the same model contract.
100
107
 
101
108
  ## Extending Tensorless
102
109
 
@@ -120,7 +127,16 @@ for the new extension, producing a `Dataset` with the appropriate `kind`.
120
127
  ### Adding a new backend/device
121
128
 
122
129
  Extend `devices/device.py`'s `_*_available()` checks and
123
- `auto_select_device()` / `get_torch_device()`.
130
+ `auto_select_device()` / `get_device()`. Add a lazy-imported backend module
131
+ and select it in `models/registry.py`, `training/trainer.py`, and
132
+ `runtime.py`. Backend models must expose the native `state_dict()` contract so
133
+ `.tl` files remain portable.
134
+
135
+ When multiple local JAX CUDA or TPU devices are visible, the JAX transformer
136
+ splits each compatible batch with `pmap` and averages gradients with a device
137
+ collective. For multi-host jobs, set `JAX_COORDINATOR_ADDRESS`,
138
+ `JAX_PROCESS_COUNT`, and `JAX_PROCESS_ID` consistently across processes before
139
+ launching. Single-device training uses the same model without collectives.
124
140
 
125
141
  See [contributing.md](contributing.md) for the contribution process
126
142
  itself (tests, PRs, etc).
@@ -12,9 +12,11 @@ create or manage this directory yourself.
12
12
  `checkpoint.pt`, containing:
13
13
 
14
14
  - `model_state_dict` — model weights
15
+ - `best_model_state_dict` — best validation weights, when validation is enabled
15
16
  - `optimizer_state_dict` — optimizer momentum/variance buffers
16
17
  - `scheduler_state_dict` — learning rate schedule position
17
18
  - `epoch`, `global_step` — where training left off
19
+ - `train_loader_epoch` — deterministic shuffle position for resumed training
18
20
  - `early_stopping_best`, `early_stopping_bad_checks` — early stopping state
19
21
  - `config` — the fully resolved training configuration used
20
22
  - `meta` — task-specific sizing info (vocab size, number of classes, etc.)
@@ -28,6 +30,12 @@ final `.tl` file — nothing about resumption depends on the original
28
30
  dataset still being on disk in the same location, only on it being
29
31
  fingerprint-identical to what was originally used.
30
32
 
33
+ Set `gradient_checkpointing=True` to reduce native transformer activation
34
+ memory. Backward recomputes each block and replays its dropout state, trading
35
+ extra compute for lower peak memory. The JAX backend uses `jax.checkpoint`
36
+ for the same option and uses device collectives for local multi-device
37
+ gradient reduction.
38
+
31
39
  ## When checkpoints are written
32
40
 
33
41
  - Every `checkpoint_every` steps (default: 50) during training, with
@@ -40,6 +48,13 @@ Writes are atomic: Tensorless writes to a temporary file in the same
40
48
  directory and renames it into place, so a crash mid-write never leaves a
41
49
  corrupt checkpoint that would block resumption.
42
50
 
51
+ JAX and MLX accelerator arrays are converted to portable NumPy arrays before
52
+ they enter a `.tl` file or checkpoint. This keeps models loadable on a CPU
53
+ machine. Set `checkpoint_shard_size_mb` above zero to store model and
54
+ optimizer state in independently readable shard files plus a small
55
+ `checkpoint.pt` manifest. The normal resume API reconstructs the state
56
+ automatically.
57
+
43
58
  When loading `.tl` files, Tensorless fills compatible fields introduced by
44
59
  older versions with safe defaults. Files created by a newer unsupported format
45
60
  version are rejected with an upgrade message instead of being partially read.
@@ -25,6 +25,7 @@ defaults.
25
25
  | `heads` | auto | Attention heads (transformer only) |
26
26
  | `ff_mult` | auto | Feed-forward expansion multiplier (transformer only) |
27
27
  | `dropout` | `0.1` | Dropout probability |
28
+ | `gradient_checkpointing` | `False` | Recompute transformer activations during backward to reduce memory use; native and JAX paths support this |
28
29
  | `max_seq_len` | `256` (text) / `1` (tabular) | Max sequence length in tokens (text tasks) |
29
30
  | `tokenizer` | `"bpe"` | Text tokenizer: `"bpe"` or `"char"` |
30
31
  | `bpe_vocab_size` | `1000` | Maximum vocabulary size when `tokenizer="bpe"` |
@@ -54,8 +55,8 @@ defaults.
54
55
 
55
56
  | Field | Default | Description |
56
57
  |---|---|---|
57
- | `device` | auto (`tpu` > `cuda` > `mps` > `cpu`) | Force a specific device |
58
- | `precision` | auto | `"fp32"`, `"fp16"`, or `"bf16"` |
58
+ | `device` | auto (`tpu` > `cuda` > `mps` > `cpu`) | Force a specific device; accelerator backends require their optional package |
59
+ | `precision` | auto | `"fp32"`, `"fp16"`, or `"bf16"`; JAX and MLX accelerator transformers use the selected dtype, fp32 loss math, and fp16 loss scaling |
59
60
 
60
61
  ## Checkpointing
61
62
 
@@ -63,6 +64,7 @@ defaults.
63
64
  |---|---|---|
64
65
  | `checkpoint_every` | `50` | Steps between checkpoint writes |
65
66
  | `checkpoint_dir` | `"<out>.ckpt"` | Checkpoint directory |
67
+ | `checkpoint_shard_size_mb` | `0` | Split model and optimizer checkpoint state into parts of approximately this size; `0` keeps the single-file format |
66
68
 
67
69
  ## Misc
68
70
 
@@ -93,3 +95,8 @@ tl.train(
93
95
  Passing an unrecognized keyword raises a `ConfigError` listing every
94
96
  valid field name, so typos are caught immediately rather than silently
95
97
  ignored.
98
+
99
+ For backend-neutral accelerator maintenance, use
100
+ `tensorless.devices.clear_memory()` and
101
+ `tensorless.devices.memory_stats()`. These helpers use JAX or MLX when
102
+ available and otherwise safely do nothing.
@@ -5,8 +5,9 @@ a commitment or timeline.
5
5
 
6
6
  ## Near-term
7
7
 
8
- - **Progress bars** for training (currently plain print-based logging)
9
- - **Multi-GPU / distributed training** for larger datasets
8
+ - **Hyperparameter search mode**: an opt-in `tl.train("./data",
9
+ search=True)` that tries a small set of configurations and keeps the
10
+ best, rather than a single heuristic choice
10
11
 
11
12
  ## Medium-term
12
13
 
@@ -15,18 +16,20 @@ a commitment or timeline.
15
16
  from scratch
16
17
  - **Additional data formats**: Parquet, Excel (`.xlsx`), image
17
18
  directories, audio
18
- - **Hyperparameter search mode**: an opt-in `tl.train("./data",
19
- search=True)` that tries a small set of configurations and keeps the
20
- best, rather than a single heuristic choice
21
19
  - **Data quality auto-fixes**: currently `tl.inspect()` only *reports*
22
20
  problems like missing values or class imbalance; a future mode could
23
21
  offer to fix them (with explicit user opt-in, consistent with "never
24
22
  silently modify user data")
25
23
 
26
- ## Long-term / exploratory
24
+ ## Implemented
27
25
 
28
- - **Alternate backends** (JAX, a lightweight NumPy-only backend) behind
29
- the same `tl.train()`/`tl.load()` API
26
+ - **Progress bars** for batch and epoch training output
27
+ - **JAX CUDA/TPU backend**, including local and multi-host data parallelism
28
+ - **MLX Apple Silicon backend** for transformer text tasks
29
+ - **NumPy CPU backend** with stacked attention, dropout, gradient checkpointing,
30
+ and sharded checkpoints
31
+
32
+ ## Long-term / exploratory
30
33
  - **Export to other formats** (ONNX, TorchScript) from a `.tl` file for
31
34
  deployment outside Python
32
35
 
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "tensorless"
7
- version = "0.4.0"
7
+ version = "0.6.0"
8
8
  description = "ML with maximum automation and minimum setup."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.9"
@@ -16,6 +16,11 @@ dependencies = [
16
16
 
17
17
  [project.optional-dependencies]
18
18
  dev = ["pytest>=7.0"]
19
+ jax = ["jax>=0.4.30"]
20
+ cuda = ["jax[cuda12]>=0.4.30"]
21
+ tpu = ["jax[tpu]>=0.4.30"]
22
+ mps = ["mlx>=0.18"]
23
+ accelerators = ["jax[cuda12]>=0.4.30", "jax[tpu]>=0.4.30", "mlx>=0.18"]
19
24
 
20
25
  [project.scripts]
21
26
  tensorless = "tensorless.cli.main:main"
@@ -179,6 +179,7 @@ def pretrain(
179
179
 
180
180
 
181
181
  def _finalize_from_checkpoint(ckpt: dict, out: str, verbose: bool) -> LoadedModel:
182
+ model_state = ckpt.get("best_model_state_dict") or ckpt["model_state_dict"]
182
183
  payload = {
183
184
  "tl_format_version": _tl_format_version,
184
185
  "tensorless_version": _tl_version,
@@ -186,7 +187,7 @@ def _finalize_from_checkpoint(ckpt: dict, out: str, verbose: bool) -> LoadedMode
186
187
  "model_type": ckpt["config"]["model_type"],
187
188
  "config": ckpt["config"],
188
189
  "meta": ckpt["meta"],
189
- "model_state_dict": ckpt["model_state_dict"],
190
+ "model_state_dict": model_state,
190
191
  "tokenizer_state": ckpt.get("tokenizer_state"),
191
192
  "preprocessor_state": ckpt.get("preprocessor_state"),
192
193
  "dataset_fingerprint": ckpt["dataset_fingerprint"],
@@ -111,6 +111,7 @@ def resolve_config(ds: Dataset, user: TrainConfig) -> ResolvedConfig:
111
111
  heads=user.heads or heads,
112
112
  ff_mult=user.ff_mult or ff_mult,
113
113
  dropout=user.dropout if user.dropout is not None else 0.1,
114
+ gradient_checkpointing=bool(user.gradient_checkpointing),
114
115
  max_seq_len=max_seq_len,
115
116
  tokenizer=tokenizer,
116
117
  bpe_vocab_size=user.bpe_vocab_size if user.bpe_vocab_size is not None else _auto_vocab_size(ds),
@@ -129,6 +130,7 @@ def resolve_config(ds: Dataset, user: TrainConfig) -> ResolvedConfig:
129
130
  precision=precision,
130
131
  checkpoint_every=user.checkpoint_every or 50,
131
132
  checkpoint_dir=checkpoint_dir,
133
+ checkpoint_shard_size_mb=user.checkpoint_shard_size_mb or 0,
132
134
  seed=user.seed,
133
135
  verbose=user.verbose,
134
136
  )
@@ -0,0 +1 @@
1
+ """Optional execution backends."""
@@ -0,0 +1,165 @@
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
+
16
+
17
+ def _jax():
18
+ try:
19
+ import jax
20
+ import jax.numpy as jnp
21
+ distributed_vars = ("JAX_COORDINATOR_ADDRESS", "JAX_PROCESS_COUNT", "JAX_PROCESS_ID")
22
+ if all(name in os.environ for name in distributed_vars) and not jax.distributed.is_initialized():
23
+ jax.distributed.initialize(
24
+ coordinator_address=os.environ["JAX_COORDINATOR_ADDRESS"],
25
+ num_processes=int(os.environ["JAX_PROCESS_COUNT"]),
26
+ process_id=int(os.environ["JAX_PROCESS_ID"]),
27
+ )
28
+ except ImportError as exc:
29
+ raise ImportError("JAX is required for the cuda/tpu backend; install tensorless[jax].") from exc
30
+ return jax, jnp
31
+
32
+
33
+ def is_available() -> bool:
34
+ try:
35
+ _jax()
36
+ return True
37
+ except ImportError:
38
+ return False
39
+
40
+
41
+ class JaxTinyTransformer(TinyTransformer):
42
+ """The native transformer parameter layout with JAX math and gradients."""
43
+
44
+ def __init__(self, *args, **kwargs):
45
+ self.precision = kwargs.pop("precision", "fp32")
46
+ super().__init__(*args, **kwargs)
47
+ d_model = kwargs.get("d_model", args[1] if len(args) > 1 else None)
48
+ self.heads = kwargs.get("heads", args[3] if len(args) > 3 else None)
49
+ if d_model is None or self.heads is None:
50
+ raise ValueError("JAX transformer requires d_model and heads")
51
+ self.d_model = d_model
52
+ self.head_dim = d_model // self.heads
53
+
54
+ def _jax_params(self):
55
+ _, jnp = _jax()
56
+ dtype = {"fp16": jnp.float16, "bf16": jnp.bfloat16}.get(self.precision, jnp.float32)
57
+ return {name: jnp.asarray(value.data, dtype=dtype) for name, value in self.named_parameters()}
58
+
59
+ def _forward_jax(self, params, input_ids, attention_mask=None):
60
+ _, jnp = _jax()
61
+
62
+ def layer_norm(x, gain, bias, eps=1e-5):
63
+ mean = jnp.mean(x, axis=-1, keepdims=True)
64
+ var = jnp.var(x, axis=-1, keepdims=True)
65
+ return (x - mean) / jnp.sqrt(var + eps) * gain + bias
66
+
67
+ input_ids = jnp.asarray(input_ids, dtype=jnp.int32)
68
+ batch, length = input_ids.shape
69
+ hidden = params["tok_emb"][input_ids] + params["pos_emb"][jnp.arange(length)]
70
+ causal = jnp.tril(jnp.ones((length, length), dtype=bool))
71
+ for index in range(len(self.blocks)):
72
+ prefix = f"blocks.{index}"
73
+ normed1 = layer_norm(hidden, params[f"{prefix}.ln1_gain"], params[f"{prefix}.ln1_bias"])
74
+ qkv = normed1 @ params[f"{prefix}.qkv_weight"] + params[f"{prefix}.qkv_bias"]
75
+ q, k, v = jnp.split(qkv, 3, axis=-1)
76
+ q = q.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
77
+ k = k.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
78
+ v = v.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
79
+ scores = q @ k.transpose(0, 1, 3, 2) / np.sqrt(self.head_dim)
80
+ scores = jnp.where(causal[None, None, :, :], scores, -1e9)
81
+ if attention_mask is not None:
82
+ mask = jnp.asarray(attention_mask) > 0
83
+ scores = jnp.where(mask[:, None, None, :], scores, -1e9)
84
+ probabilities = jnp.exp(scores - scores.max(axis=-1, keepdims=True))
85
+ probabilities /= probabilities.sum(axis=-1, keepdims=True)
86
+ context = probabilities @ v
87
+ context = context.transpose(0, 2, 1, 3).reshape(batch, length, self.d_model)
88
+ attention = context @ params[f"{prefix}.out_weight"] + params[f"{prefix}.out_bias"]
89
+ residual = hidden + attention
90
+ normed2 = layer_norm(residual, params[f"{prefix}.ln2_gain"], params[f"{prefix}.ln2_bias"])
91
+ ff_pre = normed2 @ params[f"{prefix}.ff1_weight"] + params[f"{prefix}.ff1_bias"]
92
+ ff_hidden = jnp.maximum(ff_pre, 0)
93
+ ff = ff_hidden @ params[f"{prefix}.ff2_weight"] + params[f"{prefix}.ff2_bias"]
94
+ hidden = residual + ff
95
+ hidden = layer_norm(hidden, params["ln_f_gain"], params["ln_f_bias"])
96
+ if self.task == "text-generation":
97
+ return hidden @ params["tok_emb"].T + params["head_bias"]
98
+ mask = jnp.ones((batch, length), dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
99
+ pooled = (hidden * mask[:, :, None]).sum(1) / jnp.maximum(mask.sum(1, keepdims=True), 1)
100
+ return pooled @ params["head_weight"] + params["head_bias"]
101
+
102
+ def forward(self, input_ids, attention_mask=None, cache=False):
103
+ _, jnp = _jax()
104
+ output = self._forward_jax(self._jax_params(), input_ids, attention_mask)
105
+ output = np.asarray(output)
106
+ return (output, None) if cache else output
107
+
108
+ def loss_and_backward(self, input_ids, target, attention_mask=None):
109
+ jax, jnp = _jax()
110
+ params = self._jax_params()
111
+ target = jnp.asarray(target, dtype=jnp.int32)
112
+
113
+ def loss_fn(current, current_ids, current_target, current_mask):
114
+ forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
115
+ logits = forward(current, current_ids, current_mask)
116
+ flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
117
+ flat_target = current_target.reshape(-1)
118
+ log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
119
+ losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
120
+ if self.task == "text-generation":
121
+ valid = flat_target != self.pad_id
122
+ return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
123
+ return jnp.mean(losses)
124
+
125
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
126
+ scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
127
+ mask = jnp.ones_like(input_ids, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
128
+ devices = jax.local_device_count()
129
+ if devices > 1 and input_ids.shape[0] % devices == 0:
130
+ per_device_loss = jax.value_and_grad(scaled_loss_fn)
131
+
132
+ def mapped_step(current, ids, labels, current_mask):
133
+ value, gradients = per_device_loss(current, ids, labels, current_mask)
134
+ return jax.lax.pmean(value, "data"), jax.tree_util.tree_map(
135
+ lambda gradient: jax.lax.pmean(gradient, "data"), gradients
136
+ )
137
+
138
+ shard = lambda value: value.reshape((devices, value.shape[0] // devices) + value.shape[1:])
139
+ loss, gradients = jax.pmap(mapped_step, axis_name="data", in_axes=(None, 0, 0, 0))(
140
+ params, shard(jnp.asarray(input_ids)), shard(target), shard(mask)
141
+ )
142
+ loss = loss[0]
143
+ gradients = jax.tree_util.tree_map(lambda gradient: gradient[0] / loss_scale, gradients)
144
+ else:
145
+ loss, gradients = jax.value_and_grad(scaled_loss_fn)(params, input_ids, target, mask)
146
+ gradients = jax.tree_util.tree_map(lambda gradient: gradient / loss_scale, gradients)
147
+ for name, parameter in self.named_parameters():
148
+ parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
149
+ return float(loss) / loss_scale
150
+
151
+ def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, eos_id: Optional[int] = None):
152
+ ids = np.asarray(input_ids, dtype=np.int64).copy()
153
+ for _ in range(max_new_tokens):
154
+ logits = self.forward(ids[:, -self.max_seq_len:])[:, -1, :] / max(temperature, 1e-5)
155
+ if top_k is not None:
156
+ k = min(top_k, logits.shape[-1])
157
+ excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
158
+ logits[np.arange(len(ids))[:, None], excluded] = -np.inf
159
+ probs = np.exp(logits - logits.max(1, keepdims=True))
160
+ probs /= probs.sum(1, keepdims=True)
161
+ next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
162
+ ids = np.concatenate([ids, next_ids], axis=1)
163
+ if eos_id is not None and np.all(next_ids == eos_id):
164
+ break
165
+ return ids
@@ -0,0 +1,133 @@
1
+ """MLX implementation of the text transformer.
2
+
3
+ MLX is imported lazily so Tensorless remains usable on non-Apple systems
4
+ without installing an accelerator framework.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from typing import Optional
10
+
11
+ import numpy as np
12
+
13
+ from ..models.transformer import TinyTransformer
14
+
15
+
16
+ def _mlx():
17
+ try:
18
+ import mlx.core as mx
19
+ except ImportError as exc:
20
+ raise ImportError("MLX is required for the mps backend; install tensorless[mps].") from exc
21
+ return mx
22
+
23
+
24
+ def is_available() -> bool:
25
+ try:
26
+ _mlx()
27
+ return True
28
+ except ImportError:
29
+ return False
30
+
31
+
32
+ class MlxTinyTransformer(TinyTransformer):
33
+ """The native Tensorless parameter layout with MLX math and gradients."""
34
+
35
+ def __init__(self, *args, **kwargs):
36
+ self.precision = kwargs.pop("precision", "fp32")
37
+ super().__init__(*args, **kwargs)
38
+ d_model = kwargs.get("d_model", args[1] if len(args) > 1 else None)
39
+ self.heads = kwargs.get("heads", args[3] if len(args) > 3 else None)
40
+ if d_model is None or self.heads is None:
41
+ raise ValueError("MLX transformer requires d_model and heads")
42
+ self.d_model = d_model
43
+ self.head_dim = d_model // self.heads
44
+
45
+ def _mlx_params(self):
46
+ mx = _mlx()
47
+ dtype = {"fp16": mx.float16, "bf16": mx.bfloat16}.get(self.precision, mx.float32)
48
+ return {name: mx.array(parameter.data, dtype=dtype) for name, parameter in self.named_parameters()}
49
+
50
+ def _forward_mlx(self, params, input_ids, attention_mask=None):
51
+ mx = _mlx()
52
+
53
+ def layer_norm(x, gain, bias, eps=1e-5):
54
+ mean = mx.mean(x, axis=-1, keepdims=True)
55
+ var = mx.var(x, axis=-1, keepdims=True)
56
+ return (x - mean) / mx.sqrt(var + eps) * gain + bias
57
+
58
+ input_ids = mx.array(input_ids, dtype=mx.int32)
59
+ batch, length = input_ids.shape
60
+ hidden = params["tok_emb"][input_ids] + params["pos_emb"][mx.arange(length)]
61
+ causal = mx.tril(mx.ones((length, length), dtype=mx.bool_))
62
+ for index in range(len(self.blocks)):
63
+ prefix = f"blocks.{index}"
64
+ normed1 = layer_norm(hidden, params[f"{prefix}.ln1_gain"], params[f"{prefix}.ln1_bias"])
65
+ qkv = normed1 @ params[f"{prefix}.qkv_weight"] + params[f"{prefix}.qkv_bias"]
66
+ q, k, v = mx.split(qkv, 3, axis=-1)
67
+ q = q.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
68
+ k = k.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
69
+ v = v.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
70
+ scores = q @ k.transpose(0, 1, 3, 2) / np.sqrt(self.head_dim)
71
+ scores = mx.where(causal[None, None, :, :], scores, -1e9)
72
+ if attention_mask is not None:
73
+ mask = mx.array(attention_mask) > 0
74
+ scores = mx.where(mask[:, None, None, :], scores, -1e9)
75
+ probabilities = mx.softmax(scores, axis=-1)
76
+ context = probabilities @ v
77
+ context = context.transpose(0, 2, 1, 3).reshape(batch, length, self.d_model)
78
+ attention = context @ params[f"{prefix}.out_weight"] + params[f"{prefix}.out_bias"]
79
+ residual = hidden + attention
80
+ normed2 = layer_norm(residual, params[f"{prefix}.ln2_gain"], params[f"{prefix}.ln2_bias"])
81
+ ff_pre = normed2 @ params[f"{prefix}.ff1_weight"] + params[f"{prefix}.ff1_bias"]
82
+ ff_hidden = mx.maximum(ff_pre, 0)
83
+ ff = ff_hidden @ params[f"{prefix}.ff2_weight"] + params[f"{prefix}.ff2_bias"]
84
+ hidden = residual + ff
85
+ hidden = layer_norm(hidden, params["ln_f_gain"], params["ln_f_bias"])
86
+ if self.task == "text-generation":
87
+ return hidden @ params["tok_emb"].T + params["head_bias"]
88
+ mask = mx.ones((batch, length)) if attention_mask is None else mx.array(attention_mask)
89
+ pooled = (hidden * mask[:, :, None]).sum(1) / mx.maximum(mask.sum(1, keepdims=True), 1)
90
+ return pooled @ params["head_weight"] + params["head_bias"]
91
+
92
+ def forward(self, input_ids, attention_mask=None, cache=False):
93
+ output = np.asarray(self._forward_mlx(self._mlx_params(), input_ids, attention_mask))
94
+ return (output, None) if cache else output
95
+
96
+ def loss_and_backward(self, input_ids, target, attention_mask=None):
97
+ mx = _mlx()
98
+ target = mx.array(target, dtype=mx.int32)
99
+ params = self._mlx_params()
100
+
101
+ def loss_fn(current):
102
+ logits = self._forward_mlx(current, input_ids, attention_mask)
103
+ flat_logits = logits.reshape((-1, logits.shape[-1])).astype(mx.float32)
104
+ flat_target = target.reshape((-1,))
105
+ log_probs = flat_logits - mx.logsumexp(flat_logits, axis=-1, keepdims=True)
106
+ losses = -mx.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
107
+ if self.task == "text-generation":
108
+ valid = flat_target != self.pad_id
109
+ return mx.sum(mx.where(valid, losses, 0.0)) / mx.maximum(mx.sum(valid), 1)
110
+ return mx.mean(losses)
111
+
112
+ loss_scale = 128.0 if self.precision == "fp16" else 1.0
113
+ scaled_loss_fn = lambda current: loss_fn(current) * loss_scale
114
+ loss, gradients = mx.value_and_grad(scaled_loss_fn)(params)
115
+ for name, parameter in self.named_parameters():
116
+ parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32) / loss_scale
117
+ return float(np.asarray(loss)) / loss_scale
118
+
119
+ def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, eos_id: Optional[int] = None):
120
+ ids = np.asarray(input_ids, dtype=np.int64).copy()
121
+ for _ in range(max_new_tokens):
122
+ logits = self.forward(ids[:, -self.max_seq_len:])[:, -1, :] / max(temperature, 1e-5)
123
+ if top_k is not None:
124
+ k = min(top_k, logits.shape[-1])
125
+ excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
126
+ logits[np.arange(len(ids))[:, None], excluded] = -np.inf
127
+ probs = np.exp(logits - logits.max(1, keepdims=True))
128
+ probs /= probs.sum(1, keepdims=True)
129
+ next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
130
+ ids = np.concatenate([ids, next_ids], axis=1)
131
+ if eos_id is not None and np.all(next_ids == eos_id):
132
+ break
133
+ return ids