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.
- {tensorless-0.4.0 → tensorless-0.6.0}/PKG-INFO +33 -2
- tensorless-0.4.0/tensorless.egg-info/PKG-INFO → tensorless-0.6.0/README.md +20 -15
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/architecture.md +19 -3
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/checkpointing.md +15 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/configuration.md +9 -2
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/roadmap.md +11 -8
- {tensorless-0.4.0 → tensorless-0.6.0}/pyproject.toml +6 -1
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/api.py +2 -1
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/auto/config.py +2 -0
- tensorless-0.6.0/tensorless/backends/__init__.py +1 -0
- tensorless-0.6.0/tensorless/backends/jax_backend.py +165 -0
- tensorless-0.6.0/tensorless/backends/mlx_backend.py +133 -0
- tensorless-0.6.0/tensorless/checkpoint/manager.py +129 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/config.py +4 -0
- tensorless-0.6.0/tensorless/devices/__init__.py +4 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/devices/device.py +28 -4
- tensorless-0.6.0/tensorless/devices/memory.py +36 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/engine.py +73 -17
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/models/mlp.py +3 -4
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/models/registry.py +17 -2
- tensorless-0.6.0/tensorless/models/transformer.py +222 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/runtime.py +7 -1
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/training/data_prep.py +8 -5
- tensorless-0.6.0/tensorless/training/trainer.py +138 -0
- tensorless-0.6.0/tensorless.egg-info/PKG-INFO +108 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/SOURCES.txt +6 -0
- tensorless-0.6.0/tensorless.egg-info/requires.txt +21 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_checkpoint_resume.py +16 -0
- tensorless-0.6.0/tests/test_jax_backend.py +42 -0
- tensorless-0.6.0/tests/test_mlx_backend.py +44 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_train_tabular.py +28 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_train_text_generation.py +8 -0
- tensorless-0.4.0/README.md +0 -63
- tensorless-0.4.0/tensorless/checkpoint/manager.py +0 -68
- tensorless-0.4.0/tensorless/devices/__init__.py +0 -3
- tensorless-0.4.0/tensorless/models/transformer.py +0 -255
- tensorless-0.4.0/tensorless/training/trainer.py +0 -317
- tensorless-0.4.0/tensorless.egg-info/requires.txt +0 -4
- {tensorless-0.4.0 → tensorless-0.6.0}/LICENSE +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/MANIFEST.in +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/api_reference.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/automatic_mode.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/cli.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/contributing.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/examples.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/inference.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/installation.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/limitations.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/quickstart.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/tl_format.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/training.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/troubleshooting.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/docs/tutorial.md +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/examples/tabular_classification_example.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/examples/tabular_regression_example.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/examples/text_classification_example.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/examples/text_generation_example.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/setup.cfg +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/_version.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/auto/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/auto/detector.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/checkpoint/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/cli/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/cli/main.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/english_grammar.txt +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/fingerprint.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/inspector.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/loader.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/data/tabular.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/errors.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/models/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/serialization/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/serialization/tl_format.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/tokenization/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/tokenization/bpe_tokenizer.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/tokenization/char_tokenizer.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/training/__init__.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless/training/early_stopping.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/dependency_links.txt +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/entry_points.txt +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tensorless.egg-info/top_level.txt +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_auto_detection.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_cli.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_data_loading.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_end_to_end.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_fingerprint.py +0 -0
- {tensorless-0.4.0 → tensorless-0.6.0}/tests/test_serialization.py +0 -0
- {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.
|
|
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
|
|
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
|
|
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
|
-
│
|
|
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) ->
|
|
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()` / `
|
|
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
|
-
- **
|
|
9
|
-
|
|
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
|
-
##
|
|
24
|
+
## Implemented
|
|
27
25
|
|
|
28
|
-
- **
|
|
29
|
-
|
|
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.
|
|
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":
|
|
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
|