Skip to content

virgil.inference

Local curvature tools: Hessians, Laplace covariances and Fisher matrices.

Objectives are negative log likelihoods of a flat 1D parameter vector, except for gaussian_fisher, which accepts any parameter pytree. Checks that need concrete values (positive errors, positive-definite matrices) are skipped inside jax.jit.

laplace_cov(values, params, model, data, *, dtype='float64')

Compute the full Laplace covariance matrix for all model parameters jointly.

Computes the inverse of the Hessian of the negative log-likelihood with respect to all parameters in params simultaneously, returning an N x N covariance matrix.

A path in params may be array-valued (e.g. "rim.az_amps"). Then values is the flat concatenation of every path's elements, in the order of params (each array flattened C-order), and N is the total number of elements (N = len(params) when all are scalars). The covariance is over those flattened elements. Array paths need a SourceModel template, whose leaves give the sizes; a callable model takes scalar parameters only.

This returns the full covariance matrix over all N elements. For the uncertainty of one parameter with the others held fixed (e.g. the flux at a fixed position), use :func:laplace_parameter_uncertainty.

Parameters:

Name Type Description Default
values array - like

1D flat parameter vector: the elements of each path in params, concatenated in order.

required
params list

List of parameter names.

required
model SourceModel or callable

Template model whose parameters at the dot-separated paths params are replaced by values, or a class/callable called as model(**dict(zip(params, values))) (see build_model).

required
data OIData

Object containing the data to be fitted.

required
dtype (float64, float32)

Precision of the calculation, as for fit: float64 by default, inside a local jax.enable_x64 context. The result is returned in JAX's precision outside it.

"float64"

Returns:

Type Description
array - like

N x N covariance matrix over the flattened parameter elements.

laplace_parameter_uncertainty(values, params, model, data, target_param)

Compute scalar Laplace uncertainty for one parameter with all others fixed.

Parameters:

Name Type Description Default
values array - like

Flat parameter vector at which to evaluate the curvature (the elements of every path in params, concatenated in order, as for laplace_cov).

required
params list[str]

Parameter paths corresponding to values.

required
model SourceModel or callable

Template model or class, as for loglike.

required
data OIData

Data to fit.

required
target_param str

The scalar parameter whose uncertainty is returned (a path whose template leaf has more than one element is rejected).

required

Returns:

Type Description
float

(d² -log L / d target²)^(-1/2). It is NaN where the curvature is not positive, i.e. away from a likelihood maximum along target_param.

fisher(values, params, model, data, ridge=0.0, *, dtype='float64')

Observed information (Hessian of -log L) at a parameter point.

At the maximum-likelihood point this approximates the Fisher matrix.

Parameters:

Name Type Description Default
values array - like

Flat parameter vector at which to evaluate the local curvature (the elements of every path in params, concatenated in order, as for laplace_cov).

required
params list[str]

Parameter paths corresponding to values.

required
model SourceModel or callable

Template model whose parameters at the dot-separated paths params are replaced by values, or a class/callable called as model(**dict(zip(params, values))) (see build_model).

required
data OIData

Observational data object.

required
ridge float

Diagonal regularization term.

0.0
dtype (float64, float32)

Precision of the calculation, as for laplace_cov.

"float64"

Returns:

Type Description
array - like

Observed information matrix, N x N for N the total number of parameter elements.

hessian_matrix(objective, x)

Return the Hessian matrix of objective evaluated at x.

x must be a flat 1D parameter vector; use jax.flatten_util.ravel_pytree to flatten a pytree first.

regularized_inverse(matrix, ridge=1e-10)

Return inv(matrix + ridge * I).

A RuntimeWarning is raised (outside jax.jit) if the regularized matrix is not positive definite, since its inverse is then not a covariance.

laplace_covariance(objective, x, ridge=1e-10)

Return the Laplace covariance, the inverse Hessian of objective.

objective is a negative log likelihood of the flat vector x, which should be at (or near) its minimum. A small ridge (default 1e-10) is added to the diagonal before inverting.

gaussian_fisher(prediction_fn, params, errors, ridge=0.0)

Return expected Fisher information for fixed independent Gaussian errors.

Parameters:

Name Type Description Default
prediction_fn callable

Function mapping the parameter pytree to a one-dimensional prediction.

required
params pytree

Parameter values at which to evaluate the local model sensitivity.

required
errors array - like

Standard deviations corresponding to the prediction vector.

required
ridge float

Diagonal regularization term.

0.0

Returns:

Type Description
tuple[array - like, callable]

Expected Fisher matrix Jᵀ Σ⁻¹ J and a function restoring a flat parameter vector to the structure of params.

fisher_projection(fmat, eps=1e-12)

Return projection matrix mapping unit-normal latent vectors to parameter steps.

If u ~ N(0, I), then x = x0 + P @ u has local covariance approximately F^{-1} for Fisher matrix F.

Parameters:

Name Type Description Default
fmat array - like

Symmetric Fisher (or observed information) matrix.

required
eps float

Eigenvalues below eps times the largest eigenvalue (or below eps itself, if every eigenvalue is zero) are raised to that floor, with a RuntimeWarning (outside jax.jit), so that flat directions get large but finite steps.

1e-12

Returns:

Type Description
array - like

P = V diag(λ^-1/2), so that P Pᵀ = F⁻¹.