Skip to content

Spectral (exp(-iHt), and other functions of a matrix)

A quantum system with Hamiltonian H, left alone for a time t, evolves into the state exp(-iHt). Any other smooth function of H -- sqrt(H), log(H), a step function -- works the same way: diagonalize H, apply the function to each energy, rotate back. dense_evolution.physics.spectral does that, with one specific fix.

The fix matters because of a trap you can fall into without noticing. The gradient of exp(-iHt) with respect to H -- the thing you need if you want to learn or optimize H -- is computed wrongly by the standard method when two of H's energies are exactly equal. Nothing crashes. The number just comes back silently wrong.

Gradient error vs eigenvalue gap

The x-axis is the gap between two energies of a random 4×4 Hamiltonian: 1e-14 on the left means they are almost identical, 1e-2 on the right means they are well separated. Each curve is the relative error of a gradient of exp(-iHt) against central finite differences of jax.scipy.linalg.expm. Red is the standard jnp.linalg.eigh gradient, green is spectral_evolve. Well separated, both sit at the 1e-10 level of the finite-difference reference. As the gap closes the red error climbs to 1e-2; the green one stays below 1e-8 everywhere.

Step 1. Recognize when you have a degenerate Hamiltonian

Two energies are degenerate when they are exactly equal -- two different states of the system happen to share the same energy. That is not a pathology: it happens in every symmetric molecule and every lattice model. Here is how to check for it:

import jax.numpy as jnp
from dense_evolution.physics import has_exact_degeneracy

H_deg = jnp.diag(jnp.array([1.0, 1.0, 2.0, 2.0]))
H_gap = jnp.diag(jnp.array([1.0, 1.001, 2.0, 2.001]))

has_exact_degeneracy(H_deg), has_exact_degeneracy(H_gap)
(True, False)

has_exact_degeneracy(H) answers a single question: are any two of H's eigenvalues closer together than 1e-8? H_deg has two pairs of exactly equal energies, so it returns True. H_gap looks identical on paper, but its smallest gap is 0.001 -- far above the threshold -- so it returns False.

This is a check, not a fix. Run it once before you decide which method to use for the gradient.

Step 2. Evolve a system forward in time

Now the actual evolution. spectral_evolve(H, t) gives you exp(-iHt) as a matrix, and it is safe to differentiate at any gap:

import jax
import jax.numpy as jnp
from dense_evolution.physics import spectral_evolve

H = jnp.diag(jnp.array([1.0, 1.0, 2.0, 2.0], dtype=jnp.complex128))

def loss(H_):
    return jnp.real(jnp.sum(spectral_evolve(H_, t=1.0)))

jax.grad(loss)(H)
[[-0.841-0.540j -0.841-0.540j -0.956-0.068j -0.956-0.068j]
 [-0.841-0.540j -0.841-0.540j -0.956-0.068j -0.956-0.068j]
 [-0.956-0.068j -0.956-0.068j -0.909+0.416j -0.909+0.416j]
 [-0.956-0.068j -0.956-0.068j -0.909+0.416j -0.909+0.416j]]

Forward, spectral_evolve(H, t) is the same matrix you would get from v @ diag(exp(-1j*w*t)) @ v.conj().T -- identical to machine precision. The difference is entirely on the backward pass, where the gradient with respect to H is built with a formula that stays correct at exact degeneracy (Step 4 shows which).

loss(H_) reduces exp(-iHt) to a single real number, so jax.grad can differentiate it. The result is a complex matrix of the same shape as H. The output has a visible block structure -- rows and columns 0, 1 share one value, and so do 2, 3 -- because those are the pairs of exactly equal energies. The gradient respects the degeneracy: it does not try to pick one basis state over the other inside a degenerate pair.

Step 3. Any function of H, not just the exponential

If the function you need is not exp(-iHt), use matrix_function_eigh directly. It takes H, plus the function f you want to apply to the eigenvalues, plus its derivative f_prime:

import jax.numpy as jnp
from dense_evolution.physics import matrix_function_eigh

H = jnp.diag(jnp.array([1.0, 1.0, 4.0, 4.0], dtype=jnp.complex128))

U = matrix_function_eigh(
    H,
    f       = lambda w: jnp.sqrt(w),
    f_prime = lambda w: 0.5 / jnp.sqrt(w),
)
U
[[1.+0.j 0.+0.j 0.+0.j 0.+0.j]
 [0.+0.j 1.+0.j 0.+0.j 0.+0.j]
 [0.+0.j 0.+0.j 2.+0.j 0.+0.j]
 [0.+0.j 0.+0.j 0.+0.j 2.+0.j]]

H has eigenvalues [1, 1, 4, 4]. f is sqrt, so the output is the diagonal matrix with sqrt applied to each energy: [1, 1, 2, 2]. f_prime is only consulted where two eigenvalues are exactly equal, as the limit of the divided difference in Step 4 -- away from a tie, the non-degenerate entries use the divided difference, not f_prime's value.

Both callables must return complex values if H is complex. Neither should be a closure over a JAX array: they are treated as static by JAX, so a captured array will not be traced through.

spectral_evolve from Step 2 is a one-line wrapper around this function, with f and f_prime filled in for the exponential.

Step 4. The formula behind the gradient

The gradient rule is the classical divided-difference formula for a matrix function's derivative (Kato 1995, Ch. II.5.6). For a perturbation dH of H:

d/deps [ exp(-i H(eps) t) ]  =  V (F o (V^dagger dH V)) V^dagger

with F the matrix of divided differences of the exponential:

F[i,j] = (exp(-i lambda_i t) - exp(-i lambda_j t)) / (lambda_i - lambda_j)   if lambda_i != lambda_j
F[i,j] = -i t exp(-i lambda_i t)                                              if lambda_i == lambda_j

The key detail is what is not in that formula: the eigenvectors V are consumed only through the projected matrix V^dagger dH V, never carried through as a differentiable output. At exact degeneracy the eigenvectors are not unique -- any orthonormal basis of the degenerate subspace works equally well, and jnp.linalg.eigh picks one arbitrarily. Differentiating through that choice is what makes the standard method wrong. By not depending on it, this formula stays correct.


Details

The actual failure mode. jnp.linalg.eigh's reverse-mode rule divides by lambda_i - lambda_j for every eigenvector pair. At an exact tie that is 0/0; JAX does not raise, it returns a finite but wrong number. Measured on a Kaggle CPU kernel (Dense-Evolution-Discovery, PR #173): standard eigh gradient error 0.98 versus Kato 4e-10 on an H with four exact doubly-degenerate eigenvalues -- several orders of magnitude, not a rounding difference. The plot at the top of this page measures the same effect across the gap range (random 4×4 Hamiltonian, seed 0, gaps 1e-14 to 1e-2, finite-difference step 1e-6).

When to use which method.

Situation Method
H has an exact degeneracy (gap below 1e-8) spectral_evolve / matrix_function_eigh
H has only near-degeneracy (smallest gap above 1e-8) plain jnp.linalg.eigh
Forward value only, no gradient needed plain jnp.linalg.eigh

Above the threshold, spectral_evolve never hurts -- it just does not help. The check from Step 1 is what makes the decision cheap.

References. Kato, T., Perturbation Theory for Linear Operators, Springer (1995), Ch. II.5.6 -- the classical divided-difference formula for matrix-function derivatives, which predates the modern quantum-chemistry literature by decades. Kasim, M. F., arXiv:2011.04366 (2020) -- the same formula stated specifically for degenerate Hermitian matrices, in a form more directly applicable to this code. Both were checked against the actual paper text, not trusted from a citation string alone.

Why custom_jvp rather than custom_vjp. The formula needs only (H, dH, w, v) at the primal point, with no compatibility condition on the perturbation direction. A VJP built on top of eigh's own eigenvector output would need one (Kasim's Eq. 4.72, confirmed present in the actual paper text). This is a structural property of the formula, not an implementation preference.

Hermitian only. The module uses eigh, not eig. f and f_prime must return complex values when f is complex. f_prime is only ever consulted at lambda_i == lambda_j; the non-degenerate entries use the divided difference, not its value, so it does not need to be accurate away from the diagonal.

Standalone precision. Both matrix_function_eigh and has_exact_degeneracy call dense_evolution.config.ensure_x64() on entry, so a fresh process that reaches for them first does not stay at JAX's float32 default.

Regression test. tests/unit/test_spectral.py includes a guard test, test_std_eigh_fails_at_degeneracy, that asserts the standard eigh gradient is wrong at degeneracy. If a future JAX release fixes this upstream, that test fails loudly -- the signal to retire spectral_evolve, not a silent pass.

spectral

Gauge-safe gradients for spectral functions of a Hermitian matrix (V f(Lambda) V^dagger, e.g. time evolution exp(-iHt)) at exact eigenvalue degeneracy.

jnp.linalg.eigh's reverse-mode gradient divides by lambda_i - lambda_j for every eigenvector pair. When two eigenvalues are exactly degenerate, this does not raise and does not always produce NaN -- it can silently return a finite, WRONG gradient, because the eigenvectors spanning a degenerate eigenspace are not themselves uniquely defined (any orthonormal basis of that subspace is an equally valid eigh output). Measured on a real Kaggle CPU kernel (Dense-Evolution-Discovery, PR #173): std eigh gradient error 0.98 vs Kato 4e-10 on an H with four exact doubly-degenerate eigenvalues -- several orders of magnitude, not a rounding difference.

REFERENCES (verified against the actual paper text, not trusted at face value from a citation string alone): Kasim, M. F., "Derivatives of partial eigendecomposition of a real symmetric matrix for degenerate cases", arXiv:2011.04366 (2020). Kato, T., "Perturbation Theory for Linear Operators", Springer (1995), Ch. II.5.6 (the classical divided-difference formula for matrix function derivatives, predating Kasim by decades).

matrix_function_eigh uses a jax.custom_jvp based on Kato's divided-difference formula for matrix functions:

d/deps [ V(eps) f(Lambda(eps)) V(eps)^dagger ] = V (F o (V^dagger dH V)) V^dagger

with F the matrix of divided differences of f:

F[i,j] = (f(lambda_i) - f(lambda_j)) / (lambda_i - lambda_j)   if lambda_i != lambda_j
F[i,j] = f'(lambda_i)                                          if lambda_i == lambda_j (incl. i == j)

This formula does not pass through eigenvectors as an intermediate OUTPUT, so it is gauge-invariant: the contribution from a degenerate block uses f'(lambda) directly, and there is no gauge choice to make -- unlike a custom_vjp built directly on top of eigh's own eigenvector output, which needs a compatibility condition on the perturbation direction (Kasim's Eq. 4.72, confirmed present in the actual paper text) that this formula does not.

WHEN TO USE THIS: - H has (or might have) exactly degenerate eigenvalues, AND - the function L(H) you differentiate depends on H through eigh, AND - the function is not trivially constant on degenerate blocks.

For a Hamiltonian with only near-degeneracy (e.g. min_gap ~1e-5, no exact tie), plain jnp.linalg.eigh is correct and faster -- use has_exact_degeneracy to check before reaching for spectral_evolve.

has_exact_degeneracy

has_exact_degeneracy(
    H: Array, tol: float = _DEGENERACY_TOL
) -> bool

True if H has at least one pair of eigenvalues closer than tol.

Diagnostic only -- call this before choosing spectral_evolve (Kato) over plain jnp.linalg.eigh (std). The threshold is the same one the JVP rule below uses internally, so this is the exact condition under which the two methods disagree.

Source code in dense_evolution/physics/spectral.py
def has_exact_degeneracy(H: jax.Array, tol: float = _DEGENERACY_TOL) -> bool:
    """True if H has at least one pair of eigenvalues closer than tol.

    Diagnostic only -- call this before choosing spectral_evolve (Kato)
    over plain jnp.linalg.eigh (std). The threshold is the same one the
    JVP rule below uses internally, so this is the exact condition under
    which the two methods disagree."""
    ensure_x64()
    w = jnp.linalg.eigvalsh(H)
    gaps = jnp.abs(jnp.diff(jnp.sort(w)))
    return bool((gaps < tol).any())

matrix_function_eigh

matrix_function_eigh(H: Array, f, f_prime) -> jax.Array

V f(Lambda) V^dagger with a gauge-safe gradient at exact degeneracy.

Parameters:

Name Type Description Default
H (n, n) Hermitian matrix.
required
f callable, lambda (array) -> array. Applied elementwise to the

eigenvalues.

required
f_prime callable, lambda (array) -> array. Analytic derivative of f,

used only in the JVP rule for degenerate blocks. Passed via nondiff_argnums since a Python closure is not a valid JAX type to trace.

required

Returns:

Type Description
(n, n) complex128 matrix.
Example

import jax.numpy as jnp H = jnp.diag(jnp.array([1.0, 1.0, 2.0, 2.0], dtype=jnp.complex128)) U = matrix_function_eigh(H, lambda w: jnp.exp(-1j * w), lambda w: -1j * jnp.exp(-1j * w))

Source code in dense_evolution/physics/spectral.py
@functools.partial(jax.custom_jvp, nondiff_argnums=(1, 2))
def matrix_function_eigh(H: jax.Array, f, f_prime) -> jax.Array:
    """V f(Lambda) V^dagger with a gauge-safe gradient at exact degeneracy.

    Parameters
    ----------
    H : (n, n) Hermitian matrix.
    f : callable, lambda (array) -> array. Applied elementwise to the
        eigenvalues.
    f_prime : callable, lambda (array) -> array. Analytic derivative of f,
        used only in the JVP rule for degenerate blocks. Passed via
        nondiff_argnums since a Python closure is not a valid JAX type to
        trace.

    Returns
    -------
    (n, n) complex128 matrix.

    Example
    -------
    >>> import jax.numpy as jnp
    >>> H = jnp.diag(jnp.array([1.0, 1.0, 2.0, 2.0], dtype=jnp.complex128))
    >>> U = matrix_function_eigh(H, lambda w: jnp.exp(-1j * w), lambda w: -1j * jnp.exp(-1j * w))
    """
    ensure_x64()
    w, v = jnp.linalg.eigh(H)
    return v @ jnp.diag(f(w)) @ v.conj().T

spectral_evolve

spectral_evolve(H: Array, t: float) -> jax.Array

exp(-i H t) with a gauge-safe gradient at exact degeneracy.

Equivalent forward to V @ diag(exp(-1j*w*t)) @ V.conj().T where (w, V) is jnp.linalg.eigh(H). Backward uses Kato's divided-difference rule (see module docstring) instead of eigh's own reverse-mode rule.

Source code in dense_evolution/physics/spectral.py
def spectral_evolve(H: jax.Array, t: float) -> jax.Array:
    """exp(-i H t) with a gauge-safe gradient at exact degeneracy.

    Equivalent forward to `V @ diag(exp(-1j*w*t)) @ V.conj().T` where
    (w, V) is `jnp.linalg.eigh(H)`. Backward uses Kato's divided-difference
    rule (see module docstring) instead of `eigh`'s own reverse-mode rule.
    """
    return matrix_function_eigh(
        H,
        f=lambda w: jnp.exp(-1j * w * t),
        f_prime=lambda w: -1j * t * jnp.exp(-1j * w * t),
    )

See also: Autodiff -- the differentiable-VQE pipeline where an H with exact degeneracy would otherwise produce a silently wrong parameter update. Observables -- the Pauli-sum Hamiltonian format an H on this page would typically be built from. ```