haliax 1.4.dev404__tar.gz → 1.4.dev405__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 (120) hide show
  1. {haliax-1.4.dev404 → haliax-1.4.dev405}/PKG-INFO +1 -1
  2. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/api.md +1 -0
  3. haliax-1.4.dev405/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/__init__.py +2 -0
  5. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/ops.py +29 -0
  6. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/core_test.py +20 -0
  7. haliax-1.4.dev404/src/haliax/__about__.py +0 -1
  8. {haliax-1.4.dev404 → haliax-1.4.dev405}/.coveragerc +0 -0
  9. {haliax-1.4.dev404 → haliax-1.4.dev405}/.flake8 +0 -0
  10. {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/publish_dev.yaml +0 -0
  11. {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/run_pre_commit.yaml +0 -0
  12. {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  13. {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/run_tests.yaml +0 -0
  14. {haliax-1.4.dev404 → haliax-1.4.dev405}/.gitignore +0 -0
  15. {haliax-1.4.dev404 → haliax-1.4.dev405}/.playbooks/add-types.md +0 -0
  16. {haliax-1.4.dev404 → haliax-1.4.dev405}/.playbooks/wrap-non-named.md +0 -0
  17. {haliax-1.4.dev404 → haliax-1.4.dev405}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev404 → haliax-1.4.dev405}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev404 → haliax-1.4.dev405}/AGENTS.md +0 -0
  20. {haliax-1.4.dev404 → haliax-1.4.dev405}/CONTRIBUTING.md +0 -0
  21. {haliax-1.4.dev404 → haliax-1.4.dev405}/LICENSE +0 -0
  22. {haliax-1.4.dev404 → haliax-1.4.dev405}/README.md +0 -0
  23. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/broadcasting.md +0 -0
  24. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/cheatsheet.md +0 -0
  25. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/css/material.css +0 -0
  26. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/css/mkdocstrings.css +0 -0
  27. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/faq.md +0 -0
  28. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/index.md +0 -0
  40. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/indexing.md +0 -0
  41. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/matmul.md +0 -0
  42. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/nn.md +0 -0
  43. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/partitioning.md +0 -0
  44. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/primer.md +0 -0
  45. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/rearrange.ipynb +0 -0
  46. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/rearrange.md +0 -0
  47. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/requirements.txt +0 -0
  48. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/scan.md +0 -0
  49. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/state-dict.md +0 -0
  50. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/tutorial.md +0 -0
  51. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/typing.md +0 -0
  52. {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/vmap.md +0 -0
  53. {haliax-1.4.dev404 → haliax-1.4.dev405}/mkdocs.yml +0 -0
  54. {haliax-1.4.dev404 → haliax-1.4.dev405}/pyproject.toml +0 -0
  55. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/__init__.py +0 -0
  56. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/compile_utils.py +0 -0
  57. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/dot.py +0 -0
  58. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/einsum.py +0 -0
  59. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/fp8.py +0 -0
  60. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/parsing.py +0 -0
  61. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/rearrange.py +0 -0
  62. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/scan.py +0 -0
  63. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/state_dict.py +0 -0
  64. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/util.py +0 -0
  65. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/axis.py +0 -0
  66. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/core.py +0 -0
  67. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/debug.py +0 -0
  68. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/field.py +0 -0
  69. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/haxtyping.py +0 -0
  70. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/hof.py +0 -0
  71. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/jax_utils.py +0 -0
  72. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/__init__.py +0 -0
  73. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/activations.py +0 -0
  74. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/attention.py +0 -0
  75. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/conv.py +0 -0
  76. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/dropout.py +0 -0
  77. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/embedding.py +0 -0
  78. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/linear.py +0 -0
  79. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/loss.py +0 -0
  80. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/mlp.py +0 -0
  81. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/normalization.py +0 -0
  82. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/pool.py +0 -0
  83. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/scan.py +0 -0
  84. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/partitioning.py +0 -0
  85. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/quantization.py +0 -0
  86. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/random.py +0 -0
  87. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/specialized_fns.py +0 -0
  88. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/state_dict.py +0 -0
  89. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/tree_util.py +0 -0
  90. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/types.py +0 -0
  91. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/util.py +0 -0
  92. {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/wrap.py +0 -0
  93. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_attention.py +0 -0
  94. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_axis.py +0 -0
  95. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_conv.py +0 -0
  96. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_debug.py +0 -0
  97. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_dot.py +0 -0
  98. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_dtype_typing.py +0 -0
  99. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_einsum.py +0 -0
  100. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_field.py +0 -0
  101. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_fp8.py +0 -0
  102. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_hof.py +0 -0
  103. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_int8.py +0 -0
  104. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_moe_linear.py +0 -0
  105. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_namedarray_typing.py +0 -0
  106. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_nn.py +0 -0
  107. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_ops.py +0 -0
  108. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_parsing.py +0 -0
  109. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_partitioning.py +0 -0
  110. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_pool.py +0 -0
  111. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_random.py +0 -0
  112. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_rearrange.py +0 -0
  113. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_scan.py +0 -0
  114. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_scatter_gather.py +0 -0
  115. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_specialized_fns.py +0 -0
  116. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_state_dict.py +0 -0
  117. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_tree_util.py +0 -0
  118. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_utils.py +0 -0
  119. {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_visualize_sharding.py +0 -0
  120. {haliax-1.4.dev404 → haliax-1.4.dev405}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev404
3
+ Version: 1.4.dev405
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/
@@ -260,6 +260,7 @@ These are all more or less directly from JAX's NumPy API.
260
260
  ::: haliax.clip
261
261
  ::: haliax.isclose
262
262
  ::: haliax.pad
263
+ ::: haliax.searchsorted
263
264
  ::: haliax.top_k
264
265
  ::: haliax.trace
265
266
  ::: haliax.tril
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev405"
@@ -80,6 +80,7 @@ from .ops import (
80
80
  unique_counts,
81
81
  unique_inverse,
82
82
  unique_all,
83
+ searchsorted,
83
84
  bincount,
84
85
  where,
85
86
  )
@@ -1056,6 +1057,7 @@ __all__ = [
1056
1057
  "unique_counts",
1057
1058
  "unique_inverse",
1058
1059
  "unique_all",
1060
+ "searchsorted",
1059
1061
  "bincount",
1060
1062
  "clip",
1061
1063
  "tril",
@@ -430,6 +430,34 @@ def unique_all(
430
430
  return values, indices, inverse, counts
431
431
 
432
432
 
433
+ def searchsorted(
434
+ a: NamedArray,
435
+ v: NamedArray | ArrayLike,
436
+ *,
437
+ side: str = "left",
438
+ sorter: NamedArray | ArrayLike | None = None,
439
+ method: str = "scan",
440
+ ) -> NamedArray:
441
+ """Named version of `jax.numpy.searchsorted`.
442
+
443
+ ``a`` and ``sorter`` (if provided) must be one-dimensional.
444
+ The returned array has the same axes as ``v``.
445
+ """
446
+
447
+ if a.ndim != 1:
448
+ raise ValueError("searchsorted only supports 1D 'a'")
449
+
450
+ if not isinstance(v, NamedArray):
451
+ v = haliax.named(v, ())
452
+
453
+ sorter_arr = None
454
+ if sorter is not None:
455
+ sorter_arr = sorter.array if isinstance(sorter, NamedArray) else jnp.asarray(sorter)
456
+
457
+ result = jnp.searchsorted(a.array, v.array, side=side, sorter=sorter_arr, method=method)
458
+ return NamedArray(result, v.axes)
459
+
460
+
433
461
  def bincount(
434
462
  x: NamedArray,
435
463
  Counts: Axis,
@@ -471,5 +499,6 @@ __all__ = [
471
499
  "unique_counts",
472
500
  "unique_inverse",
473
501
  "unique_all",
502
+ "searchsorted",
474
503
  "bincount",
475
504
  ]
@@ -238,6 +238,26 @@ def test_cumsum_etc():
238
238
  assert hax.argsort(named1, axis=Width).axes == (Height, Width, Depth)
239
239
 
240
240
 
241
+ def test_searchsorted():
242
+ A = hax.Axis("a", 5)
243
+ V = hax.Axis("v", 4)
244
+
245
+ a = hax.named([1, 3, 5, 7, 9], axis=A)
246
+ v = hax.named([0, 3, 6, 10], axis=V)
247
+
248
+ result = hax.searchsorted(a, v)
249
+ assert jnp.all(result.array == jnp.searchsorted(a.array, v.array))
250
+ assert result.axes == (V,)
251
+
252
+ unsorted = hax.named([5, 1, 3, 7, 4], axis=A)
253
+ sorter = hax.argsort(unsorted, axis=A)
254
+ result = hax.searchsorted(unsorted, v, sorter=sorter, side="right")
255
+ assert jnp.all(
256
+ result.array == jnp.searchsorted(unsorted.array, v.array, sorter=sorter.array, side="right")
257
+ )
258
+ assert result.axes == (V,)
259
+
260
+
241
261
  def test_rearrange():
242
262
  H, W, D, C = hax.make_axes(H=2, W=3, D=4, C=5)
243
263
 
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev404"
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
File without changes