Source code for pytc.vmc.loss

"""Loss functions for VMC optimization.

This module contains various loss functions used in variational Monte Carlo
optimization, including energy minimization and variance minimization.
"""

import functools
import jax
import jax.numpy as jnp
from typing import Callable, Optional
from .sharding import get_vmap_fn


[docs] def make_energy_loss( ansatz, optimizer_type: str = "adam", cost_fn: Optional[Callable] = None, clip_multiplier: float = 5.0, use_custom_jvp: bool = True, max_vmap_batch_size: int = 0, mesh: Optional[jax.sharding.Mesh] = None, ): """Factory to create energy-based loss function for VMC optimization. This creates a loss function that computes the average local energy, with optional energy clipping and custom JVP for memory efficiency. Args: ansatz: Wavefunction object with local_energy method optimizer_type: Type of optimizer ("adam", "sgd", "kfac", etc.) cost_fn: Optional cost function to apply to energies (defaults to mean) clip_multiplier: Multiplier for energy clipping range (clips to mean ± multiplier * std) use_custom_jvp: Whether to use custom JVP for memory-efficient gradients max_vmap_batch_size: If 0, use standard vmap everywhere. If >0, use folx.batched_vmap (or sharded_batched_vmap if multi-device). mesh: Optional device mesh for sharding. Returns: Loss function with signature (params, batch_data) -> (loss, AuxData) where AuxData is a namedtuple with (mean_energy, energy_std, clipped_energies, diff) """ if cost_fn is None: cost_fn = jnp.mean vmap_impl = get_vmap_fn(max_vmap_batch_size, mesh) batch_local_energy = vmap_impl( lambda w, p: ansatz.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) batch_network = vmap_impl( lambda w, p: ansatz(w, p)[0][1], # Returns log_psi in_axes=(0, None), out_axes=0 ) if use_custom_jvp: from collections import namedtuple AuxData = namedtuple('AuxData', ['mean_energy', 'energy_std', 'clipped_energies', 'diff']) @jax.custom_jvp def loss_fn(params, batch_data): """Energy loss function with custom JVP. Args: params: [jastrow_params, linear_coeffs] batch_data: Either walkers (Optax) or (walkers, None) (KFAC) Returns: loss: Scalar loss value (mean energy or custom cost) aux: AuxData namedtuple with (mean_energy, energy_std, clipped_energies, diff) Can be indexed as aux[0], aux[1] for backward compatibility """ # Extract walkers from batch. The custom-JVP path always # evaluates with the factory ansatz (primal AND tangent); a # batch-carried ansatz in the (walkers, ansatz) convention is # honored only by the non-JVP variant below. Mixing the two # here would make the primal energies and the JVP tangent # disagree. if isinstance(batch_data, tuple) and len(batch_data) == 2: walkers = batch_data[0] else: walkers = batch_data energies = batch_local_energy(walkers, params) mean_energy = jnp.mean(energies) energy_std = jnp.mean(jnp.abs(energies - mean_energy)) clipped_energies = jnp.clip( energies, mean_energy - clip_multiplier * energy_std, mean_energy + clip_multiplier * energy_std ) # Compute diff for gradient computation (store for reuse in JVP) mean_clipped = jnp.mean(clipped_energies) diff = clipped_energies - mean_clipped cost = cost_fn(clipped_energies) return cost, AuxData(mean_energy, energy_std, clipped_energies, diff) @loss_fn.defjvp def loss_fn_jvp(primals, tangents): """Custom JVP for memory-efficient VMC gradients. This implements the VMC gradient estimator: ∇<E> = <(E_L - <E>) * ∇log|ψ|> Key optimization: Reuses energies and diff from forward pass! 1. Forward pass already computed clipped_energies and diff 2. Extract diff from auxiliary data (no need to recompute energies!) 3. Single jvp call on batch_log_psi 4. Gradient = dot(diff, log_psi_tangent) / N By only differentiating log|ψ| instead of the full local_energy, we avoid materializing gradients of non-parameter-dependent terms. """ params, batch_data = primals params_tangent, _ = tangents cost_primal, aux_data = loss_fn(params, batch_data) mean_energy = aux_data.mean_energy energy_std = aux_data.energy_std clipped_energies = aux_data.clipped_energies diff = aux_data.diff # Extract walkers from batch_data, mirroring the forward pass. # The tangent uses the factory ansatz, matching the primal # energies above (see the forward-pass note). if isinstance(batch_data, tuple) and len(batch_data) == 2: walkers = batch_data[0] else: walkers = batch_data # batch_network takes (walkers, params) but we only differentiate # params, so curry it to a function of params alone. def log_psi_fn(p): # Use standard vmap with checkpointing for correct global gradients return vmap_impl( jax.checkpoint(lambda w, p: ansatz(w, p)[0][1]), in_axes=(0, None), out_axes=0 )(walkers, p) log_psi_primal, log_psi_tangent = jax.jvp( log_psi_fn, (params,), (params_tangent,) ) # ∇⟨E⟩ = ⟨(E_L - ⟨E⟩) * ∇log|ψ|⟩ = dot(diff, ∇log|ψ|) / N n_walkers = diff.shape[0] cost_tangent = jnp.dot(diff, log_psi_tangent) / n_walkers # Tangent aux_data: use zeros for cached values (they're not differentiated) tangent_aux = AuxData(0.0, 0.0, jnp.zeros_like(clipped_energies), jnp.zeros_like(diff)) return (cost_primal, aux_data), (cost_tangent, tangent_aux) return loss_fn else: def loss_fn(params, batch_data): """Energy loss function (standard autodiff). Args: params: [jastrow_params, linear_coeffs] batch_data: Either walkers (Optax) or (walkers, None) (KFAC) Returns: loss: Scalar loss value (mean energy or custom cost) aux: Tuple of (mean_energy, energy_std) """ if isinstance(batch_data, tuple) and len(batch_data) == 2: walkers, ansatz_arg = batch_data ansatz_dynamic = ansatz_arg else: walkers = batch_data ansatz_dynamic = ansatz batch_local_energy = vmap_impl( lambda w, p: ansatz_dynamic.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) energies = batch_local_energy(walkers, params) mean_energy = jnp.mean(energies) energy_std = jnp.mean(jnp.abs(energies - mean_energy)) clipped_energies = jnp.clip( energies, mean_energy - clip_multiplier * energy_std, mean_energy + clip_multiplier * energy_std ) cost = cost_fn(clipped_energies) return cost, (mean_energy, energy_std) return loss_fn
[docs] def make_variance_loss( ansatz, optimizer_type: str = "adam", use_custom_jvp: bool = True, max_vmap_batch_size: int = 0, clip_multiplier: float = 5.0, mesh: Optional[jax.sharding.Mesh] = None, ): """Factory to create variance-based loss function for reference variance optimization. This minimizes the variance of local energies with respect to a reference determinant, which can improve the quality of the Jastrow factor. Uses vmap (or batched_vmap for memory efficiency) following the same pattern as make_energy_loss. Args: ansatz: Wavefunction object with local_energy method optimizer_type: Type of optimizer ("adam", "sgd", "kfac", etc.) use_custom_jvp: Whether to use custom JVP for memory-efficient gradients max_vmap_batch_size: If 0, use standard vmap. If >0, use folx.batched_vmap for memory efficiency. Recommended batch size: 10-50. clip_multiplier: Multiplier for energy clipping range (clips to mean ± multiplier * mean absolute deviation (MAD) of the local energy). Set to 0 to disable clipping. Default 5.0, matching make_energy_loss. mesh: Optional device mesh for sharding. Returns: Loss function with signature (params, batch_data) -> (variance, (mean_energy, energy_mad)) """ vmap_impl = get_vmap_fn(max_vmap_batch_size, mesh) batch_local_energy = vmap_impl( lambda w, p: ansatz.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) batch_network = vmap_impl( lambda w, p: ansatz(w, p)[0][1], # Returns log_psi in_axes=(0, None), out_axes=0 ) if use_custom_jvp: @jax.custom_jvp def loss_fn(params, batch_data): """Variance loss function with custom JVP. Args: params: [jastrow_params, linear_coeffs] batch_data: Either walkers (Optax) or (walkers, None) (KFAC) Returns: variance: Sample variance of local energies aux: Tuple of (mean_energy, energy_std) """ if isinstance(batch_data, tuple): walkers = batch_data[0] else: walkers = batch_data energies = batch_local_energy(walkers, params) e_mean = jnp.mean(energies) e_std = jnp.mean(jnp.abs(energies - e_mean)) if clip_multiplier > 0: energies = jnp.clip( energies, e_mean - clip_multiplier * e_std, e_mean + clip_multiplier * e_std, ) # Recompute mean after clipping for a consistent variance e_mean = jnp.mean(energies) n_walkers = energies.shape[0] variance = jnp.sum((energies - e_mean)**2) / (n_walkers - 1) if n_walkers > 1 else 0.0 return variance, (e_mean, jnp.std(energies)) @loss_fn.defjvp def loss_fn_jvp(primals, tangents): """Custom JVP for variance minimization. """ params, batch_data = primals params_tangent, _ = tangents jastrow_params_tangent, linear_coeffs_tangent = params_tangent if isinstance(batch_data, tuple) and len(batch_data) == 2: walkers, ansatz_arg = batch_data ansatz_dynamic = ansatz_arg else: walkers = batch_data ansatz_dynamic = ansatz if ansatz_dynamic is None: raise ValueError("Ansatz must be provided either in make_variance_loss or in batch_data") energies = batch_local_energy(walkers, params) e_mean = jnp.mean(energies) e_std = jnp.mean(jnp.abs(energies - e_mean)) # Clip energies (must match the forward pass exactly) if clip_multiplier > 0: energies = jnp.clip( energies, e_mean - clip_multiplier * e_std, e_mean + clip_multiplier * e_std, ) e_mean = jnp.mean(energies) n_walkers = energies.shape[0] variance = jnp.sum((energies - e_mean)**2) / (n_walkers - 1) if n_walkers > 1 else 0.0 aux_data = (e_mean, jnp.std(energies)) def compute_energies(p): return vmap_impl( jax.checkpoint(lambda w, p: ansatz_dynamic.local_energy(w, p)[0]), in_axes=(0, None), out_axes=0 )(walkers, p) _, energy_tangent = jax.jvp( compute_energies, (params,), (params_tangent,) ) # Variance gradient: ∇var = 2 * mean((E_L - ⟨E⟩) * ∇E_L) energy_diff = energies - e_mean if n_walkers > 1: variance_tangent = 2.0 * jnp.dot(energy_diff, energy_tangent) / (n_walkers - 1) else: variance_tangent = 0.0 return (variance, aux_data), (variance_tangent, aux_data) return loss_fn else: def loss_fn(params, batch_data): """Variance loss function (standard autodiff). Args: params: [jastrow_params, linear_coeffs] batch_data: Either walkers (Optax) or (walkers, None) (KFAC) Returns: variance: Sample variance of local energies aux: Tuple of (mean_energy, energy_std) """ if isinstance(batch_data, tuple) and len(batch_data) == 2: walkers, ansatz_arg = batch_data ansatz_dynamic = ansatz_arg else: walkers = batch_data ansatz_dynamic = ansatz batch_local_energy = vmap_impl( lambda w, p: ansatz_dynamic.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) energies = batch_local_energy(walkers, params) e_mean = jnp.mean(energies) e_std = jnp.mean(jnp.abs(energies - e_mean)) if clip_multiplier > 0: energies = jnp.clip( energies, e_mean - clip_multiplier * e_std, e_mean + clip_multiplier * e_std, ) e_mean = jnp.mean(energies) n_walkers = energies.shape[0] variance = jnp.sum((energies - e_mean)**2) / (n_walkers - 1) if n_walkers > 1 else 0.0 return variance, (e_mean, jnp.std(energies)) return loss_fn