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.
Files changed (118) hide show
  1. {haliax-1.4.dev395 → haliax-1.4.dev396}/PKG-INFO +1 -1
  2. haliax-1.4.dev396/docs/vmap.md +38 -0
  3. haliax-1.4.dev396/src/haliax/__about__.py +1 -0
  4. haliax-1.4.dev395/docs/vmap.md +0 -9
  5. haliax-1.4.dev395/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev395 → haliax-1.4.dev396}/.coveragerc +0 -0
  7. {haliax-1.4.dev395 → haliax-1.4.dev396}/.flake8 +0 -0
  8. {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev395 → haliax-1.4.dev396}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev395 → haliax-1.4.dev396}/.gitignore +0 -0
  13. {haliax-1.4.dev395 → haliax-1.4.dev396}/.playbooks/add-types.md +0 -0
  14. {haliax-1.4.dev395 → haliax-1.4.dev396}/.playbooks/wrap-non-named.md +0 -0
  15. {haliax-1.4.dev395 → haliax-1.4.dev396}/.pre-commit-config.yaml +0 -0
  16. {haliax-1.4.dev395 → haliax-1.4.dev396}/.readthedocs.yaml +0 -0
  17. {haliax-1.4.dev395 → haliax-1.4.dev396}/AGENTS.md +0 -0
  18. {haliax-1.4.dev395 → haliax-1.4.dev396}/CONTRIBUTING.md +0 -0
  19. {haliax-1.4.dev395 → haliax-1.4.dev396}/LICENSE +0 -0
  20. {haliax-1.4.dev395 → haliax-1.4.dev396}/README.md +0 -0
  21. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/api.md +0 -0
  22. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/broadcasting.md +0 -0
  23. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/cheatsheet.md +0 -0
  24. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/css/material.css +0 -0
  25. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/css/mkdocstrings.css +0 -0
  26. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/faq.md +0 -0
  27. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/data_parallel_mesh.png +0 -0
  28. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  29. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_1d.png +0 -0
  30. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_1d_zero.png +0 -0
  31. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d.png +0 -0
  32. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  33. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  34. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  35. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  36. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/figures/device_mesh_2d_zero.png +0 -0
  37. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/fp8.md +0 -0
  38. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/index.md +0 -0
  39. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/indexing.md +0 -0
  40. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/matmul.md +0 -0
  41. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/nn.md +0 -0
  42. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/partitioning.md +0 -0
  43. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/primer.md +0 -0
  44. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/rearrange.ipynb +0 -0
  45. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/rearrange.md +0 -0
  46. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/requirements.txt +0 -0
  47. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/scan.md +0 -0
  48. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev395 → haliax-1.4.dev396}/docs/typing.md +0 -0
  51. {haliax-1.4.dev395 → haliax-1.4.dev396}/mkdocs.yml +0 -0
  52. {haliax-1.4.dev395 → haliax-1.4.dev396}/pyproject.toml +0 -0
  53. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/__init__.py +0 -0
  54. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/__init__.py +0 -0
  55. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/compile_utils.py +0 -0
  56. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/dot.py +0 -0
  57. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/einsum.py +0 -0
  58. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/fp8.py +0 -0
  59. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/parsing.py +0 -0
  60. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/rearrange.py +0 -0
  61. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/scan.py +0 -0
  62. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/state_dict.py +0 -0
  63. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/_src/util.py +0 -0
  64. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/axis.py +0 -0
  65. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/core.py +0 -0
  66. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/debug.py +0 -0
  67. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/haxtyping.py +0 -0
  68. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/hof.py +0 -0
  69. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/jax_utils.py +0 -0
  70. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/__init__.py +0 -0
  71. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/activations.py +0 -0
  72. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/attention.py +0 -0
  73. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/conv.py +0 -0
  74. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/dropout.py +0 -0
  75. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/embedding.py +0 -0
  76. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/linear.py +0 -0
  77. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/loss.py +0 -0
  78. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/mlp.py +0 -0
  79. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/normalization.py +0 -0
  80. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/pool.py +0 -0
  81. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/nn/scan.py +0 -0
  82. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/ops.py +0 -0
  83. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/partitioning.py +0 -0
  84. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/quantization.py +0 -0
  85. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/random.py +0 -0
  86. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/specialized_fns.py +0 -0
  87. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/state_dict.py +0 -0
  88. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/tree_util.py +0 -0
  89. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/types.py +0 -0
  90. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/util.py +0 -0
  91. {haliax-1.4.dev395 → haliax-1.4.dev396}/src/haliax/wrap.py +0 -0
  92. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/core_test.py +0 -0
  93. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_attention.py +0 -0
  94. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_axis.py +0 -0
  95. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_conv.py +0 -0
  96. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_debug.py +0 -0
  97. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_dot.py +0 -0
  98. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_dtype_typing.py +0 -0
  99. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_einsum.py +0 -0
  100. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_fp8.py +0 -0
  101. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_hof.py +0 -0
  102. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_int8.py +0 -0
  103. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_namedarray_typing.py +0 -0
  104. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_nn.py +0 -0
  105. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_ops.py +0 -0
  106. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_parsing.py +0 -0
  107. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_partitioning.py +0 -0
  108. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_pool.py +0 -0
  109. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_random.py +0 -0
  110. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_rearrange.py +0 -0
  111. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_scan.py +0 -0
  112. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_scatter_gather.py +0 -0
  113. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_specialized_fns.py +0 -0
  114. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_state_dict.py +0 -0
  115. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_tree_util.py +0 -0
  116. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_utils.py +0 -0
  117. {haliax-1.4.dev395 → haliax-1.4.dev396}/tests/test_visualize_sharding.py +0 -0
  118. {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.dev395
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"
@@ -1,9 +0,0 @@
1
- ## Vectorization
2
-
3
-
4
- This primitive is also used by [`Stacked.vmap`](scan.md#apply-blocks-in-parallel-with-vmap)
5
- to apply an entire stack of blocks in parallel.
6
-
7
- (This is a work in progress. Please contact dlwh for more information.)
8
-
9
- ::: haliax.vmap
@@ -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