haliax 1.4.dev363__tar.gz → 1.4.dev365__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.dev363 → haliax-1.4.dev365}/.pre-commit-config.yaml +16 -16
  2. haliax-1.4.dev365/AGENTS.md +79 -0
  3. {haliax-1.4.dev363 → haliax-1.4.dev365}/PKG-INFO +1 -1
  4. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/api.md +0 -1
  5. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/nn.md +0 -1
  6. haliax-1.4.dev365/docs/typing.md +100 -0
  7. {haliax-1.4.dev363 → haliax-1.4.dev365}/pyproject.toml +21 -0
  8. haliax-1.4.dev365/src/haliax/__about__.py +1 -0
  9. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/__init__.py +21 -19
  10. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/dot.py +3 -4
  11. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/fp8.py +0 -1
  12. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/state_dict.py +3 -4
  13. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/axis.py +306 -147
  14. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/core.py +172 -129
  15. haliax-1.4.dev365/src/haliax/haxtyping.py +138 -0
  16. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/hof.py +7 -4
  17. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/activations.py +1 -1
  18. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/attention.py +3 -53
  19. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/conv.py +14 -5
  20. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/embedding.py +2 -4
  21. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/pool.py +10 -11
  22. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/scan.py +3 -3
  23. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/partitioning.py +3 -3
  24. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/random.py +35 -106
  25. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/wrap.py +7 -9
  26. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/core_test.py +44 -4
  27. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_attention.py +0 -23
  28. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_axis.py +106 -15
  29. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_dot.py +1 -1
  30. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_dtype_typing.py +1 -1
  31. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_namedarray_typing.py +7 -7
  32. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_random.py +35 -21
  33. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_scatter_gather.py +0 -1
  34. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_tree_util.py +0 -1
  35. haliax-1.4.dev363/docs/typing.md +0 -61
  36. haliax-1.4.dev363/src/haliax/__about__.py +0 -1
  37. haliax-1.4.dev363/src/haliax/typing.py +0 -88
  38. {haliax-1.4.dev363 → haliax-1.4.dev365}/.coveragerc +0 -0
  39. {haliax-1.4.dev363 → haliax-1.4.dev365}/.flake8 +0 -0
  40. {haliax-1.4.dev363 → haliax-1.4.dev365}/.github/workflows/publish_dev.yaml +0 -0
  41. {haliax-1.4.dev363 → haliax-1.4.dev365}/.github/workflows/run_pre_commit.yaml +0 -0
  42. {haliax-1.4.dev363 → haliax-1.4.dev365}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  43. {haliax-1.4.dev363 → haliax-1.4.dev365}/.github/workflows/run_tests.yaml +0 -0
  44. {haliax-1.4.dev363 → haliax-1.4.dev365}/.gitignore +0 -0
  45. {haliax-1.4.dev363 → haliax-1.4.dev365}/.readthedocs.yaml +0 -0
  46. {haliax-1.4.dev363 → haliax-1.4.dev365}/CONTRIBUTING.md +0 -0
  47. {haliax-1.4.dev363 → haliax-1.4.dev365}/LICENSE +0 -0
  48. {haliax-1.4.dev363 → haliax-1.4.dev365}/README.md +0 -0
  49. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/broadcasting.md +0 -0
  50. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/cheatsheet.md +0 -0
  51. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/css/material.css +0 -0
  52. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/css/mkdocstrings.css +0 -0
  53. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/faq.md +0 -0
  54. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/data_parallel_mesh.png +0 -0
  55. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  56. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_1d.png +0 -0
  57. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_1d_zero.png +0 -0
  58. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_2d.png +0 -0
  59. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  60. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  61. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  62. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  63. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/figures/device_mesh_2d_zero.png +0 -0
  64. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/fp8.md +0 -0
  65. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/index.md +0 -0
  66. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/indexing.md +0 -0
  67. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/matmul.md +0 -0
  68. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/partitioning.md +0 -0
  69. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/rearrange.ipynb +0 -0
  70. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/rearrange.md +0 -0
  71. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/requirements.txt +0 -0
  72. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/scan.md +0 -0
  73. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/state-dict.md +1 -1
  74. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/tutorial.md +0 -0
  75. {haliax-1.4.dev363 → haliax-1.4.dev365}/docs/vmap.md +0 -0
  76. {haliax-1.4.dev363 → haliax-1.4.dev365}/mkdocs.yml +0 -0
  77. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/__init__.py +0 -0
  78. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/compile_utils.py +0 -0
  79. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/einsum.py +0 -0
  80. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/parsing.py +0 -0
  81. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/rearrange.py +0 -0
  82. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/scan.py +0 -0
  83. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/_src/util.py +0 -0
  84. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/debug.py +0 -0
  85. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/jax_utils.py +0 -0
  86. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/__init__.py +0 -0
  87. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/dropout.py +0 -0
  88. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/linear.py +0 -0
  89. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/loss.py +0 -0
  90. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/mlp.py +0 -0
  91. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/nn/normalization.py +0 -0
  92. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/ops.py +0 -0
  93. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/quantization.py +0 -0
  94. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/specialized_fns.py +0 -0
  95. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/state_dict.py +0 -0
  96. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/tree_util.py +0 -0
  97. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/types.py +0 -0
  98. {haliax-1.4.dev363 → haliax-1.4.dev365}/src/haliax/util.py +0 -0
  99. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_conv.py +0 -0
  100. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_debug.py +0 -0
  101. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_einsum.py +0 -0
  102. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_fp8.py +0 -0
  103. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_hof.py +0 -0
  104. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_int8.py +0 -0
  105. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_nn.py +0 -0
  106. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_ops.py +0 -0
  107. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_parsing.py +0 -0
  108. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_partitioning.py +0 -0
  109. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_pool.py +0 -0
  110. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_rearrange.py +0 -0
  111. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_scan.py +0 -0
  112. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_specialized_fns.py +0 -0
  113. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_state_dict.py +0 -0
  114. {haliax-1.4.dev363 → haliax-1.4.dev365}/tests/test_utils.py +0 -0
@@ -2,7 +2,7 @@
2
2
  # See https://pre-commit.com/hooks.html for more hooks
3
3
  exclude: ".git"
4
4
  default_stages:
5
- - commit
5
+ - pre-commit
6
6
  fail_fast: true
7
7
 
8
8
  repos:
@@ -16,24 +16,24 @@ repos:
16
16
  - id: check-merge-conflict
17
17
  - id: check-added-large-files
18
18
 
19
- - repo: https://github.com/psf/black
20
- rev: 22.3.0
21
- hooks:
22
- - id: black
23
-
24
- - repo: https://github.com/timothycrosley/isort
25
- rev: 5.11.5
26
- hooks:
27
- - id: isort
19
+ - repo: https://github.com/astral-sh/ruff-pre-commit
20
+ rev: v0.11.10
21
+ hooks:
22
+ - id: ruff
23
+ args: [ --fix, --exit-non-zero-on-fix ]
28
24
 
29
- - repo: https://github.com/PyCQA/flake8
30
- rev: 3.9.2
31
- hooks:
32
- - id: flake8
33
- additional_dependencies: [flake8-isort]
25
+ #- repo: local
26
+ # hooks:
27
+ # - id: ty-check
28
+ # name: ty-check
29
+ # language: python
30
+ # entry: ty check
31
+ # pass_filenames: false
32
+ # args: [--python=.venv/]
33
+ # additional_dependencies: [ty]
34
34
 
35
35
  - repo: https://github.com/pre-commit/mirrors-mypy
36
- rev: 'v1.5.1'
36
+ rev: 'v1.16.1'
37
37
  hooks:
38
38
  - id: mypy
39
39
  args: [--ignore-missing-imports, --check-untyped-defs]
@@ -0,0 +1,79 @@
1
+ # Haliax LLM Agent Guidelines
2
+
3
+ This document summarizes important conventions for contributing code or documentation to the Haliax
4
+ repository. Follow these notes when implementing new features or fixing bugs.
5
+
6
+ ## General Guidelines
7
+
8
+ * **Get better.** Whenever you discover something missing from these guidelines, or the requester
9
+ suggests a better way to do something, please update this document. The goal is to make it easier for
10
+ everyone to contribute and maintain the codebase. Generally speaking, you should add bullets or new sections.
11
+ Be sure to do this when directed to. For example, if directed that you should never relax tolerances in
12
+ floating point tests, add that to the list.
13
+ * **Playbooks.** Sometimes, there are repeatable tasks (e.g. porting models) for which we follow a standard set of steps.
14
+ Please reference `.playbooks/` to see what playbooks are available, or see the list below. If you want to add a playbook
15
+ write a markdown doc named e.g. `.playbooks/port-models.md` and add a pointer to it in the list below.
16
+
17
+ ## Playbook
18
+
19
+ - At the moment, there are no playbooks available. If you have a repeatable task that you think
20
+ should be documented, please create a new markdown file in `.playbooks/` and add it to the list above.
21
+
22
+ ## Code Style
23
+
24
+ * **Python version**: the project targets Python >=3.10.
25
+ * **Formatting and Linting**: We use `ruff` via `pre-commit`.
26
+ * **Typing**: the code base uses `mypy` for static type checking. `mypy` is run by pre‑commit and the
27
+ configuration is found in `pyproject.toml`.
28
+ * **Run `pre-commit run --all-files`** before committing. The CI workflows run the same checks.
29
+ * **Doc Strings**: All public functions, classes, and modules should have docstrings, unless
30
+ their purpose is painfully obvious. Use
31
+ [Google style](https://google.github.io/styleguide/pyguide.html#38-comments-and-docstrings) for
32
+ consistency.
33
+ * **Commenting**: Use comments to explain why something is done a certain way, especially if it is not
34
+ immediately obvious. Avoid commenting on every line of code; focus on the intent and purpose of
35
+ complex logic. Demarcating logical groups of code with comments is encouraged, unless it is better
36
+ to refactor the code into smaller functions or classes.
37
+ * **Mkdocs**: We use [Mkdocs](https://www.mkdocs.org/) for documentation. The main documentation is in
38
+ the `docs` directory. Use Markdown for writing docs, and follow the existing structure. When linking to
39
+ symbols, prefer using mkdocs-style links (e.g. With a custom title: `[full.path.object2][]` or
40
+ `[Object 1][full.path.object1]`)
41
+ * **Documentation**: When adding new features, ensure that the documentation is updated accordingly.
42
+ This includes updating the Mkdocs files and any relevant docstrings. If you add a new module or
43
+ significant functionality, consider adding a dedicated section in the documentation.
44
+
45
+ ## Testing
46
+
47
+ * Tests are executed with `pytest`. The default workflow runs
48
+ `pytest tests -m "not entry and not slow and not ray"`.
49
+ * In general, never relax tolerances in floating point tests unless specifically discussed with the
50
+ team. Use `assert_allclose` with appropriate tolerances for numerical comparisons. We typically use
51
+ 1e-4 for more complex modules, and 1e-5 for simpler ones.
52
+ * Tests should be reasonably fast. Mark long-running tests with @pytest.mark.slow so they are excluded from the default suite.
53
+ * Always mark tests that depend on pytorch with `@skip_if_no_torch` to ensure they are skipped
54
+ when PyTorch is not available. This is particularly important for tests that require PyTorch-specific
55
+ functionality.
56
+
57
+
58
+ ## Design Preferences
59
+
60
+ * **Generic code**: many utilities are written with Python generics and dataclasses. Where possible,
61
+ write reusable functions or classes that operate over TypeVars instead of hard coding concrete types.
62
+ * **Configurations**: configuration files are dataclasses loaded via `draccus`. Keep configs
63
+ declarative and typed.
64
+ * **Reproducibility**: Levanter aims for deterministic training where possible. Avoid sources of
65
+ nondeterminism unless explicitly required.
66
+ * Prefer Stacked with fold or scan over writing custom loops, for better compile times and gradient checkpointing support
67
+
68
+ ## Library conventions
69
+ - Haliax revolves around `NamedArray` and explicit `Axis` objects. Prefer APIs that accept
70
+ axes or axis names rather than hard‑coding positional dimensions.
71
+ - Utilities should be written so they work with arbitrary axis names. Avoid relying on
72
+ fixed axis orders when possible.
73
+ - Use the provided modules in `haliax.nn` or Equinox when building neural network layers.
74
+ - Type annotations can use named shapes shorthand provided in `haliax.haxtyping`: `ht.f32[NamedArray, "batch"]`
75
+ for a float32 array with a "batch" axis, or `ht.Float[NamedArray, "batch"]` for any floating point dtype.
76
+
77
+ ## Documentation
78
+ - Public functions and modules require docstrings. If behavior is non‑obvious,
79
+ add examples in `docs/`.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev363
3
+ Version: 1.4.dev365
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
 
@@ -55,7 +55,6 @@ We don't provide an explicit attention module, but we do provide an attention fu
55
55
 
56
56
  :::haliax.nn.attention.dot_product_attention
57
57
  :::haliax.nn.attention.dot_product_attention_weights
58
- :::haliax.nn.attention.self_attention
59
58
 
60
59
  ### Masks
61
60
  ::: haliax.nn.attention.causal_mask
@@ -0,0 +1,100 @@
1
+ from haliax import NamedArrayfrom haliax import NamedArray
2
+
3
+ # NamedArray Type Annotations
4
+
5
+ Haliax supports an extension to [`jaxtyping`](https://docs.kidger.site/jaxtyping/)
6
+ that allows you to annotate functions and methods that take or return
7
+ [`NamedArray`][haliax.core.NamedArray] objects. If you are familiar with
8
+ [`jaxtyping`](https://docs.kidger.site/jaxtyping/), the syntax is very similar.
9
+ In fact, for non-NamedArrays, it is exactly the same.
10
+
11
+ ```python
12
+ from haliax import NamedArray
13
+ import haliax.haxtyping as ht
14
+
15
+ def foo(x: ht.Float[NamedArray, "batch embed ..."]):
16
+ ...
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(Float[NamedArray, "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.haxtyping as ht
36
+
37
+ def foo(x: ht.f32[NamedArray, "batch"]):
38
+ ...
39
+
40
+ def bar(x: ht.i32[NamedArray, "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[NamedArray, "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.haxtyping as ht
58
+
59
+ arr = zeros({"batch": 4})
60
+ assert arr.matches_axes(ht.f32["batch"]) # dtype and axes both match
61
+ ```
62
+
63
+ ## FAQ
64
+
65
+ ### Why not use `NamedArray` directly in type annotations?
66
+
67
+ Using `NamedArray` directly in type annotations doesn't work well with
68
+ type checkers like `mypy` or `pyright`. These tools expect types to be
69
+ subscripted with other types or forward references (which are strings).
70
+ Using `NamedArray` directly would lead to type errors.
71
+
72
+ ### Why not use `jaxtyping` directly?
73
+
74
+ While `jaxtyping` is a powerful library for type annotations in JAX, it does not
75
+ support `NamedArray` objects directly. The `haliax.haxtyping` module extends
76
+ `jaxtyping` to include `NamedArray` support, allowing you to annotate functions
77
+ and methods that take or return `NamedArray` objects with specific axes and dtypes.
78
+
79
+ ### Why do I have to specify the `NamedArray` type in the annotation?
80
+
81
+ I hate this, but it's the only way to get type checkers like `mypy` and `pyright` to understand that the type is
82
+ a `NamedArray`. Underneath the hood, during type checking, `jaxtyping.Float` (and `haxtyping.Float`) are
83
+ essentially type aliases of [`Annotated`](https://docs.python.org/3/library/typing.html#typing.Annotated)
84
+ with the `NamedArray` type. There's no other way I could find to get type checkers to understand that the type is a
85
+ `NamedArray` or to accept strings like `"batch embed ..."` as valid type annotations.
86
+
87
+ ### How do I use single axes in type annotations with flake or ruff.
88
+
89
+ Like `jaxtyping`, you need to prepend a space before the axis name to use single axes in type annotations with
90
+ flake or ruff. For example, to use a single axis named `batch`, you would write:
91
+
92
+ ```python
93
+ def foo(x: ht.Float[NamedArray, " batch"]):
94
+ ...
95
+ ```
96
+
97
+ Then suppress F722 in your linter to suppress that error.
98
+
99
+ See the [jaxtyping documentation](https://docs.kidger.site/jaxtyping/faq/#flake8-or-ruff-are-throwing-an-error) for more
100
+ details on the workaround.
@@ -71,3 +71,24 @@ 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
+ [tool.ruff.lint]
76
+ ignore = [ "E203", "E501", "W605", "F821", "E266", "F722", "E731", "E741" ]
77
+
78
+ [tool.setuptools.package-data]
79
+ "haliax" = ["*.pyi"]
80
+
81
+ [tool.coverage.report]
82
+ exclude_also = [
83
+ "def __repr__",
84
+ "if self.debug:",
85
+ "if settings.DEBUG",
86
+ "raise AssertionError",
87
+ "raise NotImplementedError",
88
+ "if 0:",
89
+ "if __name__ == .__main__.:",
90
+ "if TYPE_CHECKING:",
91
+ "class .*\\bProtocol\\):",
92
+ "@(abc\\.)?abstractmethod",
93
+ "[.][.][.]"
94
+ ]
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev365"
@@ -1,4 +1,4 @@
1
- import typing
1
+ import typing as t
2
2
  from typing import Optional, Sequence
3
3
 
4
4
  import jax
@@ -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,13 +39,11 @@ 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
- Named,
44
45
  NamedArray,
45
- NamedArrayAxes,
46
- NamedArrayAxesSpec,
47
- NamedOrNumeric,
46
+ NamedArrayAxes, NamedArrayAxesSpec, NamedOrNumeric,
48
47
  are_shape_checks_enabled,
49
48
  broadcast_arrays,
50
49
  broadcast_axis,
@@ -64,6 +63,7 @@ from .core import (
64
63
  unflatten_axis,
65
64
  updated_slice,
66
65
  )
66
+ from .haxtyping import Named
67
67
  from .hof import fold, map, scan, vmap
68
68
  from .jax_utils import tree_checkpoint_name
69
69
  from .ops import clip, isclose, pad_left, trace, tril, triu, where
@@ -81,8 +81,8 @@ from .wrap import (
81
81
  )
82
82
 
83
83
 
84
- T = typing.TypeVar("T")
85
- A = typing.TypeVar("A", Scalar, NamedArray, jnp.ndarray)
84
+ T = t.TypeVar("T")
85
+ A = t.TypeVar("A", Scalar, NamedArray, jnp.ndarray)
86
86
 
87
87
 
88
88
  # creation routines
@@ -105,8 +105,8 @@ def full(shape: AxisSpec, fill_value: T, dtype: Optional[DTypeLike] = None) -> N
105
105
  if isinstance(shape, Axis):
106
106
  return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
107
107
  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))
108
+ x_shape = to_jax_shape(shape)
109
+ return NamedArray(jnp.full(shape=x_shape, fill_value=fill_value, dtype=dtype), shape)
110
110
 
111
111
 
112
112
  def zeros_like(a: NamedArray, dtype=None) -> NamedArray:
@@ -146,16 +146,11 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
146
146
  ```
147
147
 
148
148
  """
149
- from haliax.jax_utils import to_jax_shape
150
- from haliax.util import ensure_tuple
151
-
152
- # if start is a tracer, we need to be a bit cleverer since arange doesn't support tracers
153
- # return NamedArray(jnp.arange(start, stop, step, dtype=dtype), (axis,))
154
149
  size = axis_size(axis)
155
150
 
156
151
  arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=size) * step + start
157
152
  arr = arr.reshape(to_jax_shape(axis))
158
- return NamedArray(arr, ensure_tuple(axis))
153
+ return NamedArray(arr, axis_spec_to_tuple(axis))
159
154
 
160
155
 
161
156
  # TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
@@ -272,7 +267,7 @@ def concatenate(axis: AxisSelector, arrays: Sequence[NamedArray]) -> NamedArray:
272
267
  if axis_index is None:
273
268
  raise ValueError(f"Axis {aname} not found in 0th array {arrays[0]}")
274
269
 
275
- axes: typing.Tuple[AxisSelector, ...] = arrays[0].axes
270
+ axes: tuple[AxisSelector, ...] = arrays[0].axes
276
271
  # we want to use the axis name for `axis`, because it's not uncommon for those to be different lengths in the arrays
277
272
  axes = axes[:axis_index] + (aname,) + axes[axis_index + 1 :]
278
273
  arrays = [a.rearrange(axes) for a in arrays]
@@ -929,9 +924,6 @@ __all__ = [
929
924
  "make_axes",
930
925
  "axis_name",
931
926
  "axis_size",
932
- "NamedArrayAxesSpec",
933
- "NamedArrayAxes",
934
- "Named",
935
927
  "NamedArray",
936
928
  "broadcast_to",
937
929
  "broadcast_axis",
@@ -1104,4 +1096,14 @@ __all__ = [
1104
1096
  "is_named_array",
1105
1097
  "tree_checkpoint_name",
1106
1098
  "ScanCheckpointPolicy",
1099
+ "quantization",
1100
+ "util",
1101
+ "einsum",
1102
+ "broadcast_arrays",
1103
+ "unflatten_axis",
1104
+ "ReductionFunction",
1105
+ "SimpleReductionFunction",
1106
+ "NamedArrayAxes",
1107
+ "NamedArrayAxesSpec",
1108
+ "Named",
1107
1109
  ]
@@ -4,7 +4,6 @@ import warnings
4
4
  from typing import Dict, Optional, Tuple
5
5
 
6
6
  import jax
7
- import jax.numpy as jnp
8
7
 
9
8
  import haliax
10
9
  from haliax.axis import (
@@ -12,6 +11,7 @@ from haliax.axis import (
12
11
  AxisSelection,
13
12
  PartialAxisSpec,
14
13
  axis_name,
14
+ axis_spec_to_shape_dict,
15
15
  eliminate_axes,
16
16
  rearrange_for_partial_order,
17
17
  union_axes,
@@ -19,7 +19,6 @@ from haliax.axis import (
19
19
  from haliax.core import NamedArray
20
20
  from haliax.jax_utils import _jittable_dg_einsum
21
21
  from haliax.types import DTypeLike, PrecisionLike
22
- from haliax.util import ensure_tuple
23
22
 
24
23
 
25
24
  # deprecated overload
@@ -140,8 +139,8 @@ def dot(
140
139
  if axis is None:
141
140
  jax_str = f"contract {', '.join(axis_name(ax) for ax in all_axes)} -> <scalar>"
142
141
  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)}"
142
+ axis = axis_spec_to_shape_dict(axis)
143
+ jax_str = f"contract {', '.join(axis)} -> {', '.join(a.name for a in output_axes)}"
145
144
 
146
145
  with jax.named_scope(jax_str):
147
146
  output = _jittable_dg_einsum(
@@ -1,4 +1,3 @@
1
- import warnings
2
1
  from functools import partial
3
2
 
4
3
  from jax import custom_jvp, custom_vjp, lax
@@ -15,8 +15,7 @@ from jaxtyping import PyTree
15
15
 
16
16
  import haliax.partitioning as partitioning
17
17
  from haliax._src.util import index_where
18
- from haliax.axis import Axis
19
- from haliax.core import NamedArray, flatten_axes, named
18
+ from haliax.core import NamedArray, named
20
19
  from haliax.jax_utils import is_jax_array_like, is_scalarish
21
20
  from haliax.tree_util import scan_aware_tree_map
22
21
 
@@ -209,14 +208,14 @@ def from_state_dict(tree: T, state_dict: StateDict, prefix: Optional[str] = None
209
208
  array = named(array, tree.axes)
210
209
  array = partitioning.auto_sharded(array)
211
210
 
212
- return array
211
+ return array # type: ignore
213
212
  elif is_jax_array_like(tree):
214
213
  if prefix is None:
215
214
  raise ValueError("Cannot extract a leaf value from a state dict without a prefix")
216
215
  # TODO: add "strict" flag so we can return None in cases where it's just missing
217
216
  return jnp.array(state_dict[prefix])
218
217
  elif tree is None:
219
- return None
218
+ return None # type: ignore
220
219
  else:
221
220
  if prefix is None:
222
221
  return tree