quantax.utils.get_distributed_P#

quantax.utils.get_distributed_P(axis: int = 0) P#

The jax.sharding.PartitionSpec (jax.P) that distributes an array along the given axis over both mesh axes.