pytc.vmc.optimization.make_opt_update_step

pytc.vmc.optimization.make_opt_update_step(loss_fn, optimizer)[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)

Returns:

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

Return type:

A JIT-compiled function with signature