gfa.inference#
Inference functions for Group Factor Analysis.
Functions
Fit a GFA model to data via variational inference. |
|
Initialize GFA model posteriors for fitting. |
|
Infer latent posterior q(X) given observations and fitted parameters. |
|
Infer loading posterior q(C). |
|
Infer ARD posterior q(alpha). |
|
Infer observation mean posterior q(d). |
|
Infer precision posterior q(phi). |
|
Compute the variational lower bound (ELBO) for a GFA model. |
|
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,
Fit a GFA model to data via variational inference.
- Parameters:
- Y
ObsStatic Observed data.
- obs_posterior
ObsParamsPosterior Observation model posterior (modified in place).
- latents_posterior
LatentsPosteriorStatic Latent posterior (modified in place).
- config
GFAFitConfigorNone, defaultNone Fitting configuration. If None, uses default
GFAFitConfig().- obs_hyperprior
ObsParamsHyperPriororNone, defaultNone Prior hyperparameters. If None, uses default
ObsParamsHyperPrior().- tracker
GFAFitTrackerorNone, defaultNone If provided, append to existing tracker (resume). If None, create fresh.
- flags
GFAFitFlagsorNone, defaultNone If provided, preserve existing flags (resume). If None, create fresh.
- max_iter
intorNone, defaultNone Override
config.max_iter. Useful for resume with different budget.- callbacks
listofCallbackorNone, defaultNone List of callback objects. See
callbacksmodule.
- Y
- Returns:
- obs_posterior
ObsParamsPosterior Fitted observation model posterior.
- latents_posterior
LatentsPosteriorStatic Fitted latent posterior.
- tracker
GFAFitTracker Quantities tracked during fitting.
- flags
GFAFitFlags Status messages about the fitting process.
- obs_posterior
- Raises:
ValueErrorIf
obs_posterior.y_dimsdoes not matchY.dims, ifobs_posterior.x_dimdoes not matchconfig.x_dim_init(fresh fits only; resumes may have pruned dimensions), or iftrackerandflagsare not both provided or bothNone.
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,
Initialize GFA model posteriors for fitting.
- Parameters:
- Returns:
- obs_posterior
ObsParamsPosterior Initialized observation model posterior.
- latents_posterior
LatentsPosteriorStatic Initialized latent posterior.
- obs_posterior
- infer_latents(
- Y: ObsStatic,
- obs_posterior: ObsParamsPosterior,
- latents_posterior: LatentsPosteriorStatic | None = None,
- Y_centered: ndarray | None = None,
Infer latent posterior q(X) given observations and fitted parameters.
- Parameters:
- Y
ObsStatic Observed data.
- obs_posterior
ObsParamsPosterior Posterior over observation parameters. Reads
C,phi,d.- latents_posterior
LatentsPosteriorStaticorNone, defaultNone If provided, update in-place and return. If
None, create and return a newLatentsPosteriorStatic.- Y_centered
ndarrayorNone, defaultNone Pre-computed
Y.data - d.mean[:, None], shape(y_dim, n_samples). If None, computed internally.
- Y
- Returns:
LatentsPosteriorStaticPosterior 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,
Infer loading posterior q(C). Updates
obs_posterior.Cin-place.- Parameters:
- Y
ObsStatic Observed data.
- obs_posterior
ObsParamsPosterior Posterior over observation parameters. Reads
alpha,phi,d; writesC.- latents_posterior
LatentsPosteriorStatic Posterior over latents. Reads
mean,moment.- XY
ndarrayorNone, defaultNone Pre-computed correlation matrix, shape
(x_dim, y_dim). Computed if not provided.
- Y
- Returns:
LoadingPosteriorReference to
obs_posterior.C(updated in-place).
- infer_ard(
- obs_posterior: ObsParamsPosterior,
- hyperprior: ObsParamsHyperPrior,
- C_norm: ndarray | None = None,
Infer ARD posterior q(alpha). Updates
obs_posterior.alphain-place.- Parameters:
- Returns:
ARDPosteriorReference to
obs_posterior.alpha(updated in-place).
- infer_obs_mean(
- Y: ObsStatic,
- obs_posterior: ObsParamsPosterior,
- latents_posterior: LatentsPosteriorStatic,
- hyperprior: ObsParamsHyperPrior,
Infer observation mean posterior q(d). Updates
obs_posterior.din-place.- Parameters:
- Y
ObsStatic Observed data.
- obs_posterior
ObsParamsPosterior Posterior over observation parameters. Reads
C,phi; writesd.- latents_posterior
LatentsPosteriorStatic Posterior over latents. Reads
mean.- hyperprior
ObsParamsHyperPrior Hyperprior parameters (
beta_d).
- Y
- Returns:
ObsMeanPosteriorReference 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,
Infer precision posterior q(phi). Updates
obs_posterior.phiin-place.- Parameters:
- Y
ObsStatic Observed data.
- obs_posterior
ObsParamsPosterior Posterior over observation parameters. Reads
C,d; writesphi.- latents_posterior
LatentsPosteriorStatic Posterior over latents. Reads
mean,moment.- hyperprior
ObsParamsHyperPrior Hyperprior parameters (
a_phi,b_phi).- d_moment
ndarrayorNone, defaultNone Pre-computed second moment of
d, shape(y_dim,).- XY
ndarrayorNone, defaultNone Pre-computed correlation matrix, shape
(x_dim, y_dim).- Y2
ndarrayorNone, defaultNone Pre-computed sample second moments, shape
(y_dim,).- Y_sum
ndarrayorNone, defaultNone Pre-computed column sums of
Y, shape(y_dim,).
- Y
- Returns:
ObsPrecPosteriorReference 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,
Compute the variational lower bound (ELBO) for a GFA model.
- Parameters:
- Y
ObsStatic Observed data.
- obs_posterior
ObsParamsPosterior Observation model posterior.
- latents_posterior
LatentsPosteriorStatic Latent posterior.
- obs_hyperprior
ObsParamsHyperPrior Hyperparameters of the prior distributions.
- consts
tupleorNone, defaultNone Constant factors in the lower bound. If
None, computed.- logdet_C
floatorNone, defaultNone Log-determinant of loading covariances. If
None, computed.- C_norm
ndarrayorNone, defaultNone Expected squared norms of loading columns. If
None, computed.- d_moment
ndarrayorNone, defaultNone Second moment of observation mean. If
None, computed.- recon
ndarrayorNone, defaultNone Expected reconstruction error per dimension, shape
(y_dim,). IfNone, derived fromobs_posterior.phi.b.
- Y
- Returns:
floatVariational lower bound.
- compute_lower_bound_constants(
- n_samples: int,
- obs_posterior: ObsParamsPosterior,
- obs_hyperprior: ObsParamsHyperPrior,
Compute constant factors in the variational lower bound.
- Parameters:
- n_samples
int Number of samples in the observed data.
- obs_posterior
ObsParamsPosterior Observation model posterior.
- obs_hyperprior
ObsParamsHyperPrior Hyperparameters of the prior distributions.
- n_samples
- Returns: