MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / _shard_xla_device_array

Function _shard_xla_device_array

imperative/python/megengine/xla/sharding.py:329–356  ·  view source on GitHub ↗
(x: xc._xla.DeviceArray, devices, indices, sharding=None)

Source from the content-addressed store, hash-verified

327
328
329def _shard_xla_device_array(x: xc._xla.DeviceArray, devices, indices, sharding=None):
330 def _as_slice_indices(arr, idx):
331 start_indices = [0] * arr.ndim
332 limit_indices = list(arr.shape)
333 removed_dims = []
334
335 tuple_idx = idx if isinstance(idx, tuple) else (idx,)
336 for dim, sub_idx in enumerate(tuple_idx):
337 if isinstance(sub_idx, int):
338 start_indices[dim] = sub_idx
339 limit_indices[dim] = sub_idx + 1
340 removed_dims.append(dim)
341 elif sub_idx == slice(None):
342 continue
343 else:
344 assert isinstance(sub_idx, slice), sub_idx
345 assert isinstance(sub_idx.start, int), sub_idx
346 assert isinstance(sub_idx.stop, int), sub_idx
347 start_indices[dim] = sub_idx.start
348 limit_indices[dim] = sub_idx.stop
349
350 return tuple(start_indices), tuple(limit_indices), tuple(removed_dims)
351
352 start_indices, limit_indices, removed_dims = unzip3(
353 _as_slice_indices(x, idx) for idx in indices
354 )
355 shards = x._multi_slice(start_indices, limit_indices, removed_dims)
356 return device_put(shards, devices)
357
358
359def _shard_mge_tensor(x, devices, indices, sharding=None):

Callers

nothing calls this directly

Calls 3

unzip3Function · 0.85
_as_slice_indicesFunction · 0.85
device_putFunction · 0.85

Tested by

no test coverage detected