haliax 1.4.dev359__tar.gz → 1.4.dev362__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.dev359 → haliax-1.4.dev362}/.github/workflows/run_tests.yaml +1 -1
  2. {haliax-1.4.dev359 → haliax-1.4.dev362}/PKG-INFO +1 -1
  3. haliax-1.4.dev362/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/activations.py +0 -1
  5. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/linear.py +40 -14
  6. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_nn.py +1 -1
  7. haliax-1.4.dev359/src/haliax/__about__.py +0 -1
  8. {haliax-1.4.dev359 → haliax-1.4.dev362}/.coveragerc +0 -0
  9. {haliax-1.4.dev359 → haliax-1.4.dev362}/.flake8 +0 -0
  10. {haliax-1.4.dev359 → haliax-1.4.dev362}/.github/workflows/publish_dev.yaml +0 -0
  11. {haliax-1.4.dev359 → haliax-1.4.dev362}/.github/workflows/run_pre_commit.yaml +0 -0
  12. {haliax-1.4.dev359 → haliax-1.4.dev362}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  13. {haliax-1.4.dev359 → haliax-1.4.dev362}/.gitignore +0 -0
  14. {haliax-1.4.dev359 → haliax-1.4.dev362}/.pre-commit-config.yaml +0 -0
  15. {haliax-1.4.dev359 → haliax-1.4.dev362}/.readthedocs.yaml +0 -0
  16. {haliax-1.4.dev359 → haliax-1.4.dev362}/CONTRIBUTING.md +0 -0
  17. {haliax-1.4.dev359 → haliax-1.4.dev362}/LICENSE +0 -0
  18. {haliax-1.4.dev359 → haliax-1.4.dev362}/README.md +0 -0
  19. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/api.md +0 -0
  20. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/broadcasting.md +0 -0
  21. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/cheatsheet.md +0 -0
  22. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/css/material.css +0 -0
  23. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/css/mkdocstrings.css +0 -0
  24. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/faq.md +0 -0
  25. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/data_parallel_mesh.png +0 -0
  26. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  27. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_1d.png +0 -0
  28. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_1d_zero.png +0 -0
  29. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d.png +0 -0
  30. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  31. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  32. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  33. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  34. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/figures/device_mesh_2d_zero.png +0 -0
  35. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/fp8.md +0 -0
  36. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/index.md +0 -0
  37. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/indexing.md +0 -0
  38. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/matmul.md +0 -0
  39. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/nn.md +0 -0
  40. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/partitioning.md +0 -0
  41. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/rearrange.ipynb +0 -0
  42. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/rearrange.md +0 -0
  43. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/requirements.txt +0 -0
  44. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/scan.md +0 -0
  45. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/state-dict.md +0 -0
  46. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/tutorial.md +0 -0
  47. {haliax-1.4.dev359 → haliax-1.4.dev362}/docs/vmap.md +0 -0
  48. {haliax-1.4.dev359 → haliax-1.4.dev362}/mkdocs.yml +0 -0
  49. {haliax-1.4.dev359 → haliax-1.4.dev362}/pyproject.toml +0 -0
  50. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/__init__.py +0 -0
  51. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/__init__.py +0 -0
  52. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/compile_utils.py +0 -0
  53. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/dot.py +0 -0
  54. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/einsum.py +0 -0
  55. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/fp8.py +0 -0
  56. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/parsing.py +0 -0
  57. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/rearrange.py +0 -0
  58. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/scan.py +0 -0
  59. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/state_dict.py +0 -0
  60. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/_src/util.py +0 -0
  61. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/axis.py +0 -0
  62. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/core.py +0 -0
  63. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/debug.py +0 -0
  64. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/hof.py +0 -0
  65. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/jax_utils.py +0 -0
  66. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/__init__.py +0 -0
  67. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/attention.py +0 -0
  68. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/conv.py +0 -0
  69. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/dropout.py +0 -0
  70. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/embedding.py +0 -0
  71. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/loss.py +0 -0
  72. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/mlp.py +0 -0
  73. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/normalization.py +0 -0
  74. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/pool.py +0 -0
  75. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/nn/scan.py +0 -0
  76. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/ops.py +0 -0
  77. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/partitioning.py +0 -0
  78. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/quantization.py +0 -0
  79. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/random.py +0 -0
  80. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/specialized_fns.py +0 -0
  81. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/state_dict.py +0 -0
  82. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/tree_util.py +0 -0
  83. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/types.py +0 -0
  84. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/util.py +0 -0
  85. {haliax-1.4.dev359 → haliax-1.4.dev362}/src/haliax/wrap.py +0 -0
  86. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/core_test.py +0 -0
  87. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_attention.py +0 -0
  88. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_axis.py +0 -0
  89. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_conv.py +0 -0
  90. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_debug.py +0 -0
  91. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_dot.py +0 -0
  92. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_einsum.py +0 -0
  93. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_fp8.py +0 -0
  94. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_hof.py +0 -0
  95. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_int8.py +0 -0
  96. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_ops.py +0 -0
  97. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_parsing.py +0 -0
  98. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_partitioning.py +0 -0
  99. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_pool.py +0 -0
  100. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_random.py +0 -0
  101. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_rearrange.py +0 -0
  102. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_scan.py +0 -0
  103. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_scatter_gather.py +0 -0
  104. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_specialized_fns.py +0 -0
  105. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_state_dict.py +0 -0
  106. {haliax-1.4.dev359 → haliax-1.4.dev362}/tests/test_tree_util.py +0 -0
  107. {haliax-1.4.dev359 → haliax-1.4.dev362}/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 jax==0.4.35 jaxlib==0.4.35 .[dev]
20
+ pip install -e .[dev]
21
21
  - name: Test with pytest
22
22
  run: |
23
23
  XLA_FLAGS=--xla_force_host_platform_device_count=8 PYTHONPATH=tests:src:. pytest tests
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev359
3
+ Version: 1.4.dev362
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.dev362"
@@ -89,7 +89,6 @@ def quick_gelu(x):
89
89
  return x * sigmoid(1.702 * x)
90
90
 
91
91
 
92
-
93
92
  def relu_squared(x: A) -> A:
94
93
  """ReLU squared activation function. jnp.square(jnp.maximum(0, x))"""
95
94
 
@@ -18,6 +18,7 @@ from ..core import NamedArray
18
18
  from ..jax_utils import named_call
19
19
  from ..partitioning import ResourceAxis
20
20
  from ..quantization import DotGeneralOp
21
+ from ..util import ensure_tuple
21
22
 
22
23
 
23
24
  class Linear(ModuleWithStateDictSerialization):
@@ -142,7 +143,10 @@ class MoELinear(eqx.Module):
142
143
  Experts: AxisSpec = eqx.static_field()
143
144
  In: Axis = eqx.static_field()
144
145
  Out: Axis = eqx.static_field()
145
- dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
146
+ # TODO: support quanitization for ragged_dot?
147
+ # dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
148
+
149
+ use_gmm: bool = eqx.field(static=True)
146
150
 
147
151
  @staticmethod
148
152
  def init(
@@ -154,6 +158,7 @@ class MoELinear(eqx.Module):
154
158
  use_bias: bool = True,
155
159
  out_first: bool = False,
156
160
  init_scale: float = 1.0,
161
+ use_gmm: bool = False,
157
162
  ) -> "MoELinear":
158
163
  """
159
164
  Args:
@@ -172,7 +177,7 @@ class MoELinear(eqx.Module):
172
177
  weight = hax.random.truncated_normal(key, joint_spec, -3, 3) * (init_scale / math.sqrt(input_size))
173
178
  bias = hax.zeros(Out) if use_bias else None
174
179
 
175
- return MoELinear(weight, bias, Experts, In, Out)
180
+ return MoELinear(weight, bias, Experts, In, Out, use_gmm=use_gmm)
176
181
 
177
182
  @named_call
178
183
  def __call__(self, inputs, group_sizes, *, key: Optional[PRNGKey] = None):
@@ -184,21 +189,42 @@ class MoELinear(eqx.Module):
184
189
  """
185
190
  del key
186
191
 
187
- inputs = inputs.rearrange((..., self.In))
188
- out_axes = hax.replace_axis(inputs.axes, self.In, self.Out)
192
+ dim_numbers = jax.lax.RaggedDotDimensionNumbers(
193
+ dot_dimension_numbers=(
194
+ # contracting
195
+ (ensure_tuple(inputs.axis_indices(self.In)), ensure_tuple(self.weight.axis_indices(self.In))),
196
+ # batch
197
+ ((), ()),
198
+ ),
199
+ # Everything other than contracting dim is ragged
200
+ lhs_ragged_dimensions=(inputs.axis_indices(hax.axis.without_axes(inputs.axes, self.In))),
201
+ rhs_group_dimensions=(self.weight.axis_indices(self.Experts),),
202
+ )
189
203
 
190
- q = _gmm(
191
- inputs,
192
- self.weight,
193
- group_sizes,
194
- out_axes,
195
- ar=hax.partitioning.physical_axis_name(self.In) == ResourceAxis.MODEL,
196
- ) # gmm((B, D), (E, D, d)) -> (B, d)
197
- q = hax.auto_sharded(q)
204
+ if self.use_gmm:
205
+ inputs = inputs.rearrange((..., self.In))
206
+ out_axes = hax.replace_axis(inputs.axes, self.In, self.Out)
207
+ q = _gmm(
208
+ inputs,
209
+ self.weight,
210
+ group_sizes,
211
+ out_axes,
212
+ ar=hax.partitioning.physical_axis_name(self.In) == ResourceAxis.MODEL,
213
+ ) # gmm((B, D), (E, D, d)) -> (B, d)
214
+ else:
215
+ q_raw = jax.lax.ragged_dot_general(
216
+ lhs=inputs.array,
217
+ rhs=self.weight.array,
218
+ group_sizes=group_sizes.rearrange((..., self.Experts)).array,
219
+ ragged_dot_dimension_numbers=dim_numbers,
220
+ )
221
+ out_axes = hax.replace_axis(inputs.axes, self.In, self.Out)
222
+ q = hax.named(q_raw, out_axes)
198
223
 
199
224
  if self.bias is not None:
200
225
  q = q + self.bias
201
- q = hax.auto_sharded(q)
226
+
227
+ q = hax.auto_sharded(q)
202
228
 
203
229
  return q
204
230
 
@@ -249,7 +275,7 @@ def gmm_sharded(lhs_: jnp.ndarray, rhs_: jnp.ndarray, group_sizes_: jnp.ndarray,
249
275
  group_sizes_,
250
276
  preferred_element_type=lhs_.dtype,
251
277
  tiling=(min(m, tile_size[0]), min(k, tile_size[1]), min(n, tile_size[2])),
252
- # interpret=True,
278
+ interpret=jax.default_backend() == "cpu",
253
279
  )
254
280
 
255
281
  if ar:
@@ -198,4 +198,4 @@ def test_relu_squared_scalar(use_jit):
198
198
  x_neg = -5.0
199
199
  expected_neg = 0.0
200
200
  actual_neg = f(x_neg)
201
- assert jnp.allclose(actual_neg, expected_neg)
201
+ assert jnp.allclose(actual_neg, expected_neg)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev359"
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