haliax 1.4.dev307__tar.gz → 1.4.dev308__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.dev307 → haliax-1.4.dev308}/PKG-INFO +1 -1
- haliax-1.4.dev308/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/__init__.py +2 -1
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/loss.py +15 -0
- haliax-1.4.dev307/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.coveragerc +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.flake8 +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.gitignore +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/LICENSE +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/README.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/api.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/css/material.css +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/faq.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/fp8.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/hof.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/index.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/indexing.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/matmul.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/nn.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/partitioning.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/rearrange.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/requirements.txt +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/tutorial.md +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/mkdocs.yml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/pyproject.toml +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/core.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/random.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/types.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/util.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/core_test.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_attention.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_axis.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_conv.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_debug.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_dot.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_hof.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_nn.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_ops.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_pool.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_random.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_scan.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev308
|
|
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.dev308"
|
|
@@ -34,7 +34,7 @@ from .conv import Conv, ConvTranspose
|
|
|
34
34
|
from .dropout import Dropout, dropout
|
|
35
35
|
from .embedding import Embedding
|
|
36
36
|
from .linear import Linear
|
|
37
|
-
from .loss import binary_cross_entropy_loss, cross_entropy_loss, cross_entropy_loss_and_log_normalizers
|
|
37
|
+
from .loss import binary_cross_entropy_loss, cross_entropy_loss, cross_entropy_loss_and_log_normalizers, reduce_loss
|
|
38
38
|
from .mlp import MLP
|
|
39
39
|
from .normalization import LayerNorm, log_softmax, logsumexp, softmax, standardize
|
|
40
40
|
from .pool import max_pool, mean_pool, min_pool
|
|
@@ -77,6 +77,7 @@ __all__ = [
|
|
|
77
77
|
"attention",
|
|
78
78
|
"one_hot",
|
|
79
79
|
"binary_cross_entropy_loss",
|
|
80
|
+
"reduce_loss",
|
|
80
81
|
"cross_entropy_loss",
|
|
81
82
|
"cross_entropy_loss_and_log_normalizers",
|
|
82
83
|
"Conv",
|
|
@@ -94,6 +94,21 @@ def binary_cross_entropy_loss(
|
|
|
94
94
|
return loss
|
|
95
95
|
|
|
96
96
|
|
|
97
|
+
def reduce_loss(
|
|
98
|
+
arr,
|
|
99
|
+
reduction: Optional[ReductionFunction] | Unspecified = UNSPECIFIED,
|
|
100
|
+
reduction_axis: Optional[AxisSelection] = None,
|
|
101
|
+
where: Optional[NamedArray] = None,
|
|
102
|
+
):
|
|
103
|
+
"""
|
|
104
|
+
Reduce a loss array according to the given reduction and reduction axis.
|
|
105
|
+
If reduction is None, the loss is not reduced.
|
|
106
|
+
If reduction is UNSPECIFIED, the default reduction is used (mean).
|
|
107
|
+
If reduction_axis is None (default), the loss is reduced over all axes.
|
|
108
|
+
"""
|
|
109
|
+
return maybe_reduce_loss(arr, reduction, reduction_axis, where)
|
|
110
|
+
|
|
111
|
+
|
|
97
112
|
def maybe_reduce_loss(
|
|
98
113
|
arr,
|
|
99
114
|
reduction: Optional[ReductionFunction] | Unspecified,
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev307"
|
|
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
|