haliax 1.4.dev363__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.dev363 → haliax-1.4.dev364}/PKG-INFO +1 -1
  2. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/api.md +0 -1
  3. {haliax-1.4.dev363 → haliax-1.4.dev364}/pyproject.toml +16 -0
  4. haliax-1.4.dev364/src/haliax/__about__.py +1 -0
  5. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/__init__.py +5 -3
  6. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/dot.py +3 -2
  7. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/axis.py +306 -147
  8. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/core.py +170 -114
  9. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/hof.py +6 -3
  10. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/activations.py +1 -1
  11. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/attention.py +3 -52
  12. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/conv.py +14 -5
  13. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/embedding.py +2 -2
  14. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/pool.py +10 -11
  15. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/partitioning.py +2 -2
  16. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/random.py +35 -103
  17. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/wrap.py +7 -6
  18. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/core_test.py +44 -4
  19. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_attention.py +0 -22
  20. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_axis.py +106 -15
  21. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_random.py +35 -21
  22. haliax-1.4.dev363/src/haliax/__about__.py +0 -1
  23. {haliax-1.4.dev363 → haliax-1.4.dev364}/.coveragerc +0 -0
  24. {haliax-1.4.dev363 → haliax-1.4.dev364}/.flake8 +0 -0
  25. {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/publish_dev.yaml +0 -0
  26. {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/run_pre_commit.yaml +0 -0
  27. {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  28. {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/run_tests.yaml +0 -0
  29. {haliax-1.4.dev363 → haliax-1.4.dev364}/.gitignore +0 -0
  30. {haliax-1.4.dev363 → haliax-1.4.dev364}/.pre-commit-config.yaml +0 -0
  31. {haliax-1.4.dev363 → haliax-1.4.dev364}/.readthedocs.yaml +0 -0
  32. {haliax-1.4.dev363 → haliax-1.4.dev364}/CONTRIBUTING.md +0 -0
  33. {haliax-1.4.dev363 → haliax-1.4.dev364}/LICENSE +0 -0
  34. {haliax-1.4.dev363 → haliax-1.4.dev364}/README.md +0 -0
  35. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/broadcasting.md +0 -0
  36. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/cheatsheet.md +0 -0
  37. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/css/material.css +0 -0
  38. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/css/mkdocstrings.css +0 -0
  39. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/faq.md +0 -0
  40. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh.png +0 -0
  41. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  42. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_1d.png +0 -0
  43. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_1d_zero.png +0 -0
  44. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d.png +0 -0
  45. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  46. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  47. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  48. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  49. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_zero.png +0 -0
  50. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/fp8.md +0 -0
  51. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/index.md +0 -0
  52. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/indexing.md +0 -0
  53. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/matmul.md +0 -0
  54. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/nn.md +0 -0
  55. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/partitioning.md +0 -0
  56. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/rearrange.ipynb +0 -0
  57. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/rearrange.md +0 -0
  58. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/requirements.txt +0 -0
  59. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/scan.md +0 -0
  60. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/state-dict.md +0 -0
  61. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/tutorial.md +0 -0
  62. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/typing.md +0 -0
  63. {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/vmap.md +0 -0
  64. {haliax-1.4.dev363 → haliax-1.4.dev364}/mkdocs.yml +0 -0
  65. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/__init__.py +0 -0
  66. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/compile_utils.py +0 -0
  67. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/einsum.py +0 -0
  68. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/fp8.py +0 -0
  69. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/parsing.py +0 -0
  70. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/rearrange.py +0 -0
  71. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/scan.py +0 -0
  72. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/state_dict.py +0 -0
  73. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/util.py +0 -0
  74. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/debug.py +0 -0
  75. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/jax_utils.py +0 -0
  76. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/__init__.py +0 -0
  77. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/dropout.py +0 -0
  78. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/linear.py +0 -0
  79. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/loss.py +0 -0
  80. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/mlp.py +0 -0
  81. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/normalization.py +0 -0
  82. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/scan.py +0 -0
  83. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/ops.py +0 -0
  84. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/quantization.py +0 -0
  85. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/typing.py +0 -0
  90. {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/util.py +0 -0
  91. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_conv.py +0 -0
  92. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_debug.py +0 -0
  93. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_dot.py +0 -0
  94. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_dtype_typing.py +0 -0
  95. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_einsum.py +0 -0
  96. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_fp8.py +0 -0
  97. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_hof.py +0 -0
  98. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_int8.py +0 -0
  99. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_namedarray_typing.py +0 -0
  100. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_nn.py +0 -0
  101. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_ops.py +0 -0
  102. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_parsing.py +0 -0
  103. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_partitioning.py +0 -0
  104. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_pool.py +0 -0
  105. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_rearrange.py +0 -0
  106. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_scan.py +0 -0
  107. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_scatter_gather.py +0 -0
  108. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_specialized_fns.py +0 -0
  109. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_state_dict.py +0 -0
  110. {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_tree_util.py +0 -0
  111. {haliax-1.4.dev363 → 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.dev363
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
 
@@ -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,6 +39,7 @@ 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 (
43
45
  Named,
@@ -105,8 +107,8 @@ def full(shape: AxisSpec, fill_value: T, dtype: Optional[DTypeLike] = None) -> N
105
107
  if isinstance(shape, Axis):
106
108
  return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
107
109
  else:
108
- x_shape = tuple(x.size for x in shape)
109
- 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)
110
112
 
111
113
 
112
114
  def zeros_like(a: NamedArray, dtype=None) -> NamedArray:
@@ -155,7 +157,7 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
155
157
 
156
158
  arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=size) * step + start
157
159
  arr = arr.reshape(to_jax_shape(axis))
158
- return NamedArray(arr, ensure_tuple(axis))
160
+ return NamedArray(arr, axis_spec_to_tuple(axis))
159
161
 
160
162
 
161
163
  # TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
@@ -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(