haliax 1.4.dev305__tar.gz → 1.4.dev306__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 (99) hide show
  1. {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/run_tests.yaml +1 -1
  2. {haliax-1.4.dev305 → haliax-1.4.dev306}/PKG-INFO +1 -1
  3. haliax-1.4.dev306/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/partitioning.py +13 -12
  5. haliax-1.4.dev305/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev305 → haliax-1.4.dev306}/.coveragerc +0 -0
  7. {haliax-1.4.dev305 → haliax-1.4.dev306}/.flake8 +0 -0
  8. {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev305 → haliax-1.4.dev306}/.gitignore +0 -0
  12. {haliax-1.4.dev305 → haliax-1.4.dev306}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev305 → haliax-1.4.dev306}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev305 → haliax-1.4.dev306}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev305 → haliax-1.4.dev306}/LICENSE +0 -0
  16. {haliax-1.4.dev305 → haliax-1.4.dev306}/README.md +0 -0
  17. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/api.md +0 -0
  18. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/broadcasting.md +0 -0
  19. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/cheatsheet.md +0 -0
  20. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/css/material.css +0 -0
  21. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/css/mkdocstrings.css +0 -0
  22. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/faq.md +0 -0
  23. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/data_parallel_mesh.png +0 -0
  24. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  25. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_1d.png +0 -0
  26. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_1d_zero.png +0 -0
  27. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d.png +0 -0
  28. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  29. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  30. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  31. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  32. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_zero.png +0 -0
  33. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/fp8.md +0 -0
  34. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/hof.md +0 -0
  35. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/index.md +0 -0
  36. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/indexing.md +0 -0
  37. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/matmul.md +0 -0
  38. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/nn.md +0 -0
  39. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/partitioning.md +0 -0
  40. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/rearrange.ipynb +0 -0
  41. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/rearrange.md +0 -0
  42. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/requirements.txt +0 -0
  43. {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/tutorial.md +0 -0
  44. {haliax-1.4.dev305 → haliax-1.4.dev306}/mkdocs.yml +0 -0
  45. {haliax-1.4.dev305 → haliax-1.4.dev306}/pyproject.toml +0 -0
  46. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/__init__.py +0 -0
  47. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/__init__.py +0 -0
  48. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/compile_utils.py +0 -0
  49. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/dot.py +0 -0
  50. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/einsum.py +0 -0
  51. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/fp8.py +0 -0
  52. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/parsing.py +0 -0
  53. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/rearrange.py +0 -0
  54. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/util.py +0 -0
  55. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/axis.py +0 -0
  56. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/core.py +0 -0
  57. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/debug.py +0 -0
  58. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/hof.py +0 -0
  59. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/jax_utils.py +0 -0
  60. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/__init__.py +0 -0
  61. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/activations.py +0 -0
  62. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/attention.py +0 -0
  63. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/conv.py +0 -0
  64. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/dropout.py +0 -0
  65. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/embedding.py +0 -0
  66. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/linear.py +0 -0
  67. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/loss.py +0 -0
  68. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/mlp.py +0 -0
  69. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/normalization.py +0 -0
  70. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/pool.py +0 -0
  71. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/scan.py +0 -0
  72. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/ops.py +0 -0
  73. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/quantization.py +0 -0
  74. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/random.py +0 -0
  75. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/specialized_fns.py +0 -0
  76. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/tree_util.py +0 -0
  77. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/types.py +0 -0
  78. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/util.py +0 -0
  79. {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/wrap.py +0 -0
  80. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/core_test.py +0 -0
  81. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_attention.py +0 -0
  82. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_axis.py +0 -0
  83. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_conv.py +0 -0
  84. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_debug.py +0 -0
  85. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_dot.py +0 -0
  86. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_einsum.py +0 -0
  87. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_fp8.py +0 -0
  88. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_hof.py +0 -0
  89. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_nn.py +0 -0
  90. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_ops.py +0 -0
  91. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_parsing.py +0 -0
  92. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_partitioning.py +0 -0
  93. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_pool.py +0 -0
  94. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_random.py +0 -0
  95. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_rearrange.py +0 -0
  96. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_scan.py +0 -0
  97. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_specialized_fns.py +0 -0
  98. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_tree_util.py +0 -0
  99. {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_utils.py +0 -0
@@ -17,7 +17,7 @@ jobs:
17
17
  run: |
18
18
  python -m pip install --upgrade pip
19
19
  pip install flake8 pytest
20
- pip install --upgrade "jax[cpu]==0.4.25" "jaxlib[cpu]==0.4.25"
20
+ pip install --upgrade "jax[cpu]==0.4.30" "jaxlib[cpu]==0.4.30"
21
21
  pip install .[dev]
22
22
  - name: Test with pytest
23
23
  run: |
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: haliax
3
- Version: 1.4.dev305
3
+ Version: 1.4.dev306
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.dev306"
@@ -128,25 +128,26 @@ def shard(x: T, mapping: Optional[ResourceMapping] = None, mesh: Optional[Mesh]
128
128
  warnings.warn("Sharding constraints are not supported in jit on metal", RuntimeWarning)
129
129
  return x
130
130
 
131
- def _do_device_put(x):
132
- if not isinstance(x, NamedArray):
133
- return x
131
+ def _do_device_put(named):
132
+ if not isinstance(named, NamedArray):
133
+ return named
134
134
 
135
- if not is_jax_array_like(x.array):
135
+ if not is_jax_array_like(named.array):
136
136
  # this happens when we filter out params for things like lora.
137
137
  # could use eqx.partition to avoid this, but eh
138
- return x
138
+ return named
139
139
 
140
- sharding = infer_resource_partitions(x, mapping, mesh=mesh, preserve_existing_shardings=False)
140
+ sharding = infer_resource_partitions(named, mapping, mesh=mesh, preserve_existing_shardings=False)
141
141
  assert isinstance(sharding, NamedSharding)
142
142
  if is_in_jit():
143
- return with_sharding_constraint(x, sharding)
144
- elif sharding.is_fully_addressable:
145
- sharded_array = jax.device_put(x.array, sharding)
146
- return NamedArray(sharded_array, x.axes)
143
+ return with_sharding_constraint(named, sharding)
144
+ # as a special case, SingleDeviceShardings are routed through jit
145
+ elif isinstance(named.array.sharding, SingleDeviceSharding):
146
+ # TODO(dlwh): this should be unnecessary in JAX soon. Check after 2024-08-01
147
+ sharded_array = jax.jit(lambda x: x, out_shardings=sharding)(named)
148
+ return sharded_array
147
149
  else:
148
- # sharded_array = jax.device_put(x.array, sharding)
149
- ret = eqx.filter_jit(lambda x: with_sharding_constraint(x, sharding))(x)
150
+ ret = jax.device_put(named, sharding)
150
151
  return ret
151
152
 
152
153
  return htu.tree_map(_do_device_put, x)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev305"
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