tensorless 0.9.2__tar.gz → 0.9.4__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.4}/PKG-INFO +1 -1
- {tensorless-0.9.2 → tensorless-0.9.4}/pyproject.toml +1 -1
- tensorless-0.9.4/tensorless/backends/jax_backend.py +518 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/backends/mlx_backend.py +20 -1
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/devices/device.py +26 -2
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/trainer.py +76 -10
- {tensorless-0.9.2 → tensorless-0.9.4/tensorless.egg-info}/PKG-INFO +1 -1
- tensorless-0.9.2/tensorless/backends/jax_backend.py +0 -238
- {tensorless-0.9.2 → tensorless-0.9.4}/LICENSE +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/MANIFEST.in +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/README.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/api_reference.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/architecture.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/automatic_mode.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/checkpointing.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/cli.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/configuration.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/contributing.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/examples.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/inference.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/installation.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/limitations.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/quickstart.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/roadmap.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/tl_format.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/training.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/troubleshooting.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/docs/tutorial.md +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/examples/tabular_classification_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/examples/tabular_regression_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/examples/text_classification_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/examples/text_generation_example.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/setup.cfg +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/_version.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/api.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/auto/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/auto/config.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/auto/detector.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/backends/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/checkpoint/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/checkpoint/manager.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/cli/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/cli/main.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/config.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/english_grammar.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/fingerprint.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/inspector.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/loader.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/data/tabular.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/devices/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/devices/memory.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/engine.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/errors.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/mlp.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/registry.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/models/transformer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/runtime.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/serialization/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/serialization/tl_format.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/tokenization/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/tokenization/bpe_tokenizer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/tokenization/char_tokenizer.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/__init__.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/data_prep.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless/training/early_stopping.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/SOURCES.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/dependency_links.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/entry_points.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/requires.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tensorless.egg-info/top_level.txt +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_auto_detection.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_checkpoint_resume.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_cli.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_data_loading.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_end_to_end.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_fingerprint.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_jax_backend.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_mlx_backend.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_serialization.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_train_tabular.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_train_text_classification.py +0 -0
- {tensorless-0.9.2 → tensorless-0.9.4}/tests/test_train_text_generation.py +0 -0
|
@@ -0,0 +1,518 @@
|
|
|
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
|
+
from ..models.mlp import TabularMLP
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _jax():
|
|
19
|
+
try:
|
|
20
|
+
import jax
|
|
21
|
+
import jax.numpy as jnp
|
|
22
|
+
distributed_vars = ("JAX_COORDINATOR_ADDRESS", "JAX_PROCESS_COUNT", "JAX_PROCESS_ID")
|
|
23
|
+
if all(name in os.environ for name in distributed_vars) and not jax.distributed.is_initialized():
|
|
24
|
+
jax.distributed.initialize(
|
|
25
|
+
coordinator_address=os.environ["JAX_COORDINATOR_ADDRESS"],
|
|
26
|
+
num_processes=int(os.environ["JAX_PROCESS_COUNT"]),
|
|
27
|
+
process_id=int(os.environ["JAX_PROCESS_ID"]),
|
|
28
|
+
)
|
|
29
|
+
except ImportError as exc:
|
|
30
|
+
raise ImportError("JAX is required for the cuda/tpu backend; install tensorless[jax].") from exc
|
|
31
|
+
return jax, jnp
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def is_available() -> bool:
|
|
35
|
+
try:
|
|
36
|
+
_jax()
|
|
37
|
+
return True
|
|
38
|
+
except ImportError:
|
|
39
|
+
return False
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def supports_fused_training(model) -> bool:
|
|
43
|
+
"""Whether `model` can use the device-resident fused AdamW/Adam train
|
|
44
|
+
step (see `_JaxFusedAdamMixin`) instead of round-tripping full
|
|
45
|
+
parameter/gradient dicts through host NumPy every step."""
|
|
46
|
+
return isinstance(model, _JaxFusedAdamMixin)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _adam_update_tree(jnp, params, m, v, grads, step_count, lr, weight_decay, decoupled):
|
|
50
|
+
"""Fused AdamW/Adam parameter update over a pytree dict of device arrays.
|
|
51
|
+
|
|
52
|
+
Runs entirely inside `jax.jit` as part of the fused train step below, so
|
|
53
|
+
params/moments never have to leave the accelerator to be updated.
|
|
54
|
+
"""
|
|
55
|
+
beta1_correction = 1.0 - 0.9 ** step_count
|
|
56
|
+
beta2_correction = 1.0 - 0.999 ** step_count
|
|
57
|
+
new_params, new_m, new_v = {}, {}, {}
|
|
58
|
+
for name, p in params.items():
|
|
59
|
+
g = grads[name]
|
|
60
|
+
if decoupled:
|
|
61
|
+
p = p * (1.0 - lr * weight_decay)
|
|
62
|
+
else:
|
|
63
|
+
g = g + weight_decay * p
|
|
64
|
+
m_i = 0.9 * m[name] + 0.1 * g
|
|
65
|
+
v_i = 0.999 * v[name] + 0.001 * g * g
|
|
66
|
+
m_hat = m_i / beta1_correction
|
|
67
|
+
v_hat = v_i / beta2_correction
|
|
68
|
+
new_params[name] = p - lr * m_hat / (jnp.sqrt(v_hat) + 1e-8)
|
|
69
|
+
new_m[name] = m_i
|
|
70
|
+
new_v[name] = v_i
|
|
71
|
+
return new_params, new_m, new_v
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class _JaxFusedAdamMixin:
|
|
75
|
+
"""Keeps parameters + Adam moment estimates resident on the GPU/TPU
|
|
76
|
+
across many training steps, instead of re-uploading the full parameter
|
|
77
|
+
dict and downloading the full gradient dict through host NumPy on
|
|
78
|
+
*every single step*.
|
|
79
|
+
|
|
80
|
+
Compiling the forward pass with `jax.jit` (see `_get_forward_jit`) fixes
|
|
81
|
+
re-tracing overhead, but each `loss_and_backward()` call still does two
|
|
82
|
+
full host<->device transfers of the whole model (params up, grads down),
|
|
83
|
+
and the NumPy `Adam`/`AdamW` optimizer in `engine.py` then runs the
|
|
84
|
+
update on the CPU. For small/medium models that transfer, not compute,
|
|
85
|
+
dominates the step time and is why GPU utilization can stay low even
|
|
86
|
+
after JIT is fixed. This mixin fuses forward+backward+optimizer-update
|
|
87
|
+
into one jitted function so state only crosses the host/device boundary
|
|
88
|
+
at explicit sync points (checkpointing, validation, end of training) --
|
|
89
|
+
see `training/trainer.py`'s use of `init_fused_state`/`sync_fused_to_host`.
|
|
90
|
+
"""
|
|
91
|
+
|
|
92
|
+
_fused_params = None
|
|
93
|
+
_fused_m = None
|
|
94
|
+
_fused_v = None
|
|
95
|
+
_fused_step = 0
|
|
96
|
+
_fused_train_step_key = None
|
|
97
|
+
|
|
98
|
+
def fused_training_supported(self) -> bool:
|
|
99
|
+
return True
|
|
100
|
+
|
|
101
|
+
def init_fused_state(self, weight_decay: float, decoupled: bool, step: int = 0):
|
|
102
|
+
_, jnp = _jax()
|
|
103
|
+
params = self._jax_params()
|
|
104
|
+
self._fused_params = params
|
|
105
|
+
self._fused_m = {name: jnp.zeros_like(value) for name, value in params.items()}
|
|
106
|
+
self._fused_v = {name: jnp.zeros_like(value) for name, value in params.items()}
|
|
107
|
+
self._fused_step = int(step)
|
|
108
|
+
self._fused_train_step_key = None
|
|
109
|
+
|
|
110
|
+
def load_fused_optimizer_moments(self, m_by_name, v_by_name):
|
|
111
|
+
"""Restore Adam moment estimates (as NumPy, from a resumed checkpoint)."""
|
|
112
|
+
_, jnp = _jax()
|
|
113
|
+
self._fused_m = {name: jnp.asarray(value, dtype=jnp.float32) for name, value in m_by_name.items()}
|
|
114
|
+
self._fused_v = {name: jnp.asarray(value, dtype=jnp.float32) for name, value in v_by_name.items()}
|
|
115
|
+
|
|
116
|
+
def sync_fused_to_host(self):
|
|
117
|
+
"""Write device-resident params/moments back to host NumPy.
|
|
118
|
+
|
|
119
|
+
Called only at checkpoint/validation/end-of-training boundaries, so
|
|
120
|
+
this cost is paid a handful of times per epoch rather than every step.
|
|
121
|
+
"""
|
|
122
|
+
if self._fused_params is None:
|
|
123
|
+
return {}, {}
|
|
124
|
+
for name, parameter in self.named_parameters():
|
|
125
|
+
parameter.data[...] = np.asarray(self._fused_params[name], dtype=np.float32)
|
|
126
|
+
m_by_name = {name: np.asarray(value, dtype=np.float32) for name, value in self._fused_m.items()}
|
|
127
|
+
v_by_name = {name: np.asarray(value, dtype=np.float32) for name, value in self._fused_v.items()}
|
|
128
|
+
return m_by_name, v_by_name
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
class JaxTinyTransformer(_JaxFusedAdamMixin, TinyTransformer):
|
|
133
|
+
"""The native transformer parameter layout with JAX math and gradients."""
|
|
134
|
+
|
|
135
|
+
def __init__(self, *args, **kwargs):
|
|
136
|
+
self.precision = kwargs.pop("precision", "fp32")
|
|
137
|
+
super().__init__(*args, **kwargs)
|
|
138
|
+
d_model = kwargs.get("d_model", args[1] if len(args) > 1 else None)
|
|
139
|
+
self.heads = kwargs.get("heads", args[3] if len(args) > 3 else None)
|
|
140
|
+
if d_model is None or self.heads is None:
|
|
141
|
+
raise ValueError("JAX transformer requires d_model and heads")
|
|
142
|
+
self.d_model = d_model
|
|
143
|
+
self.head_dim = d_model // self.heads
|
|
144
|
+
|
|
145
|
+
def _jax_params(self):
|
|
146
|
+
_, jnp = _jax()
|
|
147
|
+
dtype = {"fp16": jnp.float16, "bf16": jnp.bfloat16}.get(self.precision, jnp.float32)
|
|
148
|
+
return {name: jnp.asarray(value.data, dtype=dtype) for name, value in self.named_parameters()}
|
|
149
|
+
|
|
150
|
+
def _forward_jax(self, params, input_ids, attention_mask=None):
|
|
151
|
+
_, jnp = _jax()
|
|
152
|
+
|
|
153
|
+
def layer_norm(x, gain, bias, eps=1e-5):
|
|
154
|
+
mean = jnp.mean(x, axis=-1, keepdims=True)
|
|
155
|
+
var = jnp.var(x, axis=-1, keepdims=True)
|
|
156
|
+
return (x - mean) / jnp.sqrt(var + eps) * gain + bias
|
|
157
|
+
|
|
158
|
+
input_ids = jnp.asarray(input_ids, dtype=jnp.int32)
|
|
159
|
+
batch, length = input_ids.shape
|
|
160
|
+
hidden = params["tok_emb"][input_ids] + params["pos_emb"][jnp.arange(length)]
|
|
161
|
+
causal = jnp.tril(jnp.ones((length, length), dtype=bool))
|
|
162
|
+
for index in range(len(self.blocks)):
|
|
163
|
+
prefix = f"blocks.{index}"
|
|
164
|
+
normed1 = layer_norm(hidden, params[f"{prefix}.ln1_gain"], params[f"{prefix}.ln1_bias"])
|
|
165
|
+
qkv = normed1 @ params[f"{prefix}.qkv_weight"] + params[f"{prefix}.qkv_bias"]
|
|
166
|
+
q, k, v = jnp.split(qkv, 3, axis=-1)
|
|
167
|
+
q = q.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
168
|
+
k = k.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
169
|
+
v = v.reshape(batch, length, self.heads, self.head_dim).transpose(0, 2, 1, 3)
|
|
170
|
+
scores = q @ k.transpose(0, 1, 3, 2) / np.sqrt(self.head_dim)
|
|
171
|
+
scores = jnp.where(causal[None, None, :, :], scores, -1e9)
|
|
172
|
+
if attention_mask is not None:
|
|
173
|
+
mask = jnp.asarray(attention_mask) > 0
|
|
174
|
+
scores = jnp.where(mask[:, None, None, :], scores, -1e9)
|
|
175
|
+
probabilities = jnp.exp(scores - scores.max(axis=-1, keepdims=True))
|
|
176
|
+
probabilities /= probabilities.sum(axis=-1, keepdims=True)
|
|
177
|
+
context = probabilities @ v
|
|
178
|
+
context = context.transpose(0, 2, 1, 3).reshape(batch, length, self.d_model)
|
|
179
|
+
attention = context @ params[f"{prefix}.out_weight"] + params[f"{prefix}.out_bias"]
|
|
180
|
+
residual = hidden + attention
|
|
181
|
+
normed2 = layer_norm(residual, params[f"{prefix}.ln2_gain"], params[f"{prefix}.ln2_bias"])
|
|
182
|
+
ff_pre = normed2 @ params[f"{prefix}.ff1_weight"] + params[f"{prefix}.ff1_bias"]
|
|
183
|
+
ff_hidden = jnp.maximum(ff_pre, 0)
|
|
184
|
+
ff = ff_hidden @ params[f"{prefix}.ff2_weight"] + params[f"{prefix}.ff2_bias"]
|
|
185
|
+
hidden = residual + ff
|
|
186
|
+
hidden = layer_norm(hidden, params["ln_f_gain"], params["ln_f_bias"])
|
|
187
|
+
if self.task == "text-generation":
|
|
188
|
+
return hidden @ params["tok_emb"].T + params["head_bias"]
|
|
189
|
+
mask = jnp.ones((batch, length), dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
|
|
190
|
+
pooled = (hidden * mask[:, :, None]).sum(1) / jnp.maximum(mask.sum(1, keepdims=True), 1)
|
|
191
|
+
return pooled @ params["head_weight"] + params["head_bias"]
|
|
192
|
+
|
|
193
|
+
def forward(self, input_ids, attention_mask=None, cache=False):
|
|
194
|
+
params = self._jax_params()
|
|
195
|
+
output = self._get_forward_jit()(params, input_ids, attention_mask)
|
|
196
|
+
output = np.asarray(output)
|
|
197
|
+
return (output, None) if cache else output
|
|
198
|
+
|
|
199
|
+
def _get_forward_jit(self):
|
|
200
|
+
# Compile once per instance and reuse. Without this, every call
|
|
201
|
+
# re-traces the whole forward pass in Python and dispatches ops to
|
|
202
|
+
# the accelerator one at a time -- which shows up as pegged CPU
|
|
203
|
+
# (tracing/dispatch overhead) with near-zero GPU utilization, even
|
|
204
|
+
# though the correct device was selected.
|
|
205
|
+
if getattr(self, "_forward_jit_fn", None) is None:
|
|
206
|
+
jax, _ = _jax()
|
|
207
|
+
self._forward_jit_fn = jax.jit(self._forward_jax)
|
|
208
|
+
return self._forward_jit_fn
|
|
209
|
+
|
|
210
|
+
def _get_loss_and_grad_jit(self):
|
|
211
|
+
if getattr(self, "_loss_and_grad_jit_fn", None) is None:
|
|
212
|
+
jax, jnp = _jax()
|
|
213
|
+
|
|
214
|
+
def loss_fn(current, current_ids, current_target, current_mask):
|
|
215
|
+
forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
|
|
216
|
+
logits = forward(current, current_ids, current_mask)
|
|
217
|
+
flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
|
|
218
|
+
flat_target = current_target.reshape(-1)
|
|
219
|
+
log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
|
|
220
|
+
losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
|
|
221
|
+
if self.task == "text-generation":
|
|
222
|
+
valid = flat_target != self.pad_id
|
|
223
|
+
return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
|
|
224
|
+
return jnp.mean(losses)
|
|
225
|
+
|
|
226
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
227
|
+
scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
|
|
228
|
+
self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(scaled_loss_fn))
|
|
229
|
+
return self._loss_and_grad_jit_fn
|
|
230
|
+
|
|
231
|
+
def _get_pmap_step(self):
|
|
232
|
+
if getattr(self, "_pmap_step_fn", None) is None:
|
|
233
|
+
jax, jnp = _jax()
|
|
234
|
+
|
|
235
|
+
def loss_fn(current, current_ids, current_target, current_mask):
|
|
236
|
+
forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
|
|
237
|
+
logits = forward(current, current_ids, current_mask)
|
|
238
|
+
flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
|
|
239
|
+
flat_target = current_target.reshape(-1)
|
|
240
|
+
log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
|
|
241
|
+
losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
|
|
242
|
+
if self.task == "text-generation":
|
|
243
|
+
valid = flat_target != self.pad_id
|
|
244
|
+
return jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
|
|
245
|
+
return jnp.mean(losses)
|
|
246
|
+
|
|
247
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
248
|
+
scaled_loss_fn = lambda *arguments: loss_fn(*arguments) * loss_scale
|
|
249
|
+
per_device_loss = jax.value_and_grad(scaled_loss_fn)
|
|
250
|
+
|
|
251
|
+
def mapped_step(current, ids, labels, current_mask):
|
|
252
|
+
value, gradients = per_device_loss(current, ids, labels, current_mask)
|
|
253
|
+
return jax.lax.pmean(value, "data"), jax.tree_util.tree_map(
|
|
254
|
+
lambda gradient: jax.lax.pmean(gradient, "data"), gradients
|
|
255
|
+
)
|
|
256
|
+
|
|
257
|
+
self._pmap_step_fn = jax.pmap(mapped_step, axis_name="data", in_axes=(None, 0, 0, 0))
|
|
258
|
+
return self._pmap_step_fn
|
|
259
|
+
|
|
260
|
+
def loss_and_backward(self, input_ids, target, attention_mask=None):
|
|
261
|
+
jax, jnp = _jax()
|
|
262
|
+
params = self._jax_params()
|
|
263
|
+
target = jnp.asarray(target, dtype=jnp.int32)
|
|
264
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
265
|
+
mask = jnp.ones_like(input_ids, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask)
|
|
266
|
+
devices = jax.local_device_count()
|
|
267
|
+
if devices > 1 and input_ids.shape[0] % devices == 0:
|
|
268
|
+
shard = lambda value: value.reshape((devices, value.shape[0] // devices) + value.shape[1:])
|
|
269
|
+
loss, gradients = self._get_pmap_step()(
|
|
270
|
+
params, shard(jnp.asarray(input_ids)), shard(target), shard(mask)
|
|
271
|
+
)
|
|
272
|
+
loss = loss[0]
|
|
273
|
+
gradients = jax.tree_util.tree_map(lambda gradient: gradient[0] / loss_scale, gradients)
|
|
274
|
+
else:
|
|
275
|
+
loss, gradients = self._get_loss_and_grad_jit()(params, input_ids, target, mask)
|
|
276
|
+
gradients = jax.tree_util.tree_map(lambda gradient: gradient / loss_scale, gradients)
|
|
277
|
+
for name, parameter in self.named_parameters():
|
|
278
|
+
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
|
|
279
|
+
return float(loss) / loss_scale
|
|
280
|
+
|
|
281
|
+
def _get_fused_train_step(self, weight_decay: float, decoupled: bool):
|
|
282
|
+
cache_key = (weight_decay, decoupled)
|
|
283
|
+
if self._fused_train_step_key != cache_key:
|
|
284
|
+
jax, jnp = _jax()
|
|
285
|
+
loss_scale = 128.0 if self.precision == "fp16" else 1.0
|
|
286
|
+
|
|
287
|
+
def loss_fn(params, ids, target, mask):
|
|
288
|
+
forward = jax.checkpoint(self._forward_jax) if self.gradient_checkpointing else self._forward_jax
|
|
289
|
+
logits = forward(params, ids, mask)
|
|
290
|
+
flat_logits = logits.reshape(-1, logits.shape[-1]).astype(jnp.float32)
|
|
291
|
+
flat_target = target.reshape(-1)
|
|
292
|
+
log_probs = jax.nn.log_softmax(flat_logits, axis=-1)
|
|
293
|
+
losses = -jnp.take_along_axis(log_probs, flat_target[:, None], axis=1).squeeze(1)
|
|
294
|
+
if self.task == "text-generation":
|
|
295
|
+
valid = flat_target != self.pad_id
|
|
296
|
+
loss = jnp.sum(jnp.where(valid, losses, 0.0)) / jnp.maximum(valid.sum(), 1)
|
|
297
|
+
else:
|
|
298
|
+
loss = jnp.mean(losses)
|
|
299
|
+
return loss * loss_scale
|
|
300
|
+
|
|
301
|
+
grad_fn = jax.value_and_grad(loss_fn)
|
|
302
|
+
|
|
303
|
+
def step(params, m, v, step_count, ids, target, mask, lr, grad_clip):
|
|
304
|
+
loss, grads = grad_fn(params, ids, target, mask)
|
|
305
|
+
grads = {name: g / loss_scale for name, g in grads.items()}
|
|
306
|
+
grad_norm = jnp.sqrt(sum(jnp.sum(g.astype(jnp.float32) ** 2) for g in grads.values()))
|
|
307
|
+
clip_scale = jnp.where(grad_clip > 0, jnp.minimum(1.0, grad_clip / (grad_norm + 1e-12)), 1.0)
|
|
308
|
+
grads = {name: g * clip_scale for name, g in grads.items()}
|
|
309
|
+
new_params, new_m, new_v = _adam_update_tree(
|
|
310
|
+
jnp, params, m, v, grads, step_count, lr, weight_decay, decoupled
|
|
311
|
+
)
|
|
312
|
+
return new_params, new_m, new_v, loss / loss_scale, grad_norm
|
|
313
|
+
|
|
314
|
+
self._fused_train_step_fn = jax.jit(step)
|
|
315
|
+
self._fused_train_step_key = cache_key
|
|
316
|
+
return self._fused_train_step_fn
|
|
317
|
+
|
|
318
|
+
def fused_train_step(self, input_ids, target, attention_mask, lr, weight_decay, grad_clip, decoupled):
|
|
319
|
+
"""One fused forward+backward+AdamW step, entirely on-device.
|
|
320
|
+
|
|
321
|
+
Call `init_fused_state()` once before the first step and
|
|
322
|
+
`sync_fused_to_host()` before reading `state_dict()`/`forward()` on
|
|
323
|
+
the host (e.g. for validation or checkpointing).
|
|
324
|
+
"""
|
|
325
|
+
if self._fused_params is None:
|
|
326
|
+
raise RuntimeError("init_fused_state() must be called before fused_train_step().")
|
|
327
|
+
jax, jnp = _jax()
|
|
328
|
+
step_fn = self._get_fused_train_step(weight_decay, decoupled)
|
|
329
|
+
self._fused_step += 1
|
|
330
|
+
ids_arr = jnp.asarray(input_ids, dtype=jnp.int32)
|
|
331
|
+
target_arr = jnp.asarray(target, dtype=jnp.int32)
|
|
332
|
+
mask_arr = jnp.ones_like(ids_arr, dtype=jnp.float32) if attention_mask is None else jnp.asarray(attention_mask, dtype=jnp.float32)
|
|
333
|
+
new_params, new_m, new_v, loss, grad_norm = step_fn(
|
|
334
|
+
self._fused_params, self._fused_m, self._fused_v,
|
|
335
|
+
jnp.asarray(self._fused_step, dtype=jnp.float32),
|
|
336
|
+
ids_arr, target_arr, mask_arr,
|
|
337
|
+
jnp.asarray(lr, dtype=jnp.float32), jnp.asarray(grad_clip, dtype=jnp.float32),
|
|
338
|
+
)
|
|
339
|
+
self._fused_params, self._fused_m, self._fused_v = new_params, new_m, new_v
|
|
340
|
+
return float(loss), float(grad_norm)
|
|
341
|
+
|
|
342
|
+
def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, top_p=.9,
|
|
343
|
+
repetition_penalty=1.1, eos_id: Optional[int] = None):
|
|
344
|
+
# Mirrors `TinyTransformer.generate` (the NumPy path) exactly, so
|
|
345
|
+
# generation quality doesn't quietly regress just because a model
|
|
346
|
+
# happens to run on cuda/tpu. Sampling itself stays in NumPy since
|
|
347
|
+
# it's inherently sequential and tiny relative to the forward pass;
|
|
348
|
+
# only the forward pass benefits from JIT.
|
|
349
|
+
ids = np.asarray(input_ids, dtype=np.int64).copy()
|
|
350
|
+
for _ in range(max_new_tokens):
|
|
351
|
+
logits = self.forward(ids[:, -self.max_seq_len:])[:, -1, :] / max(temperature, 1e-5)
|
|
352
|
+
if repetition_penalty and repetition_penalty > 1.0:
|
|
353
|
+
for row, sequence in enumerate(ids):
|
|
354
|
+
seen = np.unique(sequence)
|
|
355
|
+
logits[row, seen] = np.where(
|
|
356
|
+
logits[row, seen] < 0,
|
|
357
|
+
logits[row, seen] * repetition_penalty,
|
|
358
|
+
logits[row, seen] / repetition_penalty,
|
|
359
|
+
)
|
|
360
|
+
if top_k is not None:
|
|
361
|
+
k = min(top_k, logits.shape[-1])
|
|
362
|
+
excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
|
|
363
|
+
logits[np.arange(len(ids))[:, None], excluded] = -np.inf
|
|
364
|
+
probs = np.exp(logits - logits.max(1, keepdims=True))
|
|
365
|
+
probs /= probs.sum(1, keepdims=True)
|
|
366
|
+
if top_p is not None and 0 < top_p < 1:
|
|
367
|
+
order = np.argsort(-probs, axis=1)
|
|
368
|
+
sorted_probs = np.take_along_axis(probs, order, axis=1)
|
|
369
|
+
cutoff = np.cumsum(sorted_probs, axis=1) > top_p
|
|
370
|
+
cutoff[:, 0] = False
|
|
371
|
+
probs[np.arange(len(ids))[:, None], order] = np.where(cutoff, 0.0, sorted_probs)
|
|
372
|
+
probs /= probs.sum(1, keepdims=True)
|
|
373
|
+
next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
|
|
374
|
+
ids = np.concatenate([ids, next_ids], axis=1)
|
|
375
|
+
if eos_id is not None and np.all(next_ids == eos_id):
|
|
376
|
+
break
|
|
377
|
+
return ids
|
|
378
|
+
|
|
379
|
+
|
|
380
|
+
class JaxTabularMLP(_JaxFusedAdamMixin, TabularMLP):
|
|
381
|
+
"""The native tabular-MLP parameter layout with JAX math and gradients.
|
|
382
|
+
|
|
383
|
+
Used automatically instead of the plain NumPy `TabularMLP` whenever the
|
|
384
|
+
resolved device is `cuda`/`tpu`, so tabular training actually runs on
|
|
385
|
+
the detected accelerator instead of silently staying on CPU.
|
|
386
|
+
"""
|
|
387
|
+
|
|
388
|
+
def __init__(self, *args, **kwargs):
|
|
389
|
+
self.precision = kwargs.pop("precision", "fp32")
|
|
390
|
+
super().__init__(*args, **kwargs)
|
|
391
|
+
|
|
392
|
+
def _jax_params(self):
|
|
393
|
+
_, jnp = _jax()
|
|
394
|
+
dtype = {"fp16": jnp.float16, "bf16": jnp.bfloat16}.get(self.precision, jnp.float32)
|
|
395
|
+
return {name: jnp.asarray(value.data, dtype=dtype) for name, value in self.named_parameters()}
|
|
396
|
+
|
|
397
|
+
def _forward_jax(self, params, numeric, categorical, dropout_masks=None):
|
|
398
|
+
jax, jnp = _jax()
|
|
399
|
+
parts = [jnp.asarray(numeric, dtype=jnp.float32)] if self.n_numeric else []
|
|
400
|
+
categorical = jnp.asarray(categorical, dtype=jnp.int32)
|
|
401
|
+
for i in range(len(self.embeddings)):
|
|
402
|
+
parts.append(params[f"embeddings.{i}"][categorical[:, i]])
|
|
403
|
+
x = jnp.concatenate(parts, axis=1) if parts else jnp.asarray(numeric, dtype=jnp.float32)
|
|
404
|
+
for i in range(len(self.weights)):
|
|
405
|
+
preactivation = x @ params[f"weights.{i}"] + params[f"biases.{i}"]
|
|
406
|
+
x = jax.nn.gelu(preactivation, approximate=True)
|
|
407
|
+
if dropout_masks is not None and dropout_masks[i] is not None:
|
|
408
|
+
x = x * dropout_masks[i] / (1.0 - self.dropout)
|
|
409
|
+
out = x @ params["head_weight"] + params["head_bias"]
|
|
410
|
+
if self.task == "regression":
|
|
411
|
+
out = out[:, 0]
|
|
412
|
+
return out
|
|
413
|
+
|
|
414
|
+
def _make_dropout_masks(self, batch_size):
|
|
415
|
+
if not self.training or self.dropout <= 0:
|
|
416
|
+
return None
|
|
417
|
+
_, jnp = _jax()
|
|
418
|
+
return [
|
|
419
|
+
jnp.asarray((np.random.random((batch_size, weight.data.shape[1])) >= self.dropout).astype(np.float32))
|
|
420
|
+
for weight in self.weights
|
|
421
|
+
]
|
|
422
|
+
|
|
423
|
+
def forward(self, numeric, categorical, cache=False):
|
|
424
|
+
params = self._jax_params()
|
|
425
|
+
batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
|
|
426
|
+
dropout_masks = self._make_dropout_masks(batch_size)
|
|
427
|
+
out = np.asarray(self._get_forward_jit()(params, numeric, categorical, dropout_masks))
|
|
428
|
+
return (out, None) if cache else out
|
|
429
|
+
|
|
430
|
+
def _get_forward_jit(self):
|
|
431
|
+
if getattr(self, "_forward_jit_fn", None) is None:
|
|
432
|
+
jax, _ = _jax()
|
|
433
|
+
self._forward_jit_fn = jax.jit(self._forward_jax)
|
|
434
|
+
return self._forward_jit_fn
|
|
435
|
+
|
|
436
|
+
def _get_loss_and_grad_jit(self):
|
|
437
|
+
if getattr(self, "_loss_and_grad_jit_fn", None) is None:
|
|
438
|
+
jax, jnp = _jax()
|
|
439
|
+
is_regression = self.task == "regression"
|
|
440
|
+
|
|
441
|
+
def loss_fn(current, numeric, categorical, dropout_masks, target_arr):
|
|
442
|
+
logits = self._forward_jax(current, numeric, categorical, dropout_masks)
|
|
443
|
+
if is_regression:
|
|
444
|
+
error = logits - target_arr
|
|
445
|
+
return jnp.mean(error * error)
|
|
446
|
+
log_probs = jax.nn.log_softmax(logits, axis=-1)
|
|
447
|
+
return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
|
|
448
|
+
|
|
449
|
+
self._loss_and_grad_jit_fn = jax.jit(jax.value_and_grad(loss_fn))
|
|
450
|
+
return self._loss_and_grad_jit_fn
|
|
451
|
+
|
|
452
|
+
def loss_and_backward(self, numeric, categorical, target):
|
|
453
|
+
jax, jnp = _jax()
|
|
454
|
+
params = self._jax_params()
|
|
455
|
+
batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
|
|
456
|
+
dropout_masks = self._make_dropout_masks(batch_size)
|
|
457
|
+
is_regression = self.task == "regression"
|
|
458
|
+
target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
|
|
459
|
+
|
|
460
|
+
loss, gradients = self._get_loss_and_grad_jit()(params, numeric, categorical, dropout_masks, target_arr)
|
|
461
|
+
for name, parameter in self.named_parameters():
|
|
462
|
+
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32)
|
|
463
|
+
return float(loss)
|
|
464
|
+
|
|
465
|
+
def _get_fused_train_step(self, weight_decay: float, decoupled: bool):
|
|
466
|
+
cache_key = (weight_decay, decoupled)
|
|
467
|
+
if self._fused_train_step_key != cache_key:
|
|
468
|
+
jax, jnp = _jax()
|
|
469
|
+
is_regression = self.task == "regression"
|
|
470
|
+
|
|
471
|
+
def loss_fn(params, numeric, categorical, dropout_masks, target_arr):
|
|
472
|
+
logits = self._forward_jax(params, numeric, categorical, dropout_masks)
|
|
473
|
+
if is_regression:
|
|
474
|
+
error = logits - target_arr
|
|
475
|
+
return jnp.mean(error * error)
|
|
476
|
+
log_probs = jax.nn.log_softmax(logits, axis=-1)
|
|
477
|
+
return -jnp.mean(jnp.take_along_axis(log_probs, target_arr[:, None], axis=1))
|
|
478
|
+
|
|
479
|
+
grad_fn = jax.value_and_grad(loss_fn)
|
|
480
|
+
|
|
481
|
+
def step(params, m, v, step_count, numeric, categorical, dropout_masks, target_arr, lr, grad_clip):
|
|
482
|
+
loss, grads = grad_fn(params, numeric, categorical, dropout_masks, target_arr)
|
|
483
|
+
grad_norm = jnp.sqrt(sum(jnp.sum(g.astype(jnp.float32) ** 2) for g in grads.values()))
|
|
484
|
+
clip_scale = jnp.where(grad_clip > 0, jnp.minimum(1.0, grad_clip / (grad_norm + 1e-12)), 1.0)
|
|
485
|
+
grads = {name: g * clip_scale for name, g in grads.items()}
|
|
486
|
+
new_params, new_m, new_v = _adam_update_tree(
|
|
487
|
+
jnp, params, m, v, grads, step_count, lr, weight_decay, decoupled
|
|
488
|
+
)
|
|
489
|
+
return new_params, new_m, new_v, loss, grad_norm
|
|
490
|
+
|
|
491
|
+
self._fused_train_step_fn = jax.jit(step)
|
|
492
|
+
self._fused_train_step_key = cache_key
|
|
493
|
+
return self._fused_train_step_fn
|
|
494
|
+
|
|
495
|
+
def fused_train_step(self, numeric, categorical, target, lr, weight_decay, grad_clip, decoupled):
|
|
496
|
+
"""One fused forward+backward+AdamW step, entirely on-device.
|
|
497
|
+
|
|
498
|
+
Call `init_fused_state()` once before the first step and
|
|
499
|
+
`sync_fused_to_host()` before reading `state_dict()`/`forward()` on
|
|
500
|
+
the host (e.g. for validation or checkpointing).
|
|
501
|
+
"""
|
|
502
|
+
if self._fused_params is None:
|
|
503
|
+
raise RuntimeError("init_fused_state() must be called before fused_train_step().")
|
|
504
|
+
jax, jnp = _jax()
|
|
505
|
+
step_fn = self._get_fused_train_step(weight_decay, decoupled)
|
|
506
|
+
self._fused_step += 1
|
|
507
|
+
batch_size = numeric.shape[0] if self.n_numeric else categorical.shape[0]
|
|
508
|
+
dropout_masks = self._make_dropout_masks(batch_size)
|
|
509
|
+
is_regression = self.task == "regression"
|
|
510
|
+
target_arr = jnp.asarray(target, dtype=jnp.float32 if is_regression else jnp.int32)
|
|
511
|
+
new_params, new_m, new_v, loss, grad_norm = step_fn(
|
|
512
|
+
self._fused_params, self._fused_m, self._fused_v,
|
|
513
|
+
jnp.asarray(self._fused_step, dtype=jnp.float32),
|
|
514
|
+
numeric, categorical, dropout_masks, target_arr,
|
|
515
|
+
jnp.asarray(lr, dtype=jnp.float32), jnp.asarray(grad_clip, dtype=jnp.float32),
|
|
516
|
+
)
|
|
517
|
+
self._fused_params, self._fused_m, self._fused_v = new_params, new_m, new_v
|
|
518
|
+
return float(loss), float(grad_norm)
|
|
@@ -117,16 +117,35 @@ class MlxTinyTransformer(TinyTransformer):
|
|
|
117
117
|
parameter.grad[...] = np.asarray(gradients[name], dtype=np.float32) / loss_scale
|
|
118
118
|
return float(np.asarray(loss)) / loss_scale
|
|
119
119
|
|
|
120
|
-
def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40,
|
|
120
|
+
def generate(self, input_ids, max_new_tokens, temperature=.8, top_k=40, top_p=.9,
|
|
121
|
+
repetition_penalty=1.1, eos_id: Optional[int] = None):
|
|
122
|
+
# Mirrors `TinyTransformer.generate` (the NumPy path) exactly, so
|
|
123
|
+
# generation quality doesn't quietly regress just because a model
|
|
124
|
+
# happens to run on mps.
|
|
121
125
|
ids = np.asarray(input_ids, dtype=np.int64).copy()
|
|
122
126
|
for _ in range(max_new_tokens):
|
|
123
127
|
logits = self.forward(ids[:, -self.max_seq_len:])[:, -1, :] / max(temperature, 1e-5)
|
|
128
|
+
if repetition_penalty and repetition_penalty > 1.0:
|
|
129
|
+
for row, sequence in enumerate(ids):
|
|
130
|
+
seen = np.unique(sequence)
|
|
131
|
+
logits[row, seen] = np.where(
|
|
132
|
+
logits[row, seen] < 0,
|
|
133
|
+
logits[row, seen] * repetition_penalty,
|
|
134
|
+
logits[row, seen] / repetition_penalty,
|
|
135
|
+
)
|
|
124
136
|
if top_k is not None:
|
|
125
137
|
k = min(top_k, logits.shape[-1])
|
|
126
138
|
excluded = np.argpartition(logits, -k, axis=1)[:, :-k]
|
|
127
139
|
logits[np.arange(len(ids))[:, None], excluded] = -np.inf
|
|
128
140
|
probs = np.exp(logits - logits.max(1, keepdims=True))
|
|
129
141
|
probs /= probs.sum(1, keepdims=True)
|
|
142
|
+
if top_p is not None and 0 < top_p < 1:
|
|
143
|
+
order = np.argsort(-probs, axis=1)
|
|
144
|
+
sorted_probs = np.take_along_axis(probs, order, axis=1)
|
|
145
|
+
cutoff = np.cumsum(sorted_probs, axis=1) > top_p
|
|
146
|
+
cutoff[:, 0] = False
|
|
147
|
+
probs[np.arange(len(ids))[:, None], order] = np.where(cutoff, 0.0, sorted_probs)
|
|
148
|
+
probs /= probs.sum(1, keepdims=True)
|
|
130
149
|
next_ids = np.array([np.random.choice(logits.shape[1], p=p) for p in probs])[:, None]
|
|
131
150
|
ids = np.concatenate([ids, next_ids], axis=1)
|
|
132
151
|
if eos_id is not None and np.all(next_ids == eos_id):
|
|
@@ -37,9 +37,33 @@ def _mps_available() -> bool:
|
|
|
37
37
|
return False
|
|
38
38
|
|
|
39
39
|
|
|
40
|
+
_BF16_CAPABLE_GPU_HINTS = (
|
|
41
|
+
"a100", "a10", "a30", "a40", "l4", "l40", "l20",
|
|
42
|
+
"h100", "h200", "h800", "b100", "b200", "gh200",
|
|
43
|
+
"rtx 30", "rtx 40", "rtx 50", "rtx a", "geforce rtx 30", "geforce rtx 40", "geforce rtx 50",
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
|
|
40
47
|
def _cuda_supports_bf16() -> bool:
|
|
41
|
-
"""
|
|
42
|
-
|
|
48
|
+
"""Best-effort detection of real (fast) bf16 tensor-core support.
|
|
49
|
+
|
|
50
|
+
Older GPUs (Turing/Volta/Pascal -- T4, V100, P100, K80, RTX 20xx) either
|
|
51
|
+
lack bf16 tensor cores or emulate bf16 slowly, so fp16 (with loss
|
|
52
|
+
scaling) is the safer/faster default there. Ampere-or-newer GPUs (A100,
|
|
53
|
+
L4, H100, RTX 30xx+, ...) run bf16 at full tensor-core speed and don't
|
|
54
|
+
need loss scaling, so we opt in automatically when the detected GPU name
|
|
55
|
+
matches a known Ampere-or-newer architecture. Anything unrecognized
|
|
56
|
+
conservatively falls back to fp16, matching the previous behavior.
|
|
57
|
+
"""
|
|
58
|
+
try:
|
|
59
|
+
import jax
|
|
60
|
+
gpus = [d for d in jax.devices() if d.platform == "gpu"]
|
|
61
|
+
if not gpus:
|
|
62
|
+
return False
|
|
63
|
+
name = getattr(gpus[0], "device_kind", "").lower()
|
|
64
|
+
return any(hint in name for hint in _BF16_CAPABLE_GPU_HINTS)
|
|
65
|
+
except (ImportError, RuntimeError):
|
|
66
|
+
return False
|
|
43
67
|
|
|
44
68
|
|
|
45
69
|
def auto_select_device(user_device: Optional[str], user_precision: Optional[str]) -> Tuple[str, str]:
|