haliax 1.4.dev300__tar.gz → 1.4.dev302__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.dev300 → haliax-1.4.dev302}/PKG-INFO +1 -1
- haliax-1.4.dev302/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/partitioning.py +5 -3
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_partitioning.py +24 -2
- haliax-1.4.dev300/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.coveragerc +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.flake8 +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.gitignore +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/LICENSE +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/README.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/api.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/css/material.css +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/faq.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/fp8.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/hof.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/index.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/indexing.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/matmul.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/nn.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/partitioning.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/rearrange.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/requirements.txt +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/docs/tutorial.md +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/mkdocs.yml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/pyproject.toml +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/core.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/random.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/types.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/util.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/core_test.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_attention.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_axis.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_conv.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_debug.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_dot.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_hof.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_nn.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_ops.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_pool.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_random.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_scan.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev300 → haliax-1.4.dev302}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev302
|
|
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.dev302"
|
|
@@ -141,9 +141,13 @@ def shard(x: T, mapping: Optional[ResourceMapping] = None, mesh: Optional[Mesh]
|
|
|
141
141
|
assert isinstance(sharding, NamedSharding)
|
|
142
142
|
if is_in_jit():
|
|
143
143
|
return with_sharding_constraint(x, sharding)
|
|
144
|
-
|
|
144
|
+
elif sharding.is_fully_addressable:
|
|
145
145
|
sharded_array = jax.device_put(x.array, sharding)
|
|
146
146
|
return NamedArray(sharded_array, x.axes)
|
|
147
|
+
else:
|
|
148
|
+
# sharded_array = jax.device_put(x.array, sharding)
|
|
149
|
+
ret = eqx.filter_jit(lambda x: with_sharding_constraint(x, sharding))(x)
|
|
150
|
+
return ret
|
|
147
151
|
|
|
148
152
|
return htu.tree_map(_do_device_put, x)
|
|
149
153
|
|
|
@@ -310,8 +314,6 @@ class _NamedJitWrapper(eqx.Module):
|
|
|
310
314
|
my_pjit_args = dict(**self._pjit_args)
|
|
311
315
|
|
|
312
316
|
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
317
|
in_resources = infer_resource_partitions(
|
|
316
318
|
(dynamic_donated, dynamic_reserved),
|
|
317
319
|
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.dev300"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|