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_walkersis split across devices. Scalar leaves are replicated automatically.- Parameters:
mesh (jax.sharding.Mesh)
axis_name (str)