Derivative Operators
autodiff
Lattice operators using JAX automatic differentiation.
This module provides gradient and Laplacian operators that use JAX's autodiff capabilities instead of manual finite differences. This:
- Reduces bugs (automatic correctness)
- Enables higher-order derivatives automatically
- Works with JAX's JIT compilation
Also includes energy functionals (Hamiltonians) for force computation via autodiff.
Note: For periodic boundaries, we still use finite differences as they're more natural. But autodiff gives us Jacobians, Hessians, etc. for free.
make_gradient_operator(function, method='forward')
Create gradient operator using autodiff.
For a scalar function \(f(\phi)\), computes \(\nabla f\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
function
|
Callable[[Array], float]
|
Scalar function of field (must return single number) |
required |
method
|
str
|
'forward' (jacfwd) or 'reverse' (jacrev) |
'forward'
|
Returns:
| Type | Description |
|---|---|
Callable[[Array], ndarray]
|
Function that computes gradient ∇f |
Example
Source code in jaxlatt/operators/autodiff.py
make_hessian_operator(function)
Create Hessian operator using autodiff.
For scalar function \(f(\phi)\), computes \(\nabla^2 f\) (matrix of second derivatives).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
function
|
Callable[[Array], float]
|
Scalar function of field |
required |
Returns:
| Type | Description |
|---|---|
Callable[[Array], ndarray]
|
Function that computes Hessian matrix |
Source code in jaxlatt/operators/autodiff.py
potential_force_autodiff(phi, potential_fn)
Compute force \(-\nabla V\) using autodiff.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Field configuration |
required |
potential_fn
|
Callable[[Array], float]
|
Function that computes total potential \(V(\phi)\) |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Force \(= -\nabla V\) at each point |
Source code in jaxlatt/operators/autodiff.py
check_gradient_correctness(manual_grad, auto_grad, rtol=1e-05, atol=1e-08)
Compare manual gradient with autodiff gradient.
Useful for validating manual implementations.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
manual_grad
|
Array
|
Manually computed gradient |
required |
auto_grad
|
Array
|
Autodiff computed gradient |
required |
rtol
|
float
|
Relative tolerance |
1e-05
|
atol
|
float
|
Absolute tolerance |
1e-08
|
Returns:
| Type | Description |
|---|---|
dict
|
Dictionary with comparison metrics |
Source code in jaxlatt/operators/autodiff.py
smooth_field_gradient_x(phi, dx, order=2)
Compute \(\partial\phi/\partial x\) using autodiff on smoothed field.
For lattice fields, we can:
- Smooth using spectral filtering
- Take autodiff derivatives
- Get sub-grid accuracy
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Field on lattice |
required |
dx
|
float
|
Lattice spacing |
required |
order
|
int
|
Interpolation order |
2
|
Returns:
| Type | Description |
|---|---|
Array
|
\(\partial\phi/\partial x\) at each lattice point |
Source code in jaxlatt/operators/autodiff.py
gauss_constraint_jacobian(phi, E, g, dx)
Compute Jacobian of Gauss constraint using autodiff.
Constraint: \(C = \nabla \cdot E + g\rho\) where \(\rho = |\phi|^2\).
Returns:
| Type | Description |
|---|---|
Array
|
\(\partial C / \partial \phi\) (useful for constraint projection) |
Source code in jaxlatt/operators/autodiff.py
energy_gradient_wrt_field(phi, pi, E, links, m, lambda_, g, dx)
Compute \(\partial H / \partial \phi\) where \(H\) is total energy.
Useful for:
- Gradient flow (finding ground states)
- Stability analysis
- Constraint enforcement
Returns:
| Type | Description |
|---|---|
Array
|
\(\partial H / \partial \phi\) using automatic differentiation |
Source code in jaxlatt/operators/autodiff.py
stability_hessian(phi, m, lambda_, dx)
Compute Hessian \(\partial^2 V / \partial \phi^2\) for stability analysis.
Eigenvalues tell us about:
- Tachyonic instabilities (negative eigenvalues)
- Oscillation frequencies (positive eigenvalues)
Returns:
| Type | Description |
|---|---|
Array
|
Hessian matrix (can be huge for 3D lattices!) |
Source code in jaxlatt/operators/autodiff.py
validate_force_implementation(phi, m, lambda_, manual_force_fn, atol=1e-06)
Validate manual force implementation against autodiff.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Test field configuration |
required |
m
|
float
|
Mass parameter |
required |
lambda_
|
float
|
Self-coupling |
required |
manual_force_fn
|
Callable
|
Your manual force function |
required |
atol
|
float
|
Absolute tolerance |
1e-06
|
Returns:
| Type | Description |
|---|---|
dict
|
Dictionary with validation results |
Source code in jaxlatt/operators/autodiff.py
scalar_kinetic_energy_functional(pi, dx)
Kinetic energy \(T = \frac{1}{2} \int d^3x \, |\pi|^2\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
pi
|
Array
|
Conjugate momentum (N, N, N) complex |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total kinetic energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
scalar_gradient_energy_functional(phi, links, dx)
Gradient energy \(\frac{1}{2} \int d^3x \, |D_i \phi|^2\) (gauge-covariant).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Scalar field (N, N, N) complex |
required |
links
|
Array
|
Gauge links (3, N, N, N) complex |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total gradient energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
scalar_potential_energy_functional(phi, m, lambda_, dx)
Potential energy \(\int d^3x \, V(\phi)\) where \(V = \frac{m^2}{2}|\phi|^2 + \frac{\lambda}{4}|\phi|^4\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Scalar field (N, N, N) complex |
required |
m
|
float
|
Mass parameter |
required |
lambda_
|
float
|
Self-coupling |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total potential energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
scalar_hamiltonian(phi, pi, links, m, lambda_, dx)
Total scalar Hamiltonian H = T + gradient + potential.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Scalar field |
required |
pi
|
Array
|
Conjugate momentum |
required |
links
|
Array
|
Gauge links (for covariant derivatives) |
required |
m
|
float
|
Mass |
required |
lambda_
|
float
|
Self-coupling |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
gauge_electric_energy_functional(E, dx)
Electric energy \(\frac{1}{2} \int d^3x \, E^2\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
E
|
Array
|
Electric field (3, N, N, N) real |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total electric energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
gauge_magnetic_energy_functional(links, dx, g)
Magnetic energy \(\frac{1}{2} \int d^3x \, B^2\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
links
|
Array
|
Gauge links (3, N, N, N) |
required |
dx
|
float
|
Lattice spacing |
required |
g
|
float
|
Gauge coupling |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total magnetic energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
gauge_hamiltonian(links, E, dx, g)
Total gauge Hamiltonian \(H = \frac{1}{2}(E^2 + B^2)\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
links
|
Array
|
Gauge links |
required |
E
|
Array
|
Electric field |
required |
dx
|
float
|
Lattice spacing |
required |
g
|
float
|
Gauge coupling |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total gauge energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
coupled_hamiltonian(phi, pi, links, E, m, lambda_, g, dx)
Total Hamiltonian for coupled scalar-gauge system.
H = H_scalar + H_gauge
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi
|
Array
|
Scalar field |
required |
pi
|
Array
|
Conjugate momentum |
required |
links
|
Array
|
Gauge links |
required |
E
|
Array
|
Electric field |
required |
m
|
float
|
Scalar mass |
required |
lambda_
|
float
|
Scalar self-coupling |
required |
g
|
float
|
Gauge coupling |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total energy (scalar) |
Source code in jaxlatt/operators/autodiff.py
make_scalar_force_from_hamiltonian(m, lambda_, dx)
Create force function \(F_\phi = -\partial H / \partial \phi^*\) via autodiff.
This guarantees that the force is the exact gradient of energy, ensuring energy conservation in symplectic integrators.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
m
|
float
|
Mass parameter |
required |
lambda_
|
float
|
Self-coupling |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
|
Function \((\phi, \mathrm{links}) \to F_\phi\) |
Source code in jaxlatt/operators/autodiff.py
verify_force_is_gradient(phi, pi, links, E, m, lambda_, g, dx, force_manual)
Verify that manual force matches -∂H/∂φ*.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
phi, pi, links, E
|
Field configuration |
required | |
m, lambda_, g
|
Parameters |
required | |
dx
|
float
|
Lattice spacing |
required |
force_manual
|
Array
|
Manually computed force |
required |
Returns:
| Type | Description |
|---|---|
dict[str, float]
|
Dict with error metrics |
Source code in jaxlatt/operators/autodiff.py
wilson_action(links, g, dx)
Wilson gauge action \(S = -(1/g^2) \Sigma \operatorname{Re}(\operatorname{Tr} U_\mathrm{plaq})\).
For U(1): \(S = -(1/g^2) \Sigma \operatorname{Re}(U_\mathrm{plaq})\).
Force on links can be computed via \(\partial S / \partial U\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
links
|
Array
|
Gauge links (3, N, N, N) |
required |
g
|
float
|
Gauge coupling |
required |
dx
|
float
|
Lattice spacing |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Wilson action (scalar) |
Source code in jaxlatt/operators/autodiff.py
compute_all_energy_components(phi, pi, links, E, m, lambda_, g, dx)
Compute all energy components using functionals.
Returns dict with scalar_kinetic, scalar_gradient, scalar_potential, electric, magnetic, and total energy.