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.
- {haliax-1.4.dev404 → haliax-1.4.dev405}/PKG-INFO +1 -1
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/api.md +1 -0
- haliax-1.4.dev405/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/__init__.py +2 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/ops.py +29 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/core_test.py +20 -0
- haliax-1.4.dev404/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.coveragerc +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.flake8 +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.gitignore +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/AGENTS.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/LICENSE +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/README.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/css/material.css +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/faq.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/fp8.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/index.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/indexing.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/matmul.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/nn.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/partitioning.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/primer.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/rearrange.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/requirements.txt +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/scan.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/state-dict.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/tutorial.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/typing.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/docs/vmap.md +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/mkdocs.yml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/pyproject.toml +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/core.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/field.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/random.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/types.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/util.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_attention.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_axis.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_conv.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_debug.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_dot.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_field.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_hof.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_int8.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_nn.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_ops.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_pool.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_random.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_scan.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_utils.py +0 -0
- {haliax-1.4.dev404 → haliax-1.4.dev405}/tests/test_visualize_sharding.py +0 -0
- {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.
|
|
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/
|
|
@@ -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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|