haliax 1.4.dev336__tar.gz → 1.4.dev337__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 (105) hide show
  1. {haliax-1.4.dev336 → haliax-1.4.dev337}/PKG-INFO +1 -1
  2. haliax-1.4.dev337/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/quantization.py +6 -1
  4. haliax-1.4.dev336/src/haliax/__about__.py +0 -1
  5. {haliax-1.4.dev336 → haliax-1.4.dev337}/.coveragerc +0 -0
  6. {haliax-1.4.dev336 → haliax-1.4.dev337}/.flake8 +0 -0
  7. {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/publish_dev.yaml +0 -0
  8. {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/run_pre_commit.yaml +0 -0
  9. {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  10. {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/run_tests.yaml +0 -0
  11. {haliax-1.4.dev336 → haliax-1.4.dev337}/.gitignore +0 -0
  12. {haliax-1.4.dev336 → haliax-1.4.dev337}/.pre-commit-config.yaml +0 -0
  13. {haliax-1.4.dev336 → haliax-1.4.dev337}/.readthedocs.yaml +0 -0
  14. {haliax-1.4.dev336 → haliax-1.4.dev337}/CONTRIBUTING.md +0 -0
  15. {haliax-1.4.dev336 → haliax-1.4.dev337}/LICENSE +0 -0
  16. {haliax-1.4.dev336 → haliax-1.4.dev337}/README.md +0 -0
  17. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/api.md +0 -0
  18. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/broadcasting.md +0 -0
  19. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/cheatsheet.md +0 -0
  20. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/css/material.css +0 -0
  21. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/css/mkdocstrings.css +0 -0
  22. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/faq.md +0 -0
  23. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/data_parallel_mesh.png +0 -0
  24. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  25. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_1d.png +0 -0
  26. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_1d_zero.png +0 -0
  27. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d.png +0 -0
  28. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  29. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  30. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  31. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  32. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_zero.png +0 -0
  33. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/fp8.md +0 -0
  34. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/hof.md +0 -0
  35. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/index.md +0 -0
  36. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/indexing.md +0 -0
  37. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/matmul.md +0 -0
  38. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/nn.md +0 -0
  39. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/partitioning.md +0 -0
  40. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/rearrange.ipynb +0 -0
  41. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/rearrange.md +0 -0
  42. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/requirements.txt +0 -0
  43. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/stacked.md +0 -0
  44. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/state-dict.md +0 -0
  45. {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/tutorial.md +0 -0
  46. {haliax-1.4.dev336 → haliax-1.4.dev337}/mkdocs.yml +0 -0
  47. {haliax-1.4.dev336 → haliax-1.4.dev337}/pyproject.toml +0 -0
  48. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/__init__.py +0 -0
  49. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/__init__.py +0 -0
  50. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/compile_utils.py +0 -0
  51. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/dot.py +0 -0
  52. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/einsum.py +0 -0
  53. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/fp8.py +0 -0
  54. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/parsing.py +0 -0
  55. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/rearrange.py +0 -0
  56. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/state_dict.py +0 -0
  57. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/util.py +0 -0
  58. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/axis.py +0 -0
  59. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/core.py +0 -0
  60. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/debug.py +0 -0
  61. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/hof.py +0 -0
  62. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/jax_utils.py +0 -0
  63. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/__init__.py +0 -0
  64. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/activations.py +0 -0
  65. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/attention.py +0 -0
  66. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/conv.py +0 -0
  67. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/dropout.py +0 -0
  68. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/embedding.py +0 -0
  69. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/linear.py +0 -0
  70. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/loss.py +0 -0
  71. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/mlp.py +0 -0
  72. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/normalization.py +0 -0
  73. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/pool.py +0 -0
  74. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/scan.py +0 -0
  75. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/ops.py +0 -0
  76. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/partitioning.py +0 -0
  77. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/random.py +0 -0
  78. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/specialized_fns.py +0 -0
  79. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/state_dict.py +0 -0
  80. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/tree_util.py +0 -0
  81. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/types.py +0 -0
  82. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/util.py +0 -0
  83. {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/wrap.py +0 -0
  84. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/core_test.py +0 -0
  85. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_attention.py +0 -0
  86. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_axis.py +0 -0
  87. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_conv.py +0 -0
  88. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_debug.py +0 -0
  89. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_dot.py +0 -0
  90. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_einsum.py +0 -0
  91. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_fp8.py +0 -0
  92. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_hof.py +0 -0
  93. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_int8.py +0 -0
  94. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_nn.py +0 -0
  95. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_ops.py +0 -0
  96. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_parsing.py +0 -0
  97. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_partitioning.py +0 -0
  98. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_pool.py +0 -0
  99. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_random.py +0 -0
  100. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_rearrange.py +0 -0
  101. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_scan.py +0 -0
  102. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_specialized_fns.py +0 -0
  103. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_state_dict.py +0 -0
  104. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_tree_util.py +0 -0
  105. {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev336
3
+ Version: 1.4.dev337
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.dev337"
@@ -14,9 +14,10 @@ import jax.random as jrandom
14
14
  from aqt.jax.v2.aqt_dot_general import DotGeneral
15
15
  from jax import numpy as jnp
16
16
  from jax.tree_util import DictKey, FlattenedIndexKey, GetAttrKey, SequenceKey
17
- from jax.typing import DTypeLike
17
+ from jaxtyping import DTypeLike, PyTree
18
18
 
19
19
  import haliax.nn as hnn
20
+ from haliax.state_dict import StateDict
20
21
  from haliax.types import PrecisionLike
21
22
 
22
23
  from ._src.fp8 import dot_general_with_precision, in_qdq, out_qdq
@@ -206,6 +207,10 @@ class Int8DotGeneralOp(OverwriteWithGradient):
206
207
  cfg = aqt_config.set_context(self.cfg, jrandom.PRNGKey(42), train_step=None)
207
208
  return cfg(lhs, rhs, dimension_numbers, precision, preferred_element_type)
208
209
 
210
+ def to_state_dict(tree: PyTree, prefix: Optional[str] = None) -> StateDict:
211
+ warnings.warn("Ignore all int8 states (if any) for now.")
212
+ return {}
213
+
209
214
 
210
215
  @dataclass(frozen=True)
211
216
  class QuantizationConfig:
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev336"
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