haliax 1.4.dev299__tar.gz → 1.4.dev301__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.dev299 → haliax-1.4.dev301}/.github/workflows/run_pre_commit.yaml +1 -1
- haliax-1.4.dev301/.github/workflows/run_quick_levanter_tests.yaml +36 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.github/workflows/run_tests.yaml +3 -6
- {haliax-1.4.dev299 → haliax-1.4.dev301}/PKG-INFO +1 -1
- haliax-1.4.dev301/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/partitioning.py +0 -2
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_partitioning.py +24 -2
- haliax-1.4.dev299/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.coveragerc +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.flake8 +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.gitignore +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/LICENSE +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/README.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/api.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/css/material.css +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/faq.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/fp8.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/hof.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/index.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/indexing.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/matmul.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/nn.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/partitioning.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/rearrange.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/requirements.txt +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/docs/tutorial.md +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/mkdocs.yml +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/pyproject.toml +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/core.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/random.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/types.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/util.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/core_test.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_attention.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_axis.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_conv.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_debug.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_dot.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_hof.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_nn.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_ops.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_pool.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_random.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_scan.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev299 → haliax-1.4.dev301}/tests/test_utils.py +0 -0
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
name: Run Levanter Tests
|
|
2
|
+
|
|
3
|
+
on: [pull_request]
|
|
4
|
+
|
|
5
|
+
jobs:
|
|
6
|
+
build:
|
|
7
|
+
|
|
8
|
+
runs-on: ubuntu-latest
|
|
9
|
+
|
|
10
|
+
steps:
|
|
11
|
+
- uses: actions/checkout@v3
|
|
12
|
+
- name: Set up Python 3.10.11
|
|
13
|
+
uses: actions/setup-python@v4
|
|
14
|
+
with:
|
|
15
|
+
python-version: 3.10.11
|
|
16
|
+
- name: Install dependencies
|
|
17
|
+
run: |
|
|
18
|
+
python -m pip install --upgrade pip
|
|
19
|
+
pip install flake8 pytest
|
|
20
|
+
pip install --upgrade "jax[cpu]==0.4.26" "jaxlib[cpu]==0.4.26"
|
|
21
|
+
|
|
22
|
+
- name: Install Levanter from source
|
|
23
|
+
run: |
|
|
24
|
+
cd ..
|
|
25
|
+
git clone https://github.com/stanford-crfm/levanter.git
|
|
26
|
+
cd levanter
|
|
27
|
+
pip install -e .
|
|
28
|
+
- name: Install Haliax on top
|
|
29
|
+
run: |
|
|
30
|
+
# install second since levanter will install a built version of haliax
|
|
31
|
+
cd ../haliax
|
|
32
|
+
pip install .[dev]
|
|
33
|
+
- name: Test levanter with pytest
|
|
34
|
+
run: |
|
|
35
|
+
cd ../levanter
|
|
36
|
+
XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:../src pytest tests -m "not entry and not slow"
|
|
@@ -1,21 +1,18 @@
|
|
|
1
1
|
name: Run Tests
|
|
2
2
|
|
|
3
|
-
on: [push]
|
|
3
|
+
on: [push, pull_request]
|
|
4
4
|
|
|
5
5
|
jobs:
|
|
6
6
|
build:
|
|
7
7
|
|
|
8
8
|
runs-on: ubuntu-latest
|
|
9
|
-
strategy:
|
|
10
|
-
matrix:
|
|
11
|
-
python-version: ["3.10.11"]
|
|
12
9
|
|
|
13
10
|
steps:
|
|
14
11
|
- uses: actions/checkout@v3
|
|
15
|
-
- name: Set up Python
|
|
12
|
+
- name: Set up Python 3.10.11
|
|
16
13
|
uses: actions/setup-python@v4
|
|
17
14
|
with:
|
|
18
|
-
python-version:
|
|
15
|
+
python-version: 3.10.11
|
|
19
16
|
- name: Install dependencies
|
|
20
17
|
run: |
|
|
21
18
|
python -m pip install --upgrade pip
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev301
|
|
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.dev301"
|
|
@@ -310,8 +310,6 @@ class _NamedJitWrapper(eqx.Module):
|
|
|
310
310
|
my_pjit_args = dict(**self._pjit_args)
|
|
311
311
|
|
|
312
312
|
if in_axis_resources is not None or axis_resources is not None:
|
|
313
|
-
if in_axis_resources is None:
|
|
314
|
-
in_axis_resources = axis_resources
|
|
315
313
|
in_resources = infer_resource_partitions(
|
|
316
314
|
(dynamic_donated, dynamic_reserved),
|
|
317
315
|
in_axis_resources,
|
|
@@ -99,7 +99,7 @@ def test_pjit_class_init_with_args():
|
|
|
99
99
|
|
|
100
100
|
devices = jax.devices()
|
|
101
101
|
with Mesh(np.array(devices).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL)):
|
|
102
|
-
mod = named_jit(ModWithArgs)(hax.ones((Dim1, Dim2)))
|
|
102
|
+
mod = named_jit(ModWithArgs)(hax.shard(hax.ones((Dim1, Dim2))))
|
|
103
103
|
assert isinstance(mod, ModWithArgs)
|
|
104
104
|
assert mod.array.array.shape == (Dim1.size, Dim2.size)
|
|
105
105
|
assert mod.array2.array.shape == (Dim3.size,)
|
|
@@ -173,7 +173,7 @@ def test_shard_with_axis_mapping_inside_jit():
|
|
|
173
173
|
|
|
174
174
|
jax.debug.inspect_array_sharding(arr.array, callback=lambda x: assert_eq(x, expected))
|
|
175
175
|
|
|
176
|
-
@named_jit(
|
|
176
|
+
@named_jit(out_axis_resources=resource_map)
|
|
177
177
|
def do_shard(x, y):
|
|
178
178
|
x = hax.shard(x, resource_map)
|
|
179
179
|
assert_inside_pjit(x, NamedSharding(mesh, PartitionSpec(None, ResourceAxis.DATA)))
|
|
@@ -293,3 +293,25 @@ def test_cross_device_sharding():
|
|
|
293
293
|
z_devices = z.array.devices()
|
|
294
294
|
|
|
295
295
|
assert set(d.platform for d in x_devices) == set(d.platform for d in z_devices)
|
|
296
|
+
|
|
297
|
+
|
|
298
|
+
def test_named_jit_no_in_axis_resources():
|
|
299
|
+
mesh = Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
|
|
300
|
+
with axis_mapping(resource_map), mesh:
|
|
301
|
+
|
|
302
|
+
class MyModule(eqx.Module):
|
|
303
|
+
array: NamedArray
|
|
304
|
+
|
|
305
|
+
def __init__(self):
|
|
306
|
+
self.array = hax.ones((Dim1, Dim2))
|
|
307
|
+
|
|
308
|
+
data = hax.ones((Dim1, Dim2))
|
|
309
|
+
data = hax.shard(data, {})
|
|
310
|
+
|
|
311
|
+
@named_jit(axis_resources=resource_map)
|
|
312
|
+
def fn(data):
|
|
313
|
+
mod = MyModule()
|
|
314
|
+
return mod.array
|
|
315
|
+
|
|
316
|
+
r = fn(data)
|
|
317
|
+
assert r.array.sharding.device_set == set(jax.devices())
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev299"
|
|
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
|