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) """ # Default cost function: mean energy if cost_fn is None: cost_fn = jnp.mean # Choose vmap implementation 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 ) # Note: We don't need batch_log_psi anymore - in the JVP we call ansatz directly on the batch if use_custom_jvp: # Internal structure to hold all data needed for gradient computation 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 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 # Compute all local energies using vmap (or batched_vmap) energies = batch_local_energy(walkers, params) # Compute statistics mean_energy = jnp.mean(energies) energy_std = jnp.mean(jnp.abs(energies - mean_energy)) # Clip energies to avoid numerical instability 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 # Compute cost cost = cost_fn(clipped_energies) # Return cost and auxiliary data (namedtuple is indexable and JIT-compatible) 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 # Run forward pass to get primal output and cached intermediate values 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 if isinstance(batch_data, tuple): walkers = batch_data[0] else: walkers = batch_data # batch_network takes (walkers, params) but we only differentiate params # So we curry it to make a function of just params # batch_network takes (walkers, params) but we only differentiate params # So we curry it to make a function of just params def log_psi_fn(p): # Use standard vmap with checkpointing for correct global gradients return vmap_impl( jax.checkpoint(lambda w, p: ansatz_dynamic(w, p)[0][1]), in_axes=(0, None), out_axes=0 )(walkers, p) # Single JVP call - now only differentiating wrt params log_psi_primal, log_psi_tangent = jax.jvp( log_psi_fn, (params,), (params_tangent,) ) # VMC gradient: single dot product! # ∇⟨E⟩ = ⟨(E_L - ⟨E⟩) * ∇log|ψ|⟩ = dot(diff, ∇log|ψ|) / N n_walkers = diff.shape[0] cost_tangent = jnp.dot(diff, log_psi_tangent) / n_walkers # Return primal cost and tangent # 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: # Standard loss without custom JVP 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) """ # Extract walkers from batch 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 # Define batch_local_energy using the current ansatz batch_local_energy = vmap_impl( lambda w, p: ansatz_dynamic.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) # Compute all local energies using vmap energies = batch_local_energy(walkers, params) # Compute statistics mean_energy = jnp.mean(energies) energy_std = jnp.mean(jnp.abs(energies - mean_energy)) # Clip energies clipped_energies = jnp.clip( energies, mean_energy - clip_multiplier * energy_std, mean_energy + clip_multiplier * energy_std ) # Compute cost 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)) """ # Choose vmap implementation vmap_impl = get_vmap_fn(max_vmap_batch_size, mesh) # Define batch_local_energy using the current ansatz batch_local_energy = vmap_impl( lambda w, p: ansatz.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) # Define batch_network using the current ansatz 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) """ # Extract walkers from batch if isinstance(batch_data, tuple): walkers = batch_data[0] else: walkers = batch_data # Compute local energies energies = batch_local_energy(walkers, params) e_mean = jnp.mean(energies) e_std = jnp.mean(jnp.abs(energies - e_mean)) # Clip energies to suppress outliers (same scheme as make_energy_loss) 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) # Sample variance: sum((E - <E>)^2) / (n - 1) 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 # Extract walkers and ansatz 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") # Forward pass energies = batch_local_energy(walkers, params) # Pass ansatz 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)) # ========== Standard Gradient Method ========== # Compute JVP of local energies def compute_energies(p): # Use standard vmap with checkpointing 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,) ) # Single JVP call - now only differentiating wrt params #log_psi_primal = batch_network(walkers, params) # 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: # Standard variance loss without custom JVP 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) """ # Extract walkers from batch 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 # Define batch_local_energy using the current ansatz batch_local_energy = vmap_impl( lambda w, p: ansatz_dynamic.local_energy(w, p)[0], in_axes=(0, None), out_axes=0 ) # Compute local energies energies = batch_local_energy(walkers, params) e_mean = jnp.mean(energies) e_std = jnp.mean(jnp.abs(energies - e_mean)) # Clip energies 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) # Sample variance: sum((E - <E>)^2) / (n - 1) 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