haliax 1.4.dev315__tar.gz → 1.4.dev316__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.dev315 → haliax-1.4.dev316}/PKG-INFO +1 -1
  2. haliax-1.4.dev316/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/mlp.py +20 -21
  4. haliax-1.4.dev315/src/haliax/__about__.py +0 -1
  5. {haliax-1.4.dev315 → haliax-1.4.dev316}/.coveragerc +0 -0
  6. {haliax-1.4.dev315 → haliax-1.4.dev316}/.flake8 +0 -0
  7. {haliax-1.4.dev315 → haliax-1.4.dev316}/.github/workflows/publish_dev.yaml +0 -0
  8. {haliax-1.4.dev315 → haliax-1.4.dev316}/.github/workflows/run_pre_commit.yaml +0 -0
  9. {haliax-1.4.dev315 → haliax-1.4.dev316}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  10. {haliax-1.4.dev315 → haliax-1.4.dev316}/.github/workflows/run_tests.yaml +0 -0
  11. {haliax-1.4.dev315 → haliax-1.4.dev316}/.gitignore +0 -0
  12. {haliax-1.4.dev315 → haliax-1.4.dev316}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev315 → haliax-1.4.dev316}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev315 → haliax-1.4.dev316}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev315 → haliax-1.4.dev316}/LICENSE +0 -0
  16. {haliax-1.4.dev315 → haliax-1.4.dev316}/README.md +0 -0
  17. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/api.md +0 -0
  18. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/broadcasting.md +0 -0
  19. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/cheatsheet.md +0 -0
  20. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/css/material.css +0 -0
  21. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/css/mkdocstrings.css +0 -0
  22. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/faq.md +0 -0
  23. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/data_parallel_mesh.png +0 -0
  24. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  25. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_1d.png +0 -0
  26. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_1d_zero.png +0 -0
  27. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_2d.png +0 -0
  28. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  29. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  30. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  31. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  32. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/figures/device_mesh_2d_zero.png +0 -0
  33. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/fp8.md +0 -0
  34. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/hof.md +0 -0
  35. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/index.md +0 -0
  36. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/indexing.md +0 -0
  37. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/matmul.md +0 -0
  38. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/nn.md +0 -0
  39. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/partitioning.md +0 -0
  40. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/rearrange.ipynb +0 -0
  41. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/rearrange.md +0 -0
  42. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/requirements.txt +0 -0
  43. {haliax-1.4.dev315 → haliax-1.4.dev316}/docs/tutorial.md +0 -0
  44. {haliax-1.4.dev315 → haliax-1.4.dev316}/mkdocs.yml +0 -0
  45. {haliax-1.4.dev315 → haliax-1.4.dev316}/pyproject.toml +0 -0
  46. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/__init__.py +0 -0
  47. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/__init__.py +0 -0
  48. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/compile_utils.py +0 -0
  49. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/dot.py +0 -0
  50. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/einsum.py +0 -0
  51. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/fp8.py +0 -0
  52. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/parsing.py +0 -0
  53. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/rearrange.py +0 -0
  54. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/_src/util.py +0 -0
  55. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/axis.py +0 -0
  56. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/core.py +0 -0
  57. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/debug.py +0 -0
  58. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/hof.py +0 -0
  59. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/jax_utils.py +0 -0
  60. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/__init__.py +0 -0
  61. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/activations.py +0 -0
  62. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/attention.py +0 -0
  63. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/conv.py +0 -0
  64. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/dropout.py +0 -0
  65. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/embedding.py +0 -0
  66. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/linear.py +0 -0
  67. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/loss.py +0 -0
  68. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/normalization.py +0 -0
  69. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/pool.py +0 -0
  70. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/nn/scan.py +0 -0
  71. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/ops.py +0 -0
  72. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/partitioning.py +0 -0
  73. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/quantization.py +0 -0
  74. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/random.py +0 -0
  75. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/specialized_fns.py +0 -0
  76. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/tree_util.py +0 -0
  77. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/types.py +0 -0
  78. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/util.py +0 -0
  79. {haliax-1.4.dev315 → haliax-1.4.dev316}/src/haliax/wrap.py +0 -0
  80. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/core_test.py +0 -0
  81. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_attention.py +0 -0
  82. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_axis.py +0 -0
  83. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_conv.py +0 -0
  84. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_debug.py +0 -0
  85. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_dot.py +0 -0
  86. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_einsum.py +0 -0
  87. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_fp8.py +0 -0
  88. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_hof.py +0 -0
  89. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_nn.py +0 -0
  90. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_ops.py +0 -0
  91. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_parsing.py +0 -0
  92. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_partitioning.py +0 -0
  93. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_pool.py +0 -0
  94. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_random.py +0 -0
  95. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_rearrange.py +0 -0
  96. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_scan.py +0 -0
  97. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_specialized_fns.py +0 -0
  98. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_tree_util.py +0 -0
  99. {haliax-1.4.dev315 → haliax-1.4.dev316}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: haliax
3
- Version: 1.4.dev315
3
+ Version: 1.4.dev316
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.dev316"
@@ -44,6 +44,7 @@ class MLP(eqx.Module):
44
44
  depth: int,
45
45
  activation: Callable = relu,
46
46
  *,
47
+ out_first: bool = True,
47
48
  use_bias: bool = True,
48
49
  use_final_bias: bool = True,
49
50
  key: PRNGKeyArray,
@@ -57,36 +58,34 @@ class MLP(eqx.Module):
57
58
 
58
59
  layers = []
59
60
 
61
+ kwargs: dict = {
62
+ "use_bias": use_bias,
63
+ "dot_general": dot_general,
64
+ "init_scale": init_scale,
65
+ "out_first": out_first,
66
+ }
67
+
68
+ last_kwargs: dict = {
69
+ "use_bias": use_final_bias,
70
+ "dot_general": dot_general,
71
+ "init_scale": init_scale,
72
+ "out_first": out_first,
73
+ }
74
+
60
75
  if depth == 0:
61
76
  # special case: no hidden layers
62
- layers.append(
63
- Linear.init(
64
- Input, Output, use_bias=use_final_bias, key=keys[0], dot_general=dot_general, init_scale=init_scale
65
- )
66
- )
77
+ layers.append(Linear.init(Input, Output, key=keys[0], **last_kwargs))
67
78
  else:
68
79
  # first hidden layer
69
- layers.append(
70
- Linear.init(
71
- Input, Width, use_bias=use_bias, key=keys[0], dot_general=dot_general, init_scale=init_scale
72
- )
73
- )
80
+ layers.append(Linear.init(Input, Width, key=keys[0], **kwargs))
74
81
  # middle hidden layers
75
82
  cur = Width
76
83
  next = Width2
77
84
  for i in range(1, depth):
78
- layers.append(
79
- Linear.init(
80
- cur, next, use_bias=use_bias, key=keys[i], dot_general=dot_general, init_scale=init_scale
81
- )
82
- )
85
+ layers.append(Linear.init(cur, next, key=keys[i], **kwargs))
83
86
  cur, next = next, cur
84
- # final hidden layer
85
- layers.append(
86
- Linear.init(
87
- cur, Output, use_bias=use_final_bias, key=keys[-1], dot_general=dot_general, init_scale=init_scale
88
- )
89
- )
87
+ # final layer
88
+ layers.append(Linear.init(cur, Output, key=keys[-1], **last_kwargs))
90
89
 
91
90
  return MLP(
92
91
  layers=tuple(layers),
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev315"
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