gfa.inference#

Inference functions for Group Factor Analysis.

Functions

fit

Fit a GFA model to data via variational inference.

init_posteriors

Initialize GFA model posteriors for fitting.

infer_latents

Infer latent posterior q(X) given observations and fitted parameters.

infer_loadings

Infer loading posterior q(C).

infer_ard

Infer ARD posterior q(alpha).

infer_obs_mean

Infer observation mean posterior q(d).

infer_obs_prec

Infer precision posterior q(phi).

compute_lower_bound

Compute the variational lower bound (ELBO) for a GFA model.

compute_lower_bound_constants

Compute constant factors in the variational lower bound.


fit(
Y: ObsStatic,
obs_posterior: ObsParamsPosterior,
latents_posterior: LatentsPosteriorStatic,
config: GFAFitConfig | None = None,
obs_hyperprior: ObsParamsHyperPrior | None = None,
tracker: GFAFitTracker | None = None,
flags: GFAFitFlags | None = None,
max_iter: int | None = None,
callbacks: list | None = None,
) tuple[ObsParamsPosterior, LatentsPosteriorStatic, GFAFitTracker, GFAFitFlags][source]#

Fit a GFA model to data via variational inference.

Parameters:
YObsStatic

Observed data.

obs_posteriorObsParamsPosterior

Observation model posterior (modified in place).

latents_posteriorLatentsPosteriorStatic

Latent posterior (modified in place).

configGFAFitConfig or None, default None

Fitting configuration. If None, uses default GFAFitConfig().

obs_hyperpriorObsParamsHyperPrior or None, default None

Prior hyperparameters. If None, uses default ObsParamsHyperPrior().

trackerGFAFitTracker or None, default None

If provided, append to existing tracker (resume). If None, create fresh.

flagsGFAFitFlags or None, default None

If provided, preserve existing flags (resume). If None, create fresh.

max_iterint or None, default None

Override config.max_iter. Useful for resume with different budget.

callbackslist of Callback or None, default None

List of callback objects. See callbacks module.

Returns:
obs_posteriorObsParamsPosterior

Fitted observation model posterior.

latents_posteriorLatentsPosteriorStatic

Fitted latent posterior.

trackerGFAFitTracker

Quantities tracked during fitting.

flagsGFAFitFlags

Status messages about the fitting process.

Raises:
ValueError

If obs_posterior.y_dims does not match Y.dims, if obs_posterior.x_dim does not match config.x_dim_init (fresh fits only; resumes may have pruned dimensions), or if tracker and flags are not both provided or both None.

Examples

Most users should use GFAModel.fit() instead of calling this directly.

Low-level usage

>>> from latents.gfa.inference import fit
>>> obs_posterior, latents_posterior, tracker, flags = fit(
...     Y, obs_posterior, latents_posterior, config=config
... )
init_posteriors(
Y: ObsStatic,
config: GFAFitConfig | None = None,
obs_hyperprior: ObsParamsHyperPrior | None = None,
) tuple[ObsParamsPosterior, LatentsPosteriorStatic][source]#

Initialize GFA model posteriors for fitting.

Parameters:
YObsStatic

Observed data.

configGFAFitConfig or None, default None

Fitting configuration. If None, uses default GFAFitConfig().

obs_hyperpriorObsParamsHyperPrior or None, default None

Prior hyperparameters. If None, uses default ObsParamsHyperPrior().

Returns:
obs_posteriorObsParamsPosterior

Initialized observation model posterior.

latents_posteriorLatentsPosteriorStatic

Initialized latent posterior.

infer_latents(
Y: ObsStatic,
obs_posterior: ObsParamsPosterior,
latents_posterior: LatentsPosteriorStatic | None = None,
Y_centered: ndarray | None = None,
) LatentsPosteriorStatic[source]#

Infer latent posterior q(X) given observations and fitted parameters.

Parameters:
YObsStatic

Observed data.

obs_posteriorObsParamsPosterior

Posterior over observation parameters. Reads C, phi, d.

latents_posteriorLatentsPosteriorStatic or None, default None

If provided, update in-place and return. If None, create and return a new LatentsPosteriorStatic.

Y_centeredndarray or None, default None

Pre-computed Y.data - d.mean[:, None], shape (y_dim, n_samples). If None, computed internally.

Returns:
LatentsPosteriorStatic

Posterior over latent variables.

Examples

Infer latents for new data using a fitted model:

>>> X_new = infer_latents(Y_new, model.obs_posterior)
>>> X_new.mean  # Posterior mean, shape (x_dim, n_samples)
infer_loadings(
Y: ObsStatic,
obs_posterior: ObsParamsPosterior,
latents_posterior: LatentsPosteriorStatic,
XY: ndarray | None = None,
) LoadingPosterior[source]#

Infer loading posterior q(C). Updates obs_posterior.C in-place.

Parameters:
YObsStatic

Observed data.

obs_posteriorObsParamsPosterior

Posterior over observation parameters. Reads alpha, phi, d; writes C.

latents_posteriorLatentsPosteriorStatic

Posterior over latents. Reads mean, moment.

XYndarray or None, default None

Pre-computed correlation matrix, shape (x_dim, y_dim). Computed if not provided.

Returns:
LoadingPosterior

Reference to obs_posterior.C (updated in-place).

infer_ard(
obs_posterior: ObsParamsPosterior,
hyperprior: ObsParamsHyperPrior,
C_norm: ndarray | None = None,
) ARDPosterior[source]#

Infer ARD posterior q(alpha). Updates obs_posterior.alpha in-place.

Parameters:
obs_posteriorObsParamsPosterior

Posterior over observation parameters. Reads C; writes alpha.

hyperpriorObsParamsHyperPrior

Hyperprior parameters (a_alpha, b_alpha).

C_normndarray or None, default None

Pre-computed squared column norms of C per group, shape (n_groups, x_dim).

Returns:
ARDPosterior

Reference to obs_posterior.alpha (updated in-place).

infer_obs_mean(
Y: ObsStatic,
obs_posterior: ObsParamsPosterior,
latents_posterior: LatentsPosteriorStatic,
hyperprior: ObsParamsHyperPrior,
) ObsMeanPosterior[source]#

Infer observation mean posterior q(d). Updates obs_posterior.d in-place.

Parameters:
YObsStatic

Observed data.

obs_posteriorObsParamsPosterior

Posterior over observation parameters. Reads C, phi; writes d.

latents_posteriorLatentsPosteriorStatic

Posterior over latents. Reads mean.

hyperpriorObsParamsHyperPrior

Hyperprior parameters (beta_d).

Returns:
ObsMeanPosterior

Reference to obs_posterior.d (updated in-place).

infer_obs_prec(
Y: ObsStatic,
obs_posterior: ObsParamsPosterior,
latents_posterior: LatentsPosteriorStatic,
hyperprior: ObsParamsHyperPrior,
d_moment: ndarray | None = None,
XY: ndarray | None = None,
Y2: ndarray | None = None,
Y_sum: ndarray | None = None,
) ObsPrecPosterior[source]#

Infer precision posterior q(phi). Updates obs_posterior.phi in-place.

Parameters:
YObsStatic

Observed data.

obs_posteriorObsParamsPosterior

Posterior over observation parameters. Reads C, d; writes phi.

latents_posteriorLatentsPosteriorStatic

Posterior over latents. Reads mean, moment.

hyperpriorObsParamsHyperPrior

Hyperprior parameters (a_phi, b_phi).

d_momentndarray or None, default None

Pre-computed second moment of d, shape (y_dim,).

XYndarray or None, default None

Pre-computed correlation matrix, shape (x_dim, y_dim).

Y2ndarray or None, default None

Pre-computed sample second moments, shape (y_dim,).

Y_sumndarray or None, default None

Pre-computed column sums of Y, shape (y_dim,).

Returns:
ObsPrecPosterior

Reference to obs_posterior.phi (updated in-place).

compute_lower_bound(
Y: ObsStatic,
obs_posterior: ObsParamsPosterior,
latents_posterior: LatentsPosteriorStatic,
obs_hyperprior: ObsParamsHyperPrior,
consts: tuple | None = None,
logdet_C: float | None = None,
C_norm: ndarray | None = None,
d_moment: ndarray | None = None,
recon: ndarray | None = None,
) float[source]#

Compute the variational lower bound (ELBO) for a GFA model.

Parameters:
YObsStatic

Observed data.

obs_posteriorObsParamsPosterior

Observation model posterior.

latents_posteriorLatentsPosteriorStatic

Latent posterior.

obs_hyperpriorObsParamsHyperPrior

Hyperparameters of the prior distributions.

conststuple or None, default None

Constant factors in the lower bound. If None, computed.

logdet_Cfloat or None, default None

Log-determinant of loading covariances. If None, computed.

C_normndarray or None, default None

Expected squared norms of loading columns. If None, computed.

d_momentndarray or None, default None

Second moment of observation mean. If None, computed.

reconndarray or None, default None

Expected reconstruction error per dimension, shape (y_dim,). If None, derived from obs_posterior.phi.b.

Returns:
float

Variational lower bound.

compute_lower_bound_constants(
n_samples: int,
obs_posterior: ObsParamsPosterior,
obs_hyperprior: ObsParamsHyperPrior,
) tuple[float, float, float, float, float, float, float, float, ndarray, ndarray][source]#

Compute constant factors in the variational lower bound.

Parameters:
n_samplesint

Number of samples in the observed data.

obs_posteriorObsParamsPosterior

Observation model posterior.

obs_hyperpriorObsParamsHyperPrior

Hyperparameters of the prior distributions.

Returns:
tuple of (float, float, float, float, float, float, float, float, ndarray, ndarray)

Constant factors: const_lik, const_d, alogb_phi, loggamma_a_phi_prior, loggamma_a_phi_post, digamma_a_phi, alogb_alpha, loggamma_a_alpha_prior, loggamma_a_alpha_post, digamma_a_alpha.