pytc.vmc.sharding.pad_walker

pytc.vmc.sharding.pad_walker(walker, target_n_walkers)[source]

Pad a Walker pytree along axis 0 to target_n_walkers.

Extra walkers are copies of the first walker (so they have valid shapes/dtypes for JIT tracing). They should be excluded from statistics after the training step.

Parameters:

target_n_walkers (int)