haliax 1.4.dev326__tar.gz → 1.4.dev330__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.
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.github/workflows/run_quick_levanter_tests.yaml +2 -2
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.github/workflows/run_tests.yaml +1 -2
- {haliax-1.4.dev326 → haliax-1.4.dev330}/PKG-INFO +4 -2
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/fp8.md +7 -7
- {haliax-1.4.dev326 → haliax-1.4.dev330}/pyproject.toml +2 -1
- haliax-1.4.dev330/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/jax_utils.py +18 -1
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/attention.py +6 -1
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/partitioning.py +6 -6
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/quantization.py +57 -17
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/random.py +1 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_fp8.py +8 -8
- haliax-1.4.dev330/tests/test_int8.py +23 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_ops.py +3 -1
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_random.py +8 -1
- haliax-1.4.dev326/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.coveragerc +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.flake8 +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.gitignore +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/LICENSE +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/README.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/api.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/css/material.css +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/faq.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/hof.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/index.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/indexing.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/matmul.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/nn.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/partitioning.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/rearrange.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/requirements.txt +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/state-dict.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/docs/tutorial.md +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/mkdocs.yml +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/core.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/types.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/util.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/core_test.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_attention.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_axis.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_conv.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_debug.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_dot.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_hof.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_nn.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_pool.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_scan.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev326 → haliax-1.4.dev330}/tests/test_utils.py +0 -0
|
@@ -17,7 +17,7 @@ jobs:
|
|
|
17
17
|
run: |
|
|
18
18
|
python -m pip install --upgrade pip
|
|
19
19
|
pip install flake8 pytest
|
|
20
|
-
pip install --upgrade "jax[cpu]==0.4.
|
|
20
|
+
pip install --upgrade "jax[cpu]==0.4.35" "jaxlib[cpu]==0.4.35"
|
|
21
21
|
|
|
22
22
|
- name: Install Levanter from source
|
|
23
23
|
run: |
|
|
@@ -29,7 +29,7 @@ jobs:
|
|
|
29
29
|
run: |
|
|
30
30
|
# install second since levanter will install a built version of haliax
|
|
31
31
|
cd ../haliax
|
|
32
|
-
pip install .[dev]
|
|
32
|
+
pip install "jax[cpu]==0.4.35" "jaxlib[cpu]==0.4.35" .[dev]
|
|
33
33
|
- name: Test levanter with pytest
|
|
34
34
|
run: |
|
|
35
35
|
cd ../levanter
|
|
@@ -17,8 +17,7 @@ jobs:
|
|
|
17
17
|
run: |
|
|
18
18
|
python -m pip install --upgrade pip
|
|
19
19
|
pip install flake8 pytest
|
|
20
|
-
pip install
|
|
21
|
-
pip install .[dev]
|
|
20
|
+
pip install jax==0.4.35 jaxlib==0.4.35 .[dev]
|
|
22
21
|
- name: Test with pytest
|
|
23
22
|
run: |
|
|
24
23
|
XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:. pytest tests
|
|
@@ -1,11 +1,12 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev330
|
|
4
4
|
Summary: Named Tensors for Legible Deep Learning in JAX
|
|
5
5
|
Project-URL: Homepage, https://github.com/stanford-crfm/haliax
|
|
6
6
|
Project-URL: Bug Tracker, https://github.com/stanford-crfm/haliax/issues/
|
|
7
7
|
Project-URL: Documentation, https://haliax.readthedocs.io/en/latest/
|
|
8
8
|
Author-email: David Hall <dlwh@cs.stanford.edu>
|
|
9
|
+
License-File: LICENSE
|
|
9
10
|
Classifier: Development Status :: 4 - Beta
|
|
10
11
|
Classifier: Intended Audience :: Science/Research
|
|
11
12
|
Classifier: License :: OSI Approved :: Apache Software License
|
|
@@ -13,6 +14,7 @@ Classifier: Operating System :: MacOS :: MacOS X
|
|
|
13
14
|
Classifier: Operating System :: POSIX :: Linux
|
|
14
15
|
Classifier: Programming Language :: Python :: 3
|
|
15
16
|
Requires-Python: >=3.10
|
|
17
|
+
Requires-Dist: aqtp>=0.8.2
|
|
16
18
|
Requires-Dist: equinox>=0.10.6
|
|
17
19
|
Requires-Dist: jaxtyping>=0.2.20
|
|
18
20
|
Requires-Dist: jmp>=0.0.4
|
|
@@ -23,7 +23,7 @@ Using FP8 with Haliax is actually pretty straightforward. To enable FP8, do this
|
|
|
23
23
|
```python
|
|
24
24
|
import haliax.quantization as haxq
|
|
25
25
|
# setup
|
|
26
|
-
module = haxq.
|
|
26
|
+
module = haxq.quantize_linear_layers(module, haxq.QuantizationConfig(fp8=True))
|
|
27
27
|
|
|
28
28
|
# if using optax. This saves a tiny amount of memory so you can skip it if you want
|
|
29
29
|
_, trainable_module = haxq.partition_for_grad_overwrite(module)
|
|
@@ -81,13 +81,13 @@ class MyModule(eqx.Module):
|
|
|
81
81
|
module = MyModule.init(key=jax.random.PRNGKey(0))
|
|
82
82
|
|
|
83
83
|
# Enable FP8
|
|
84
|
-
module = hax.quantization.
|
|
84
|
+
module = hax.quantization.quantize_linear_layers(module, QuantizationConfig(fp8=True))
|
|
85
85
|
|
|
86
86
|
# Enable FP8 for a specific layer
|
|
87
|
-
from haliax.quantization import
|
|
87
|
+
from haliax.quantization import QuantizationConfig
|
|
88
88
|
|
|
89
|
-
config =
|
|
90
|
-
module = hax.quantization.
|
|
89
|
+
config = QuantizationConfig(targets=["up_proj"], fp8=True)
|
|
90
|
+
module = hax.quantization.quantize_linear_layers(module, config)
|
|
91
91
|
|
|
92
92
|
# Train step
|
|
93
93
|
grads = eqx.filter_grad(loss_fn)(module, data)
|
|
@@ -96,7 +96,7 @@ updates, opt_state = opt.update(grads, opt_state, params=module) # or however y
|
|
|
96
96
|
module = hax.quantization.apply_updates(module, updates, grads)
|
|
97
97
|
```
|
|
98
98
|
|
|
99
|
-
That's it! Just a few lines of code to
|
|
99
|
+
That's it! Just a few lines of code to enable FP8. The `quantize_linear_layers` function will transform your module to use FP8
|
|
100
100
|
for linear layers (or a subset if you want), and the combo of `partition_for_grad_overwrite` and `apply_updates` function will apply the updates to the module
|
|
101
101
|
in a way that is compatible with FP8.
|
|
102
102
|
|
|
@@ -138,7 +138,7 @@ gradient and stores it in the gradient.
|
|
|
138
138
|
|
|
139
139
|
## Functions
|
|
140
140
|
|
|
141
|
-
::: haliax.quantization.
|
|
141
|
+
::: haliax.quantization.quantize_linear_layers
|
|
142
142
|
::: haliax.quantization.partition_for_grad_overwrite
|
|
143
143
|
::: haliax.quantization.apply_updates
|
|
144
144
|
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev330"
|
|
@@ -16,6 +16,16 @@ import haliax
|
|
|
16
16
|
from haliax.types import PrecisionLike
|
|
17
17
|
|
|
18
18
|
|
|
19
|
+
try:
|
|
20
|
+
# jax v0.5.1 or newer
|
|
21
|
+
from jax._src.numpy import (
|
|
22
|
+
einsum as jax_einsum, # pylint: disable=g-import-not-at-top # pytype: disable=import-error
|
|
23
|
+
)
|
|
24
|
+
except ImportError:
|
|
25
|
+
# jax v0.5.0 or older
|
|
26
|
+
from jax._src.numpy import lax_numpy as jax_einsum # pylint: disable=g-import-not-at-top
|
|
27
|
+
|
|
28
|
+
|
|
19
29
|
F = typing.TypeVar("F", bound=Callable[..., Any])
|
|
20
30
|
|
|
21
31
|
|
|
@@ -169,6 +179,13 @@ def _jittable_dg_einsum(
|
|
|
169
179
|
preferred_element_type: DTypeLike | None = None,
|
|
170
180
|
_dot_general: Callable[..., Array] = jax.lax.dot_general,
|
|
171
181
|
) -> Array:
|
|
182
|
+
"""
|
|
183
|
+
So we want to pass around a jittable dot_general module, but JAX's builtin version doesn't support this.
|
|
184
|
+
So we copy over the implementation of jax.numpy.einsum and modify thing so that it is jittable (via
|
|
185
|
+
eqx.filter_jit)
|
|
186
|
+
|
|
187
|
+
More or less copied from AQT
|
|
188
|
+
"""
|
|
172
189
|
operands = (subscripts, *operands)
|
|
173
190
|
if out is not None:
|
|
174
191
|
raise NotImplementedError("The 'out' argument to jnp.einsum is not supported.")
|
|
@@ -191,7 +208,7 @@ def _jittable_dg_einsum(
|
|
|
191
208
|
|
|
192
209
|
contractions = tuple((a, frozenset(b), c) for a, b, c, *_ in contractions)
|
|
193
210
|
|
|
194
|
-
einsum = eqx.filter_jit(
|
|
211
|
+
einsum = eqx.filter_jit(jax_einsum._einsum, inline=True)
|
|
195
212
|
if spec is not None:
|
|
196
213
|
einsum = jax.named_call(einsum, name=spec)
|
|
197
214
|
return einsum(operands, contractions, precision, preferred_element_type, _dot_general) # type: ignore[operator]
|
|
@@ -33,6 +33,7 @@ def dot_product_attention_weights(
|
|
|
33
33
|
bias: Optional[NamedArray] = None,
|
|
34
34
|
attention_dtype: Optional[jnp.dtype] = None,
|
|
35
35
|
precision: PrecisionLike = None,
|
|
36
|
+
scaling_factor: Optional[float] = None,
|
|
36
37
|
) -> NamedArray:
|
|
37
38
|
"""
|
|
38
39
|
NamedArray version of dot product attention. Computes the logits for the attention weights. Note that the
|
|
@@ -46,12 +47,16 @@ def dot_product_attention_weights(
|
|
|
46
47
|
:param bias: Optional[NamedArray] broadcast compatible with (KeySize, QPos, KPos). Should be float
|
|
47
48
|
:param attention_dtype: Optional dtype to use for attention
|
|
48
49
|
:param precision: PrecisionLike for dot product. See precision argument to jax.lax.dot_general
|
|
50
|
+
:param scaling_factor: Optional float as scaling factor for attention score. Default to 1/sqrt(D)
|
|
49
51
|
:return: NamedArray of shape (QPos, KPos)
|
|
50
52
|
"""
|
|
51
53
|
# cf https://github.com/google/flax/blob/509bf97ea272e130d932920f45307ac98947d994/flax/linen/attention.py#L40
|
|
52
54
|
|
|
53
55
|
orig_dtype = query.dtype
|
|
54
|
-
|
|
56
|
+
if scaling_factor is None:
|
|
57
|
+
scaling_factor = 1.0 / jnp.sqrt(query.axis_size(Key))
|
|
58
|
+
|
|
59
|
+
query = query * scaling_factor
|
|
55
60
|
|
|
56
61
|
if attention_dtype is not None:
|
|
57
62
|
query = query.astype(attention_dtype)
|
|
@@ -8,7 +8,7 @@ from typing import Callable, ContextManager, Mapping, Optional, ParamSpec, Seque
|
|
|
8
8
|
|
|
9
9
|
import equinox as eqx
|
|
10
10
|
import jax
|
|
11
|
-
from equinox import module_update_wrapper
|
|
11
|
+
from equinox import is_array, module_update_wrapper
|
|
12
12
|
from jax.lax import with_sharding_constraint
|
|
13
13
|
from jax.sharding import Mesh, NamedSharding, PartitionSpec, SingleDeviceSharding
|
|
14
14
|
from jaxtyping import PyTree
|
|
@@ -20,7 +20,7 @@ from .axis import Axis, AxisSelection, AxisSelector
|
|
|
20
20
|
from .core import NamedArray
|
|
21
21
|
from .jax_utils import Static, is_in_jit, is_jax_array_like, is_on_mac_metal
|
|
22
22
|
from .tree_util import hashable_combine, hashable_partition
|
|
23
|
-
from .util import StringHolderEnum, ensure_tuple
|
|
23
|
+
from .util import StringHolderEnum, ensure_tuple
|
|
24
24
|
|
|
25
25
|
|
|
26
26
|
PhysicalAxisSpec = Union[(str), Sequence[str]]
|
|
@@ -274,7 +274,7 @@ class _NamedJitWrapper(eqx.Module):
|
|
|
274
274
|
if out_axis_resources is None:
|
|
275
275
|
out_axis_resources = axis_resources
|
|
276
276
|
|
|
277
|
-
dynamic_argspec, static_argspec = hashable_partition((args, kwargs),
|
|
277
|
+
dynamic_argspec, static_argspec = hashable_partition((args, kwargs), is_array)
|
|
278
278
|
dynamic = (self._dynamic_fun, dynamic_argspec)
|
|
279
279
|
|
|
280
280
|
donate_args = self._donate_args
|
|
@@ -436,7 +436,7 @@ def named_jit(
|
|
|
436
436
|
**pjit_args,
|
|
437
437
|
)
|
|
438
438
|
|
|
439
|
-
dynamic_fun, static_fun = hashable_partition(fn,
|
|
439
|
+
dynamic_fun, static_fun = hashable_partition(fn, is_array)
|
|
440
440
|
|
|
441
441
|
wrapper = _NamedJitWrapper(
|
|
442
442
|
fn,
|
|
@@ -514,7 +514,7 @@ def _named_pjit_cache(fun_names, **jitkwargs) -> WrappedCallable:
|
|
|
514
514
|
fun = hashable_combine(dynamic_fun, static_fun)
|
|
515
515
|
args, kwargs = hashable_combine(dynamic_spec, static_spec)
|
|
516
516
|
out = fun(*args, **kwargs)
|
|
517
|
-
out_dynamic, out_static = hashable_partition(out,
|
|
517
|
+
out_dynamic, out_static = hashable_partition(out, is_array)
|
|
518
518
|
return out_dynamic, Static(out_static)
|
|
519
519
|
|
|
520
520
|
fun_name, fun_qualname = fun_names
|
|
@@ -543,7 +543,7 @@ def _cached_filter_eval_shape(fun, *args, **kwargs):
|
|
|
543
543
|
eval_shape is surprisingly expensive, so we cache it. We use this for named_pjit for evaluating resource partitions
|
|
544
544
|
of the output.
|
|
545
545
|
"""
|
|
546
|
-
dynamic, static = hashable_partition((fun, args, kwargs),
|
|
546
|
+
dynamic, static = hashable_partition((fun, args, kwargs), is_array)
|
|
547
547
|
if static not in _eval_shape_cache:
|
|
548
548
|
_eval_shape_cache[static] = eqx.filter_eval_shape(fun, *args, **kwargs)
|
|
549
549
|
|
|
@@ -7,10 +7,12 @@ import warnings
|
|
|
7
7
|
from dataclasses import dataclass
|
|
8
8
|
from typing import Optional, Protocol, TypeVar
|
|
9
9
|
|
|
10
|
+
import aqt.jax.v2.config as aqt_config
|
|
10
11
|
import equinox as eqx
|
|
11
12
|
import jax
|
|
13
|
+
import jax.random as jrandom
|
|
14
|
+
from aqt.jax.v2.aqt_dot_general import DotGeneral
|
|
12
15
|
from jax import numpy as jnp
|
|
13
|
-
from jax._src.tree_util import BuiltInKeyEntry
|
|
14
16
|
from jax.tree_util import DictKey, FlattenedIndexKey, GetAttrKey, SequenceKey
|
|
15
17
|
from jax.typing import DTypeLike
|
|
16
18
|
|
|
@@ -184,23 +186,63 @@ class Fp8DotGeneralOp(OverwriteWithGradient):
|
|
|
184
186
|
return y
|
|
185
187
|
|
|
186
188
|
|
|
189
|
+
class Int8DotGeneralOp(OverwriteWithGradient):
|
|
190
|
+
|
|
191
|
+
cfg: DotGeneral
|
|
192
|
+
|
|
193
|
+
@classmethod
|
|
194
|
+
def init(cls):
|
|
195
|
+
cfg = aqt_config.config_v3()
|
|
196
|
+
return cls(cfg)
|
|
197
|
+
|
|
198
|
+
def __call__(
|
|
199
|
+
self,
|
|
200
|
+
lhs,
|
|
201
|
+
rhs,
|
|
202
|
+
dimension_numbers,
|
|
203
|
+
precision,
|
|
204
|
+
preferred_element_type=None,
|
|
205
|
+
):
|
|
206
|
+
cfg = aqt_config.set_context(self.cfg, jrandom.PRNGKey(42), train_step=None)
|
|
207
|
+
return cfg(lhs, rhs, dimension_numbers, precision, preferred_element_type)
|
|
208
|
+
|
|
209
|
+
|
|
187
210
|
@dataclass(frozen=True)
|
|
188
|
-
class
|
|
189
|
-
amax_history_length: int = 1024
|
|
190
|
-
compute_dtype: DTypeLike = None
|
|
211
|
+
class QuantizationConfig:
|
|
191
212
|
targets: Optional[list[str] | str] = dataclasses.field(default=None)
|
|
192
213
|
"""
|
|
193
214
|
If provided, only modules with names in this list will be quantized. If a single string, will be treated as a regex
|
|
194
215
|
"""
|
|
195
216
|
|
|
217
|
+
amax_history_length: int = 1024
|
|
218
|
+
compute_dtype: DTypeLike = None
|
|
219
|
+
|
|
220
|
+
fp8: bool = False
|
|
221
|
+
int8: bool = False
|
|
222
|
+
|
|
223
|
+
def __post_init__(self):
|
|
224
|
+
assert not (self.fp8 and self.int8), "Cannot use FP8 and INT8 quantization at the same time."
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def quantize_linear_layers(tree: T, config: QuantizationConfig) -> T:
|
|
228
|
+
"""
|
|
229
|
+
Converts a module tree to use FP8/INT8 quantization.
|
|
230
|
+
"""
|
|
231
|
+
if config.fp8:
|
|
232
|
+
return _quantize_linear_layers(tree, config, Fp8DotGeneralOp, config.amax_history_length, config.compute_dtype)
|
|
233
|
+
elif config.int8:
|
|
234
|
+
return _quantize_linear_layers(tree, config, Int8DotGeneralOp)
|
|
235
|
+
else:
|
|
236
|
+
warnings.warn("Both fp8 and int8 are set to False. `quantize_linear_layers()` is no-op.")
|
|
237
|
+
return tree
|
|
238
|
+
|
|
196
239
|
|
|
197
|
-
def
|
|
240
|
+
def _quantize_linear_layers(tree: T, config: QuantizationConfig, dot_general_cls, *args, **kwargs) -> T:
|
|
198
241
|
"""
|
|
199
|
-
|
|
200
|
-
Linear modules that have a name that matches the targets (if provided) will be converted to use FP8.
|
|
242
|
+
Linear modules that have a name that matches the targets (if provided) will be converted to quantized version.
|
|
201
243
|
(If targets is None, all linear modules will be converted.)
|
|
202
244
|
|
|
203
|
-
This essentially goes through and adds
|
|
245
|
+
This essentially goes through and adds corresponding DotGeneralOp to the Linear modules.
|
|
204
246
|
"""
|
|
205
247
|
|
|
206
248
|
def _is_special_module(module):
|
|
@@ -212,31 +254,29 @@ def fp8_linear_layers(tree: T, config: Fp8Config = Fp8Config()) -> T:
|
|
|
212
254
|
return functools.reduce(lambda ctor, batch_axis: vmap(ctor, batch_axis), reversed(batch_dims), ctor)
|
|
213
255
|
|
|
214
256
|
# TODO: test scanlayers for dg
|
|
215
|
-
def
|
|
257
|
+
def quantize_module(path_prefix, batch_dims: tuple[Axis, ...], path, module: T) -> T:
|
|
216
258
|
path = path_prefix + path
|
|
217
259
|
if isinstance(module, hnn.Stacked):
|
|
218
260
|
new_inner = jax.tree_util.tree_map_with_path(
|
|
219
|
-
functools.partial(
|
|
261
|
+
functools.partial(quantize_module, path_prefix + (GetAttrKey("stacked"),), batch_dims + (module.Block,)), # type: ignore
|
|
220
262
|
module.stacked,
|
|
221
263
|
is_leaf=_is_special_module,
|
|
222
264
|
)
|
|
223
265
|
return dataclasses.replace(module, stacked=new_inner) # type: ignore
|
|
224
266
|
elif isinstance(module, hnn.Linear):
|
|
225
|
-
if
|
|
226
|
-
vmapped_dg = _batchify_ctor(
|
|
227
|
-
config.amax_history_length, config.compute_dtype
|
|
228
|
-
)
|
|
267
|
+
if _matches_target(path, config):
|
|
268
|
+
vmapped_dg = _batchify_ctor(dot_general_cls.init, batch_dims)(*args, **kwargs)
|
|
229
269
|
module = dataclasses.replace(module, dot_general=vmapped_dg) # type: ignore
|
|
230
270
|
return module
|
|
231
271
|
else:
|
|
232
272
|
return module
|
|
233
273
|
|
|
234
274
|
return jax.tree_util.tree_map_with_path(
|
|
235
|
-
lambda p, m:
|
|
275
|
+
lambda p, m: quantize_module((), (), p, m), tree, is_leaf=_is_special_module
|
|
236
276
|
)
|
|
237
277
|
|
|
238
278
|
|
|
239
|
-
def
|
|
279
|
+
def _matches_target(key_path, config: QuantizationConfig) -> bool:
|
|
240
280
|
if not key_path:
|
|
241
281
|
key = ""
|
|
242
282
|
else:
|
|
@@ -253,7 +293,7 @@ def _matches_target_fp8(key_path, config: Fp8Config) -> bool:
|
|
|
253
293
|
return re.match(config.targets, key_path_str) is not None
|
|
254
294
|
|
|
255
295
|
|
|
256
|
-
def _key_path_to_str(key_path: tuple
|
|
296
|
+
def _key_path_to_str(key_path: tuple) -> str:
|
|
257
297
|
out = ""
|
|
258
298
|
for k in key_path:
|
|
259
299
|
match k:
|
|
@@ -24,6 +24,7 @@ def uniform(
|
|
|
24
24
|
minval = broadcast_to(minval, shape).array
|
|
25
25
|
maxval = broadcast_to(maxval, shape).array
|
|
26
26
|
jax_shape = _to_jax_shape(shape)
|
|
27
|
+
print(jax_shape, minval, maxval)
|
|
27
28
|
jax_array = jrandom.uniform(key=key, shape=jax_shape, dtype=dtype, minval=minval, maxval=maxval)
|
|
28
29
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
29
30
|
|
|
@@ -10,11 +10,11 @@ import haliax as hax
|
|
|
10
10
|
from haliax._src.fp8 import compute_scale
|
|
11
11
|
from haliax.nn import Linear
|
|
12
12
|
from haliax.quantization import (
|
|
13
|
-
Fp8Config,
|
|
14
13
|
Fp8DotGeneralOp,
|
|
14
|
+
QuantizationConfig,
|
|
15
15
|
apply_updates,
|
|
16
|
-
fp8_linear_layers,
|
|
17
16
|
partition_for_grad_overwrite,
|
|
17
|
+
quantize_linear_layers,
|
|
18
18
|
)
|
|
19
19
|
|
|
20
20
|
|
|
@@ -34,7 +34,7 @@ def test_fp8_is_reasonable():
|
|
|
34
34
|
assert output.shape == fp8_output.shape
|
|
35
35
|
assert output.dtype == fp8_output.dtype
|
|
36
36
|
|
|
37
|
-
assert_trees_all_close(output.array, fp8_output.array, atol=
|
|
37
|
+
assert_trees_all_close(output.array, fp8_output.array, atol=2e-2, rtol=5e-2)
|
|
38
38
|
|
|
39
39
|
|
|
40
40
|
# https://github.com/google/flax/blob/6f2b08e024c2fd2f8cec42a6c82408cb35412319/tests/linen/linen_test.py#L1222
|
|
@@ -138,7 +138,7 @@ def test_layer_splicing():
|
|
|
138
138
|
Output = hax.Axis("Output", 32)
|
|
139
139
|
mlp = hax.nn.MLP.init(Input, Output, Hidden, 3, key=init_key, init_scale=0.1)
|
|
140
140
|
|
|
141
|
-
mlp_q =
|
|
141
|
+
mlp_q = quantize_linear_layers(mlp, QuantizationConfig(fp8=True))
|
|
142
142
|
for layer in mlp_q.layers:
|
|
143
143
|
assert isinstance(layer.dot_general, Fp8DotGeneralOp)
|
|
144
144
|
|
|
@@ -148,14 +148,14 @@ def test_layer_splicing():
|
|
|
148
148
|
chex.assert_trees_all_close(output.array, output_q.array, atol=1e-3, rtol=1e-3)
|
|
149
149
|
assert not jnp.allclose(output_q.array, 0) # don't want them to all underflow
|
|
150
150
|
|
|
151
|
-
mlp_q =
|
|
151
|
+
mlp_q = quantize_linear_layers(mlp, QuantizationConfig(targets="layers.0", fp8=True))
|
|
152
152
|
for i, layer in enumerate(mlp_q.layers):
|
|
153
153
|
if i == 0:
|
|
154
154
|
assert isinstance(layer.dot_general, Fp8DotGeneralOp)
|
|
155
155
|
else:
|
|
156
156
|
assert not isinstance(layer.dot_general, Fp8DotGeneralOp)
|
|
157
157
|
|
|
158
|
-
mlp_q =
|
|
158
|
+
mlp_q = quantize_linear_layers(mlp, QuantizationConfig(targets=["0", "1"], fp8=True))
|
|
159
159
|
for i, layer in enumerate(mlp_q.layers):
|
|
160
160
|
if i < 2:
|
|
161
161
|
assert isinstance(layer.dot_general, Fp8DotGeneralOp)
|
|
@@ -193,7 +193,7 @@ def test_fp8ize_stacking():
|
|
|
193
193
|
In = hax.Axis("In", 16)
|
|
194
194
|
Out = hax.Axis("Out", 32)
|
|
195
195
|
tformer = Tformer.init(In, Out, key=jrandom.PRNGKey(0))
|
|
196
|
-
tformer_q =
|
|
196
|
+
tformer_q = quantize_linear_layers(tformer, QuantizationConfig(fp8=True))
|
|
197
197
|
|
|
198
198
|
# want to be sure this vmaps the dot_general to the right places
|
|
199
199
|
dg = tformer_q.blocks.stacked.up_proj.dot_general
|
|
@@ -204,7 +204,7 @@ def test_fp8ize_stacking():
|
|
|
204
204
|
assert isinstance(dg, Fp8DotGeneralOp)
|
|
205
205
|
|
|
206
206
|
# just stack the up_proj
|
|
207
|
-
tformer_q =
|
|
207
|
+
tformer_q = quantize_linear_layers(tformer, QuantizationConfig(targets=["up_proj"], fp8=True))
|
|
208
208
|
dg = tformer_q.blocks.stacked.up_proj.dot_general
|
|
209
209
|
assert isinstance(dg, Fp8DotGeneralOp)
|
|
210
210
|
dg = tformer_q.blocks.stacked.down_proj.dot_general
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
import jax.random as jrandom
|
|
2
|
+
from chex import assert_trees_all_close
|
|
3
|
+
|
|
4
|
+
import haliax as hax
|
|
5
|
+
from haliax.nn import Linear
|
|
6
|
+
from haliax.quantization import Int8DotGeneralOp
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def test_int8_is_reasonable():
|
|
10
|
+
In = hax.Axis("In", 8)
|
|
11
|
+
Out = hax.Axis("Out", 8)
|
|
12
|
+
linear = Linear.init(In, Out, key=jrandom.PRNGKey(0), init_scale=0.1)
|
|
13
|
+
|
|
14
|
+
int8_linear = Linear.init(In, Out, key=jrandom.PRNGKey(0), dot_general=Int8DotGeneralOp.init(), init_scale=0.1)
|
|
15
|
+
|
|
16
|
+
input = hax.random.normal(jrandom.PRNGKey(3), In)
|
|
17
|
+
output = linear(input)
|
|
18
|
+
int8_output = int8_linear(input)
|
|
19
|
+
|
|
20
|
+
assert output.shape == int8_output.shape
|
|
21
|
+
assert output.dtype == int8_output.dtype
|
|
22
|
+
|
|
23
|
+
assert_trees_all_close(output.array, int8_output.array, atol=1e-2, rtol=5e-2)
|
|
@@ -150,7 +150,9 @@ def test_where(use_jit):
|
|
|
150
150
|
Volume = hax.Axis("Volume", Height.size * Width.size * Depth.size)
|
|
151
151
|
named7 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
|
|
152
152
|
named8, named9, named10 = hax_where(named7 > 0.5, fill_value=-1, new_axis=Volume)
|
|
153
|
-
|
|
153
|
+
unnamed_7 = named7.array
|
|
154
|
+
unnamed_8, unnamed_9, unnamed_10 = jnp.where(unnamed_7 > 0.5, size=Volume.size, fill_value=-1)
|
|
155
|
+
assert jnp.all(unnamed_8 == named8.array)
|
|
154
156
|
|
|
155
157
|
|
|
156
158
|
def test_clip():
|
|
@@ -41,8 +41,15 @@ def test_uniform_with_bounds_broadcast():
|
|
|
41
41
|
assert hax.all(u >= lb)
|
|
42
42
|
assert hax.all(u <= 0.5)
|
|
43
43
|
|
|
44
|
+
# for lb and ub, broadcast the raw arrays for jax
|
|
45
|
+
lb_raw = lb.array.reshape(-1, 1)
|
|
46
|
+
ub_raw = ub.array.reshape(1, -1)
|
|
47
|
+
|
|
48
|
+
lb_raw = jnp.broadcast_to(lb_raw, (Height.size, Width.size))
|
|
49
|
+
ub_raw = jnp.broadcast_to(ub_raw, (Height.size, Width.size))
|
|
50
|
+
|
|
44
51
|
check_gen_is_equal(
|
|
45
|
-
lambda k, s: jax.random.uniform(k, shape=s, minval=
|
|
52
|
+
lambda k, s: jax.random.uniform(k, shape=s, minval=lb_raw, maxval=ub_raw),
|
|
46
53
|
lambda k, s: hax.random.uniform(k, s, minval=lb, maxval=ub),
|
|
47
54
|
)
|
|
48
55
|
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev326"
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|