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