"""Tests for the Ansatz class."""
import unittest
import numpy as np
import jax
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
from pyscf import gto, scf
from pytc.ansatz import SlaterJastrow, SlaterDet
from pytc.jastrow import Poly
from pytc.vmc.walker import Walker
[docs]
def create_test_walker(positions, det):
"""Helper to create a Walker for testing.
Creates unbatched walker - positions should have shape (n_electrons, 3).
"""
n_alpha = det.n_alpha
n_beta = det.n_beta
n_electrons = n_alpha + n_beta
is_batched = positions.ndim == 3
if is_batched:
batch_size = positions.shape[0]
return Walker(
positions=positions,
det_up=(jnp.ones(batch_size), jnp.zeros(batch_size)),
det_down=(jnp.ones(batch_size), jnp.zeros(batch_size)),
slater_up=jnp.zeros((batch_size, n_alpha, n_alpha)),
slater_down=jnp.zeros((batch_size, n_beta, n_beta)),
inv_up=jnp.zeros((batch_size, n_alpha, n_alpha)),
inv_down=jnp.zeros((batch_size, n_beta, n_beta)),
grad_up=jnp.zeros((batch_size, n_alpha, n_alpha, 3)),
grad_down=jnp.zeros((batch_size, n_beta, n_beta, 3)),
lap_up=jnp.zeros((batch_size, n_alpha, n_alpha)),
lap_down=jnp.zeros((batch_size, n_beta, n_beta)),
move_mask=jnp.ones((batch_size, n_electrons), dtype=bool),
log_psi=jnp.zeros((batch_size,)),
psi_sign=jnp.zeros((batch_size,)),
log_jastrow=jnp.zeros((batch_size,)),
)
else:
return Walker(
positions=positions,
det_up=(jnp.array(1.0), jnp.array(0.0)),
det_down=(jnp.array(1.0), jnp.array(0.0)),
slater_up=jnp.zeros((n_alpha, n_alpha)),
slater_down=jnp.zeros((n_beta, n_beta)),
inv_up=jnp.zeros((n_alpha, n_alpha)),
inv_down=jnp.zeros((n_beta, n_beta)),
grad_up=jnp.zeros((n_alpha, n_alpha, 3)),
grad_down=jnp.zeros((n_beta, n_beta, 3)),
lap_up=jnp.zeros((n_alpha, n_alpha)),
lap_down=jnp.zeros((n_beta, n_beta)),
move_mask=jnp.ones(n_electrons, dtype=bool),
log_psi=jnp.array(0.0),
psi_sign=jnp.array(0.0),
log_jastrow=jnp.array(0.0),
)
[docs]
class TestAnsatzH2(unittest.TestCase):
"""Test cases for Ansatz class using real H2 molecule."""
[docs]
def setUp(self):
"""Set up H2 molecule and compute RHF."""
self.mol = gto.M(
atom='H 0 0 0; H 0 0 0.742',
basis='sto-3g',
unit='bohr',
spin=2
)
self.mf = scf.RHF(self.mol)
self.mf.kernel()
self.det = SlaterDet.create(self.mol, self.mf.mo_coeff)
self.jastrow_params = jnp.array([0.5])
self.jastrow = Poly() # No params in constructor
self.linear_coeffs = jnp.array([1.0])
self.ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det])
# Test positions: two electrons slightly offset from nuclei (unbatched)
self.test_pos = jnp.array([
[0.0, 0.1, 0.0], # electron 1 near first H
[0.0, 0.1, 0.742], # electron 2 near second H
]) # Shape: (2, 3) - unbatched
self.params = (self.jastrow_params, self.linear_coeffs)
[docs]
def test_wavefunction_evaluation(self):
"""Test full wavefunction evaluation for H2."""
walker = create_test_walker(self.test_pos, self.det)
psi_values, updated_walker = self.ansatz(walker, self.params)
# psi_values is now (sign, log|psi|) tuple
psi_sign, psi_logabs = psi_values
psi_val = psi_sign * jnp.exp(psi_logabs)
self.assertTrue(np.isreal(psi_val)) # Unbatched walker
self.assertNotEqual(float(psi_val), 0.0) # Unbatched walker
# Test that moving electrons far apart gives smaller absolute value
far_pos = jnp.array([
[0.0, 0.0, -5.0],
[0.0, 0.0, 5.0],
]) # Shape: (2, 3) - unbatched
far_walker = create_test_walker(far_pos, self.det)
far_psi_values, _ = self.ansatz(far_walker, self.params)
far_psi_sign, far_psi_logabs = far_psi_values
far_psi_val = far_psi_sign * jnp.exp(far_psi_logabs)
self.assertLess(abs(float(far_psi_val)), abs(float(psi_val)))
[docs]
def test_jastrow_parameter_sensitivity(self):
"""Test sensitivity to Jastrow parameter changes."""
walker = create_test_walker(self.test_pos, self.det)
from pytc.ansatz.det import value_and_grad
_, populated_walker = value_and_grad(self.det, walker)
psi_values_original, _ = self.ansatz(populated_walker, self.params)
psi_sign_orig, psi_logabs_orig = psi_values_original
value_original = psi_sign_orig * jnp.exp(psi_logabs_orig)
new_params = (jnp.array([2.0]), self.linear_coeffs)
psi_values_new, _ = self.ansatz(populated_walker, new_params)
psi_sign_new, psi_logabs_new = psi_values_new
value_new = psi_sign_new * jnp.exp(psi_logabs_new)
self.assertNotAlmostEqual(float(value_original), float(value_new))
[docs]
def test_antisymmetry(self):
"""Test that wavefunction is antisymmetric under electron exchange."""
walker = create_test_walker(self.test_pos, self.det)
psi_values1, _ = self.ansatz(walker, self.params)
psi_sign1, psi_logabs1 = psi_values1
value1 = psi_sign1 * jnp.exp(psi_logabs1)
# Note: For H2 in RHF, we need to swap within same spin block to see antisymmetry
# First electron is spin-up, second is spin-down, so swapping won't show antisymmetry
spin_up_pos = jnp.array([[
[0.0, 0.1, 0.0], # first spin-up electron
[0.0, 0.1, 1.0], # second spin-up electron
]]) # Shape: (1, 2, 3)
walker1 = create_test_walker(spin_up_pos, self.det)
batch_ansatz = jax.vmap(self.ansatz, in_axes=(0, None))
psi_values1, _ = batch_ansatz(walker1, self.params)
psi_sign1, psi_logabs1 = psi_values1
value1 = psi_sign1 * jnp.exp(psi_logabs1)
swapped_pos = spin_up_pos[:, ::-1, :] # Swap along electron dimension
walker2 = create_test_walker(swapped_pos, self.det)
psi_values2, _ = batch_ansatz(walker2, self.params)
psi_sign2, psi_logabs2 = psi_values2
value2 = psi_sign2 * jnp.exp(psi_logabs2)
np.testing.assert_allclose(value1[0], -value2[0])
[docs]
def test_jastrow_terms(self):
"""Test computation of Jastrow gradient and laplacian terms."""
from pytc.vmc.hamiltonian import compute_jastrow_terms
grad_J, lap_J = compute_jastrow_terms(self.ansatz, self.test_pos, self.jastrow_params)
self.assertEqual(grad_J.shape, (2, 3)) # (n_electrons, xyz)
self.assertEqual(lap_J.shape, (2,)) # (n_electrons,)
# Gradients should be opposite for electrons near equilibrium
np.testing.assert_allclose(grad_J[0], -grad_J[1], rtol=1e-5)
[docs]
def test_jastrow_terms_analytical(self):
"""Test Jastrow gradient and laplacian against analytical values.
For a simple polynomial Jastrow factor u(r_ij) = a*r_ij with parameter a=0.5,
we can derive the analytical expressions for gradient and laplacian.
"""
# Use a simple Jastrow with u(r_ij) = 0.5*r_ij
simple_jastrow = Poly()
simple_ansatz = SlaterJastrow.create(self.mol, simple_jastrow, [self.det])
simple_jastrow_params = jnp.array([0.5])
# Use simple positions for easier analytical calculation
# Two electrons along the x-axis at positions 0 and 1
positions = jnp.array([
[0.0, 0.0, 0.0], # first electron at origin
[1.0, 0.0, 0.0] # second electron at x=1
])
# For u(r) = 0.5*r, where r = |r_i - r_j|
# The gradient with respect to r_i depends on the convention:
# In our implementation, we get:
# ∇_i u(r_ij) = 0.5 * (r_i - r_j)/|r_i - r_j|
# For our positions:
# ∇_1 u(r_12) = 0.5 * ([0,0,0] - [1,0,0])/1 = [-0.5, 0, 0]
# ∇_2 u(r_21) = 0.5 * ([1,0,0] - [0,0,0])/1 = [0.5, 0, 0]
from pytc.vmc.hamiltonian import compute_jastrow_terms
grad_J_over_J, lap_J_over_J = compute_jastrow_terms(simple_ansatz, positions, simple_jastrow_params)
expected_grad = jnp.array([
[-0.5, 0.0, 0.0], # gradient for electron 1
[0.5, 0.0, 0.0] # gradient for electron 2
])
expected_lap = jnp.array([1.25, 1.25]) # laplacian for electrons 1 and 2
np.testing.assert_allclose(grad_J_over_J, expected_grad, rtol=1e-5)
np.testing.assert_allclose(lap_J_over_J, expected_lap, rtol=1e-5)
different_jastrow_params = jnp.array([2.0]) # Parameter a=2.0
different_ansatz = SlaterJastrow.create(self.mol, simple_jastrow, [self.det])
grad_J_over_J_2, lap_J_over_J_2 = compute_jastrow_terms(different_ansatz, positions, different_jastrow_params)
np.testing.assert_allclose(grad_J_over_J_2, 4.0 * expected_grad, rtol=1e-5)
#np.testing.assert_allclose(lap_J_over_J_2, 4.0 * expected_lap, rtol=1e-5)
three_electron_pos = jnp.array([
[0.0, 0.0, 0.0], # at origin
[1.0, 0.0, 0.0], # along x-axis
[0.0, 1.0, 0.0] # along y-axis
])
grad_J_over_J_3, lap_J_over_J_3 = compute_jastrow_terms(simple_ansatz, three_electron_pos, simple_jastrow_params)
# For 3 electron case, the expected gradients and laplacians are:
# For electron 1 at [0,0,0]:
# Gradient wrt electron 2 at [1,0,0]:
# Vector = [0,0,0] - [1,0,0] = [-1,0,0]
# ∇₁u(r₁₂) = 0.5 * [-1,0,0]/1 = [-0.5, 0, 0]
# Gradient wrt electron 3 at [0,1,0]:
# Vector = [0,0,0] - [0,1,0] = [0,-1,0]
# ∇₁u(r₁₃) = 0.5 * [0,-1,0]/1 = [0, -0.5, 0]
# Total: ∇₁J/J = [-0.5, -0.5, 0]
# For electron 2 at [1,0,0]:
# Gradient wrt electron 1 at [0,0,0]:
# Vector = [1,0,0] - [0,0,0] = [1,0,0]
# ∇₂u(r₂₁) = 0.5 * [1,0,0]/1 = [0.5, 0, 0]
# Gradient wrt electron 3 at [0,1,0]:
# Vector = [1,0,0] - [0,1,0] = [1,-1,0]
# Distance = √2
# ∇₂u(r₂₃) = 0.5 * [1,-1,0]/√2 ≈ [0.35, -0.35, 0]
# Total: ∇₂J/J = [0.85, -0.35, 0]
# For electron 3 at [0,1,0]:
# Gradient wrt electron 1 at [0,0,0]:
# Vector = [0,1,0] - [0,0,0] = [0,1,0]
# ∇₃u(r₃₁) = 0.5 * [0,1,0]/1 = [0, 0.5, 0]
# Gradient wrt electron 2 at [1,0,0]:
# Vector = [0,1,0] - [1,0,0] = [-1,1,0]
# Distance = √2
# ∇₃u(r₃₂) = 0.5 * [-1,1,0]/√2 ≈ [-0.35, 0.35, 0]
# Total: ∇₃J/J = [-0.35, 0.85, 0]
expected_grad_3 = jnp.array([
[-0.5, -0.5, 0.0], # gradient for electron 1
[0.8535534, -0.35355338, 0.0], # gradient for electron 2
[-0.35355338, 0.8535534, 0.0] # gradient for electron 3
])
expected_lap_3 = jnp.array([2.5, 2.56066, 2.56066])
np.testing.assert_allclose(grad_J_over_J_3, expected_grad_3, rtol=1e-5)
np.testing.assert_allclose(lap_J_over_J_3, expected_lap_3, rtol=1e-5)
[docs]
def test_potential_matrix_values(self):
"""Test potential matrix calculations with analytical values."""
# For our H2 molecule with atoms at [0,0,0] and [0,0,0.742]
# and test positions at [[0.0, 0.1, 0.0], [0.0, 0.1, 0.742]]
# Electron-Nuclear potential for electron 1 at [0.0, 0.1, 0.0]:
# Distance to H atom 1 (at [0,0,0]): 0.1
# Distance to H atom 2 (at [0,0,0.742]): sqrt(0.01 + 0.550564) ≈ 0.748
# V_en_1 = -1/0.1 - 1/0.748 ≈ -11.34
# Electron-Nuclear potential for electron 2 at [0.0, 0.1, 0.742]:
# Distance to H atom 1: 0.748
# Distance to H atom 2: 0.1
# V_en_2 = -1/0.748 - 1/0.1 ≈ -11.34
# Electron-Electron potential:
# Distance between electrons: 0.742
# V_ee = 1/0.742 ≈ 1.35
slater_up, slater_down = self.det.matrix(self.test_pos)
from pytc.vmc.hamiltonian import compute_potential_matrix, compute_jastrow_terms
B_pot_up, B_pot_down = compute_potential_matrix(
self.ansatz,
self.test_pos, # Unbatched positions
slater_up,
slater_down
)
n_alpha = self.det.n_alpha
n_beta = self.det.n_beta
potentials = []
if n_alpha > 0: # If we have alpha electrons
pot_e1 = float(jnp.asarray(B_pot_up[0, 0] / slater_up[0, 0])) # Convert to scalar
potentials.append(pot_e1)
if n_beta > 0: # If we have beta electrons
pot_e2 = float(jnp.asarray(B_pot_down[0, 0] / slater_down[0, 0]))
potentials.append(pot_e2)
self.assertGreater(len(potentials), 0, "No potentials were calculated")
atom_coords = self.mol.atom_coords()
atom_charges = self.mol.atom_charges()
def compute_nuclear_pot(pos):
dists = jnp.linalg.norm(pos - atom_coords, axis=1)
return -jnp.sum(atom_charges / (dists + 1e-10))
positions = self.test_pos
n_electrons = len(positions)
for i in range(n_electrons):
e_n = compute_nuclear_pot(positions[i])
e_e_sum = 0.0
for j in range(n_electrons):
if i != j: # Skip self-interaction
e_e_dist = jnp.linalg.norm(positions[i] - positions[j])
e_e_sum += 1.0 / (e_e_dist + 1e-10)
expected_e_n = -1/0.1 - 1/np.sqrt(0.1**2+0.742**2) # From 1/0.1 + 1/0.748
expected_e_e = 1/0.742 # From 1/0.742
self.assertAlmostEqual(float(e_n), expected_e_n, delta=1e-6)
self.assertAlmostEqual(float(e_e_sum), expected_e_e, delta=1e-6)
expected_total = expected_e_n + expected_e_e/2.
if i < len(potentials):
self.assertAlmostEqual(potentials[i], expected_total, delta=1e-6)
[docs]
def test_param_gradient(self):
"""Test parameter gradient calculation."""
# For our current setup:
# - Jastrow with parameter a=0.5: exp(0.5*a*r_ij)
# - Single determinant
# - Test positions at [0.0, 0.1, 0.0] and [0.0, 0.1, 0.742]
electron_dist = jnp.linalg.norm(self.test_pos[0] - self.test_pos[1])
self.assertAlmostEqual(float(electron_dist), 0.742, places=3)
def wf_value(param):
param_tuple = (jnp.array([param]), self.linear_coeffs)
walker = create_test_walker(self.test_pos, self.det)
from pytc.ansatz.det import value_and_grad
_, populated_walker = value_and_grad(self.det, walker)
psi_values, _ = self.ansatz(populated_walker, param_tuple)
psi_sign, psi_logabs = psi_values
return psi_sign * jnp.exp(psi_logabs)
param_grad = jax.grad(wf_value)(0.5)
# Calculate expected gradient analytically:
# For Poly Jastrow with a single parameter a, the implementation is:
# J = exp(0.5 * sum_ij param * |r_i - r_j|)
#
# For two electrons:
walker = create_test_walker(self.test_pos, self.det)
from pytc.ansatz.det import value_and_grad
_, populated_walker = value_and_grad(self.det, walker)
psi_values, _ = self.ansatz(populated_walker, (self.jastrow_params, self.linear_coeffs))
psi_sign, psi_logabs = psi_values
current_wf = psi_sign * jnp.exp(psi_logabs)
dj_da = electron_dist # The full electron distance
expected_grad = float(current_wf * dj_da)
self.assertAlmostEqual(float(param_grad), expected_grad, places=8)
different_param = 1.0
different_jastrow_params = jnp.array([different_param])
different_params = (different_jastrow_params, self.linear_coeffs)
walker_diff = create_test_walker(self.test_pos, self.det)
psi_values_diff, _ = self.ansatz(walker_diff, different_params)
psi_sign_diff, psi_logabs_diff = psi_values_diff
different_wf = psi_sign_diff * jnp.exp(psi_logabs_diff)
# The dJ/da is the same (electron_dist), but the wavefunction value is different
different_expected_grad = float(different_wf * dj_da)
different_param_grad = jax.grad(wf_value)(different_param)
self.assertAlmostEqual(float(different_param_grad), different_expected_grad, places=8)
self.assertGreater(different_param, 0.5) # Parameter increased
ratio = different_wf / current_wf
self.assertGreater(ratio, 1.0) # Wavefunction increased
[docs]
class TestLocalEnergyWithWalker(unittest.TestCase):
"""Test local_energy function with Walker dataclass."""
[docs]
def setUp(self):
"""Set up H2 molecule for testing."""
self.mol = gto.M(
atom='H 0 0 0; H 0 0 0.742',
basis='sto-3g',
unit='bohr'
)
self.mf = scf.RHF(self.mol)
self.mf.kernel()
self.det = SlaterDet.create(self.mol, self.mf.mo_coeff)
self.jastrow = Poly()
self.ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det])
self.jastrow_params = jnp.array([0.5])
self.linear_coeffs = jnp.array([1.0])
self.params = (self.jastrow_params, self.linear_coeffs)
[docs]
def test_local_energy_with_walker(self):
"""Test that local_energy works with Walker and returns updated walker."""
from pytc.vmc.walker import Walker
positions = jnp.array([
[[0.0, 0.1, 0.0], [0.0, 0.1, 0.742]],
[[0.1, 0.0, 0.0], [0.1, 0.0, 0.742]]
])
n_walkers = 2
walker = Walker(
positions=positions,
slater_up=jnp.zeros((n_walkers, 1, 1)),
slater_down=jnp.zeros((n_walkers, 1, 1)),
inv_up=jnp.zeros((n_walkers, 1, 1)),
inv_down=jnp.zeros((n_walkers, 1, 1)),
det_up=(jnp.ones((n_walkers,)), jnp.zeros((n_walkers,))),
det_down=(jnp.ones((n_walkers,)), jnp.zeros((n_walkers,))),
grad_up=jnp.zeros((n_walkers, 1, 1, 3)),
grad_down=jnp.zeros((n_walkers, 1, 1, 3)),
lap_up=jnp.zeros((n_walkers, 1, 1)),
lap_down=jnp.zeros((n_walkers, 1, 1)),
move_mask=jnp.ones((n_walkers, 2), dtype=bool),
log_psi=jnp.zeros((n_walkers,)),
psi_sign=jnp.zeros((n_walkers,)),
log_jastrow=jnp.zeros((n_walkers,)),
)
# The ansatz method is written for single walkers, so we need to vmap it
batch_ansatz = jax.vmap(self.ansatz, in_axes=(0, None))
psi_values, walker = batch_ansatz(walker, self.params)
print(f"Wavefunction log values: {psi_values[1]}")
# local_energy works with single walkers, so vmap over batch
batch_local_energy = jax.vmap(
lambda w, p: self.ansatz.local_energy(w, p),
in_axes=(0, None)
)
energies, updated_walker = batch_local_energy(walker, self.params)
self.assertEqual(energies.shape, (n_walkers,))
self.assertTrue(jnp.all(jnp.isfinite(energies)))
self.assertFalse(jnp.allclose(updated_walker.grad_up, 0.0))
self.assertFalse(jnp.allclose(updated_walker.lap_up, 0.0))
self.assertTrue(jnp.all(jnp.isfinite(energies)))
self.assertTrue(jnp.all(jnp.abs(energies) < 100.0)) # Should be reasonable magnitude
[docs]
class TestMultiDetEvaluation(unittest.TestCase):
"""Test multi-determinant combination logic in eval_sj."""
[docs]
def setUp(self):
self.mol = gto.M(
atom='H 0 0 0; H 0 0 0.742',
basis='sto3g',
unit='bohr',
)
self.mf = scf.RHF(self.mol)
self.mf.kernel()
self.det = SlaterDet.create(self.mol, self.mf.mo_coeff)
self.jastrow = Poly()
self.jastrow_params = jnp.array([0.5])
self.test_pos = jnp.array([
[0.0, 0.1, 0.0],
[0.0, 0.1, 0.742],
])
[docs]
def test_multi_det_two_identical(self):
"""Two identical dets with coeffs summing to 1 should equal single det with coeff 1."""
single_ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det])
multi_ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det, self.det])
walker = create_test_walker(self.test_pos, self.det)
single_params = (self.jastrow_params, jnp.array([1.0]))
single_psi, _ = single_ansatz(walker, single_params)
single_val = single_psi[0] * jnp.exp(single_psi[1])
multi_params = (self.jastrow_params, jnp.array([0.3, 0.7]))
multi_psi, _ = multi_ansatz(walker, multi_params)
multi_val = multi_psi[0] * jnp.exp(multi_psi[1])
np.testing.assert_allclose(float(multi_val), float(single_val), rtol=1e-10)
[docs]
def test_multi_det_value_matches_manual(self):
"""Multi-det combination value matches manual sum(coeffs * det_values)."""
from pytc.ansatz.det import value_and_grad
from pytc.ansatz.sj import compute_jastrow_log_value
coeffs = jnp.array([0.3, 0.7])
ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det, self.det])
walker = create_test_walker(self.test_pos, self.det)
psi, _ = ansatz(walker, (self.jastrow_params, coeffs))
actual_val = psi[0] * jnp.exp(psi[1])
det_val, _ = value_and_grad(self.det, walker)
det_scalar = det_val[0] * jnp.exp(det_val[1])
log_j = compute_jastrow_log_value(ansatz, self.test_pos, self.jastrow_params)
expected_val = jnp.exp(log_j) * jnp.sum(coeffs * det_scalar)
np.testing.assert_allclose(float(actual_val), float(expected_val), rtol=1e-10)
[docs]
def test_multi_det_cancellation(self):
"""Coeffs [1, -1] with identical dets should give near-zero value."""
ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det, self.det])
walker = create_test_walker(self.test_pos, self.det)
cancel_params = (self.jastrow_params, jnp.array([1.0, -1.0]))
psi, _ = ansatz(walker, cancel_params)
val = psi[0] * jnp.exp(psi[1])
self.assertAlmostEqual(float(val), 0.0, places=80)
[docs]
def test_multi_det_walker_cache(self):
"""Multi-det eval_sj caches psi values in the walker."""
ansatz = SlaterJastrow.create(self.mol, self.jastrow, [self.det, self.det])
walker = create_test_walker(self.test_pos, self.det)
psi, updated_walker = ansatz(walker, (self.jastrow_params, jnp.array([1.0, 0.5])))
psi_sign, psi_logabs = psi
self.assertAlmostEqual(float(updated_walker.log_psi), float(psi_logabs), places=10)
self.assertAlmostEqual(float(updated_walker.psi_sign), float(psi_sign), places=10)
if __name__ == '__main__':
unittest.main()