pytc.vmc.sharding._assemble_sharded_from_local_pytrees

pytc.vmc.sharding._assemble_sharded_from_local_pytrees(local_pytrees, mesh, axis_name='walkers')[source]

Assemble per-device local pytrees into a globally sharded pytree.

Parameters:
  • mesh (jax.sharding.Mesh)

  • axis_name (str)