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