pytc.vmc.sharding.initialize_walkers_sharded

pytc.vmc.sharding.initialize_walkers_sharded(ansatz, n_walkers, mesh, initial_walkers=None, key=None)[source]

Initialize walkers independently per device and return globally sharded walkers.

This avoids creating a full (n_walkers, ...) walker tensor on one GPU before sharding, which is important for very large walker counts.

Parameters:
  • n_walkers (int)

  • mesh (jax.sharding.Mesh)