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

Function copy_recursively

axlearn/common/utils.py:1100–1136  ·  view source on GitHub ↗

Sets target[path] = source[path]. Args: source: The source tree. target: The target tree. path: The sequence of keys used to recursively index the nested tensor. If `isinstance(path, str)`, it will be split into sequence of strings based on the `d

(
    *,
    source: NestedTensor,
    target: NestedTensor,
    path: Union[str, Sequence[str]],
    separator: str = "/",
)

Source from the content-addressed store, hash-verified

1098 global_index_to_local_shard = index_to_shard(local_shards, value.shape)
1099 # A mapping from dim -> local slices.
1100 dim_to_local_slices: list[list[slice]] = [[] for _ in range(value.ndim)]
1101
1102 # For each dim, we bucket global slices into the corresponding local slice index.
1103 for dim, slices in enumerate(dim_to_local_slices):
1104 global_to_local = {}
1105 for global_index in global_index_to_local_shard:
1106 global_slice = global_index[dim]
1107 size = global_slice[1] - global_slice[0]
1108 if global_slice not in global_to_local:
1109 # This exploits the fact that global slices are already sorted.
1110 bucket_idx = len(global_to_local)
1111 global_to_local[global_slice] = bucket_idx
1112 else:
1113 # Along a given dim, global slices can appear multiples times if the array is
1114 # replicated along a different dim, in which case the local slice is the same.
1115 bucket_idx = global_to_local[global_slice]
1116
1117 # We require uniform sharding along each dim.
1118 assert not slices or (slices[-1].stop - slices[-1].start) == size
1119 start = bucket_idx * size
1120 slices.append(slice(start, start + size))
1121
1122 # The local shape can be inferred from the last offset along each dim.
1123 local_shape = tuple(local_slices[-1].stop for local_slices in dim_to_local_slices)
1124
1125 # Build the final output array.
1126 output = np.empty(local_shape, dtype=value.dtype)
1127 for local_index, local_shard in zip(
1128 zip(*dim_to_local_slices), global_index_to_local_shard.values()
1129 ):
1130 output[tuple(local_index)] = local_shard.data
1131 return output
1132
1133 return jax.tree.map(get_local_array, global_arrays)
1134
1135
1136def get_recursively(
1137 x: NestedTensor, path: Union[str, Sequence[str]], separator: Optional[str] = "/"
1138) -> NestedTensor:
1139 """Recursively indexes through the nested tensor.

Callers 1

test_copy_recursivelyMethod · 0.90

Calls

no outgoing calls

Tested by 1

test_copy_recursivelyMethod · 0.72