Returns a list of _ShardInfo for addressable shards that need to be saved. If replica count for the shards are greater than 0, all replicas will save slices of the shard provided that any dim of the shard is divisible by the replica count. If no such dim exists, we fallback to only repl
(
arr_inp: Tensor, *, max_data_shard_degree: int, shard_threshold_bytes: int
)
| 186 | |
| 187 | |
| 188 | def _get_shard_infos( |
| 189 | arr_inp: Tensor, *, max_data_shard_degree: int, shard_threshold_bytes: int |
| 190 | ) -> list[_ShardInfo]: |
| 191 | """Returns a list of _ShardInfo for addressable shards that need to be saved. |
| 192 | |
| 193 | If replica count for the shards are greater than 0, all replicas will save slices of the |
| 194 | shard provided that any dim of the shard is divisible by the replica count. If no such |
| 195 | dim exists, we fallback to only replica 0 saving the shard. |
| 196 | """ |
| 197 | shard_infos: list[_ShardInfo] = [] |
| 198 | replica_count_map = _num_replicas_per_shard(arr_inp) |
| 199 | for shard in arr_inp.addressable_shards: |
| 200 | replica_count = replica_count_map[_slices_to_tuple(shard.index)] |
| 201 | assert replica_count > 0 |
| 202 | shard_degree = ( |
| 203 | min(replica_count, max_data_shard_degree) |
| 204 | if max_data_shard_degree > 0 |
| 205 | else replica_count |
| 206 | ) |
| 207 | should_skip = ( |
| 208 | shard_degree == 1 |
| 209 | or shard.data.nbytes < shard_threshold_bytes |
| 210 | or shard.replica_id >= shard_degree |
| 211 | ) |
| 212 | for axis, size in enumerate(shard.data.shape): |
| 213 | # Find the first dim divisible by partial replication size. |
| 214 | if should_skip or size % shard_degree != 0: |
| 215 | continue |
| 216 | part_size = size // shard_degree |
| 217 | slice_obj = shard.index[axis] |
| 218 | assert slice_obj.step is None |
| 219 | start_offset = shard.replica_id * part_size |
| 220 | end_offset = start_offset + part_size |
| 221 | # When an axis of a tensor is not sharded, the slice object corresponding to |
| 222 | # that axis will be (None, None, None). |
| 223 | slice_start = slice_obj.start or 0 |
| 224 | shard_infos.append( |
| 225 | _ShardInfo( |
| 226 | shard.data, |
| 227 | shard.index[:axis] |
| 228 | + (slice(slice_start + start_offset, slice_start + end_offset),) |
| 229 | + shard.index[axis + 1 :], |
| 230 | (start_offset, end_offset, axis), |
| 231 | shard_degree, |
| 232 | ) |
| 233 | ) |
| 234 | break |
| 235 | else: |
| 236 | # We only have 1 replica or shard is not evenly divisible across replicas. |
| 237 | # Assign replica=0 only. |
| 238 | if shard.replica_id == 0: |
| 239 | shard_infos.append(_ShardInfo(shard.data, shard.index, None, 1)) |
| 240 | return shard_infos |
| 241 | |
| 242 | |
| 243 | def _transfer_to_host(data: Tensor) -> Tensor: |