haliax 1.4.dev294__tar.gz → 1.4.dev296__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 (98) hide show
  1. {haliax-1.4.dev294 → haliax-1.4.dev296}/PKG-INFO +2 -4
  2. {haliax-1.4.dev294 → haliax-1.4.dev296}/README.md +1 -3
  3. haliax-1.4.dev296/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/partitioning.py +1 -12
  5. haliax-1.4.dev294/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev294 → haliax-1.4.dev296}/.coveragerc +0 -0
  7. {haliax-1.4.dev294 → haliax-1.4.dev296}/.flake8 +0 -0
  8. {haliax-1.4.dev294 → haliax-1.4.dev296}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev294 → haliax-1.4.dev296}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev294 → haliax-1.4.dev296}/.github/workflows/run_tests.yaml +0 -0
  11. {haliax-1.4.dev294 → haliax-1.4.dev296}/.gitignore +0 -0
  12. {haliax-1.4.dev294 → haliax-1.4.dev296}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev294 → haliax-1.4.dev296}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev294 → haliax-1.4.dev296}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev294 → haliax-1.4.dev296}/LICENSE +0 -0
  16. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/api.md +0 -0
  17. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/broadcasting.md +0 -0
  18. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/cheatsheet.md +0 -0
  19. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/css/material.css +0 -0
  20. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/css/mkdocstrings.css +0 -0
  21. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/faq.md +0 -0
  22. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/data_parallel_mesh.png +0 -0
  23. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  24. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_1d.png +0 -0
  25. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_1d_zero.png +0 -0
  26. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_2d.png +0 -0
  27. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  28. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  29. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  30. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  31. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/figures/device_mesh_2d_zero.png +0 -0
  32. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/fp8.md +0 -0
  33. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/hof.md +0 -0
  34. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/index.md +0 -0
  35. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/indexing.md +0 -0
  36. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/matmul.md +0 -0
  37. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/nn.md +0 -0
  38. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/partitioning.md +0 -0
  39. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/rearrange.ipynb +0 -0
  40. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/rearrange.md +0 -0
  41. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/requirements.txt +0 -0
  42. {haliax-1.4.dev294 → haliax-1.4.dev296}/docs/tutorial.md +0 -0
  43. {haliax-1.4.dev294 → haliax-1.4.dev296}/mkdocs.yml +0 -0
  44. {haliax-1.4.dev294 → haliax-1.4.dev296}/pyproject.toml +0 -0
  45. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/__init__.py +0 -0
  46. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/__init__.py +0 -0
  47. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/compile_utils.py +0 -0
  48. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/dot.py +0 -0
  49. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/einsum.py +0 -0
  50. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/fp8.py +0 -0
  51. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/parsing.py +0 -0
  52. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/rearrange.py +0 -0
  53. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/_src/util.py +0 -0
  54. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/axis.py +0 -0
  55. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/core.py +0 -0
  56. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/debug.py +0 -0
  57. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/hof.py +0 -0
  58. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/jax_utils.py +0 -0
  59. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/__init__.py +0 -0
  60. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/activations.py +0 -0
  61. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/attention.py +0 -0
  62. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/conv.py +0 -0
  63. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/dropout.py +0 -0
  64. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/embedding.py +0 -0
  65. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/linear.py +0 -0
  66. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/loss.py +0 -0
  67. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/mlp.py +0 -0
  68. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/normalization.py +0 -0
  69. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/pool.py +0 -0
  70. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/nn/scan.py +0 -0
  71. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/ops.py +0 -0
  72. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/quantization.py +0 -0
  73. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/random.py +0 -0
  74. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/specialized_fns.py +0 -0
  75. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/tree_util.py +0 -0
  76. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/types.py +0 -0
  77. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/util.py +0 -0
  78. {haliax-1.4.dev294 → haliax-1.4.dev296}/src/haliax/wrap.py +0 -0
  79. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/core_test.py +0 -0
  80. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_attention.py +0 -0
  81. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_axis.py +0 -0
  82. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_conv.py +0 -0
  83. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_debug.py +0 -0
  84. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_dot.py +0 -0
  85. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_einsum.py +0 -0
  86. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_fp8.py +0 -0
  87. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_hof.py +0 -0
  88. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_nn.py +0 -0
  89. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_ops.py +0 -0
  90. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_parsing.py +0 -0
  91. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_partitioning.py +0 -0
  92. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_pool.py +0 -0
  93. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_random.py +0 -0
  94. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_rearrange.py +0 -0
  95. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_scan.py +0 -0
  96. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_specialized_fns.py +0 -0
  97. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_tree_util.py +0 -0
  98. {haliax-1.4.dev294 → haliax-1.4.dev296}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: haliax
3
- Version: 1.4.dev294
3
+ Version: 1.4.dev296
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/
@@ -67,8 +67,6 @@ please see the [Haliax tutorial](https://colab.research.google.com/drive/1TiTcQQ
67
67
  (We use the excellent [Equinox](https://github.com/patrick-kidger/equinox) library for its module system and tree transformations.)
68
68
 
69
69
  ```python
70
- import haliax.nn.normalization
71
- import haliax.nn.activations
72
70
  import equinox as eqx
73
71
  import jax
74
72
  import jax.numpy as jnp
@@ -90,7 +88,7 @@ def attention_scores(Key, KPos, query, key, mask):
90
88
  scores -= 1E9 * (1.0 - mask)
91
89
 
92
90
  # convert to probabilities
93
- scores = haliax.nn.normalization.softmax(scores, KPos)
91
+ scores = haliax.nn.softmax(scores, KPos)
94
92
  return scores
95
93
 
96
94
 
@@ -35,8 +35,6 @@ please see the [Haliax tutorial](https://colab.research.google.com/drive/1TiTcQQ
35
35
  (We use the excellent [Equinox](https://github.com/patrick-kidger/equinox) library for its module system and tree transformations.)
36
36
 
37
37
  ```python
38
- import haliax.nn.normalization
39
- import haliax.nn.activations
40
38
  import equinox as eqx
41
39
  import jax
42
40
  import jax.numpy as jnp
@@ -58,7 +56,7 @@ def attention_scores(Key, KPos, query, key, mask):
58
56
  scores -= 1E9 * (1.0 - mask)
59
57
 
60
58
  # convert to probabilities
61
- scores = haliax.nn.normalization.softmax(scores, KPos)
59
+ scores = haliax.nn.softmax(scores, KPos)
62
60
  return scores
63
61
 
64
62
 
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev296"
@@ -3,7 +3,6 @@ import functools
3
3
  import threading
4
4
  import typing
5
5
  import warnings
6
- from itertools import chain
7
6
  from math import prod
8
7
  from typing import Callable, ContextManager, Mapping, Optional, ParamSpec, Sequence, TypeVar, Union
9
8
 
@@ -586,17 +585,7 @@ def sharding_for_axis(
586
585
  def pspec_for_axis(axis: AxisSelection, mapping: Optional[ResourceMapping] = None) -> PartitionSpec:
587
586
  """Get the PartitionSpec for a single axis"""
588
587
  axis = ensure_tuple(axis)
589
- phys_axes = []
590
- for a in axis:
591
- pa = physical_axis_name(a, mapping)
592
- if pa is None or isinstance(pa, str):
593
- phys_axes.append(pa)
594
- else:
595
- # I have no way to resolve the mypy check :)
596
- for i in pa:
597
- phys_axes.append(i)
598
-
599
- return PartitionSpec(*phys_axes)
588
+ return PartitionSpec(*(physical_axis_name(a, mapping) for a in axis))
600
589
 
601
590
 
602
591
  def round_axis_for_partitioning(axis: Axis, mapping: Optional[ResourceMapping] = None) -> Axis:
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev294"
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