pytc.vmc.sharding.shard_walker

pytc.vmc.sharding.shard_walker(walker, mesh, axis_name='walkers')[source]

Place a Walker (or any pytree) with sharding along axis 0.

Every leaf whose leading dimension equals n_walkers is split across devices. Scalar leaves are replicated automatically.

Parameters:
  • mesh (jax.sharding.Mesh)

  • axis_name (str)