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 = "/",
)
| 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 | |
| 1136 | def get_recursively( |
| 1137 | x: NestedTensor, path: Union[str, Sequence[str]], separator: Optional[str] = "/" |
| 1138 | ) -> NestedTensor: |
| 1139 | """Recursively indexes through the nested tensor. |
no outgoing calls