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)