haliax 1.4.dev354__tar.gz → 1.4.dev355__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 (107) hide show
  1. {haliax-1.4.dev354 → haliax-1.4.dev355}/PKG-INFO +1 -1
  2. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/indexing.md +39 -0
  3. haliax-1.4.dev355/src/haliax/__about__.py +1 -0
  4. haliax-1.4.dev354/src/haliax/__about__.py +0 -1
  5. {haliax-1.4.dev354 → haliax-1.4.dev355}/.coveragerc +0 -0
  6. {haliax-1.4.dev354 → haliax-1.4.dev355}/.flake8 +0 -0
  7. {haliax-1.4.dev354 → haliax-1.4.dev355}/.github/workflows/publish_dev.yaml +0 -0
  8. {haliax-1.4.dev354 → haliax-1.4.dev355}/.github/workflows/run_pre_commit.yaml +0 -0
  9. {haliax-1.4.dev354 → haliax-1.4.dev355}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  10. {haliax-1.4.dev354 → haliax-1.4.dev355}/.github/workflows/run_tests.yaml +0 -0
  11. {haliax-1.4.dev354 → haliax-1.4.dev355}/.gitignore +0 -0
  12. {haliax-1.4.dev354 → haliax-1.4.dev355}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev354 → haliax-1.4.dev355}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev354 → haliax-1.4.dev355}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev354 → haliax-1.4.dev355}/LICENSE +0 -0
  16. {haliax-1.4.dev354 → haliax-1.4.dev355}/README.md +0 -0
  17. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/api.md +0 -0
  18. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/broadcasting.md +0 -0
  19. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/cheatsheet.md +0 -0
  20. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/css/material.css +0 -0
  21. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/css/mkdocstrings.css +0 -0
  22. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/faq.md +0 -0
  23. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/data_parallel_mesh.png +0 -0
  24. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  25. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_1d.png +0 -0
  26. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_1d_zero.png +0 -0
  27. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_2d.png +0 -0
  28. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  29. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  30. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  31. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  32. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/figures/device_mesh_2d_zero.png +0 -0
  33. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/fp8.md +0 -0
  34. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/index.md +0 -0
  35. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/matmul.md +0 -0
  36. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/nn.md +0 -0
  37. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/partitioning.md +0 -0
  38. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/rearrange.ipynb +0 -0
  39. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/rearrange.md +0 -0
  40. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/requirements.txt +0 -0
  41. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/scan.md +0 -0
  42. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/state-dict.md +0 -0
  43. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/tutorial.md +0 -0
  44. {haliax-1.4.dev354 → haliax-1.4.dev355}/docs/vmap.md +0 -0
  45. {haliax-1.4.dev354 → haliax-1.4.dev355}/mkdocs.yml +0 -0
  46. {haliax-1.4.dev354 → haliax-1.4.dev355}/pyproject.toml +0 -0
  47. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/__init__.py +0 -0
  48. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/__init__.py +0 -0
  49. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/compile_utils.py +0 -0
  50. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/dot.py +0 -0
  51. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/einsum.py +0 -0
  52. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/fp8.py +0 -0
  53. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/parsing.py +0 -0
  54. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/rearrange.py +0 -0
  55. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/scan.py +0 -0
  56. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/state_dict.py +0 -0
  57. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/_src/util.py +0 -0
  58. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/axis.py +0 -0
  59. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/core.py +0 -0
  60. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/debug.py +0 -0
  61. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/hof.py +0 -0
  62. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/jax_utils.py +0 -0
  63. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/__init__.py +0 -0
  64. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/activations.py +0 -0
  65. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/attention.py +0 -0
  66. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/conv.py +0 -0
  67. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/dropout.py +0 -0
  68. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/embedding.py +0 -0
  69. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/linear.py +0 -0
  70. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/loss.py +0 -0
  71. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/mlp.py +0 -0
  72. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/normalization.py +0 -0
  73. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/pool.py +0 -0
  74. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/nn/scan.py +0 -0
  75. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/ops.py +0 -0
  76. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/partitioning.py +0 -0
  77. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/quantization.py +0 -0
  78. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/random.py +0 -0
  79. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/specialized_fns.py +0 -0
  80. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/state_dict.py +0 -0
  81. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/tree_util.py +0 -0
  82. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/types.py +0 -0
  83. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/util.py +0 -0
  84. {haliax-1.4.dev354 → haliax-1.4.dev355}/src/haliax/wrap.py +0 -0
  85. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/core_test.py +0 -0
  86. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_attention.py +0 -0
  87. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_axis.py +0 -0
  88. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_conv.py +0 -0
  89. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_debug.py +0 -0
  90. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_dot.py +0 -0
  91. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_einsum.py +0 -0
  92. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_fp8.py +0 -0
  93. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_hof.py +0 -0
  94. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_int8.py +0 -0
  95. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_nn.py +0 -0
  96. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_ops.py +0 -0
  97. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_parsing.py +0 -0
  98. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_partitioning.py +0 -0
  99. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_pool.py +0 -0
  100. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_random.py +0 -0
  101. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_rearrange.py +0 -0
  102. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_scan.py +0 -0
  103. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_scatter_gather.py +0 -0
  104. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_specialized_fns.py +0 -0
  105. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_state_dict.py +0 -0
  106. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_tree_util.py +0 -0
  107. {haliax-1.4.dev354 → haliax-1.4.dev355}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev354
3
+ Version: 1.4.dev355
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/
@@ -287,3 +287,42 @@ operation more effectively.)
287
287
 
288
288
  It's worth emphasizing that these functions are typically compiled to scatter-add and friends (as appropriate).
289
289
  This is the preferred way to do scatter/gather operations in JAX, as well as in Haliax.
290
+
291
+ ## Scatter/Gather
292
+
293
+ Haliax supports scatter/gather semantics in its indexing operations. When an axis
294
+ is indexed by another NamedArray (or a 1-D JAX array), the values of that axis
295
+ are gathered according to the index array and the axes of the indexer are
296
+ inserted into the result.
297
+
298
+ ```python
299
+ import haliax as hax
300
+ import jax.numpy as jnp
301
+
302
+ B, S, V = Axis("batch", 4), Axis("seq", 3), Axis("vocab", 7)
303
+ x = hax.arange((B, S, V))
304
+ idx = hax.arange((B, S), dtype=jnp.int32) % V.size
305
+
306
+ out = x["vocab", idx]
307
+ ```
308
+
309
+ Here `out` has axes `(B, S)` and its values match `jax.numpy.take_along_axis`
310
+ on the underlying ndarray.
311
+
312
+ For scatter-style updates where each batch writes to a different position, use
313
+ [`updated_slice`][haliax.updated_slice]:
314
+
315
+ ```python
316
+ Batch = hax.Axis("batch", 2)
317
+ Seq = hax.Axis("seq", 5)
318
+ New = hax.Axis("seq", 2)
319
+
320
+ cache = hax.zeros((Batch, Seq), dtype=int)
321
+ lengths = hax.named([1, 3], axis=Batch)
322
+ kv = hax.named([[1, 2], [3, 4]], axis=(Batch, New))
323
+
324
+ result = updated_slice(cache, {"seq": lengths}, kv)
325
+ ```
326
+
327
+ This inserts `[1, 2]` starting at position `1` in batch `0` and `[3, 4]` starting
328
+ at position `3` in batch `1`.
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev355"
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev354"
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
File without changes