haliax 1.4.dev382__tar.gz → 1.4.dev386__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.dev382 → haliax-1.4.dev386}/PKG-INFO +1 -1
  2. haliax-1.4.dev386/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/__init__.py +2 -1
  4. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/ops.py +147 -1
  5. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_ops.py +109 -0
  6. haliax-1.4.dev382/src/haliax/__about__.py +0 -1
  7. {haliax-1.4.dev382 → haliax-1.4.dev386}/.coveragerc +0 -0
  8. {haliax-1.4.dev382 → haliax-1.4.dev386}/.flake8 +0 -0
  9. {haliax-1.4.dev382 → haliax-1.4.dev386}/.github/workflows/publish_dev.yaml +0 -0
  10. {haliax-1.4.dev382 → haliax-1.4.dev386}/.github/workflows/run_pre_commit.yaml +0 -0
  11. {haliax-1.4.dev382 → haliax-1.4.dev386}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  12. {haliax-1.4.dev382 → haliax-1.4.dev386}/.github/workflows/run_tests.yaml +0 -0
  13. {haliax-1.4.dev382 → haliax-1.4.dev386}/.gitignore +0 -0
  14. {haliax-1.4.dev382 → haliax-1.4.dev386}/.playbooks/add-types.md +0 -0
  15. {haliax-1.4.dev382 → haliax-1.4.dev386}/.pre-commit-config.yaml +0 -0
  16. {haliax-1.4.dev382 → haliax-1.4.dev386}/.readthedocs.yaml +0 -0
  17. {haliax-1.4.dev382 → haliax-1.4.dev386}/AGENTS.md +0 -0
  18. {haliax-1.4.dev382 → haliax-1.4.dev386}/CONTRIBUTING.md +0 -0
  19. {haliax-1.4.dev382 → haliax-1.4.dev386}/LICENSE +0 -0
  20. {haliax-1.4.dev382 → haliax-1.4.dev386}/README.md +0 -0
  21. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/api.md +0 -0
  22. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/broadcasting.md +0 -0
  23. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/cheatsheet.md +0 -0
  24. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/css/material.css +0 -0
  25. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/css/mkdocstrings.css +0 -0
  26. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/faq.md +0 -0
  27. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/data_parallel_mesh.png +0 -0
  28. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  29. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_1d.png +0 -0
  30. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_1d_zero.png +0 -0
  31. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_2d.png +0 -0
  32. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  33. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  34. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  35. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  36. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/figures/device_mesh_2d_zero.png +0 -0
  37. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/fp8.md +0 -0
  38. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/index.md +0 -0
  39. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/indexing.md +0 -0
  40. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/matmul.md +0 -0
  41. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/nn.md +0 -0
  42. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/partitioning.md +0 -0
  43. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/rearrange.ipynb +0 -0
  44. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/rearrange.md +0 -0
  45. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/requirements.txt +0 -0
  46. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/scan.md +0 -0
  47. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/state-dict.md +0 -0
  48. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/tutorial.md +0 -0
  49. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/typing.md +0 -0
  50. {haliax-1.4.dev382 → haliax-1.4.dev386}/docs/vmap.md +0 -0
  51. {haliax-1.4.dev382 → haliax-1.4.dev386}/mkdocs.yml +0 -0
  52. {haliax-1.4.dev382 → haliax-1.4.dev386}/pyproject.toml +0 -0
  53. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/__init__.py +0 -0
  54. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/compile_utils.py +0 -0
  55. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/dot.py +0 -0
  56. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/einsum.py +0 -0
  57. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/fp8.py +0 -0
  58. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/parsing.py +0 -0
  59. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/rearrange.py +0 -0
  60. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/scan.py +0 -0
  61. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/state_dict.py +0 -0
  62. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/_src/util.py +0 -0
  63. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/axis.py +0 -0
  64. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/core.py +0 -0
  65. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/debug.py +0 -0
  66. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/haxtyping.py +0 -0
  67. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/hof.py +0 -0
  68. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/jax_utils.py +0 -0
  69. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/__init__.py +0 -0
  70. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/activations.py +0 -0
  71. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/attention.py +0 -0
  72. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/conv.py +0 -0
  73. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/dropout.py +0 -0
  74. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/embedding.py +0 -0
  75. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/linear.py +0 -0
  76. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/loss.py +0 -0
  77. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/mlp.py +0 -0
  78. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/normalization.py +0 -0
  79. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/pool.py +0 -0
  80. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/nn/scan.py +0 -0
  81. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/partitioning.py +0 -0
  82. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/quantization.py +0 -0
  83. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/random.py +0 -0
  84. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/specialized_fns.py +0 -0
  85. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/state_dict.py +0 -0
  86. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/tree_util.py +0 -0
  87. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/types.py +0 -0
  88. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/util.py +0 -0
  89. {haliax-1.4.dev382 → haliax-1.4.dev386}/src/haliax/wrap.py +0 -0
  90. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/core_test.py +0 -0
  91. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_attention.py +0 -0
  92. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_axis.py +0 -0
  93. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_conv.py +0 -0
  94. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_debug.py +0 -0
  95. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_dot.py +0 -0
  96. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_dtype_typing.py +0 -0
  97. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_einsum.py +0 -0
  98. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_fp8.py +0 -0
  99. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_hof.py +0 -0
  100. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_int8.py +0 -0
  101. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_namedarray_typing.py +0 -0
  102. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_nn.py +0 -0
  103. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_parsing.py +0 -0
  104. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_partitioning.py +0 -0
  105. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_pool.py +0 -0
  106. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_random.py +0 -0
  107. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_rearrange.py +0 -0
  108. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_scan.py +0 -0
  109. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_scatter_gather.py +0 -0
  110. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_specialized_fns.py +0 -0
  111. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_state_dict.py +0 -0
  112. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_tree_util.py +0 -0
  113. {haliax-1.4.dev382 → haliax-1.4.dev386}/tests/test_utils.py +0 -0
  114. {haliax-1.4.dev382 → haliax-1.4.dev386}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev382
3
+ Version: 1.4.dev386
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.dev386"
@@ -66,7 +66,7 @@ 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, where
69
+ from .ops import clip, isclose, pad_left, pad, trace, tril, triu, unique, where
70
70
  from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
71
71
  from .specialized_fns import top_k
72
72
  from .types import Scalar
@@ -1034,6 +1034,7 @@ __all__ = [
1034
1034
  "vmap",
1035
1035
  "trace",
1036
1036
  "where",
1037
+ "unique",
1037
1038
  "clip",
1038
1039
  "tril",
1039
1040
  "triu",
@@ -3,6 +3,9 @@ from typing import Mapping, Optional, Union
3
3
 
4
4
  import jax
5
5
  import jax.numpy as jnp
6
+ from jaxtyping import ArrayLike
7
+
8
+ import haliax
6
9
 
7
10
  from .axis import Axis, AxisSelector, axis_name
8
11
  from .core import NamedArray, NamedOrNumeric, broadcast_arrays, broadcast_arrays_and_return_axes, named
@@ -196,4 +199,147 @@ def raw_array_or_scalar(x: NamedOrNumeric):
196
199
  return x
197
200
 
198
201
 
199
- __all__ = ["trace", "where", "tril", "triu", "isclose", "pad_left", "pad", "clip"]
202
+ @typing.overload
203
+ def unique(
204
+ array: NamedArray, Unique: Axis, *, axis: AxisSelector | None = None, fill_value: ArrayLike | None = None
205
+ ) -> NamedArray:
206
+ ...
207
+
208
+
209
+ @typing.overload
210
+ def unique(
211
+ array: NamedArray,
212
+ Unique: Axis,
213
+ *,
214
+ return_index: typing.Literal[True],
215
+ axis: AxisSelector | None = None,
216
+ fill_value: ArrayLike | None = None,
217
+ ) -> tuple[NamedArray, NamedArray]:
218
+ ...
219
+
220
+
221
+ @typing.overload
222
+ def unique(
223
+ array: NamedArray,
224
+ Unique: Axis,
225
+ *,
226
+ return_inverse: typing.Literal[True],
227
+ axis: AxisSelector | None = None,
228
+ fill_value: ArrayLike | None = None,
229
+ ) -> tuple[NamedArray, NamedArray]:
230
+ ...
231
+
232
+
233
+ @typing.overload
234
+ def unique(
235
+ array: NamedArray,
236
+ Unique: Axis,
237
+ *,
238
+ return_counts: typing.Literal[True],
239
+ axis: AxisSelector | None = None,
240
+ fill_value: ArrayLike | None = None,
241
+ ) -> tuple[NamedArray, NamedArray]:
242
+ ...
243
+
244
+
245
+ @typing.overload
246
+ def unique(
247
+ array: NamedArray,
248
+ Unique: Axis,
249
+ *,
250
+ return_index: bool = False,
251
+ return_inverse: bool = False,
252
+ return_counts: bool = False,
253
+ axis: AxisSelector | None = None,
254
+ fill_value: ArrayLike | None = None,
255
+ ) -> NamedArray | tuple[NamedArray, ...]:
256
+ ...
257
+
258
+
259
+ def unique(
260
+ array: NamedArray,
261
+ Unique: Axis,
262
+ *,
263
+ return_index: bool = False,
264
+ return_inverse: bool = False,
265
+ return_counts: bool = False,
266
+ axis: AxisSelector | None = None,
267
+ fill_value: ArrayLike | None = None,
268
+ ) -> NamedArray | tuple[NamedArray, ...]:
269
+ """
270
+ Like jnp.unique, but with named axes.
271
+
272
+ Args:
273
+ array: The input array.
274
+ Unique: The name of the axis that will be created to hold the unique values.
275
+ fill_value: The value to use for the fill_value argument of jnp.unique
276
+ axis: The axis along which to find unique values.
277
+ return_index: If True, return the indices of the unique values.
278
+ return_inverse: If True, return the indices of the input array that would reconstruct the unique values.
279
+ """
280
+ size = Unique.size
281
+
282
+ is_multireturn = return_index or return_inverse or return_counts
283
+
284
+ kwargs = dict(
285
+ size=size,
286
+ fill_value=fill_value,
287
+ return_index=return_index,
288
+ return_inverse=return_inverse,
289
+ return_counts=return_counts,
290
+ )
291
+
292
+ if axis is not None:
293
+ axis_index = array._lookup_indices(axis)
294
+ if axis_index is None:
295
+ raise ValueError(f"Axis {axis} not found in array. Available axes: {array.axes}")
296
+ out = jnp.unique(array.array, axis=axis_index, **kwargs)
297
+ else:
298
+ out = jnp.unique(array.array, **kwargs)
299
+
300
+ if is_multireturn:
301
+ unique = out[0]
302
+ next_index = 1
303
+ if return_index:
304
+ index = out[next_index]
305
+ next_index += 1
306
+ if return_inverse:
307
+ inverse = out[next_index]
308
+ next_index += 1
309
+ if return_counts:
310
+ counts = out[next_index]
311
+ next_index += 1
312
+ else:
313
+ unique = out
314
+
315
+ ret = []
316
+
317
+ if axis is not None:
318
+ out_axes = haliax.axis.replace_axis(array.axes, axis, Unique)
319
+ else:
320
+ out_axes = (Unique,)
321
+
322
+ unique_values = haliax.named(unique, out_axes)
323
+ if not is_multireturn:
324
+ return unique_values
325
+
326
+ ret.append(unique_values)
327
+
328
+ if return_index:
329
+ ret.append(haliax.named(index, Unique))
330
+
331
+ if return_inverse:
332
+ if axis is not None:
333
+ assert axis_index is not None
334
+ inverse = haliax.named(inverse, array.axes[axis_index])
335
+ else:
336
+ inverse = haliax.named(inverse, array.axes)
337
+ ret.append(inverse)
338
+
339
+ if return_counts:
340
+ ret.append(haliax.named(counts, Unique))
341
+
342
+ return tuple(ret)
343
+
344
+
345
+ __all__ = ["trace", "where", "tril", "triu", "isclose", "pad_left", "pad", "clip", "unique"]
@@ -250,3 +250,112 @@ def test_pad():
250
250
  assert padded.axes[0].size == Height.size + 3
251
251
  assert padded.axes[1].size == Width.size + 1
252
252
  assert jnp.all(expected == padded.array)
253
+
254
+
255
+ def test_unique():
256
+ # named version of this test:
257
+ # >>> M = jnp.array([[1, 2],
258
+ # ... [2, 3],
259
+ # ... [1, 2]])
260
+ # >>> jnp.unique(M)
261
+ # Array([1, 2, 3], dtype=int32)
262
+
263
+ Height = Axis("Height", 3)
264
+ Width = Axis("Width", 2)
265
+
266
+ named1 = hax.named([[1, 2], [2, 3], [1, 2]], (Height, Width))
267
+
268
+ U = Axis("U", 3)
269
+
270
+ named2 = hax.unique(named1, U)
271
+
272
+ assert jnp.all(jnp.equal(named2.array, jnp.array([1, 2, 3])))
273
+
274
+ # If you pass an ``axis`` keyword, you can find unique *slices* of the array along
275
+ # that axis:
276
+ #
277
+ # >>> jnp.unique(M, axis=0)
278
+ # Array([[1, 2],
279
+ # [2, 3]], dtype=int32)
280
+
281
+ U2 = Axis("U2", 2)
282
+ named3 = hax.unique(named1, U2, axis=Height)
283
+ assert jnp.all(jnp.equal(named3.array, jnp.array([[1, 2], [2, 3]])))
284
+
285
+ # >>> x = jnp.array([3, 4, 1, 3, 1])
286
+ # >>> values, indices = jnp.unique(x, return_index=True)
287
+ # >>> print(values)
288
+ # [1 3 4]
289
+ # >>> print(indices)
290
+ # [2 0 1]
291
+ # >>> jnp.all(values == x[indices])
292
+ # Array(True, dtype=bool)
293
+
294
+ x = hax.named([3, 4, 1, 3, 1], ("Height",))
295
+ U3 = Axis("U3", 3)
296
+ values, indices = hax.unique(x, U3, return_index=True)
297
+
298
+ assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
299
+ assert jnp.all(jnp.equal(indices.array, jnp.array([2, 0, 1])))
300
+
301
+ assert jnp.all(jnp.equal(values.array, x[{"Height": indices}].array))
302
+
303
+ # If you set ``return_inverse=True``, then ``unique`` returns the indices within the
304
+ # unique values for every entry in the input array:
305
+ #
306
+ # >>> x = jnp.array([3, 4, 1, 3, 1])
307
+ # >>> values, inverse = jnp.unique(x, return_inverse=True)
308
+ # >>> print(values)
309
+ # [1 3 4]
310
+ # >>> print(inverse)
311
+ # [1 2 0 1 0]
312
+ # >>> jnp.all(values[inverse] == x)
313
+ # Array(True, dtype=bool)
314
+
315
+ values, inverse = hax.unique(x, U3, return_inverse=True)
316
+
317
+ assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
318
+ assert jnp.all(jnp.equal(inverse.array, jnp.array([1, 2, 0, 1, 0])))
319
+
320
+ # In multiple dimensions, the input can be reconstructed using
321
+ # :func:`jax.numpy.take`:
322
+ #
323
+ # >>> values, inverse = jnp.unique(M, axis=0, return_inverse=True)
324
+ # >>> jnp.all(jnp.take(values, inverse, axis=0) == M)
325
+ # Array(True, dtype=bool)
326
+ #
327
+
328
+ values, inverse = hax.unique(named1, U3, axis=Height, return_inverse=True)
329
+
330
+ assert jnp.all((values[{"U3": inverse}] == named1).array)
331
+
332
+ # **Returning counts**
333
+ # If you set ``return_counts=True``, then ``unique`` returns the number of occurrences
334
+ # within the input for every unique value:
335
+ #
336
+ # >>> x = jnp.array([3, 4, 1, 3, 1])
337
+ # >>> values, counts = jnp.unique(x, return_counts=True)
338
+ # >>> print(values)
339
+ # [1 3 4]
340
+ # >>> print(counts)
341
+ # [2 2 1]
342
+ #
343
+ # For multi-dimensional arrays, this also returns a 1D array of counts
344
+ # indicating number of occurrences along the specified axis:
345
+ #
346
+ # >>> values, counts = jnp.unique(M, axis=0, return_counts=True)
347
+ # >>> print(values)
348
+ # [[1 2]
349
+ # [2 3]]
350
+ # >>> print(counts)
351
+ # [2 1]
352
+
353
+ values, counts = hax.unique(x, U3, return_counts=True)
354
+
355
+ assert jnp.all(jnp.equal(values.array, jnp.array([1, 3, 4])))
356
+ assert jnp.all(jnp.equal(counts.array, jnp.array([2, 2, 1])))
357
+
358
+ values, counts = hax.unique(named1, U2, axis=Height, return_counts=True)
359
+
360
+ assert jnp.all(jnp.equal(values.array, jnp.array([[1, 2], [2, 3]])))
361
+ assert jnp.all(jnp.equal(counts.array, jnp.array([2, 1])))
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev382"
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