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

Function _num_replicas_per_shard

axlearn/common/array_serialization.py:178–185  ·  view source on GitHub ↗

Gets the global replication count for each unique shard.

(arr: Tensor)

Source from the content-addressed store, hash-verified

176
177
178def _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
188def _get_shard_infos(

Calls 1

_slices_to_tupleFunction · 0.85