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 |
required |
params
|
list
|
List of parameter names. |
required |
model
|
SourceModel or callable
|
Template model whose parameters at the dot-separated paths |
required |
data
|
OIData
|
Object containing the data to be fitted. |
required |
dtype
|
(float64, float32)
|
Precision of the calculation, as for |
"float64"
|
Returns:
| Type | Description |
|---|---|
array - like
|
|
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 |
required |
params
|
list[str]
|
Parameter paths corresponding to |
required |
model
|
SourceModel or callable
|
Template model or class, as for |
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
|
|
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 |
required |
params
|
list[str]
|
Parameter paths corresponding to |
required |
model
|
SourceModel or callable
|
Template model whose parameters at the dot-separated paths |
required |
data
|
OIData
|
Observational data object. |
required |
ridge
|
float
|
Diagonal regularization term. |
0.0
|
dtype
|
(float64, float32)
|
Precision of the calculation, as for |
"float64"
|
Returns:
| Type | Description |
|---|---|
array - like
|
Observed information matrix, |
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 |
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 |
1e-12
|
Returns:
| Type | Description |
|---|---|
array - like
|
|