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.
Files changed (99) hide show
  1. {haliax-1.4.dev307 → haliax-1.4.dev308}/PKG-INFO +1 -1
  2. haliax-1.4.dev308/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/__init__.py +2 -1
  4. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/loss.py +15 -0
  5. haliax-1.4.dev307/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev307 → haliax-1.4.dev308}/.coveragerc +0 -0
  7. {haliax-1.4.dev307 → haliax-1.4.dev308}/.flake8 +0 -0
  8. {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev307 → haliax-1.4.dev308}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev307 → haliax-1.4.dev308}/.gitignore +0 -0
  13. {haliax-1.4.dev307 → haliax-1.4.dev308}/.pre-commit-config.yaml +0 -0
  14. {haliax-1.4.dev307 → haliax-1.4.dev308}/.readthedocs.yaml +0 -0
  15. {haliax-1.4.dev307 → haliax-1.4.dev308}/CONTRIBUTING.md +0 -0
  16. {haliax-1.4.dev307 → haliax-1.4.dev308}/LICENSE +0 -0
  17. {haliax-1.4.dev307 → haliax-1.4.dev308}/README.md +0 -0
  18. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/api.md +0 -0
  19. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/broadcasting.md +0 -0
  20. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/cheatsheet.md +0 -0
  21. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/css/material.css +0 -0
  22. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/css/mkdocstrings.css +0 -0
  23. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/faq.md +0 -0
  24. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/data_parallel_mesh.png +0 -0
  25. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  26. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_1d.png +0 -0
  27. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_1d_zero.png +0 -0
  28. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d.png +0 -0
  29. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  30. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  31. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  32. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  33. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/figures/device_mesh_2d_zero.png +0 -0
  34. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/fp8.md +0 -0
  35. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/hof.md +0 -0
  36. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/index.md +0 -0
  37. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/indexing.md +0 -0
  38. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/matmul.md +0 -0
  39. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/nn.md +0 -0
  40. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/partitioning.md +0 -0
  41. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/rearrange.ipynb +0 -0
  42. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/rearrange.md +0 -0
  43. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/requirements.txt +0 -0
  44. {haliax-1.4.dev307 → haliax-1.4.dev308}/docs/tutorial.md +0 -0
  45. {haliax-1.4.dev307 → haliax-1.4.dev308}/mkdocs.yml +0 -0
  46. {haliax-1.4.dev307 → haliax-1.4.dev308}/pyproject.toml +0 -0
  47. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/__init__.py +0 -0
  48. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/__init__.py +0 -0
  49. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/compile_utils.py +0 -0
  50. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/dot.py +0 -0
  51. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/einsum.py +0 -0
  52. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/fp8.py +0 -0
  53. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/parsing.py +0 -0
  54. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/rearrange.py +0 -0
  55. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/_src/util.py +0 -0
  56. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/axis.py +0 -0
  57. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/core.py +0 -0
  58. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/debug.py +0 -0
  59. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/hof.py +0 -0
  60. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/jax_utils.py +0 -0
  61. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/activations.py +0 -0
  62. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/attention.py +0 -0
  63. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/conv.py +0 -0
  64. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/dropout.py +0 -0
  65. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/embedding.py +0 -0
  66. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/linear.py +0 -0
  67. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/mlp.py +0 -0
  68. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/normalization.py +0 -0
  69. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/pool.py +0 -0
  70. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/nn/scan.py +0 -0
  71. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/ops.py +0 -0
  72. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/partitioning.py +0 -0
  73. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/quantization.py +0 -0
  74. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/random.py +0 -0
  75. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/specialized_fns.py +0 -0
  76. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/tree_util.py +0 -0
  77. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/types.py +0 -0
  78. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/util.py +0 -0
  79. {haliax-1.4.dev307 → haliax-1.4.dev308}/src/haliax/wrap.py +0 -0
  80. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/core_test.py +0 -0
  81. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_attention.py +0 -0
  82. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_axis.py +0 -0
  83. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_conv.py +0 -0
  84. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_debug.py +0 -0
  85. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_dot.py +0 -0
  86. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_einsum.py +0 -0
  87. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_fp8.py +0 -0
  88. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_hof.py +0 -0
  89. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_nn.py +0 -0
  90. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_ops.py +0 -0
  91. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_parsing.py +0 -0
  92. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_partitioning.py +0 -0
  93. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_pool.py +0 -0
  94. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_random.py +0 -0
  95. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_rearrange.py +0 -0
  96. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_scan.py +0 -0
  97. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_specialized_fns.py +0 -0
  98. {haliax-1.4.dev307 → haliax-1.4.dev308}/tests/test_tree_util.py +0 -0
  99. {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.dev307
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