pytc.vmc.optimization.make_opt_update_step

pytc.vmc.optimization.make_opt_update_step(loss_fn, optimizer, gradient_mask=None)[source]

Factory to create a JIT-compilable optimizer step for Optax optimizers.

Parameters:
  • loss_fn – Loss function with signature (params, walkers) -> (loss, aux_data) where aux_data is a tuple of auxiliary outputs

  • optimizer – Optax optimizer (e.g., optax.adam)

  • gradient_mask – Optional boolean PyTree matching the parameters. Gradients at False leaves are zeroed before the optimizer update.

Returns:

opt_step(ansatz, params, walkers, opt_state, key) -> (params, opt_state, loss, aux_data)

Return type:

A JIT-compiled function with signature