quantax.utils.get_distributed_sharding#
- quantax.utils.get_distributed_sharding() NamedSharding#
The sharding that splits an array’s first dimension evenly across all devices in
jax.devices().- Returns:
A
jax.sharding.NamedShardingcombiningmake_mesh()withget_distributed_P().