Source code for pytc.vmc.hamiltonian

"""Hamiltonian and local energy computation functions for VMC."""

import jax
import jax.numpy as jnp


[docs] def compute_jastrow_terms(sj, elec_coords, jastrow_params): """Compute ∇J/J and ∇²J/J with explicit parameters.""" n_electrons = elec_coords.shape[0] # Fully vectorized implementation (O(N^2) parallelism) # Optimized for symmetric Jastrow factors (u(r1, r2) = u(r2, r1)) # 1. Create all pairs (i, j) # Broadcast to (N, N, 3) r1 = elec_coords[:, None, :] # (N, 1, 3) - represents i r2 = elec_coords[None, :, :] # (1, N, 3) - represents j # 2. Compute gradients and laplacians for all pairs # We only compute gradients w.r.t first argument (r_i) # Due to symmetry, sum_j grad_1(ri, rj) == sum_i grad_2(ri, rj) # And specifically for the total gradient on electron k: # grad_k U = sum_{j!=k} grad_1(rk, rj) def compute_pair_grads(r_i, r_j): g1, l1 = sj.jastrow.get_log_grads_r1(r_i, r_j, jastrow_params) return g1, l1 # vmap over j (inner), then i (outer) inner_vmap = jax.vmap(compute_pair_grads, in_axes=(None, 0)) outer_vmap = jax.vmap(inner_vmap, in_axes=(0, None)) # Compute for all pairs # g1s: (N, N, 3), l1s: (N, N) g1s, l1s = outer_vmap(elec_coords, elec_coords) # 3. Mask diagonal (i == j) mask = 1.0 - jnp.eye(n_electrons) # Expand mask for gradients (N, N, 1) mask_grad = mask[:, :, None] g1s = g1s * mask_grad l1s = l1s * mask # 4. Sum over j to get values for each electron i sum_g1 = jnp.sum(g1s, axis=1) sum_l1 = jnp.sum(l1s, axis=1) # 5. Result # grad_k U = sum_{j!=k} grad_1(rk, rj) # The factor of 0.5 from the definition U = 0.5 * sum u(ri, rj) cancels with the # fact that we have two identical sums (one for i=k, one for j=k). # So we just take the sum over j of grad_1. grad_J_over_J = sum_g1 lap_sum = sum_l1 grad_squared = jnp.sum(grad_J_over_J**2, axis=1) lap_J_over_J = lap_sum + grad_squared return grad_J_over_J, lap_J_over_J
[docs] def compute_potential_matrix(sj, elec_coords, slater_alpha, slater_beta): """Compute potential energy part of B matrix.""" n_alpha = sj.dets[0].n_alpha n_electrons = len(elec_coords) atom_coords = sj.atom_coords atom_charges = sj.atom_charges def e_n_potential(r): r_reshaped = r[:, jnp.newaxis, :] diff = r_reshaped - atom_coords[jnp.newaxis, :, :] dists = jnp.linalg.norm(diff, axis=2) potentials = -atom_charges[jnp.newaxis, :] / (dists + 1e-10) return jnp.sum(potentials, axis=1) alpha_coords = jnp.take(elec_coords, jnp.arange(n_alpha), axis=0) beta_coords = jnp.take(elec_coords, jnp.arange(n_alpha, n_electrons), axis=0) V_en_alpha = e_n_potential(alpha_coords)[:, None] V_en_beta = e_n_potential(beta_coords)[:, None] def pairwise_distance(r_i, r_j): diff = r_i - r_j dist = jnp.sqrt(jnp.sum(diff**2) + 1e-10) return 1.0 / dist ee_vmap_inner = jax.vmap(pairwise_distance, in_axes=(None, 0)) ee_vmap_outer = jax.vmap(ee_vmap_inner, in_axes=(0, None)) all_e_e_pot = jax.jit(ee_vmap_outer)(elec_coords, elec_coords) mask = 1.0 - jnp.eye(n_electrons) all_e_e_pot = all_e_e_pot * mask e_e_pot = 0.5 * jnp.sum(all_e_e_pot, axis=1) V_ee_alpha = e_e_pot[:n_alpha, None] V_ee_beta = e_e_pot[n_alpha:, None] B_alpha = (V_en_alpha + V_ee_alpha) * slater_alpha B_beta = (V_en_beta + V_ee_beta) * slater_beta return B_alpha, B_beta
[docs] def compute_single_walker_energy(sj, walker, jastrow_params): """Compute energy for a single walker. Args: sj: SlaterJastrow ansatz object walker: Walker object containing positions and Slater matrices jastrow_params: Jastrow parameters Returns: Local energy value """ n_alpha = sj.dets[0].n_alpha # Compute Jastrow terms internally grad_J_over_J, lap_J_over_J = compute_jastrow_terms(sj, walker.positions, jastrow_params) grad_J_alpha = grad_J_over_J[:n_alpha] grad_J_beta = grad_J_over_J[n_alpha:] lap_J_alpha = lap_J_over_J[:n_alpha] lap_J_beta = lap_J_over_J[n_alpha:] B_kin_alpha = -0.5 * ( walker.lap_up + 2 * jnp.einsum('ik,ijk->ij', grad_J_alpha, walker.grad_up) + jnp.multiply(lap_J_alpha[:, None], walker.slater_up) ) B_kin_beta = -0.5 * ( walker.lap_down + 2 * jnp.einsum('ik,ijk->ij', grad_J_beta, walker.grad_down) + jnp.multiply(lap_J_beta[:, None], walker.slater_down) ) B_pot_alpha, B_pot_beta = compute_potential_matrix( sj, walker.positions, walker.slater_up, walker.slater_down ) E_L = (jnp.trace(walker.inv_up @ (B_kin_alpha + B_pot_alpha)) + jnp.trace(walker.inv_down @ (B_kin_beta + B_pot_beta))) E_L = E_L + sj.ion_ion_potential return jnp.real(E_L)
[docs] def eval_local_energy(sj, walker, params): """Evaluate local energy for a SlaterJastrow ansatz. Args: sj: SlaterJastrow ansatz object walker: Walker object params: Tuple of (jastrow_params, linear_coeffs) Returns: Tuple of (energy, walker) """ jastrow_params, linear_coeffs = params energy = compute_single_walker_energy(sj, walker, jastrow_params) return energy, walker