haliax 1.4.dev381__tar.gz → 1.4.dev382__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 (114) hide show
  1. {haliax-1.4.dev381 → haliax-1.4.dev382}/PKG-INFO +3 -3
  2. {haliax-1.4.dev381 → haliax-1.4.dev382}/README.md +2 -2
  3. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/indexing.md +1 -1
  4. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/nn.md +1 -1
  5. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/state-dict.md +1 -1
  6. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/typing.md +1 -1
  7. haliax-1.4.dev382/src/haliax/__about__.py +1 -0
  8. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/conv.py +1 -1
  9. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/linear.py +1 -1
  10. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_rearrange.py +1 -1
  11. haliax-1.4.dev381/src/haliax/__about__.py +0 -1
  12. {haliax-1.4.dev381 → haliax-1.4.dev382}/.coveragerc +0 -0
  13. {haliax-1.4.dev381 → haliax-1.4.dev382}/.flake8 +0 -0
  14. {haliax-1.4.dev381 → haliax-1.4.dev382}/.github/workflows/publish_dev.yaml +0 -0
  15. {haliax-1.4.dev381 → haliax-1.4.dev382}/.github/workflows/run_pre_commit.yaml +0 -0
  16. {haliax-1.4.dev381 → haliax-1.4.dev382}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  17. {haliax-1.4.dev381 → haliax-1.4.dev382}/.github/workflows/run_tests.yaml +0 -0
  18. {haliax-1.4.dev381 → haliax-1.4.dev382}/.gitignore +0 -0
  19. {haliax-1.4.dev381 → haliax-1.4.dev382}/.playbooks/add-types.md +0 -0
  20. {haliax-1.4.dev381 → haliax-1.4.dev382}/.pre-commit-config.yaml +0 -0
  21. {haliax-1.4.dev381 → haliax-1.4.dev382}/.readthedocs.yaml +0 -0
  22. {haliax-1.4.dev381 → haliax-1.4.dev382}/AGENTS.md +0 -0
  23. {haliax-1.4.dev381 → haliax-1.4.dev382}/CONTRIBUTING.md +0 -0
  24. {haliax-1.4.dev381 → haliax-1.4.dev382}/LICENSE +0 -0
  25. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/api.md +0 -0
  26. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/broadcasting.md +0 -0
  27. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/cheatsheet.md +0 -0
  28. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/css/material.css +0 -0
  29. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/css/mkdocstrings.css +0 -0
  30. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/faq.md +0 -0
  31. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/data_parallel_mesh.png +0 -0
  32. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  33. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_1d.png +0 -0
  34. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_1d_zero.png +0 -0
  35. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_2d.png +0 -0
  36. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  37. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  38. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  39. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  40. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_zero.png +0 -0
  41. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/fp8.md +0 -0
  42. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/index.md +0 -0
  43. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/matmul.md +0 -0
  44. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/partitioning.md +0 -0
  45. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/rearrange.ipynb +0 -0
  46. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/rearrange.md +0 -0
  47. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/requirements.txt +0 -0
  48. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/scan.md +0 -0
  49. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev381 → haliax-1.4.dev382}/docs/vmap.md +0 -0
  51. {haliax-1.4.dev381 → haliax-1.4.dev382}/mkdocs.yml +0 -0
  52. {haliax-1.4.dev381 → haliax-1.4.dev382}/pyproject.toml +0 -0
  53. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/__init__.py +0 -0
  54. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/__init__.py +0 -0
  55. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/compile_utils.py +0 -0
  56. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/dot.py +0 -0
  57. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/einsum.py +0 -0
  58. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/fp8.py +0 -0
  59. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/parsing.py +0 -0
  60. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/rearrange.py +0 -0
  61. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/scan.py +0 -0
  62. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/state_dict.py +0 -0
  63. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/_src/util.py +0 -0
  64. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/axis.py +0 -0
  65. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/core.py +0 -0
  66. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/debug.py +0 -0
  67. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/haxtyping.py +0 -0
  68. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/hof.py +0 -0
  69. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/jax_utils.py +0 -0
  70. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/__init__.py +0 -0
  71. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/activations.py +0 -0
  72. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/attention.py +0 -0
  73. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/dropout.py +0 -0
  74. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/embedding.py +0 -0
  75. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/loss.py +0 -0
  76. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/mlp.py +0 -0
  77. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/normalization.py +0 -0
  78. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/pool.py +0 -0
  79. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/nn/scan.py +0 -0
  80. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/ops.py +0 -0
  81. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/partitioning.py +0 -0
  82. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/quantization.py +0 -0
  83. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/random.py +0 -0
  84. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/specialized_fns.py +0 -0
  85. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/state_dict.py +0 -0
  86. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/tree_util.py +0 -0
  87. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/types.py +0 -0
  88. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/util.py +0 -0
  89. {haliax-1.4.dev381 → haliax-1.4.dev382}/src/haliax/wrap.py +0 -0
  90. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/core_test.py +0 -0
  91. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_attention.py +0 -0
  92. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_axis.py +0 -0
  93. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_conv.py +0 -0
  94. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_debug.py +0 -0
  95. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_dot.py +0 -0
  96. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_dtype_typing.py +0 -0
  97. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_einsum.py +0 -0
  98. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_fp8.py +0 -0
  99. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_hof.py +0 -0
  100. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_int8.py +0 -0
  101. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_namedarray_typing.py +0 -0
  102. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_nn.py +0 -0
  103. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_ops.py +0 -0
  104. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_parsing.py +0 -0
  105. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_partitioning.py +0 -0
  106. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_pool.py +0 -0
  107. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_random.py +0 -0
  108. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_scan.py +0 -0
  109. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_scatter_gather.py +0 -0
  110. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_specialized_fns.py +0 -0
  111. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_state_dict.py +0 -0
  112. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_tree_util.py +0 -0
  113. {haliax-1.4.dev381 → haliax-1.4.dev382}/tests/test_utils.py +0 -0
  114. {haliax-1.4.dev381 → haliax-1.4.dev382}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev381
3
+ Version: 1.4.dev382
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/
@@ -33,14 +33,14 @@ Description-Content-Type: text/markdown
33
33
  <a href="">
34
34
  <img alt="License" src="https://img.shields.io/github/license/stanford-crfm/haliax?color=blue" />
35
35
  </a>
36
- <a href="https://https://pypi.org/project/haliax/">
36
+ <a href="https://pypi.org/project/haliax/">
37
37
  <img alt="PyPI" src="https://img.shields.io/pypi/v/haliax?color=blue" />
38
38
  </a>
39
39
 
40
40
  > *Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me.*<br/>
41
41
  > — Patrick Rothfuss, *The Name of the Wind*
42
42
 
43
- Haliax is a [JAX](https:://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
43
+ Haliax is a [JAX](https://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
44
44
  Named tensors improve the **legibility** and **compositionality** of tensor programs by using named axes instead of positional indices
45
45
  as typically used in NumPy, PyTorch, etc.
46
46
 
@@ -10,14 +10,14 @@
10
10
  <a href="">
11
11
  <img alt="License" src="https://img.shields.io/github/license/stanford-crfm/haliax?color=blue" />
12
12
  </a>
13
- <a href="https://https://pypi.org/project/haliax/">
13
+ <a href="https://pypi.org/project/haliax/">
14
14
  <img alt="PyPI" src="https://img.shields.io/pypi/v/haliax?color=blue" />
15
15
  </a>
16
16
 
17
17
  > *Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me.*<br/>
18
18
  > — Patrick Rothfuss, *The Name of the Wind*
19
19
 
20
- Haliax is a [JAX](https:://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
20
+ Haliax is a [JAX](https://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
21
21
  Named tensors improve the **legibility** and **compositionality** of tensor programs by using named axes instead of positional indices
22
22
  as typically used in NumPy, PyTorch, etc.
23
23
 
@@ -1,7 +1,7 @@
1
1
  # Indexing and Slicing
2
2
 
3
3
  Haliax supports Numpy-style indexing, including so-called [Advanced Indexing](https://numpy.org/doc/stable/user/basics.indexing.html#advanced-indexing),
4
- though the syntax is necessarily different. Most forms of indexing are supporting, except we don't support indexing with
4
+ though the syntax is necessarily different. Most forms of indexing are supported, except we don't support indexing with
5
5
  booleans right now. (JAX doesn't support indexing with non-constant bool arrays anyway,
6
6
  so I don't think it's worth the effort to implement it in Haliax.)
7
7
 
@@ -6,7 +6,7 @@
6
6
  Haliax provides a small number of neural network modules that are compatible with Equinox, though
7
7
  they naturally all use [haliax.NamedArray][]. (We welcome PRs for more modules! Nothing too exotic though.)
8
8
 
9
- The most interesting of these modules is [haliax.nn.Stacked][], which allows you to create homogenous "stacks"
9
+ The most interesting of these modules is [haliax.nn.Stacked][], which allows you to create homogeneous "stacks"
10
10
  of the same module (e.g. transformer blocks), which is a common pattern in deep learning.
11
11
 
12
12
  ### Linear
@@ -226,7 +226,7 @@ any Axis members to match the new shape.
226
226
  ::: haliax.state_dict.save_state_dict
227
227
  ::: haliax.state_dict.load_state_dict
228
228
 
229
- ### Converting betweewn State Dicts and Modules
229
+ ### Converting between State Dicts and Modules
230
230
 
231
231
  ::: haliax.state_dict.from_state_dict
232
232
  ::: haliax.state_dict.to_state_dict
@@ -1,4 +1,4 @@
1
- from haliax import NamedArrayfrom haliax import NamedArray
1
+ from haliax import NamedArray
2
2
 
3
3
  # NamedArray Type Annotations
4
4
 
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev382"
@@ -213,7 +213,7 @@ class Conv(_ConvBase):
213
213
  return x
214
214
 
215
215
  def _do_conv(self, inputs):
216
- # _do_conv expects there ot be a single __batch__ dimension
216
+ # _do_conv expects there to be a single __batch__ dimension
217
217
  output_axes = _compute_output_axes(inputs, "__batch__", self.In, self.Out)
218
218
 
219
219
  batch_index = _index_of_name(inputs.axes, "__batch__")
@@ -143,7 +143,7 @@ class MoELinear(eqx.Module):
143
143
  Experts: AxisSpec = eqx.field(static=True)
144
144
  In: Axis = eqx.field(static=True)
145
145
  Out: Axis = eqx.field(static=True)
146
- # TODO: support quanitization for ragged_dot?
146
+ # TODO: support quantization for ragged_dot?
147
147
  # dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
148
148
 
149
149
  use_gmm: bool = eqx.field(static=True)
@@ -293,7 +293,7 @@ def test_examples():
293
293
  r = einops_rearrange(z, "{B (H: h1 h) (W: w1 w) C} -> (B: B h1 w1) ... (C: C h w) ", h1=2, w1=2)
294
294
  assert r.axes == (Axis("B", B.size * 2 * 2), D, Axis("C", C.size * sH.size * sW.size))
295
295
  # unet attention reordering:
296
- # postional: (qkv heads c) h w -> qkv heads c (h w)
296
+ # positional: (qkv heads c) h w -> qkv heads c (h w)
297
297
  # named: { (embed: qkv heads c) h w } -> qkv heads c (pos: h w)
298
298
  Embed = Axis("embed", 3 * 4 * C.size)
299
299
  attn = hax.random.randint(PRNGKey(0), (Embed, H, W), 0, 255)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev381"
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes