Initial commit. Stage one work done.
This commit is contained in:
100
tests/test_guards.py
Normal file
100
tests/test_guards.py
Normal file
@@ -0,0 +1,100 @@
|
||||
"""NaN-safety and differentiability tests for the guard layer.
|
||||
|
||||
The regression tests establish numeric parity; these establish that funxc
|
||||
stays finite (values *and* gradients) at the singular corners libxc's grids
|
||||
avoid: rho -> 0 tails, sigma = 0, full spin polarization, and combinations.
|
||||
"""
|
||||
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
import funxc
|
||||
|
||||
GGA_IDS = [fid for fid in funxc.REGISTRY if funxc.REGISTRY[fid].family == "gga"]
|
||||
LDA_IDS = [fid for fid in funxc.REGISTRY if funxc.REGISTRY[fid].family == "lda"]
|
||||
|
||||
# Adversarial polarized points:
|
||||
# (rho_up, rho_dn, sigma_uu, sigma_ud, sigma_dd, tau_up, tau_dn)
|
||||
EDGE_POINTS = [
|
||||
(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0), # vacuum
|
||||
(1e-300, 0.0, 0.0, 0.0, 0.0, 1e-300, 0.0), # denormal tail
|
||||
(0.3, 0.0, 7e-3, 0.0, 0.0, 2.9e-3, 0.0), # fully polarized, tau ~ tau_W
|
||||
(0.3, 0.3, 0.0, 0.0, 0.0, 0.0, 0.0), # sigma = tau = 0 exactly
|
||||
(0.3, 1e-16, 1e-2, 0.0, 1e-30, 5e-2, 1e-30), # one channel below thr
|
||||
(1e-14, 1e-14, 1e-25, -1e-25, 1e-25, 1e-25, 1e-25), # negative sigma_ud
|
||||
(1e3, 1e3, 1e8, 1e8, 1e8, 1e8, 1e8), # high density, huge gradient
|
||||
(0.5, 0.5, 1e12, 0.0, 1e12, 1e-5, 1e-5), # s -> inf, tau << tau_W (FHC)
|
||||
]
|
||||
|
||||
|
||||
def _edge_args(spec, f):
|
||||
rho = np.array([[p[0], p[1]] for p in EDGE_POINTS])
|
||||
sigma = np.array([[p[2], p[3], p[4]] for p in EDGE_POINTS])
|
||||
tau = np.array([[p[5], p[6]] for p in EDGE_POINTS])
|
||||
lapl = np.zeros_like(tau)
|
||||
if spec.family == "lda":
|
||||
return (rho,)
|
||||
if spec.family == "gga":
|
||||
return (rho, sigma)
|
||||
return (rho, sigma, lapl, tau)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fid", sorted(funxc.REGISTRY))
|
||||
def test_values_and_gradients_finite_at_edges(fid):
|
||||
spec = funxc.REGISTRY[fid]
|
||||
f = funxc.functional(fid, polarized=True)
|
||||
out = f.exc_vxc(*_edge_args(spec, f))
|
||||
for key, val in out.items():
|
||||
assert np.all(np.isfinite(val)), f"{fid}: non-finite {key} at edge points"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"fid",
|
||||
["LDA_X", "LDA_C_PW", "GGA_X_PBE", "GGA_C_PBE", "GGA_C_LYP",
|
||||
"GGA_XC_B97_D", "MGGA_X_LTA"],
|
||||
)
|
||||
def test_second_derivatives_finite(fid):
|
||||
"""fxc-level derivatives must also stay finite (autodiff twice)."""
|
||||
spec = funxc.REGISTRY[fid]
|
||||
f = funxc.functional(fid, polarized=True)
|
||||
if spec.family == "lda":
|
||||
hess = jax.hessian(f.energy_density, argnums=0)
|
||||
args = lambda p: (jnp.array(p[:2]),) # noqa: E731
|
||||
elif spec.family == "gga":
|
||||
hess = jax.hessian(f.energy_density, argnums=(0, 1))
|
||||
args = lambda p: (jnp.array(p[:2]), jnp.array(p[2:5])) # noqa: E731
|
||||
else:
|
||||
hess = jax.hessian(f.energy_density, argnums=(0, 1, 2, 3))
|
||||
args = lambda p: ( # noqa: E731
|
||||
jnp.array(p[:2]), jnp.array(p[2:5]), jnp.zeros(2), jnp.array(p[5:7])
|
||||
)
|
||||
for p in EDGE_POINTS:
|
||||
h = hess(*args(p))
|
||||
leaves = jax.tree.leaves(h)
|
||||
assert all(jnp.all(jnp.isfinite(x)) for x in leaves), (
|
||||
f"{fid}: non-finite second derivative at {p}"
|
||||
)
|
||||
|
||||
|
||||
def test_vacuum_outputs_are_zero():
|
||||
f = funxc.functional("GGA_X_PBE", polarized=True)
|
||||
out = f.exc_vxc(np.zeros((1, 2)), np.zeros((1, 3)))
|
||||
for key, val in out.items():
|
||||
assert np.all(np.asarray(val) == 0.0), f"nonzero {key} in vacuum"
|
||||
|
||||
|
||||
def test_jit_and_vmap_compose():
|
||||
"""The per-point eps must survive user-side jit/vmap/grad composition."""
|
||||
f = funxc.functional("GGA_C_PBE", polarized=True)
|
||||
g = jax.jit(jax.vmap(jax.grad(f.energy_density, argnums=(0, 1))))
|
||||
rho = jnp.array([[0.3, 0.2], [1.0, 1.0]])
|
||||
sigma = jnp.array([[0.01, 0.0, 0.02], [0.1, 0.05, 0.1]])
|
||||
vrho, vsigma = g(rho, sigma)
|
||||
assert vrho.shape == (2, 2) and vsigma.shape == (2, 3)
|
||||
assert jnp.all(jnp.isfinite(vrho)) and jnp.all(jnp.isfinite(vsigma))
|
||||
|
||||
|
||||
def test_float64_active():
|
||||
assert jnp.asarray(1.0).dtype == jnp.float64
|
||||
Reference in New Issue
Block a user