haliax 1.4.dev368__tar.gz → 1.4.dev371__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 (114) hide show
  1. {haliax-1.4.dev368 → haliax-1.4.dev371}/.github/workflows/run_tests.yaml +3 -4
  2. {haliax-1.4.dev368 → haliax-1.4.dev371}/AGENTS.md +8 -4
  3. {haliax-1.4.dev368 → haliax-1.4.dev371}/PKG-INFO +1 -1
  4. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/scan.md +11 -0
  5. haliax-1.4.dev371/docs/vmap.md +9 -0
  6. haliax-1.4.dev371/src/haliax/__about__.py +1 -0
  7. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/scan.py +91 -1
  8. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_scan.py +79 -0
  9. haliax-1.4.dev368/docs/vmap.md +0 -5
  10. haliax-1.4.dev368/src/haliax/__about__.py +0 -1
  11. {haliax-1.4.dev368 → haliax-1.4.dev371}/.coveragerc +0 -0
  12. {haliax-1.4.dev368 → haliax-1.4.dev371}/.flake8 +0 -0
  13. {haliax-1.4.dev368 → haliax-1.4.dev371}/.github/workflows/publish_dev.yaml +0 -0
  14. {haliax-1.4.dev368 → haliax-1.4.dev371}/.github/workflows/run_pre_commit.yaml +0 -0
  15. {haliax-1.4.dev368 → haliax-1.4.dev371}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  16. {haliax-1.4.dev368 → haliax-1.4.dev371}/.gitignore +0 -0
  17. {haliax-1.4.dev368 → haliax-1.4.dev371}/.playbooks/add-types.md +0 -0
  18. {haliax-1.4.dev368 → haliax-1.4.dev371}/.pre-commit-config.yaml +0 -0
  19. {haliax-1.4.dev368 → haliax-1.4.dev371}/.readthedocs.yaml +0 -0
  20. {haliax-1.4.dev368 → haliax-1.4.dev371}/CONTRIBUTING.md +0 -0
  21. {haliax-1.4.dev368 → haliax-1.4.dev371}/LICENSE +0 -0
  22. {haliax-1.4.dev368 → haliax-1.4.dev371}/README.md +0 -0
  23. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/api.md +0 -0
  24. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/broadcasting.md +0 -0
  25. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/cheatsheet.md +0 -0
  26. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/css/material.css +0 -0
  27. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/css/mkdocstrings.css +0 -0
  28. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/faq.md +0 -0
  29. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/data_parallel_mesh.png +0 -0
  30. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  31. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_1d.png +0 -0
  32. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_1d_zero.png +0 -0
  33. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_2d.png +0 -0
  34. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  35. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  36. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  37. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  38. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/figures/device_mesh_2d_zero.png +0 -0
  39. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/fp8.md +0 -0
  40. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/index.md +0 -0
  41. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/indexing.md +0 -0
  42. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/matmul.md +0 -0
  43. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/nn.md +0 -0
  44. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/partitioning.md +0 -0
  45. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/rearrange.ipynb +0 -0
  46. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/rearrange.md +0 -0
  47. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/requirements.txt +0 -0
  48. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev368 → haliax-1.4.dev371}/docs/typing.md +0 -0
  51. {haliax-1.4.dev368 → haliax-1.4.dev371}/mkdocs.yml +0 -0
  52. {haliax-1.4.dev368 → haliax-1.4.dev371}/pyproject.toml +0 -0
  53. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/__init__.py +0 -0
  54. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/__init__.py +0 -0
  55. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/compile_utils.py +0 -0
  56. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/dot.py +0 -0
  57. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/einsum.py +0 -0
  58. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/fp8.py +0 -0
  59. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/parsing.py +0 -0
  60. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/rearrange.py +0 -0
  61. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/scan.py +0 -0
  62. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/state_dict.py +0 -0
  63. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/_src/util.py +0 -0
  64. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/axis.py +0 -0
  65. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/core.py +0 -0
  66. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/debug.py +0 -0
  67. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/haxtyping.py +0 -0
  68. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/hof.py +0 -0
  69. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/jax_utils.py +0 -0
  70. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/__init__.py +0 -0
  71. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/activations.py +0 -0
  72. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/attention.py +0 -0
  73. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/conv.py +0 -0
  74. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/dropout.py +0 -0
  75. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/embedding.py +0 -0
  76. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/linear.py +0 -0
  77. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/loss.py +0 -0
  78. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/mlp.py +0 -0
  79. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/normalization.py +0 -0
  80. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/nn/pool.py +0 -0
  81. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/ops.py +0 -0
  82. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/partitioning.py +0 -0
  83. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/quantization.py +0 -0
  84. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/random.py +0 -0
  85. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/util.py +0 -0
  90. {haliax-1.4.dev368 → haliax-1.4.dev371}/src/haliax/wrap.py +0 -0
  91. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/core_test.py +0 -0
  92. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_attention.py +0 -0
  93. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_axis.py +0 -0
  94. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_conv.py +0 -0
  95. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_debug.py +0 -0
  96. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_dot.py +0 -0
  97. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_dtype_typing.py +0 -0
  98. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_einsum.py +0 -0
  99. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_fp8.py +0 -0
  100. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_hof.py +0 -0
  101. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_int8.py +0 -0
  102. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_namedarray_typing.py +0 -0
  103. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_nn.py +0 -0
  104. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_ops.py +0 -0
  105. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_parsing.py +0 -0
  106. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_partitioning.py +0 -0
  107. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_pool.py +0 -0
  108. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_random.py +0 -0
  109. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_rearrange.py +0 -0
  110. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_scatter_gather.py +0 -0
  111. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_specialized_fns.py +0 -0
  112. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_state_dict.py +0 -0
  113. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_tree_util.py +0 -0
  114. {haliax-1.4.dev368 → haliax-1.4.dev371}/tests/test_utils.py +0 -0
@@ -15,9 +15,8 @@ jobs:
15
15
  python-version: 3.10.11
16
16
  - name: Install dependencies
17
17
  run: |
18
- python -m pip install --upgrade pip
19
- pip install flake8 pytest
20
- pip install -e .[dev]
18
+ python -m pip install uv
19
+ uv sync
21
20
  - name: Test with pytest
22
21
  run: |
23
- XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:. pytest tests
22
+ XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:. uv run pytest tests
@@ -12,12 +12,11 @@ repository. Follow these notes when implementing new features or fixing bugs.
12
12
  floating point tests, add that to the list.
13
13
  * **Playbooks.** Sometimes, there are repeatable tasks (e.g. porting models) for which we follow a standard set of steps.
14
14
  Please reference `.playbooks/` to see what playbooks are available, or see the list below. If you want to add a playbook
15
- write a markdown doc named e.g. `.playbooks/port-models.md` and add a pointer to it in the list below.
15
+ write a markdown doc named e.g. `.playbooks/add-types.md` and add a pointer to it in the list below.
16
16
 
17
17
  ## Playbook
18
18
 
19
- - At the moment, there are no playbooks available. If you have a repeatable task that you think
20
- should be documented, please create a new markdown file in `.playbooks/` and add it to the list above.
19
+ - Adding Haliax-style tensor typing annotations are described in @.playbooks/add-types.md
21
20
 
22
21
  ## Code Style
23
22
 
@@ -45,7 +44,7 @@ repository. Follow these notes when implementing new features or fixing bugs.
45
44
 
46
45
  ## Testing
47
46
 
48
- * Tests are executed with `pytest`. The default workflow runs `uv run pytest tests`.
47
+ * Tests are executed with `pytest`. The default workflow runs ` XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:. uv run pytest tests`.
49
48
  * In general, never relax tolerances in floating point tests unless specifically discussed with the
50
49
  team. Use `assert_allclose` with appropriate tolerances for numerical comparisons. We typically use
51
50
  1e-4 for more complex modules, and 1e-5 for simpler ones.
@@ -58,6 +57,11 @@ repository. Follow these notes when implementing new features or fixing bugs.
58
57
 
59
58
  * **Generic code**: many utilities are written with Python generics and dataclasses. Where possible,
60
59
  write reusable functions or classes that operate over TypeVars instead of hard coding concrete types.
60
+ * **Configurations**: configuration files are dataclasses loaded via `draccus`. Keep configs
61
+ declarative and typed.
62
+ * **Reproducibility**: Levanter aims for deterministic training where possible. Avoid sources of
63
+ nondeterminism unless explicitly required.
64
+ * Prefer Stacked with fold or scan over writing custom loops, for better compile times and gradient checkpointing support
61
65
 
62
66
  ## Library conventions
63
67
  - Haliax revolves around `NamedArray` and explicit `Axis` objects. Prefer APIs that accept
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev368
3
+ Version: 1.4.dev371
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/
@@ -405,6 +405,17 @@ blocks = Stacked.init(Layers, Gpt2Block)(
405
405
  Any NamedArray passed to the Stacked init will have its Layers axis (if present) vmapped over. Any
406
406
  JAX array will have its first axis vmapped over.
407
407
 
408
+ #### Apply Blocks in Parallel with `vmap`
409
+
410
+ Sometimes you may want to apply each block independently, without feeding the
411
+ output of one block into the next. `Stacked.vmap` does exactly that: it uses
412
+ [`haliax.vmap`][] to broadcast the initial value to every block and evaluates
413
+ them in parallel, returning the stack of outputs.
414
+
415
+ ```python
416
+ y = stacked.vmap(x)
417
+ ```
418
+
408
419
 
409
420
  #### Fold Blocks vs Scan Blocks
410
421
 
@@ -0,0 +1,9 @@
1
+ ## Vectorization
2
+
3
+
4
+ This primitive is also used by [`Stacked.vmap`](scan.md#apply-blocks-in-parallel-with-vmap)
5
+ to apply an entire stack of blocks in parallel.
6
+
7
+ (This is a work in progress. Please contact dlwh for more information.)
8
+
9
+ ::: haliax.vmap
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev371"
@@ -2,7 +2,7 @@ import dataclasses
2
2
  import functools
3
3
  import re
4
4
  import warnings
5
- from typing import Any, Dict, Generic, Optional, Protocol, Sequence, Type, TypeVar, cast
5
+ from typing import Any, Callable, Dict, Generic, Optional, Protocol, Sequence, Type, TypeVar, cast, overload, ParamSpec
6
6
 
7
7
  import equinox as eqx
8
8
  import jax
@@ -20,8 +20,21 @@ from ..axis import Axis
20
20
 
21
21
  M = TypeVar("M", bound=eqx.Module)
22
22
  M_co = TypeVar("M_co", bound=eqx.Module, covariant=True)
23
+ M_contra = TypeVar("M_contra", bound=eqx.Module, contravariant=True)
23
24
  S = TypeVar("S", bound=eqx.Module)
24
25
  T = TypeVar("T")
26
+ CarryT = TypeVar("CarryT")
27
+ OutputT_co = TypeVar("OutputT_co", covariant=True)
28
+ P = ParamSpec("P")
29
+
30
+ class FoldFunction(Protocol[M_contra, CarryT]):
31
+ def __call__(self, module: M_contra, carry: CarryT) -> CarryT:
32
+ ...
33
+
34
+
35
+ class ScanFunction(Protocol[M_contra, CarryT, OutputT_co]):
36
+ def __call__(self, module: M_contra, carry: CarryT) -> tuple[CarryT, OutputT_co]:
37
+ ...
25
38
 
26
39
 
27
40
  class ModuleInit(Protocol[M_co]):
@@ -233,6 +246,8 @@ class Stacked(ModuleWithStateDictSerialization, Generic[M]):
233
246
  output, while "scan" is the same as a for loop that accumulates a list of outputs as well as the final output.
234
247
 
235
248
  Stacked also supports gradient checkpointing, which is useful for very large models that don't fit in memory.
249
+ If your blocks are independent of each other you can instead use :py:meth:`Stacked.vmap`
250
+ to apply every block in parallel.
236
251
 
237
252
  Typically only one of "fold" or "scan" can be used with a given Stacked module, depending on the what the module
238
253
  returns: if the module returns a single output, use "fold"; if the module returns a sequence of outputs and
@@ -382,6 +397,81 @@ class Stacked(ModuleWithStateDictSerialization, Generic[M]):
382
397
 
383
398
  return do_fold(init, *args, **kwargs)
384
399
 
400
+ @overload
401
+ def fold_via(self, fn: FoldFunction[M, CarryT]) -> Callable[[CarryT], CarryT]:
402
+ ...
403
+
404
+ @overload
405
+ def fold_via(self, fn: Callable[[M, CarryT], CarryT]) -> Callable[[CarryT], CarryT]:
406
+ ...
407
+
408
+ def fold_via(self, fn: Callable[..., CarryT]):
409
+ """Return a function that folds over the stack using ``fn``.
410
+
411
+ ``fn`` should take a block and a carry and return a new carry. The
412
+ returned function mirrors :func:`haliax.fold` over the block axis.
413
+ """
414
+
415
+ def do_block(carry: CarryT, block: M) -> CarryT:
416
+ return fn(block, carry)
417
+
418
+ def do_fold(init: CarryT) -> CarryT:
419
+ return haliax.fold(do_block, self.Block, remat=self.gradient_checkpointing)(init, self.stacked)
420
+
421
+ return do_fold
422
+
423
+ @overload
424
+ def scan_via(self, fn: ScanFunction[M, CarryT, OutputT_co]) -> Callable[[CarryT], tuple[CarryT, OutputT_co]]:
425
+ ...
426
+
427
+ @overload
428
+ def scan_via(self, fn: Callable[[M, CarryT], tuple[CarryT, OutputT_co]]) -> Callable[[CarryT], tuple[CarryT, OutputT_co]]:
429
+ ...
430
+
431
+ def scan_via(self, fn: Callable[..., tuple[CarryT, OutputT_co]]):
432
+ """Return a function that scans over the stack using ``fn``.
433
+
434
+ ``fn`` should take a block and a carry and return ``(carry, output)``.
435
+ Semantics match :func:`haliax.scan` over the block axis.
436
+ """
437
+
438
+ def do_block(carry: CarryT, block: M) -> tuple[CarryT, OutputT_co]:
439
+ return fn(block, carry)
440
+
441
+ def do_scan(init: CarryT) -> tuple[CarryT, OutputT_co]:
442
+ return haliax.scan(do_block, self.Block, remat=self.gradient_checkpointing)(init, self.stacked)
443
+
444
+ return do_scan
445
+
446
+ def vmap(self, init, *extra_args, **extra_kwargs):
447
+ """Apply each block independently using :func:`haliax.vmap`.
448
+
449
+ This maps ``init`` through every block in parallel, so each block
450
+ receives the same ``init`` but its own parameters. Extra ``args`` and
451
+ ``kwargs`` are also mapped over the block axis by default.
452
+
453
+ Returns the stacked outputs of each block.
454
+ """
455
+
456
+ if haliax.is_named_array(init):
457
+ init = init.broadcast_axis(self.Block)
458
+ elif haliax.jax_utils.is_jax_array_like(init):
459
+ init = jnp.broadcast_to(init, (self.Block.size,) + init.shape)
460
+ else:
461
+ init = tuple(init for _ in range(self.Block.size))
462
+
463
+ arg_spec = (0, 0) + (0,) * len(extra_args)
464
+ kwarg_spec = {k: 0 for k in extra_kwargs}
465
+
466
+ do_vmap = haliax.vmap(
467
+ Stacked._do_block,
468
+ self.Block,
469
+ default=0,
470
+ args=arg_spec,
471
+ kwargs=kwarg_spec,
472
+ )
473
+ return do_vmap(init, self.stacked, *extra_args, **extra_kwargs)
474
+
385
475
  @staticmethod
386
476
  def _do_block(carry, block, *extra_args, **extra_kwargs):
387
477
  return block(carry, *extra_args, **extra_kwargs)
@@ -45,6 +45,30 @@ def test_unstacked():
45
45
  assert hax.all(module.array == m.stacked.array[i])
46
46
 
47
47
 
48
+ def test_vmap():
49
+ class Module(eqx.Module):
50
+ weight: hax.NamedArray
51
+
52
+ def __call__(self, x):
53
+ return x + self.weight
54
+
55
+ @staticmethod
56
+ def init(weight):
57
+ return Module(weight=weight)
58
+
59
+ Block = hax.Axis("block", 4)
60
+ E = hax.Axis("E", 10)
61
+
62
+ weights = hax.random.uniform(jax.random.PRNGKey(0), (Block, E))
63
+ m = Stacked.init(Block, Module)(weight=weights)
64
+
65
+ x = hax.random.uniform(jax.random.PRNGKey(1), (E,))
66
+ y = m.vmap(x)
67
+
68
+ assert y.axes == (Block, E)
69
+ assert hax.all(y == weights + x)
70
+
71
+
48
72
  def test_seq_and_stacked_give_same_results():
49
73
  class Module(eqx.Module):
50
74
  named: hax.NamedArray
@@ -285,3 +309,58 @@ def test_checkpoint_carries(name, policy, expected_scan_shapes, check_offloading
285
309
 
286
310
  assert target is not None, f"Could not find named value for {name}"
287
311
  assert found_saved, f"Could not find offloaded value for {name}"
312
+
313
+
314
+ def test_fold_via():
315
+ class Module(eqx.Module):
316
+ w: hax.NamedArray
317
+
318
+ def __call__(self, x):
319
+ return x + self.w
320
+
321
+ def intermediate(self, x):
322
+ return x + 2 * self.w
323
+
324
+ @staticmethod
325
+ def init(named):
326
+ return Module(w=named)
327
+
328
+ Block = hax.Axis("block", 3)
329
+ E = hax.Axis("E", 5)
330
+
331
+ named = hax.random.uniform(jax.random.PRNGKey(0), (Block, E))
332
+ m = Stacked.init(Block, Module)(named=named)
333
+
334
+ x = hax.random.uniform(jax.random.PRNGKey(1), (E,))
335
+ result = m.fold_via(Module.intermediate)(x)
336
+
337
+ expected = x + 2 * hax.sum(named, Block)
338
+ assert hax.all(hax.isclose(result, expected))
339
+
340
+
341
+ def test_scan_via():
342
+ class Module(eqx.Module):
343
+ w: hax.NamedArray
344
+
345
+ def with_output(self, x):
346
+ out = x + self.w
347
+ return out, 2 * self.w
348
+
349
+ @staticmethod
350
+ def init(named):
351
+ return Module(w=named)
352
+
353
+ Block = hax.Axis("block", 4)
354
+ E = hax.Axis("E", 6)
355
+
356
+ named = hax.random.uniform(jax.random.PRNGKey(0), (Block, E))
357
+ m = Stacked.init(Block, Module)(named=named)
358
+
359
+ x = hax.random.uniform(jax.random.PRNGKey(1), (E,))
360
+ carry, outs = m.scan_via(Module.with_output)(x)
361
+
362
+ expected_carry = x + hax.sum(named, Block)
363
+ expected_outs = 2 * named
364
+
365
+ assert hax.all(hax.isclose(carry, expected_carry))
366
+ assert hax.all(hax.isclose(outs, expected_outs))
@@ -1,5 +0,0 @@
1
- ## Vectorization
2
-
3
- (This is a work in progress. Please contact dlwh for more information.)
4
-
5
- ::: haliax.vmap
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev368"
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