haliax 1.4.dev362__tar.gz → 1.4.dev364__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 (111) hide show
  1. {haliax-1.4.dev362 → haliax-1.4.dev364}/PKG-INFO +1 -1
  2. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/api.md +0 -1
  3. haliax-1.4.dev364/docs/typing.md +61 -0
  4. {haliax-1.4.dev362 → haliax-1.4.dev364}/mkdocs.yml +1 -0
  5. {haliax-1.4.dev362 → haliax-1.4.dev364}/pyproject.toml +16 -0
  6. haliax-1.4.dev364/src/haliax/__about__.py +1 -0
  7. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/__init__.py +11 -3
  8. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/dot.py +3 -2
  9. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/axis.py +306 -147
  10. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/core.py +343 -115
  11. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/hof.py +6 -3
  12. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/activations.py +1 -1
  13. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/attention.py +3 -52
  14. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/conv.py +14 -5
  15. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/embedding.py +2 -2
  16. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/pool.py +10 -11
  17. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/partitioning.py +2 -2
  18. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/random.py +35 -103
  19. haliax-1.4.dev364/src/haliax/typing.py +88 -0
  20. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/wrap.py +7 -6
  21. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/core_test.py +44 -4
  22. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_attention.py +0 -22
  23. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_axis.py +106 -15
  24. haliax-1.4.dev364/tests/test_dtype_typing.py +41 -0
  25. haliax-1.4.dev364/tests/test_namedarray_typing.py +64 -0
  26. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_random.py +35 -21
  27. haliax-1.4.dev362/src/haliax/__about__.py +0 -1
  28. {haliax-1.4.dev362 → haliax-1.4.dev364}/.coveragerc +0 -0
  29. {haliax-1.4.dev362 → haliax-1.4.dev364}/.flake8 +0 -0
  30. {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/publish_dev.yaml +0 -0
  31. {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/run_pre_commit.yaml +0 -0
  32. {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  33. {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/run_tests.yaml +0 -0
  34. {haliax-1.4.dev362 → haliax-1.4.dev364}/.gitignore +0 -0
  35. {haliax-1.4.dev362 → haliax-1.4.dev364}/.pre-commit-config.yaml +0 -0
  36. {haliax-1.4.dev362 → haliax-1.4.dev364}/.readthedocs.yaml +0 -0
  37. {haliax-1.4.dev362 → haliax-1.4.dev364}/CONTRIBUTING.md +0 -0
  38. {haliax-1.4.dev362 → haliax-1.4.dev364}/LICENSE +0 -0
  39. {haliax-1.4.dev362 → haliax-1.4.dev364}/README.md +0 -0
  40. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/broadcasting.md +0 -0
  41. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/cheatsheet.md +0 -0
  42. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/css/material.css +0 -0
  43. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/css/mkdocstrings.css +0 -0
  44. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/faq.md +0 -0
  45. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh.png +0 -0
  46. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  47. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_1d.png +0 -0
  48. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_1d_zero.png +0 -0
  49. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d.png +0 -0
  50. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  51. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  52. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  53. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  54. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_zero.png +0 -0
  55. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/fp8.md +0 -0
  56. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/index.md +0 -0
  57. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/indexing.md +0 -0
  58. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/matmul.md +0 -0
  59. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/nn.md +0 -0
  60. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/partitioning.md +0 -0
  61. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/rearrange.ipynb +0 -0
  62. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/rearrange.md +0 -0
  63. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/requirements.txt +0 -0
  64. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/scan.md +0 -0
  65. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/state-dict.md +0 -0
  66. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/tutorial.md +0 -0
  67. {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/vmap.md +0 -0
  68. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/__init__.py +0 -0
  69. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/compile_utils.py +0 -0
  70. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/einsum.py +0 -0
  71. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/fp8.py +0 -0
  72. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/parsing.py +0 -0
  73. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/rearrange.py +0 -0
  74. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/scan.py +0 -0
  75. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/state_dict.py +0 -0
  76. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/util.py +0 -0
  77. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/debug.py +0 -0
  78. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/jax_utils.py +0 -0
  79. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/__init__.py +0 -0
  80. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/dropout.py +0 -0
  81. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/linear.py +0 -0
  82. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/loss.py +0 -0
  83. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/mlp.py +0 -0
  84. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/normalization.py +0 -0
  85. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/scan.py +0 -0
  86. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/ops.py +0 -0
  87. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/quantization.py +0 -0
  88. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/specialized_fns.py +0 -0
  89. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/state_dict.py +0 -0
  90. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/tree_util.py +0 -0
  91. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/types.py +0 -0
  92. {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/util.py +0 -0
  93. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_conv.py +0 -0
  94. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_debug.py +0 -0
  95. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_dot.py +0 -0
  96. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_einsum.py +0 -0
  97. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_fp8.py +0 -0
  98. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_hof.py +0 -0
  99. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_int8.py +0 -0
  100. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_nn.py +0 -0
  101. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_ops.py +0 -0
  102. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_parsing.py +0 -0
  103. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_partitioning.py +0 -0
  104. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_pool.py +0 -0
  105. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_rearrange.py +0 -0
  106. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_scan.py +0 -0
  107. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_scatter_gather.py +0 -0
  108. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_specialized_fns.py +0 -0
  109. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_state_dict.py +0 -0
  110. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_tree_util.py +0 -0
  111. {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev362
3
+ Version: 1.4.dev364
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/
@@ -35,7 +35,6 @@ Occasionally, an axis size can be inferred in some circumstances but not others.
35
35
  ::: haliax.axis.eliminate_axes
36
36
  ::: haliax.axis.without_axes
37
37
  ::: haliax.axis.selects_axis
38
- ::: haliax.axis.overlapping_axes
39
38
  ::: haliax.axis.is_axis_compatible
40
39
 
41
40
 
@@ -0,0 +1,61 @@
1
+ # NamedArray Type Annotations
2
+
3
+ Haliax supports a lightweight syntax for specifying the axes of a `NamedArray`
4
+ in type annotations. Internally, `Named[...]` expands to
5
+ `Annotated[NamedArray, axes]`, so it works well with static type checkers like
6
+ ``mypy``. The syntax mirrors normal indexing with axis names. Some examples:
7
+
8
+ ```python
9
+ from haliax import Named
10
+
11
+ arr: Named["batch", "embed"]
12
+ arr: Named["batch embed ..."] # starts with these axes
13
+ arr: Named["... embed"] # ends with this axis
14
+ arr: Named["batch ... embed"] # contains these axes in order
15
+ arr: Named[{"batch", "embed"}] # has exactly these axes, order ignored
16
+ arr: Named[{"batch", "embed", ...}] # has at least these axes
17
+ ```
18
+
19
+ At runtime you can verify that a `NamedArray` conforms to a particular
20
+ annotation using `matches_axes`:
21
+
22
+ ```python
23
+ if not arr.matches_axes(Named["batch embed ..."]):
24
+ raise ValueError("unexpected axes")
25
+ ```
26
+
27
+ ## DType-aware annotations
28
+
29
+ Sometimes it is useful to express both the axes **and** the dtype in the type
30
+ annotation. The :mod:`haliax.typing` module defines symbolic types for all of
31
+ JAX's common dtypes that can be indexed just like ``Named``. In documentation
32
+ examples we'll use ``import haliax.typing as ht``:
33
+
34
+ ```python
35
+ import haliax.typing as ht
36
+
37
+ def foo(x: ht.f32["batch"]):
38
+ ...
39
+
40
+ def bar(x: ht.i32["batch"]):
41
+ ...
42
+ ```
43
+
44
+ For convenience the module also provides aggregate categories ``Float``,
45
+ ``Complex``, ``Int`` and ``UInt`` that match any floating point, complex,
46
+ signed integer or unsigned integer dtype respectively:
47
+
48
+ ```python
49
+ def baz(x: ht.Float["batch"]):
50
+ ...
51
+ ```
52
+
53
+ At runtime ``matches_axes`` also checks the dtype when one is present:
54
+
55
+ ```python
56
+ from haliax import Axis, zeros
57
+ import haliax.typing as ht
58
+
59
+ arr = zeros({"batch": 4})
60
+ assert arr.matches_axes(ht.f32["batch"]) # dtype and axes both match
61
+ ```
@@ -90,6 +90,7 @@ nav:
90
90
  - Indexing and Slicing: 'indexing.md'
91
91
  - Rearrange: 'rearrange.md'
92
92
  - Matrix Multiplication: 'matmul.md'
93
+ - Type Annotations: 'typing.md'
93
94
  - Higher Order Functions:
94
95
  - Scan and Fold: 'scan.md'
95
96
  - Vectorization: 'vmap.md'
@@ -71,3 +71,19 @@ src_paths = ["src", "tests"]
71
71
  "Homepage" = "https://github.com/stanford-crfm/haliax"
72
72
  "Bug Tracker" = "https://github.com/stanford-crfm/haliax/issues/"
73
73
  "Documentation" = "https://haliax.readthedocs.io/en/latest/"
74
+
75
+
76
+ [tool.coverage.report]
77
+ exclude_also = [
78
+ "def __repr__",
79
+ "if self.debug:",
80
+ "if settings.DEBUG",
81
+ "raise AssertionError",
82
+ "raise NotImplementedError",
83
+ "if 0:",
84
+ "if __name__ == .__main__.:",
85
+ "if TYPE_CHECKING:",
86
+ "class .*\\bProtocol\\):",
87
+ "@(abc\\.)?abstractmethod",
88
+ "[.][.][.]"
89
+ ]
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev364"
@@ -29,6 +29,7 @@ from .axis import (
29
29
  AxisSpec,
30
30
  axis_name,
31
31
  axis_size,
32
+ axis_spec_to_tuple,
32
33
  concat_axes,
33
34
  dblock,
34
35
  ds,
@@ -38,9 +39,13 @@ from .axis import (
38
39
  replace_axis,
39
40
  resolve_axis,
40
41
  selects_axis,
42
+ to_jax_shape,
41
43
  )
42
44
  from .core import (
45
+ Named,
43
46
  NamedArray,
47
+ NamedArrayAxes,
48
+ NamedArrayAxesSpec,
44
49
  NamedOrNumeric,
45
50
  are_shape_checks_enabled,
46
51
  broadcast_arrays,
@@ -102,8 +107,8 @@ def full(shape: AxisSpec, fill_value: T, dtype: Optional[DTypeLike] = None) -> N
102
107
  if isinstance(shape, Axis):
103
108
  return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
104
109
  else:
105
- x_shape = tuple(x.size for x in shape)
106
- return NamedArray(jnp.full(shape=x_shape, fill_value=fill_value, dtype=dtype), tuple(shape))
110
+ x_shape = to_jax_shape(shape)
111
+ return NamedArray(jnp.full(shape=x_shape, fill_value=fill_value, dtype=dtype), shape)
107
112
 
108
113
 
109
114
  def zeros_like(a: NamedArray, dtype=None) -> NamedArray:
@@ -152,7 +157,7 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
152
157
 
153
158
  arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=size) * step + start
154
159
  arr = arr.reshape(to_jax_shape(axis))
155
- return NamedArray(arr, ensure_tuple(axis))
160
+ return NamedArray(arr, axis_spec_to_tuple(axis))
156
161
 
157
162
 
158
163
  # TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
@@ -926,6 +931,9 @@ __all__ = [
926
931
  "make_axes",
927
932
  "axis_name",
928
933
  "axis_size",
934
+ "NamedArrayAxesSpec",
935
+ "NamedArrayAxes",
936
+ "Named",
929
937
  "NamedArray",
930
938
  "broadcast_to",
931
939
  "broadcast_axis",
@@ -12,6 +12,7 @@ from haliax.axis import (
12
12
  AxisSelection,
13
13
  PartialAxisSpec,
14
14
  axis_name,
15
+ axis_spec_to_shape_dict,
15
16
  eliminate_axes,
16
17
  rearrange_for_partial_order,
17
18
  union_axes,
@@ -140,8 +141,8 @@ def dot(
140
141
  if axis is None:
141
142
  jax_str = f"contract {', '.join(axis_name(ax) for ax in all_axes)} -> <scalar>"
142
143
  else:
143
- axis = ensure_tuple(axis)
144
- jax_str = f"contract {', '.join(axis_name(ax) for ax in axis)} -> {', '.join(a.name for a in output_axes)}"
144
+ axis = axis_spec_to_shape_dict(axis)
145
+ jax_str = f"contract {', '.join(axis)} -> {', '.join(a.name for a in output_axes)}"
145
146
 
146
147
  with jax.named_scope(jax_str):
147
148
  output = _jittable_dg_einsum(