haliax 1.4.dev303__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.dev303 → haliax-1.4.dev306}/.github/workflows/run_tests.yaml +1 -1
  2. {haliax-1.4.dev303 → haliax-1.4.dev306}/PKG-INFO +1 -1
  3. haliax-1.4.dev306/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/partitioning.py +15 -10
  5. haliax-1.4.dev303/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev303 → haliax-1.4.dev306}/.coveragerc +0 -0
  7. {haliax-1.4.dev303 → haliax-1.4.dev306}/.flake8 +0 -0
  8. {haliax-1.4.dev303 → haliax-1.4.dev306}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev303 → haliax-1.4.dev306}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev303 → haliax-1.4.dev306}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev303 → haliax-1.4.dev306}/.gitignore +0 -0
  12. {haliax-1.4.dev303 → haliax-1.4.dev306}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev303 → haliax-1.4.dev306}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev303 → haliax-1.4.dev306}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev303 → haliax-1.4.dev306}/LICENSE +0 -0
  16. {haliax-1.4.dev303 → haliax-1.4.dev306}/README.md +0 -0
  17. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/api.md +0 -0
  18. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/broadcasting.md +0 -0
  19. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/cheatsheet.md +0 -0
  20. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/css/material.css +0 -0
  21. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/css/mkdocstrings.css +0 -0
  22. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/faq.md +0 -0
  23. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/data_parallel_mesh.png +0 -0
  24. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  25. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_1d.png +0 -0
  26. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_1d_zero.png +0 -0
  27. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_2d.png +0 -0
  28. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  29. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  30. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  31. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  32. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_zero.png +0 -0
  33. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/fp8.md +0 -0
  34. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/hof.md +0 -0
  35. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/index.md +0 -0
  36. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/indexing.md +0 -0
  37. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/matmul.md +0 -0
  38. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/nn.md +0 -0
  39. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/partitioning.md +0 -0
  40. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/rearrange.ipynb +0 -0
  41. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/rearrange.md +0 -0
  42. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/requirements.txt +0 -0
  43. {haliax-1.4.dev303 → haliax-1.4.dev306}/docs/tutorial.md +0 -0
  44. {haliax-1.4.dev303 → haliax-1.4.dev306}/mkdocs.yml +0 -0
  45. {haliax-1.4.dev303 → haliax-1.4.dev306}/pyproject.toml +0 -0
  46. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/__init__.py +0 -0
  47. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/__init__.py +0 -0
  48. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/compile_utils.py +0 -0
  49. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/dot.py +0 -0
  50. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/einsum.py +0 -0
  51. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/fp8.py +0 -0
  52. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/parsing.py +0 -0
  53. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/rearrange.py +0 -0
  54. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/_src/util.py +0 -0
  55. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/axis.py +0 -0
  56. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/core.py +0 -0
  57. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/debug.py +0 -0
  58. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/hof.py +0 -0
  59. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/jax_utils.py +0 -0
  60. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/__init__.py +0 -0
  61. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/activations.py +0 -0
  62. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/attention.py +0 -0
  63. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/conv.py +0 -0
  64. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/dropout.py +0 -0
  65. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/embedding.py +0 -0
  66. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/linear.py +0 -0
  67. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/loss.py +0 -0
  68. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/mlp.py +0 -0
  69. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/normalization.py +0 -0
  70. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/pool.py +0 -0
  71. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/nn/scan.py +0 -0
  72. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/ops.py +0 -0
  73. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/quantization.py +0 -0
  74. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/random.py +0 -0
  75. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/specialized_fns.py +0 -0
  76. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/tree_util.py +0 -0
  77. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/types.py +0 -0
  78. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/util.py +0 -0
  79. {haliax-1.4.dev303 → haliax-1.4.dev306}/src/haliax/wrap.py +0 -0
  80. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/core_test.py +0 -0
  81. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_attention.py +0 -0
  82. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_axis.py +0 -0
  83. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_conv.py +0 -0
  84. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_debug.py +0 -0
  85. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_dot.py +0 -0
  86. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_einsum.py +0 -0
  87. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_fp8.py +0 -0
  88. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_hof.py +0 -0
  89. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_nn.py +0 -0
  90. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_ops.py +0 -0
  91. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_parsing.py +0 -0
  92. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_partitioning.py +0 -0
  93. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_pool.py +0 -0
  94. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_random.py +0 -0
  95. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_rearrange.py +0 -0
  96. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_scan.py +0 -0
  97. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_specialized_fns.py +0 -0
  98. {haliax-1.4.dev303 → haliax-1.4.dev306}/tests/test_tree_util.py +0 -0
  99. {haliax-1.4.dev303 → 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.dev303
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,22 +128,27 @@ 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)
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
144
149
  else:
145
- sharded_array = jax.device_put(x.array, sharding)
146
- return NamedArray(sharded_array, x.axes)
150
+ ret = jax.device_put(named, sharding)
151
+ return ret
147
152
 
148
153
  return htu.tree_map(_do_device_put, x)
149
154
 
@@ -309,7 +314,7 @@ class _NamedJitWrapper(eqx.Module):
309
314
  output_shape = _cached_filter_eval_shape(self._fn, *args, **kwargs)
310
315
  my_pjit_args = dict(**self._pjit_args)
311
316
 
312
- if in_axis_resources is not None or axis_resources is not None:
317
+ if in_axis_resources is not None:
313
318
  in_resources = infer_resource_partitions(
314
319
  (dynamic_donated, dynamic_reserved),
315
320
  in_axis_resources,
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev303"
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