tensorless 0.9.2__tar.gz → 0.9.3__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {tensorless-0.9.2/tensorless.egg-info → tensorless-0.9.3}/PKG-INFO +1 -1
- {tensorless-0.9.2 → tensorless-0.9.3}/pyproject.toml +1 -1
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/backends/jax_backend.py +89 -36
- {tensorless-0.9.2 → tensorless-0.9.3/tensorless.egg-info}/PKG-INFO +1 -1
- {tensorless-0.9.2 → tensorless-0.9.3}/LICENSE +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/MANIFEST.in +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/README.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/api_reference.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/architecture.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/automatic_mode.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/checkpointing.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/cli.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/configuration.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/contributing.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/examples.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/inference.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/installation.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/limitations.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/quickstart.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/roadmap.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/tl_format.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/training.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/troubleshooting.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/docs/tutorial.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/examples/tabular_classification_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/examples/tabular_regression_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/examples/text_classification_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/examples/text_generation_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/setup.cfg +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/_version.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/api.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/auto/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/auto/config.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/auto/detector.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/backends/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/backends/mlx_backend.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/checkpoint/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/checkpoint/manager.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/cli/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/cli/main.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/config.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/english_grammar.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/fingerprint.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/inspector.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/loader.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/data/tabular.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/devices/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/devices/device.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/devices/memory.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/engine.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/errors.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/mlp.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/registry.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/models/transformer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/runtime.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/serialization/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/serialization/tl_format.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/tokenization/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/tokenization/bpe_tokenizer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/tokenization/char_tokenizer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/data_prep.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/early_stopping.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless/training/trainer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/SOURCES.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/dependency_links.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/entry_points.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/requires.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tensorless.egg-info/top_level.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_auto_detection.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_checkpoint_resume.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_cli.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_data_loading.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_end_to_end.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_fingerprint.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_jax_backend.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_mlx_backend.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_serialization.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_train_tabular.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_train_text_classification.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.3}/tests/test_train_text_generation.py +0 -0
|
@@ -101,33 +101,61 @@ class JaxTinyTransformer(TinyTransformer):
|
|
|
101
101
|
return pooled @ params["head_weight"] + params["head_bias"]
|
|
102
102
|
|
|
103
103
|
def forward(self, input_ids, attention_mask=None, cache=False):
|
|
104
|
-
|
|
105
|
-
output = self.
|
|
104
|
+
params = self._jax_params()
|
|
105
|
+
output = self._get_forward_jit()(params, input_ids, attention_mask)
|
|
106
106
|
output = np.asarray(output)
|
|
107
107
|
return (output, None) if cache else output
|
|
108
108
|
|
|
109
|
-
def
|
|
110
|
-
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
|
|
127
|
-
|
|
128
|
-
|
|
129
|
-
|
|
130
|
-
|
|
109
|
+
def _get_forward_jit(self):
|
|
110
|
+
# Compile once per instance and reuse. Without this, every call
|
|
111
|
+
# re-traces the whole forward pass in Python and dispatches ops to
|
|
112
|
+
# the accelerator one at a time -- which shows up as pegged CPU
|
|
113
|
+
# (tracing/dispatch overhead) with near-zero GPU utilization, even
|
|
114
|
+
# though the correct device was selected.
|
|
115
|
+
if getattr(self, "_forward_jit_fn", None) is None:
|
|
116
|
+
jax, _ = _jax()
|
|
117
|
+
self._forward_jit_fn = jax.jit(self._forward_jax)
|
|
118
|
+
return self._forward_jit_fn
|
|
119
|
+
|
|
120
|
+
def _get_loss_and_grad_jit(self):
|
|
121
|
+
if getattr(self, "_loss_and_grad_jit_fn", None) is None:
|
|
122
|
+
jax, jnp = _jax()
|
|
123
|
+
|
|
124
|
+
def loss_fn(current, current_ids, current_target, current_mask):
|
|
125
|
+
forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
|
|
126
|
+
logits = forward(current, current_ids, current_mask)
|
|
127
|
+
flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
|
|
128
|
+
flat_target = current_target.reshape(-1)
|
|
129
|
+
log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
|
|
130
|
+
losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
|
|
131
|
+
if self.task == "text-generation":
|
|
132
|
+
valid = flat_target != self.pad_id
|
|
133
|
+
return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
|
|
134
|
+
return jnp.mean(losses)
|
|
135
|
+
|
|
136
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
137
|
+
scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
|
|
138
|
+
self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(scaled_loss_fn))
|
|
139
|
+
return self._loss_and_grad_jit_fn
|
|
140
|
+
|
|
141
|
+
def _get_pmap_step(self):
|
|
142
|
+
if getattr(self, "_pmap_step_fn", None) is None:
|
|
143
|
+
jax, jnp = _jax()
|
|
144
|
+
|
|
145
|
+
def loss_fn(current, current_ids, current_target, current_mask):
|
|
146
|
+
forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
|
|
147
|
+
logits = forward(current, current_ids, current_mask)
|
|
148
|
+
flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
|
|
149
|
+
flat_target = current_target.reshape(-1)
|
|
150
|
+
log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
|
|
151
|
+
losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
|
|
152
|
+
if self.task == "text-generation":
|
|
153
|
+
valid = flat_target != self.pad_id
|
|
154
|
+
return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
|
|
155
|
+
return jnp.mean(losses)
|
|
156
|
+
|
|
157
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
158
|
+
scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
|
|
131
159
|
per_device_loss = jax.value_and_grad(scaled_loss_fn)
|
|
132
160
|
|
|
133
161
|
def mapped_step(current, ids, labels, current_mask):
|
|
@@ -136,14 +164,25 @@ class JaxTinyTransformer(TinyTransformer):
|
|
|
136
164
|
lambda gradient: jax.lax.pmean(gradient, "data"), gradients
|
|
137
165
|
)
|
|
138
166
|
|
|
167
|
+
self._pmap_step_fn = jax.pmap(mapped_step, axis_name="data", in_axes=(None, 0, 0, 0))
|
|
168
|
+
return self._pmap_step_fn
|
|
169
|
+
|
|
170
|
+
def loss_and_backward(self, input_ids, target, attention_mask=None):
|
|
171
|
+
jax, jnp = _jax()
|
|
172
|
+
params = self._jax_params()
|
|
173
|
+
target = jnp.asarray(target, dtype=jnp.int32)
|
|
174
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
175
|
+
mask = jnp.ones_like(input_ids, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
|
|
176
|
+
devices = jax.local_device_count()
|
|
177
|
+
if devices > 1 and input_ids.shape[0] % devices == 0:
|
|
139
178
|
shard = lambda value: value.reshape((devices, value.shape[0] // devices) + value.shape[1:])
|
|
140
|
-
loss, gradients =
|
|
179
|
+
loss, gradients = self._get_pmap_step()(
|
|
141
180
|
params, shard(jnp.asarray(input_ids)), shard(target), shard(mask)
|
|
142
181
|
)
|
|
143
182
|
loss = loss[0]
|
|
144
183
|
gradients = jax.tree_util.tree_map(lambda gradient: gradient[0] / loss_scale, gradients)
|
|
145
184
|
else:
|
|
146
|
-
loss, gradients =
|
|
185
|
+
loss, gradients = self._get_loss_and_grad_jit()(params, input_ids, target, mask)
|
|
147
186
|
gradients = jax.tree_util.tree_map(lambda gradient: gradient / loss_scale, gradients)
|
|
148
187
|
for name, parameter in self.named_parameters():
|
|
149
188
|
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
|
|
@@ -213,9 +252,31 @@ class JaxTabularMLP(TabularMLP):
|
|
|
213
252
|
params = self._jax_params()
|
|
214
253
|
batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
|
|
215
254
|
dropout_masks = self._make_dropout_masks(batch_size)
|
|
216
|
-
out = np.asarray(self.
|
|
255
|
+
out = np.asarray(self._get_forward_jit()(params, numeric, categorical, dropout_masks))
|
|
217
256
|
return (out, None) if cache else out
|
|
218
257
|
|
|
258
|
+
def _get_forward_jit(self):
|
|
259
|
+
if getattr(self, "_forward_jit_fn", None) is None:
|
|
260
|
+
jax, _ = _jax()
|
|
261
|
+
self._forward_jit_fn = jax.jit(self._forward_jax)
|
|
262
|
+
return self._forward_jit_fn
|
|
263
|
+
|
|
264
|
+
def _get_loss_and_grad_jit(self):
|
|
265
|
+
if getattr(self, "_loss_and_grad_jit_fn", None) is None:
|
|
266
|
+
jax, jnp = _jax()
|
|
267
|
+
is_regression = self.task == "regression"
|
|
268
|
+
|
|
269
|
+
def loss_fn(current, numeric, categorical, dropout_masks, target_arr):
|
|
270
|
+
logits = self._forward_jax(current, numeric, categorical, dropout_masks)
|
|
271
|
+
if is_regression:
|
|
272
|
+
error = logits - target_arr
|
|
273
|
+
return jnp.mean(error * error)
|
|
274
|
+
log_probs = jax.nn.log_softmax(logits, axis=-1)
|
|
275
|
+
return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
|
|
276
|
+
|
|
277
|
+
self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(loss_fn))
|
|
278
|
+
return self._loss_and_grad_jit_fn
|
|
279
|
+
|
|
219
280
|
def loss_and_backward(self, numeric, categorical, target):
|
|
220
281
|
jax, jnp = _jax()
|
|
221
282
|
params = self._jax_params()
|
|
@@ -224,15 +285,7 @@ class JaxTabularMLP(TabularMLP):
|
|
|
224
285
|
is_regression = self.task == "regression"
|
|
225
286
|
target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
|
|
226
287
|
|
|
227
|
-
|
|
228
|
-
logits = self._forward_jax(current, numeric, categorical, dropout_masks)
|
|
229
|
-
if is_regression:
|
|
230
|
-
error = logits - target_arr
|
|
231
|
-
return jnp.mean(error * error)
|
|
232
|
-
log_probs = jax.nn.log_softmax(logits, axis=-1)
|
|
233
|
-
return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
|
|
234
|
-
|
|
235
|
-
loss, gradients = jax.value_and_grad(loss_fn)(params)
|
|
288
|
+
loss, gradients = self._get_loss_and_grad_jit()(params, numeric, categorical, dropout_masks, target_arr)
|
|
236
289
|
for name, parameter in self.named_parameters():
|
|
237
290
|
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
|
|
238
291
|
return float(loss)
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|