haliax 1.4.dev395__tar.gz → 1.4.dev396__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.dev395 → haliax-1.4.dev396}/PKG-INFO +1 -1
- haliax-1.4.dev396/docs/vmap.md +38 -0
- haliax-1.4.dev396/src/haliax/__about__.py +1 -0
- haliax-1.4.dev395/docs/vmap.md +0 -9
- haliax-1.4.dev395/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.coveragerc +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.flake8 +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.gitignore +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/AGENTS.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/LICENSE +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/README.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/api.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/css/material.css +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/faq.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/fp8.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/index.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/indexing.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/matmul.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/nn.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/partitioning.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/primer.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/rearrange.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/requirements.txt +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/scan.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/state-dict.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/tutorial.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/typing.md +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/mkdocs.yml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/pyproject.toml +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/core.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/random.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/types.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/util.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/core_test.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_attention.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_axis.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_conv.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_debug.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_dot.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_hof.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_int8.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_nn.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_ops.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_pool.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_random.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_scan.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_utils.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev395 → haliax-1.4.dev396}/uv.lock +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev396
|
|
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,38 @@
|
|
|
1
|
+
## Vectorization with `haliax.vmap`
|
|
2
|
+
|
|
3
|
+
`haliax.vmap` is a [`NamedArray`][haliax.NamedArray] aware wrapper around
|
|
4
|
+
[`jax.vmap`][jax.vmap]. Instead of supplying positional axis numbers you pass
|
|
5
|
+
the [`Axis`][haliax.Axis] (or axis name) you want to map over. Any
|
|
6
|
+
`NamedArray` containing that axis is mapped in parallel and the axis is
|
|
7
|
+
reinserted in the output. Regular JAX arrays can be mapped as well by
|
|
8
|
+
providing a `default` spec or per‑argument overrides.
|
|
9
|
+
|
|
10
|
+
Unlike vanilla `jax.vmap`, you may supply **one or more axes**. When multiple
|
|
11
|
+
axes are given, the function is vmapped over each axis in turn (innermost first).
|
|
12
|
+
If an axis isn't already present in the array you must also specify its size,
|
|
13
|
+
either by passing an `Axis` object (`Axis("batch", 4)`) or a mapping such as
|
|
14
|
+
`{"batch": 4}` so the new dimension can be inserted.
|
|
15
|
+
|
|
16
|
+
### Basic Example
|
|
17
|
+
|
|
18
|
+
```python
|
|
19
|
+
import haliax as hax
|
|
20
|
+
|
|
21
|
+
Batch = hax.Axis("batch", 4)
|
|
22
|
+
|
|
23
|
+
def double(x):
|
|
24
|
+
return x * 2
|
|
25
|
+
|
|
26
|
+
x = hax.arange(Batch)
|
|
27
|
+
y = hax.vmap(double, Batch)(x)
|
|
28
|
+
```
|
|
29
|
+
|
|
30
|
+
The result `y` has the same `Batch` axis as `x`, and each element was processed
|
|
31
|
+
in parallel. With JAX you would write `jax.vmap(double)(x.array)` and manually
|
|
32
|
+
specify `in_axes`, but Haliax handles the axis automatically.
|
|
33
|
+
|
|
34
|
+
For applying many modules in parallel see
|
|
35
|
+
[`Stacked.vmap`](scan.md#apply-blocks-in-parallel-with-vmap) which builds on this
|
|
36
|
+
primitive.
|
|
37
|
+
|
|
38
|
+
::: haliax.vmap
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev396"
|
haliax-1.4.dev395/docs/vmap.md
DELETED
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev395"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|