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.
- {haliax-1.4.dev336 → haliax-1.4.dev337}/PKG-INFO +1 -1
- haliax-1.4.dev337/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/quantization.py +6 -1
- haliax-1.4.dev336/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.coveragerc +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.flake8 +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.gitignore +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/LICENSE +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/README.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/api.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/css/material.css +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/faq.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/fp8.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/hof.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/index.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/indexing.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/matmul.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/nn.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/partitioning.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/rearrange.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/requirements.txt +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/stacked.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/state-dict.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/docs/tutorial.md +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/mkdocs.yml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/pyproject.toml +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/core.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/random.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/types.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/util.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/core_test.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_attention.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_axis.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_conv.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_debug.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_dot.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_hof.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_int8.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_nn.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_ops.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_pool.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_random.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_scan.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev336 → haliax-1.4.dev337}/tests/test_tree_util.py +0 -0
- {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.
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|