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