Constrains all leaf values in the input batch, then (optionally) dispatches examples to a subset along the batch axis. The dispatchings are applied to all nested dicts which contain a special dispatching key in their root. This is deprecated in favor of `axlearn.common.input_dispat
(
input_batch: NestedTensor, *, batch_axis_names: Union[str, Sequence[str]] = "data"
)
| 761 | return thread_resources.env.physical_mesh |
| 762 | return mesh |
| 763 | |
| 764 | |
| 765 | def replicate_to_local_data(x: NestedTensor) -> NestedTensor: |
| 766 | """Replicates and converts Tensors in `x` to local DeviceArrays. |
| 767 | |
| 768 | Args: |
| 769 | x: The tensor to replicate. |
| 770 | |
| 771 | Returns: |
| 772 | Replicated tensor. |
| 773 | """ |
| 774 | return multihost_utils.process_allgather(x, tiled=True) |
| 775 | |
| 776 | |
| 777 | def complete_partition_spec_tree( |
| 778 | treedef: jax.tree_util.PyTreeDef, partition_spec_tree: NestedTree |
| 779 | ) -> NestedTree: |
| 780 | """Adapted from flatten_axes(), but with a simplified API and more error logging and messages. |
| 781 | |
| 782 | Original: |
| 783 | https://github.com/google/jax/blob/cdf4177f9219c9c0f70e243a097740e46138fc35/jax/_src/api_util.py#L277-L315 |
| 784 | |
| 785 | Args: |
| 786 | treedef: The tree structure of data to be partitioned according to `partition_spec_tree`. |
| 787 | partition_spec_tree: A nested structure with PartitionSpecs or ParamPartitionSpecs and Nones |
| 788 | at the leaves. Must be a tree prefix of `treedef`. |
| 789 | |
| 790 | Returns: |
| 791 | A complete tree of PartitionSpecs or ParamPartitionSpecs and Nones that have the exact same |
| 792 | structure as `treedef`. |
| 793 | |
| 794 | Raises: |
| 795 | ValueError: If an unsupported type is encountered, or if partition_spec_tree is not a tree |
| 796 | prefix of treedef. |
| 797 | """ |
| 798 | proxy = object() |
| 799 | dummy = jax.tree_util.tree_unflatten(treedef, [object()] * treedef.num_leaves) |
| 800 | axes = [] |
| 801 | |
| 802 | def replace_none_with_proxy(tree): |
| 803 | if tree is None: |
| 804 | return proxy |
| 805 | if isinstance(tree, PartitionSpec) or dataclasses.is_dataclass(tree): |
| 806 | return tree |
| 807 | if is_named_tuple(tree): |
| 808 | return type(tree)(*[replace_none_with_proxy(x) for x in tree]) |