pytc.vmc.sharding.get_walker_sharding¶ pytc.vmc.sharding.get_walker_sharding(mesh, axis_name='walkers')[source]¶ NamedSharding that partitions the leading (walker) dimension. Parameters: mesh (jax.sharding.Mesh) axis_name (str)