MCPcopy Create free account
hub / github.com/apple/axlearn / dispatch_input_batch

Function dispatch_input_batch

axlearn/common/utils.py:763–805  ·  view source on GitHub ↗

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"
)

Source from the content-addressed store, hash-verified

761 return thread_resources.env.physical_mesh
762 return mesh
763
764
765def 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
777def 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])

Calls 3

with_sharding_constraintFunction · 0.85
traverse_and_dispatchFunction · 0.85
mapMethod · 0.80