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

Function _get_shard_infos

axlearn/common/array_serialization.py:188–240  ·  view source on GitHub ↗

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
)

Source from the content-addressed store, hash-verified

186
187
188def _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
243def _transfer_to_host(data: Tensor) -> Tensor:

Callers 2

_verify_shard_infoMethod · 0.90
_async_serializeFunction · 0.85

Calls 3

_num_replicas_per_shardFunction · 0.85
_slices_to_tupleFunction · 0.85
_ShardInfoClass · 0.85

Tested by 1

_verify_shard_infoMethod · 0.72