haliax 1.4.dev386__tar.gz → 1.4.dev388__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 (115) hide show
  1. haliax-1.4.dev388/.playbooks/wrap-non-named.md +49 -0
  2. {haliax-1.4.dev386 → haliax-1.4.dev388}/AGENTS.md +1 -0
  3. {haliax-1.4.dev386 → haliax-1.4.dev388}/PKG-INFO +1 -1
  4. haliax-1.4.dev388/src/haliax/__about__.py +1 -0
  5. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/__init__.py +19 -1
  6. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/ops.py +103 -1
  7. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_ops.py +48 -0
  8. haliax-1.4.dev386/src/haliax/__about__.py +0 -1
  9. {haliax-1.4.dev386 → haliax-1.4.dev388}/.coveragerc +0 -0
  10. {haliax-1.4.dev386 → haliax-1.4.dev388}/.flake8 +0 -0
  11. {haliax-1.4.dev386 → haliax-1.4.dev388}/.github/workflows/publish_dev.yaml +0 -0
  12. {haliax-1.4.dev386 → haliax-1.4.dev388}/.github/workflows/run_pre_commit.yaml +0 -0
  13. {haliax-1.4.dev386 → haliax-1.4.dev388}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  14. {haliax-1.4.dev386 → haliax-1.4.dev388}/.github/workflows/run_tests.yaml +0 -0
  15. {haliax-1.4.dev386 → haliax-1.4.dev388}/.gitignore +0 -0
  16. {haliax-1.4.dev386 → haliax-1.4.dev388}/.playbooks/add-types.md +0 -0
  17. {haliax-1.4.dev386 → haliax-1.4.dev388}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev386 → haliax-1.4.dev388}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev386 → haliax-1.4.dev388}/CONTRIBUTING.md +0 -0
  20. {haliax-1.4.dev386 → haliax-1.4.dev388}/LICENSE +0 -0
  21. {haliax-1.4.dev386 → haliax-1.4.dev388}/README.md +0 -0
  22. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/api.md +0 -0
  23. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/broadcasting.md +0 -0
  24. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/cheatsheet.md +0 -0
  25. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/css/material.css +0 -0
  26. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/css/mkdocstrings.css +0 -0
  27. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/faq.md +0 -0
  28. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/index.md +0 -0
  40. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/indexing.md +0 -0
  41. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/matmul.md +0 -0
  42. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/nn.md +0 -0
  43. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/partitioning.md +0 -0
  44. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/rearrange.ipynb +0 -0
  45. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/rearrange.md +0 -0
  46. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/requirements.txt +0 -0
  47. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/scan.md +0 -0
  48. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/typing.md +0 -0
  51. {haliax-1.4.dev386 → haliax-1.4.dev388}/docs/vmap.md +0 -0
  52. {haliax-1.4.dev386 → haliax-1.4.dev388}/mkdocs.yml +0 -0
  53. {haliax-1.4.dev386 → haliax-1.4.dev388}/pyproject.toml +0 -0
  54. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/__init__.py +0 -0
  55. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/compile_utils.py +0 -0
  56. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/dot.py +0 -0
  57. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/einsum.py +0 -0
  58. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/fp8.py +0 -0
  59. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/parsing.py +0 -0
  60. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/rearrange.py +0 -0
  61. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/scan.py +0 -0
  62. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/state_dict.py +0 -0
  63. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/_src/util.py +0 -0
  64. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/axis.py +0 -0
  65. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/core.py +0 -0
  66. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/debug.py +0 -0
  67. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/haxtyping.py +0 -0
  68. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/hof.py +0 -0
  69. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/jax_utils.py +0 -0
  70. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/__init__.py +0 -0
  71. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/activations.py +0 -0
  72. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/attention.py +0 -0
  73. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/conv.py +0 -0
  74. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/dropout.py +0 -0
  75. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/embedding.py +0 -0
  76. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/linear.py +0 -0
  77. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/loss.py +0 -0
  78. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/mlp.py +0 -0
  79. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/normalization.py +0 -0
  80. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/pool.py +0 -0
  81. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/nn/scan.py +0 -0
  82. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/partitioning.py +0 -0
  83. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/quantization.py +0 -0
  84. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/random.py +0 -0
  85. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/util.py +0 -0
  90. {haliax-1.4.dev386 → haliax-1.4.dev388}/src/haliax/wrap.py +0 -0
  91. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/core_test.py +0 -0
  92. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_attention.py +0 -0
  93. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_axis.py +0 -0
  94. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_conv.py +0 -0
  95. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_debug.py +0 -0
  96. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_dot.py +0 -0
  97. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_dtype_typing.py +0 -0
  98. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_einsum.py +0 -0
  99. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_fp8.py +0 -0
  100. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_hof.py +0 -0
  101. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_int8.py +0 -0
  102. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_namedarray_typing.py +0 -0
  103. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_nn.py +0 -0
  104. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_parsing.py +0 -0
  105. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_partitioning.py +0 -0
  106. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_pool.py +0 -0
  107. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_random.py +0 -0
  108. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_rearrange.py +0 -0
  109. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_scan.py +0 -0
  110. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_scatter_gather.py +0 -0
  111. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_specialized_fns.py +0 -0
  112. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_state_dict.py +0 -0
  113. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_tree_util.py +0 -0
  114. {haliax-1.4.dev386 → haliax-1.4.dev388}/tests/test_utils.py +0 -0
  115. {haliax-1.4.dev386 → haliax-1.4.dev388}/uv.lock +0 -0
@@ -0,0 +1,49 @@
1
+ # Wrapping Functions with NamedArray Support
2
+
3
+ This playbook explains how to convert a regular JAX function that works on unnamed arrays into a Haliax function that accepts `NamedArray` inputs and returns `NamedArray` outputs.
4
+
5
+ ## When is wrapping needed?
6
+ Many JAX primitives only operate on regular arrays. To integrate them in Haliax you should provide a thin wrapper that handles axis metadata. Simple elementwise operations and reductions have helper utilities.
7
+
8
+ ## Elemwise Unary
9
+ For a unary function that acts elementwise (e.g. `jnp.abs`):
10
+
11
+ ```python
12
+ from haliax import wrap_elemwise_unary
13
+
14
+ def abs(a):
15
+ return wrap_elemwise_unary(jnp.abs, a)
16
+ ```
17
+
18
+ This preserves axis order and dtype.
19
+
20
+ ## Elemwise Binary
21
+ For binary operations (e.g. `jnp.add`), decorate a function with `wrap_elemwise_binary`:
22
+
23
+ ```python
24
+ from haliax import wrap_elemwise_binary
25
+
26
+ @wrap_elemwise_binary
27
+ def add(x1, x2):
28
+ return jnp.add(x1, x2)
29
+ ```
30
+
31
+ Broadcasting between `NamedArray`s is handled automatically.
32
+
33
+ ## Reductions
34
+ Reductions require choosing axes to eliminate. Use `wrap_reduction_call`:
35
+
36
+ ```python
37
+ from haliax import wrap_reduction_call
38
+
39
+ def sum(a, axis=None):
40
+ return wrap_reduction_call(jnp.sum, a, axis)
41
+ ```
42
+
43
+ `axis` can be an `AxisSelector` or tuple. The wrapper returns a `NamedArray` with those axes removed.
44
+
45
+ ## Harder Cases
46
+ Some functions need bespoke handling. For example `jnp.unique` returns several arrays and may change shape unpredictably. There is no generic helper, so you will need to manually map between `NamedArray` axes and the outputs. Use the lower level utilities in `haliax.wrap` for broadcasting and axis lookup.
47
+
48
+ ## Testing
49
+ Add tests to ensure that named and unnamed calls produce the same results and that axis names are preserved or removed correctly.
@@ -17,6 +17,7 @@ repository. Follow these notes when implementing new features or fixing bugs.
17
17
  ## Playbook
18
18
 
19
19
  - Adding Haliax-style tensor typing annotations are described in @.playbooks/add-types.md
20
+ - Wrapping standard JAX functions so they operate on `NamedArray` is explained in @.playbooks/wrap-non-named.md
20
21
 
21
22
  ## Code Style
22
23
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev386
3
+ Version: 1.4.dev388
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 @@
1
+ __version__ = "1.4.dev388"
@@ -66,7 +66,21 @@ from .core import (
66
66
  from .haxtyping import Named
67
67
  from .hof import fold, map, scan, vmap
68
68
  from .jax_utils import tree_checkpoint_name
69
- from .ops import clip, isclose, pad_left, pad, trace, tril, triu, unique, where
69
+ from .ops import (
70
+ clip,
71
+ isclose,
72
+ pad_left,
73
+ pad,
74
+ trace,
75
+ tril,
76
+ triu,
77
+ unique,
78
+ unique_values,
79
+ unique_counts,
80
+ unique_inverse,
81
+ unique_all,
82
+ where,
83
+ )
70
84
  from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
71
85
  from .specialized_fns import top_k
72
86
  from .types import Scalar
@@ -1035,6 +1049,10 @@ __all__ = [
1035
1049
  "trace",
1036
1050
  "where",
1037
1051
  "unique",
1052
+ "unique_values",
1053
+ "unique_counts",
1054
+ "unique_inverse",
1055
+ "unique_all",
1038
1056
  "clip",
1039
1057
  "tril",
1040
1058
  "triu",
@@ -342,4 +342,106 @@ def unique(
342
342
  return tuple(ret)
343
343
 
344
344
 
345
- __all__ = ["trace", "where", "tril", "triu", "isclose", "pad_left", "pad", "clip", "unique"]
345
+ def unique_values(
346
+ array: NamedArray,
347
+ Unique: Axis,
348
+ *,
349
+ axis: AxisSelector | None = None,
350
+ fill_value: ArrayLike | None = None,
351
+ ) -> NamedArray:
352
+ """Shortcut for :func:`unique` that returns only unique values."""
353
+
354
+ return typing.cast(
355
+ NamedArray,
356
+ unique(
357
+ array,
358
+ Unique,
359
+ axis=axis,
360
+ fill_value=fill_value,
361
+ ),
362
+ )
363
+
364
+
365
+ def unique_counts(
366
+ array: NamedArray,
367
+ Unique: Axis,
368
+ *,
369
+ axis: AxisSelector | None = None,
370
+ fill_value: ArrayLike | None = None,
371
+ ) -> tuple[NamedArray, NamedArray]:
372
+ """Shortcut for :func:`unique` that also returns counts."""
373
+
374
+ values, counts = typing.cast(
375
+ tuple[NamedArray, NamedArray],
376
+ unique(
377
+ array,
378
+ Unique,
379
+ return_counts=True,
380
+ axis=axis,
381
+ fill_value=fill_value,
382
+ ),
383
+ )
384
+ return values, counts
385
+
386
+
387
+ def unique_inverse(
388
+ array: NamedArray,
389
+ Unique: Axis,
390
+ *,
391
+ axis: AxisSelector | None = None,
392
+ fill_value: ArrayLike | None = None,
393
+ ) -> tuple[NamedArray, NamedArray]:
394
+ """Shortcut for :func:`unique` that also returns inverse indices."""
395
+
396
+ values, inverse = typing.cast(
397
+ tuple[NamedArray, NamedArray],
398
+ unique(
399
+ array,
400
+ Unique,
401
+ return_inverse=True,
402
+ axis=axis,
403
+ fill_value=fill_value,
404
+ ),
405
+ )
406
+ return values, inverse
407
+
408
+
409
+ def unique_all(
410
+ array: NamedArray,
411
+ Unique: Axis,
412
+ *,
413
+ axis: AxisSelector | None = None,
414
+ fill_value: ArrayLike | None = None,
415
+ ) -> tuple[NamedArray, NamedArray, NamedArray, NamedArray]:
416
+ """Shortcut for :func:`unique` returning values, indices, inverse, and counts."""
417
+
418
+ values, indices, inverse, counts = typing.cast(
419
+ tuple[NamedArray, NamedArray, NamedArray, NamedArray],
420
+ unique(
421
+ array,
422
+ Unique,
423
+ return_index=True,
424
+ return_inverse=True,
425
+ return_counts=True,
426
+ axis=axis,
427
+ fill_value=fill_value,
428
+ ),
429
+ )
430
+ return values, indices, inverse, counts
431
+
432
+
433
+ __all__ = [
434
+ "trace",
435
+ "where",
436
+ "tril",
437
+ "triu",
438
+ "isclose",
439
+ "pad_left",
440
+ "pad",
441
+ "clip",
442
+ "unique",
443
+ "unique_values",
444
+ "unique_counts",
445
+ "unique_inverse",
446
+ "unique_all",
447
+ ]
@@ -1,4 +1,5 @@
1
1
  from typing import Callable
2
+ import typing
2
3
 
3
4
  import jax.numpy as jnp
4
5
  import pytest
@@ -359,3 +360,50 @@ def test_unique():
359
360
 
360
361
  assert jnp.all(jnp.equal(values.array, jnp.array([[1, 2], [2, 3]])))
361
362
  assert jnp.all(jnp.equal(counts.array, jnp.array([2, 1])))
363
+
364
+
365
+ def test_unique_shortcuts():
366
+ Height = Axis("Height", 3)
367
+ Width = Axis("Width", 2)
368
+
369
+ arr2d = hax.named([[1, 2], [2, 3], [1, 2]], (Height, Width))
370
+ U = Axis("U", 3)
371
+
372
+ # unique_values
373
+ uv = hax.unique_values(arr2d, U)
374
+ uv_expected = hax.unique(arr2d, U)
375
+ assert jnp.all(uv.array == uv_expected.array)
376
+
377
+ # unique_counts
378
+ vc, cc = hax.unique_counts(arr2d, U)
379
+ vc_exp, cc_exp = hax.unique(arr2d, U, return_counts=True)
380
+ assert jnp.all(vc.array == vc_exp.array)
381
+ assert jnp.all(cc.array == cc_exp.array)
382
+
383
+ # unique_inverse
384
+ Height1 = Axis("Height1", 5)
385
+ arr1d = hax.named([3, 4, 1, 3, 1], (Height1,))
386
+ U2 = Axis("U2", 3)
387
+ vi, ii = hax.unique_inverse(arr1d, U2)
388
+ vi_exp, ii_exp = hax.unique(arr1d, U2, return_inverse=True)
389
+ assert jnp.all(vi.array == vi_exp.array)
390
+ assert jnp.all(ii.array == ii_exp.array)
391
+
392
+ # unique_all
393
+ U3 = Axis("U3", 2)
394
+ va, ia, ina, ca = hax.unique_all(arr2d, U3, axis=Height)
395
+ va_exp, ia_exp, ina_exp, ca_exp = typing.cast(
396
+ tuple[NamedArray, NamedArray, NamedArray, NamedArray],
397
+ hax.unique(
398
+ arr2d,
399
+ U3,
400
+ axis=Height,
401
+ return_index=True,
402
+ return_inverse=True,
403
+ return_counts=True,
404
+ ),
405
+ )
406
+ assert jnp.all(va.array == va_exp.array)
407
+ assert jnp.all(ia.array == ia_exp.array)
408
+ assert jnp.all(ina.array == ina_exp.array)
409
+ assert jnp.all(ca.array == ca_exp.array)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev386"
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