(devices, indices, args, shardings=None)
| 371 | |
| 372 | |
| 373 | def shard_args(devices, indices, args, shardings=None): |
| 374 | def _shard_arg(arg, devices, arg_indices, sharding=None): |
| 375 | arg = canonicalize_arg(arg) |
| 376 | return shard_arg_handlers[type(arg)](arg, devices, arg_indices, sharding) |
| 377 | |
| 378 | if shardings is None: |
| 379 | return [_shard_arg(arg, devices, indices[i]) for i, arg in enumerate(args)] |
| 380 | else: |
| 381 | return [ |
| 382 | _shard_arg(arg, devices, indices[i], shardings[i]) |
| 383 | for i, arg in enumerate(args) |
| 384 | ] |
| 385 | |
| 386 | |
| 387 | @functools.lru_cache() |
nothing calls this directly
no test coverage detected