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

Function _as_slice_indices

imperative/python/megengine/xla/sharding.py:330–350  ·  view source on GitHub ↗
(arr, idx)

Source from the content-addressed store, hash-verified

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

Callers 1

_shard_xla_device_arrayFunction · 0.85

Calls 2

listFunction · 0.85
appendMethod · 0.45

Tested by

no test coverage detected