(x, devices, indices, sharding=None)
| 321 | |
| 322 | |
| 323 | def _shard_nparray(x, devices, indices, sharding=None): |
| 324 | if x.shape == (): |
| 325 | return device_put([x] * len(devices), devices) |
| 326 | return device_put([x[i] for i in indices], devices) |
| 327 | |
| 328 | |
| 329 | def _shard_xla_device_array(x: xc._xla.DeviceArray, devices, indices, sharding=None): |
nothing calls this directly
no test coverage detected