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