haliax 1.4.dev359__tar.gz → 1.4.dev362__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.dev359 → haliax-1.4.dev362}/.github/workflows/run_tests.yaml +1 -1
- {haliax-1.4.dev359 → haliax-1.4.dev362}/PKG-INFO +1 -1
- haliax-1.4.dev362/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/activations.py +0 -1
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/linear.py +40 -14
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_nn.py +1 -1
- haliax-1.4.dev359/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.coveragerc +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.flake8 +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.gitignore +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/LICENSE +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/README.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/api.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/css/material.css +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/faq.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/fp8.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/index.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/indexing.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/matmul.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/nn.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/partitioning.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/rearrange.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/requirements.txt +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/scan.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/state-dict.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/tutorial.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/vmap.md +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/mkdocs.yml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/pyproject.toml +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/core.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/random.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/types.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/util.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/core_test.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_attention.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_axis.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_conv.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_debug.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_dot.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_hof.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_int8.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_ops.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_pool.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_random.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_scan.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev359 → haliax-1.4.dev362}/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
|
|
20
|
+
pip install -e .[dev]
|
|
21
21
|
- name: Test with pytest
|
|
22
22
|
run: |
|
|
23
23
|
XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:. pytest tests
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev362
|
|
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/
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev362"
|
|
@@ -18,6 +18,7 @@ from ..core import NamedArray
|
|
|
18
18
|
from ..jax_utils import named_call
|
|
19
19
|
from ..partitioning import ResourceAxis
|
|
20
20
|
from ..quantization import DotGeneralOp
|
|
21
|
+
from ..util import ensure_tuple
|
|
21
22
|
|
|
22
23
|
|
|
23
24
|
class Linear(ModuleWithStateDictSerialization):
|
|
@@ -142,7 +143,10 @@ class MoELinear(eqx.Module):
|
|
|
142
143
|
Experts: AxisSpec = eqx.static_field()
|
|
143
144
|
In: Axis = eqx.static_field()
|
|
144
145
|
Out: Axis = eqx.static_field()
|
|
145
|
-
|
|
146
|
+
# TODO: support quanitization for ragged_dot?
|
|
147
|
+
# dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
|
|
148
|
+
|
|
149
|
+
use_gmm: bool = eqx.field(static=True)
|
|
146
150
|
|
|
147
151
|
@staticmethod
|
|
148
152
|
def init(
|
|
@@ -154,6 +158,7 @@ class MoELinear(eqx.Module):
|
|
|
154
158
|
use_bias: bool = True,
|
|
155
159
|
out_first: bool = False,
|
|
156
160
|
init_scale: float = 1.0,
|
|
161
|
+
use_gmm: bool = False,
|
|
157
162
|
) -> "MoELinear":
|
|
158
163
|
"""
|
|
159
164
|
Args:
|
|
@@ -172,7 +177,7 @@ class MoELinear(eqx.Module):
|
|
|
172
177
|
weight = hax.random.truncated_normal(key, joint_spec, -3, 3) * (init_scale / math.sqrt(input_size))
|
|
173
178
|
bias = hax.zeros(Out) if use_bias else None
|
|
174
179
|
|
|
175
|
-
return MoELinear(weight, bias, Experts, In, Out)
|
|
180
|
+
return MoELinear(weight, bias, Experts, In, Out, use_gmm=use_gmm)
|
|
176
181
|
|
|
177
182
|
@named_call
|
|
178
183
|
def __call__(self, inputs, group_sizes, *, key: Optional[PRNGKey] = None):
|
|
@@ -184,21 +189,42 @@ class MoELinear(eqx.Module):
|
|
|
184
189
|
"""
|
|
185
190
|
del key
|
|
186
191
|
|
|
187
|
-
|
|
188
|
-
|
|
192
|
+
dim_numbers = jax.lax.RaggedDotDimensionNumbers(
|
|
193
|
+
dot_dimension_numbers=(
|
|
194
|
+
# contracting
|
|
195
|
+
(ensure_tuple(inputs.axis_indices(self.In)), ensure_tuple(self.weight.axis_indices(self.In))),
|
|
196
|
+
# batch
|
|
197
|
+
((), ()),
|
|
198
|
+
),
|
|
199
|
+
# Everything other than contracting dim is ragged
|
|
200
|
+
lhs_ragged_dimensions=(inputs.axis_indices(hax.axis.without_axes(inputs.axes, self.In))),
|
|
201
|
+
rhs_group_dimensions=(self.weight.axis_indices(self.Experts),),
|
|
202
|
+
)
|
|
189
203
|
|
|
190
|
-
|
|
191
|
-
inputs
|
|
192
|
-
self.
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
|
|
196
|
-
|
|
197
|
-
|
|
204
|
+
if self.use_gmm:
|
|
205
|
+
inputs = inputs.rearrange((..., self.In))
|
|
206
|
+
out_axes = hax.replace_axis(inputs.axes, self.In, self.Out)
|
|
207
|
+
q = _gmm(
|
|
208
|
+
inputs,
|
|
209
|
+
self.weight,
|
|
210
|
+
group_sizes,
|
|
211
|
+
out_axes,
|
|
212
|
+
ar=hax.partitioning.physical_axis_name(self.In) == ResourceAxis.MODEL,
|
|
213
|
+
) # gmm((B, D), (E, D, d)) -> (B, d)
|
|
214
|
+
else:
|
|
215
|
+
q_raw = jax.lax.ragged_dot_general(
|
|
216
|
+
lhs=inputs.array,
|
|
217
|
+
rhs=self.weight.array,
|
|
218
|
+
group_sizes=group_sizes.rearrange((..., self.Experts)).array,
|
|
219
|
+
ragged_dot_dimension_numbers=dim_numbers,
|
|
220
|
+
)
|
|
221
|
+
out_axes = hax.replace_axis(inputs.axes, self.In, self.Out)
|
|
222
|
+
q = hax.named(q_raw, out_axes)
|
|
198
223
|
|
|
199
224
|
if self.bias is not None:
|
|
200
225
|
q = q + self.bias
|
|
201
|
-
|
|
226
|
+
|
|
227
|
+
q = hax.auto_sharded(q)
|
|
202
228
|
|
|
203
229
|
return q
|
|
204
230
|
|
|
@@ -249,7 +275,7 @@ def gmm_sharded(lhs_: jnp.ndarray, rhs_: jnp.ndarray, group_sizes_: jnp.ndarray,
|
|
|
249
275
|
group_sizes_,
|
|
250
276
|
preferred_element_type=lhs_.dtype,
|
|
251
277
|
tiling=(min(m, tile_size[0]), min(k, tile_size[1]), min(n, tile_size[2])),
|
|
252
|
-
|
|
278
|
+
interpret=jax.default_backend() == "cpu",
|
|
253
279
|
)
|
|
254
280
|
|
|
255
281
|
if ar:
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev359"
|
|
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
|
|
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
|