Source code for pytc.test.test_kmat

"""Tests for JAX implementation of kinetic matrix elements."""

import unittest
import numpy as np
import jax
# Enable float64 support
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
from pytc.legacy.kmat import calc_K1 as calc_K1_numpy, calc_K3 as calc_K3_numpy
from pytc.kmat import calc_K1, calc_K3
from pytc.jastrow import Poly


[docs] class TestKmat(unittest.TestCase): """Test JAX implementation of K matrix elements."""
[docs] def setUp(self): """Set up test fixtures.""" rng = np.random.RandomState(42) # Create multiple test systems with different sizes self.test_configs = [ # Small system {'Nb': 2, 'N_grid': 3, 'name': 'small'}, # Medium system {'Nb': 4, 'N_grid': 10, 'name': 'medium'}, # Larger system {'Nb': 6, 'N_grid': 20, 'name': 'large'} ] for config in self.test_configs: Nb, N_grid = config['Nb'], config['N_grid'] # Create test data for each configuration config['grid_points'] = rng.randn(N_grid, 3) config['weights'] = rng.rand(N_grid) # Random weights # Generate orbitals (phi) and gradients (grad_phi) config['phi'] = rng.randn(Nb, N_grid) config['grad_phi'] = rng.randn(Nb, N_grid, 3) # Compute paired densities for NumPy reference (which expects pairs) # phi_paired_ij = phi_i * phi_j config['phi_paired'] = np.einsum('in,jn->ijn', config['phi'], config['phi']).reshape(Nb * Nb, N_grid) # grad_phi_paired_ij = grad_phi_i * phi_j # Note: This matches JAX calc_K1 logic (grad on first index) config['grad_phi_paired'] = np.einsum('ind,jn->ijnd', config['grad_phi'], config['phi']).reshape(Nb * Nb, N_grid, 3) # Create Jastrow factors self.params = jnp.array([1.0]) self.jastrow_jax = Poly() class PolyNumpy: """NumPy implementation to match original implementation.""" def __init__(self, params): self.params = params def grad(self, r1, r2): """Numpy gradient computation handling both single and batched inputs.""" # Handle single point inputs if r1.ndim == 1: r1 = r1[None, :] if r2.ndim == 1: r2 = r2[None, :] diff = r1[:, None, :] - r2[None, :, :] r12 = np.sqrt(np.sum(diff * diff, axis=-1) + 1e-10) # Match epsilon grad = diff / r12[..., None] grad = grad * self.params[0] # Return single point result without batch dimensions if grad.shape[0] == 1 and grad.shape[1] == 1: return grad[0, 0] return grad self.jastrow_numpy = PolyNumpy(self.params)
[docs] def test_K1_shapes(self): """Test K1 output shapes for different input sizes.""" for config in self.test_configs: with self.subTest(size=config['name']): Nb = config['Nb'] result = calc_K1( jnp.asarray(config['phi']), jnp.asarray(config['grad_phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']) ) self.assertEqual(result.shape, (Nb * Nb, Nb * Nb))
[docs] def test_K1_against_numpy_all_sizes(self): """Compare JAX K1 implementation against numpy for different sizes.""" for config in self.test_configs: with self.subTest(size=config['name']): k1_jax_raw = calc_K1( jnp.asarray(config['phi']), jnp.asarray(config['grad_phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']) ) # JAX returns (Nb, Nb, Nb, Nb) flattened to (Nb^2, Nb^2) # This matches NumPy (Nb^2, Nb^2) k1_jax = k1_jax_raw k1_numpy = calc_K1_numpy( config['phi_paired'], config['grad_phi_paired'], self.jastrow_numpy, config['grid_points'], config['weights'] ) np.testing.assert_allclose( np.asarray(k1_jax), k1_numpy, rtol=1e-5, atol=1e-5, err_msg=f"JAX and numpy K1 don't match for {config['name']} system" )
[docs] def test_K3_shapes(self): """Test K3 output shapes for different input sizes.""" for config in self.test_configs: with self.subTest(size=config['name']): Nb = config['Nb'] result = calc_K3( jnp.asarray(config['phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']) ) self.assertEqual(result.shape, (Nb * Nb, Nb * Nb))
[docs] def test_K3_against_numpy_all_sizes(self): """Compare JAX K3 implementation against numpy for different sizes.""" for config in self.test_configs: with self.subTest(size=config['name']): k3_jax = calc_K3( jnp.asarray(config['phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']) ) k3_numpy = calc_K3_numpy( config['phi_paired'], self.jastrow_numpy, config['grid_points'], config['weights'] ) np.testing.assert_allclose( np.asarray(k3_jax), k3_numpy, rtol=1e-5, atol=1e-5, err_msg=f"JAX and numpy K3 don't match for {config['name']} system" )
[docs] def test_batch_size_handling(self): """Test different batch sizes produce same results.""" config = self.test_configs[-1] # Use largest system batch_sizes = [1, 5, 10, 20] # Get reference result with default batch size ref_k1 = calc_K1( jnp.asarray(config['phi']), jnp.asarray(config['grad_phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']) ) ref_k3 = calc_K3( jnp.asarray(config['phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']) ) for batch_size in batch_sizes: with self.subTest(batch_size=batch_size): # Test K1 k1 = calc_K1( jnp.asarray(config['phi']), jnp.asarray(config['grad_phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']), batch_size=batch_size ) np.testing.assert_allclose(k1, ref_k1, rtol=1e-5, atol=1e-5) # Test K3 k3 = calc_K3( jnp.asarray(config['phi']), self.jastrow_jax, self.params, # Add params argument jnp.asarray(config['grid_points']), jnp.asarray(config['weights']), batch_size=batch_size ) np.testing.assert_allclose(k3, ref_k3, rtol=1e-5, atol=1e-5)
[docs] def test_single_point_gradient(self): """Test single point gradient computation matches between JAX and NumPy.""" r1 = np.array([0., 0., 0.]) r2 = np.array([1., 0., 0.]) grad_jax = self.jastrow_jax.grad_r(r1, r2, self.params) grad_numpy = self.jastrow_numpy.grad(r1, r2) np.testing.assert_allclose( np.asarray(grad_jax), grad_numpy, rtol=1e-5, atol=1e-5, err_msg="Single point gradients don't match" )
[docs] class TestStreamingContractionParity(unittest.TestCase): """contract_K1_minus_K2_isdf / contract_K3_isdf_streaming must match the resident JIT on the same inputs — regardless of whether U is host numpy or device jax, and regardless of panel_size."""
[docs] def setUp(self): rng = np.random.RandomState(123) self.n_fused = 97 # deliberately not a multiple of any panel_size below self.n_orb = 11 self.Np = self.Nq = self.Nr = self.Ns = self.n_orb self.phi = jnp.asarray(rng.randn(self.n_orb, self.n_fused)) self.grad_phi = jnp.asarray(rng.randn(self.n_orb, self.n_fused, 3)) self.U1 = jnp.asarray(rng.randn(self.n_fused, self.n_fused, 3)) self.U3 = jnp.asarray(rng.randn(self.n_fused, self.n_fused)) self.rbs = 16
def _reference_K1_minus_K2(self): from pytc.kmat import contract_K1_minus_K2_isdf_streaming as contract_K1_minus_K2_isdf_jit return np.asarray(contract_K1_minus_K2_isdf_jit( self.phi, self.phi, self.phi, self.phi, self.grad_phi, self.grad_phi, self.U1, self.rbs, )) def _reference_K3(self): from pytc.kmat import contract_K3_isdf_jit return np.asarray(contract_K3_isdf_jit( self.phi, self.phi, self.phi, self.phi, self.U3, self.rbs, ))
[docs] def test_K1_minus_K2_resident_fast_path_matches_jit(self): from pytc.kmat import contract_K1_minus_K2_isdf_streaming as contract_K1_minus_K2_isdf ref = self._reference_K1_minus_K2() out = np.asarray(contract_K1_minus_K2_isdf( self.phi, self.phi, self.phi, self.phi, self.grad_phi, self.grad_phi, self.U1, self.rbs, panel_size=None, )) np.testing.assert_allclose(out, ref, atol=1e-14, rtol=0)
[docs] def test_K1_minus_K2_streaming_host_matches_jit(self): from pytc.kmat import contract_K1_minus_K2_isdf_streaming as contract_K1_minus_K2_isdf ref = self._reference_K1_minus_K2() U1_host = np.asarray(self.U1) # explicitly on host for panel_size in (16, 32, 48): # none divides n_fused=97 evenly out = np.asarray(contract_K1_minus_K2_isdf( self.phi, self.phi, self.phi, self.phi, self.grad_phi, self.grad_phi, U1_host, self.rbs, panel_size=panel_size, )) np.testing.assert_allclose( out, ref, atol=1e-12, rtol=0, err_msg=f"panel_size={panel_size}", )
[docs] def test_K1_minus_K2_streaming_device_matches_jit(self): from pytc.kmat import contract_K1_minus_K2_isdf_streaming as contract_K1_minus_K2_isdf ref = self._reference_K1_minus_K2() for panel_size in (16, 32, 48): out = np.asarray(contract_K1_minus_K2_isdf( self.phi, self.phi, self.phi, self.phi, self.grad_phi, self.grad_phi, self.U1, self.rbs, panel_size=panel_size, )) np.testing.assert_allclose( out, ref, atol=1e-12, rtol=0, err_msg=f"panel_size={panel_size}", )
[docs] def test_K1_isdf_streaming_host_matches_jit(self): from pytc.kmat import contract_K1_isdf_jit, contract_K1_isdf_streaming ref = np.asarray(contract_K1_isdf_jit( self.phi, self.phi, self.phi, self.phi, self.grad_phi, self.U1, self.rbs, )) U1_host = np.asarray(self.U1) for panel_size in (16, 32, 48, None): out = np.asarray(contract_K1_isdf_streaming( self.phi, self.phi, self.phi, self.phi, self.grad_phi, U1_host, self.rbs, panel_size=panel_size, )) np.testing.assert_allclose( out, ref, atol=1e-12, rtol=0, err_msg=f"panel_size={panel_size}", )
[docs] def test_K1_antisym_pq_matches_legacy_transpose(self): """In-kernel antisym must equal the legacy ``k12 - k12.T(1,0,2,3)`` path (resident & all panel sizes, host & device U1).""" from pytc.kmat import (contract_K1_isdf_jit, contract_K1_antisym_pq_isdf_streaming) k12 = np.asarray(contract_K1_isdf_jit( self.phi, self.phi, self.phi, self.phi, self.grad_phi, self.U1, self.rbs, )) ref = k12 - k12.transpose(1, 0, 2, 3) for U1_in, label in ((self.U1, "device"), (np.asarray(self.U1), "host")): for panel_size in (None, 16, 32, 48): out = np.asarray(contract_K1_antisym_pq_isdf_streaming( self.phi, self.phi, self.phi, self.grad_phi, U1_in, self.rbs, panel_size=panel_size, )) np.testing.assert_allclose( out, ref, atol=1e-12, rtol=0, err_msg=f"U1={label}, panel_size={panel_size}", )
[docs] def test_K3_streaming_host_matches_jit(self): from pytc.kmat import contract_K3_isdf_streaming ref = self._reference_K3() U3_host = np.asarray(self.U3) for panel_size in (16, 32, 48, None): out = np.asarray(contract_K3_isdf_streaming( self.phi, self.phi, self.phi, self.phi, U3_host, self.rbs, panel_size=panel_size, )) np.testing.assert_allclose( out, ref, atol=1e-12, rtol=0, err_msg=f"panel_size={panel_size}", )
if __name__ == '__main__': unittest.main()