(x: xc._xla.DeviceArray, devices, indices, sharding=None)
| 327 | |
| 328 | |
| 329 | def _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 | |
| 359 | def _shard_mge_tensor(x, devices, indices, sharding=None): |
nothing calls this directly
no test coverage detected