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)