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)