haliax 1.4.dev298__tar.gz → 1.4.dev302__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.dev298 → haliax-1.4.dev302}/.github/workflows/run_pre_commit.yaml +1 -1
  2. haliax-1.4.dev302/.github/workflows/run_quick_levanter_tests.yaml +36 -0
  3. {haliax-1.4.dev298 → haliax-1.4.dev302}/.github/workflows/run_tests.yaml +3 -6
  4. {haliax-1.4.dev298 → haliax-1.4.dev302}/PKG-INFO +1 -1
  5. haliax-1.4.dev302/src/haliax/__about__.py +1 -0
  6. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/partitioning.py +5 -3
  7. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_partitioning.py +24 -2
  8. haliax-1.4.dev298/src/haliax/__about__.py +0 -1
  9. {haliax-1.4.dev298 → haliax-1.4.dev302}/.coveragerc +0 -0
  10. {haliax-1.4.dev298 → haliax-1.4.dev302}/.flake8 +0 -0
  11. {haliax-1.4.dev298 → haliax-1.4.dev302}/.github/workflows/publish_dev.yaml +0 -0
  12. {haliax-1.4.dev298 → haliax-1.4.dev302}/.gitignore +0 -0
  13. {haliax-1.4.dev298 → haliax-1.4.dev302}/.pre-commit-config.yaml +0 -0
  14. {haliax-1.4.dev298 → haliax-1.4.dev302}/.readthedocs.yaml +0 -0
  15. {haliax-1.4.dev298 → haliax-1.4.dev302}/CONTRIBUTING.md +0 -0
  16. {haliax-1.4.dev298 → haliax-1.4.dev302}/LICENSE +0 -0
  17. {haliax-1.4.dev298 → haliax-1.4.dev302}/README.md +0 -0
  18. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/api.md +0 -0
  19. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/broadcasting.md +0 -0
  20. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/cheatsheet.md +0 -0
  21. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/css/material.css +0 -0
  22. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/css/mkdocstrings.css +0 -0
  23. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/faq.md +0 -0
  24. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/data_parallel_mesh.png +0 -0
  25. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  26. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_1d.png +0 -0
  27. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_1d_zero.png +0 -0
  28. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_2d.png +0 -0
  29. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  30. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  31. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  32. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  33. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/figures/device_mesh_2d_zero.png +0 -0
  34. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/fp8.md +0 -0
  35. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/hof.md +0 -0
  36. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/index.md +0 -0
  37. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/indexing.md +0 -0
  38. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/matmul.md +0 -0
  39. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/nn.md +0 -0
  40. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/partitioning.md +0 -0
  41. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/rearrange.ipynb +0 -0
  42. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/rearrange.md +0 -0
  43. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/requirements.txt +0 -0
  44. {haliax-1.4.dev298 → haliax-1.4.dev302}/docs/tutorial.md +0 -0
  45. {haliax-1.4.dev298 → haliax-1.4.dev302}/mkdocs.yml +0 -0
  46. {haliax-1.4.dev298 → haliax-1.4.dev302}/pyproject.toml +0 -0
  47. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/__init__.py +1 -1
  48. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/__init__.py +0 -0
  49. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/compile_utils.py +0 -0
  50. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/dot.py +0 -0
  51. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/einsum.py +0 -0
  52. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/fp8.py +0 -0
  53. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/parsing.py +0 -0
  54. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/rearrange.py +0 -0
  55. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/_src/util.py +0 -0
  56. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/axis.py +0 -0
  57. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/core.py +0 -0
  58. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/debug.py +0 -0
  59. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/hof.py +0 -0
  60. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/jax_utils.py +0 -0
  61. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/__init__.py +0 -0
  62. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/activations.py +0 -0
  63. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/attention.py +0 -0
  64. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/conv.py +0 -0
  65. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/dropout.py +0 -0
  66. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/embedding.py +0 -0
  67. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/linear.py +0 -0
  68. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/loss.py +0 -0
  69. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/mlp.py +0 -0
  70. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/normalization.py +0 -0
  71. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/pool.py +0 -0
  72. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/nn/scan.py +0 -0
  73. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/ops.py +0 -0
  74. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/quantization.py +0 -0
  75. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/random.py +0 -0
  76. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/specialized_fns.py +0 -0
  77. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/tree_util.py +0 -0
  78. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/types.py +0 -0
  79. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/util.py +0 -0
  80. {haliax-1.4.dev298 → haliax-1.4.dev302}/src/haliax/wrap.py +0 -0
  81. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/core_test.py +0 -0
  82. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_attention.py +0 -0
  83. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_axis.py +0 -0
  84. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_conv.py +0 -0
  85. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_debug.py +0 -0
  86. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_dot.py +0 -0
  87. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_einsum.py +0 -0
  88. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_fp8.py +0 -0
  89. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_hof.py +0 -0
  90. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_nn.py +0 -0
  91. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_ops.py +0 -0
  92. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_parsing.py +0 -0
  93. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_pool.py +0 -0
  94. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_random.py +0 -0
  95. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_rearrange.py +0 -0
  96. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_scan.py +0 -0
  97. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_specialized_fns.py +0 -0
  98. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_tree_util.py +0 -0
  99. {haliax-1.4.dev298 → haliax-1.4.dev302}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  name: Pre-Commit
2
2
 
3
- on: [push]
3
+ on: [push, pull_request]
4
4
 
5
5
  jobs:
6
6
  build:
@@ -0,0 +1,36 @@
1
+ name: Run Levanter Tests
2
+
3
+ on: [pull_request]
4
+
5
+ jobs:
6
+ build:
7
+
8
+ runs-on: ubuntu-latest
9
+
10
+ steps:
11
+ - uses: actions/checkout@v3
12
+ - name: Set up Python 3.10.11
13
+ uses: actions/setup-python@v4
14
+ with:
15
+ python-version: 3.10.11
16
+ - name: Install dependencies
17
+ run: |
18
+ python -m pip install --upgrade pip
19
+ pip install flake8 pytest
20
+ pip install --upgrade "jax[cpu]==0.4.26" "jaxlib[cpu]==0.4.26"
21
+
22
+ - name: Install Levanter from source
23
+ run: |
24
+ cd ..
25
+ git clone https://github.com/stanford-crfm/levanter.git
26
+ cd levanter
27
+ pip install -e .
28
+ - name: Install Haliax on top
29
+ run: |
30
+ # install second since levanter will install a built version of haliax
31
+ cd ../haliax
32
+ pip install .[dev]
33
+ - name: Test levanter with pytest
34
+ run: |
35
+ cd ../levanter
36
+ XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:../src pytest tests -m "not entry and not slow"
@@ -1,21 +1,18 @@
1
1
  name: Run Tests
2
2
 
3
- on: [push]
3
+ on: [push, pull_request]
4
4
 
5
5
  jobs:
6
6
  build:
7
7
 
8
8
  runs-on: ubuntu-latest
9
- strategy:
10
- matrix:
11
- python-version: ["3.10.11"]
12
9
 
13
10
  steps:
14
11
  - uses: actions/checkout@v3
15
- - name: Set up Python ${{ matrix.python-version }}
12
+ - name: Set up Python 3.10.11
16
13
  uses: actions/setup-python@v4
17
14
  with:
18
- python-version: ${{ matrix.python-version }}
15
+ python-version: 3.10.11
19
16
  - name: Install dependencies
20
17
  run: |
21
18
  python -m pip install --upgrade pip
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: haliax
3
- Version: 1.4.dev298
3
+ Version: 1.4.dev302
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.dev302"
@@ -141,9 +141,13 @@ def shard(x: T, mapping: Optional[ResourceMapping] = None, mesh: Optional[Mesh]
141
141
  assert isinstance(sharding, NamedSharding)
142
142
  if is_in_jit():
143
143
  return with_sharding_constraint(x, sharding)
144
- else:
144
+ elif sharding.is_fully_addressable:
145
145
  sharded_array = jax.device_put(x.array, sharding)
146
146
  return NamedArray(sharded_array, x.axes)
147
+ else:
148
+ # sharded_array = jax.device_put(x.array, sharding)
149
+ ret = eqx.filter_jit(lambda x: with_sharding_constraint(x, sharding))(x)
150
+ return ret
147
151
 
148
152
  return htu.tree_map(_do_device_put, x)
149
153
 
@@ -310,8 +314,6 @@ class _NamedJitWrapper(eqx.Module):
310
314
  my_pjit_args = dict(**self._pjit_args)
311
315
 
312
316
  if in_axis_resources is not None or axis_resources is not None:
313
- if in_axis_resources is None:
314
- in_axis_resources = axis_resources
315
317
  in_resources = infer_resource_partitions(
316
318
  (dynamic_donated, dynamic_reserved),
317
319
  in_axis_resources,
@@ -99,7 +99,7 @@ def test_pjit_class_init_with_args():
99
99
 
100
100
  devices = jax.devices()
101
101
  with Mesh(np.array(devices).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL)):
102
- mod = named_jit(ModWithArgs)(hax.ones((Dim1, Dim2)))
102
+ mod = named_jit(ModWithArgs)(hax.shard(hax.ones((Dim1, Dim2))))
103
103
  assert isinstance(mod, ModWithArgs)
104
104
  assert mod.array.array.shape == (Dim1.size, Dim2.size)
105
105
  assert mod.array2.array.shape == (Dim3.size,)
@@ -173,7 +173,7 @@ def test_shard_with_axis_mapping_inside_jit():
173
173
 
174
174
  jax.debug.inspect_array_sharding(arr.array, callback=lambda x: assert_eq(x, expected))
175
175
 
176
- @named_jit(in_axis_resources={}, out_axis_resources=resource_map)
176
+ @named_jit(out_axis_resources=resource_map)
177
177
  def do_shard(x, y):
178
178
  x = hax.shard(x, resource_map)
179
179
  assert_inside_pjit(x, NamedSharding(mesh, PartitionSpec(None, ResourceAxis.DATA)))
@@ -293,3 +293,25 @@ def test_cross_device_sharding():
293
293
  z_devices = z.array.devices()
294
294
 
295
295
  assert set(d.platform for d in x_devices) == set(d.platform for d in z_devices)
296
+
297
+
298
+ def test_named_jit_no_in_axis_resources():
299
+ mesh = Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
300
+ with axis_mapping(resource_map), mesh:
301
+
302
+ class MyModule(eqx.Module):
303
+ array: NamedArray
304
+
305
+ def __init__(self):
306
+ self.array = hax.ones((Dim1, Dim2))
307
+
308
+ data = hax.ones((Dim1, Dim2))
309
+ data = hax.shard(data, {})
310
+
311
+ @named_jit(axis_resources=resource_map)
312
+ def fn(data):
313
+ mod = MyModule()
314
+ return mod.array
315
+
316
+ r = fn(data)
317
+ assert r.array.sharding.device_set == set(jax.devices())
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev298"
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
@@ -56,8 +56,8 @@ from .core import (
56
56
  unflatten_axis,
57
57
  updated_slice,
58
58
  )
59
- from .jax_utils import filter_checkpoint
60
59
  from .hof import fold, map, scan, vmap
60
+ from .jax_utils import filter_checkpoint
61
61
  from .ops import clip, isclose, pad_left, trace, tril, triu, where
62
62
  from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
63
63
  from .specialized_fns import top_k