"""Tests for JAX SimpleJastrow implementation."""
import unittest
import numpy as np
import jax
import jax.numpy as jnp
from pytc.jastrow import Poly
# Enable float64 support
jax.config.update("jax_enable_x64", True)
[docs]
def numerical_gradient_params(jastrow, r1, r2, params, eps=1e-4):
"""Compute numerical gradient with respect to parameters for single points."""
grad = jnp.zeros_like(params)
for i in range(len(params)):
# Forward step
params_plus = params.at[i].add(eps)
params_minus = params.at[i].add(-eps)
j_plus = jastrow._compute(r1, r2, params_plus)
j_minus = jastrow._compute(r1, r2, params_minus)
# Central difference
grad = grad.at[i].set((j_plus - j_minus) / (2 * eps))
return grad
[docs]
def numerical_gradient_r1(jastrow, r1, r2, params, eps=1e-7):
"""Compute numerical gradient with respect to r1 for single points."""
grad = jnp.zeros_like(r1)
for j in range(3): # x, y, z components
# Forward step
r1_plus = r1.at[j].add(eps)
r1_minus = r1.at[j].add(-eps)
j_plus = jastrow._compute(r1_plus, r2, params)
j_minus = jastrow._compute(r1_minus, r2, params)
# Central difference
grad = grad.at[j].set((j_plus - j_minus) / (2 * eps))
return grad
[docs]
class TestSimpleJastrowJAX(unittest.TestCase):
"""Test cases for SimpleJastrowJAX class."""
[docs]
def setUp(self):
self.params = jnp.array([1.0])
self.jastrow = Poly() # No params in constructor
[docs]
def test_single_point_evaluation(self):
"""Test single point Jastrow evaluation."""
r1 = jnp.array([0., 0., 0.])
r2 = jnp.array([1., 0., 0.])
value = self.jastrow._compute(r1, r2, self.params)
self.assertTrue(jnp.isfinite(value))
np.testing.assert_allclose(float(value), 1.0, rtol=1e-8)
[docs]
def test_batch_evaluation(self):
"""Test batched Jastrow evaluation."""
r1 = jnp.array([[0., 0., 0.], [1., 1., 1.]]) # (2, 3)
r2 = jnp.array([[1., 0., 0.]]) # (1, 3) - single point for r2
values = self.jastrow._compute(r1, r2, self.params)
self.assertEqual(values.shape, (2,)) # Changed from (2, 1)
self.assertTrue(jnp.all(jnp.isfinite(values)))
[docs]
def test_param_gradient(self):
"""Test parameter gradient computation for single points."""
r1 = jnp.array([0., 0., 0.])
r2 = jnp.array([1., 0., 0.])
grad_analytical = self.jastrow.grad_params(r1, r2, self.params)
grad_numerical = numerical_gradient_params(self.jastrow, r1, r2, self.params)
np.testing.assert_allclose(
grad_analytical, grad_numerical,
rtol=1e-5, atol=1e-5,
err_msg="Single point parameter gradients don't match"
)
[docs]
def test_position_gradient(self):
"""Test position gradient computation for single points."""
r1 = jnp.array([0., 0., 0.])
r2 = jnp.array([1., 0., 0.])
grad_analytical = self.jastrow.grad_r(r1, r2, self.params)
grad_numerical = numerical_gradient_r1(self.jastrow, r1, r2, self.params)
np.testing.assert_allclose(
grad_analytical, grad_numerical,
rtol=1e-5, atol=1e-5,
err_msg="Single point position gradients don't match"
)
[docs]
def test_batch_consistency(self):
"""Test that batched results match single point computations."""
r1_single = jnp.array([0., 0., 0.])
r2_single = jnp.array([1., 0., 0.])
r1_batch = jnp.array([[0., 0., 0.]])
r2_batch = jnp.array([[1., 0., 0.]])
# Compare raw _compute values
single_u = self.jastrow._compute(r1_single, r2_single, self.params)
batch_u = self.jastrow._compute(r1_batch[0], r2_batch[0], self.params)
np.testing.assert_allclose(
single_u, batch_u,
rtol=1e-10, atol=1e-10,
err_msg="Batch and single point u values don't match"
)
# Compare exp(u) values from __call__
single_J = self.jastrow(r1_single, r2_single, self.params)
batch_J = self.jastrow(r1_batch, r2_batch, self.params)
np.testing.assert_allclose(
single_J, batch_J[0],
rtol=1e-10, atol=1e-10,
err_msg="Batch and single point J values don't match"
)
if __name__ == '__main__':
unittest.main()