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)