Gets the global replication count for each unique shard.
(arr: Tensor)
| 176 | |
| 177 | |
| 178 | def _num_replicas_per_shard(arr: Tensor) -> dict[tuple[_SliceTuple, ...], int]: |
| 179 | """Gets the global replication count for each unique shard.""" |
| 180 | replica_count = defaultdict(int) |
| 181 | for slices in arr.sharding.devices_indices_map(arr.shape).values(): |
| 182 | # Slice() object is not hashable before python 3.12. |
| 183 | # Manually convert it to hashable types. |
| 184 | replica_count[_slices_to_tuple(slices)] += 1 |
| 185 | return dict(replica_count) |
| 186 | |
| 187 | |
| 188 | def _get_shard_infos( |