haliax 1.4.dev293__tar.gz → 1.4.dev295__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.dev293 → haliax-1.4.dev295}/PKG-INFO +2 -4
  2. {haliax-1.4.dev293 → haliax-1.4.dev295}/README.md +1 -3
  3. haliax-1.4.dev295/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/partitioning.py +12 -1
  5. haliax-1.4.dev293/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev293 → haliax-1.4.dev295}/.coveragerc +0 -0
  7. {haliax-1.4.dev293 → haliax-1.4.dev295}/.flake8 +0 -0
  8. {haliax-1.4.dev293 → haliax-1.4.dev295}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev293 → haliax-1.4.dev295}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev293 → haliax-1.4.dev295}/.github/workflows/run_tests.yaml +0 -0
  11. {haliax-1.4.dev293 → haliax-1.4.dev295}/.gitignore +0 -0
  12. {haliax-1.4.dev293 → haliax-1.4.dev295}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev293 → haliax-1.4.dev295}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev293 → haliax-1.4.dev295}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev293 → haliax-1.4.dev295}/LICENSE +0 -0
  16. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/api.md +0 -0
  17. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/broadcasting.md +0 -0
  18. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/cheatsheet.md +0 -0
  19. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/css/material.css +0 -0
  20. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/css/mkdocstrings.css +0 -0
  21. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/faq.md +0 -0
  22. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/data_parallel_mesh.png +0 -0
  23. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  24. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_1d.png +0 -0
  25. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_1d_zero.png +0 -0
  26. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_2d.png +0 -0
  27. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  28. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  29. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  30. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  31. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/figures/device_mesh_2d_zero.png +0 -0
  32. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/fp8.md +0 -0
  33. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/hof.md +0 -0
  34. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/index.md +0 -0
  35. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/indexing.md +0 -0
  36. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/matmul.md +0 -0
  37. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/nn.md +0 -0
  38. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/partitioning.md +0 -0
  39. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/rearrange.ipynb +0 -0
  40. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/rearrange.md +0 -0
  41. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/requirements.txt +0 -0
  42. {haliax-1.4.dev293 → haliax-1.4.dev295}/docs/tutorial.md +0 -0
  43. {haliax-1.4.dev293 → haliax-1.4.dev295}/mkdocs.yml +0 -0
  44. {haliax-1.4.dev293 → haliax-1.4.dev295}/pyproject.toml +0 -0
  45. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/__init__.py +0 -0
  46. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/__init__.py +0 -0
  47. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/compile_utils.py +0 -0
  48. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/dot.py +0 -0
  49. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/einsum.py +0 -0
  50. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/fp8.py +0 -0
  51. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/parsing.py +0 -0
  52. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/rearrange.py +0 -0
  53. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/_src/util.py +0 -0
  54. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/axis.py +0 -0
  55. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/core.py +0 -0
  56. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/debug.py +0 -0
  57. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/hof.py +0 -0
  58. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/jax_utils.py +0 -0
  59. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/__init__.py +0 -0
  60. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/activations.py +0 -0
  61. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/attention.py +0 -0
  62. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/conv.py +0 -0
  63. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/dropout.py +0 -0
  64. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/embedding.py +0 -0
  65. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/linear.py +0 -0
  66. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/loss.py +0 -0
  67. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/mlp.py +0 -0
  68. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/normalization.py +0 -0
  69. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/pool.py +0 -0
  70. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/nn/scan.py +0 -0
  71. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/ops.py +0 -0
  72. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/quantization.py +0 -0
  73. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/random.py +0 -0
  74. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/specialized_fns.py +0 -0
  75. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/tree_util.py +0 -0
  76. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/types.py +0 -0
  77. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/util.py +0 -0
  78. {haliax-1.4.dev293 → haliax-1.4.dev295}/src/haliax/wrap.py +0 -0
  79. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/core_test.py +0 -0
  80. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_attention.py +0 -0
  81. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_axis.py +0 -0
  82. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_conv.py +0 -0
  83. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_debug.py +0 -0
  84. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_dot.py +0 -0
  85. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_einsum.py +0 -0
  86. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_fp8.py +0 -0
  87. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_hof.py +0 -0
  88. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_nn.py +0 -0
  89. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_ops.py +0 -0
  90. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_parsing.py +0 -0
  91. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_partitioning.py +0 -0
  92. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_pool.py +0 -0
  93. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_random.py +0 -0
  94. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_rearrange.py +0 -0
  95. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_scan.py +0 -0
  96. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_specialized_fns.py +0 -0
  97. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_tree_util.py +0 -0
  98. {haliax-1.4.dev293 → haliax-1.4.dev295}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: haliax
3
- Version: 1.4.dev293
3
+ Version: 1.4.dev295
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.dev295"
@@ -3,6 +3,7 @@ import functools
3
3
  import threading
4
4
  import typing
5
5
  import warnings
6
+ from itertools import chain
6
7
  from math import prod
7
8
  from typing import Callable, ContextManager, Mapping, Optional, ParamSpec, Sequence, TypeVar, Union
8
9
 
@@ -585,7 +586,17 @@ def sharding_for_axis(
585
586
  def pspec_for_axis(axis: AxisSelection, mapping: Optional[ResourceMapping] = None) -> PartitionSpec:
586
587
  """Get the PartitionSpec for a single axis"""
587
588
  axis = ensure_tuple(axis)
588
- return PartitionSpec(*(physical_axis_name(a, mapping) for a in 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)
589
600
 
590
601
 
591
602
  def round_axis_for_partitioning(axis: Axis, mapping: Optional[ResourceMapping] = None) -> Axis:
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev293"
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