Source code for latents.state.priors

"""Prior distributions for state model parameters."""

from __future__ import annotations

from dataclasses import dataclass, field

import numpy as np

from latents.state.realizations import LatentsRealization


[docs] class LatentsPriorStatic: """Static latent prior: X ~ N(0, I). The standard GFA prior assumes independent standard normal latents. Examples -------- >>> prior = LatentsPriorStatic() >>> rng = np.random.default_rng(42) >>> X = prior.sample(x_dim=5, n_samples=100, rng=rng) >>> X.data.shape (5, 100) """
[docs] def sample( self, x_dim: int, n_samples: int, rng: np.random.Generator, ) -> LatentsRealization: """Sample X ~ N(0, I). Parameters ---------- x_dim : int Number of latent dimensions. n_samples : int Number of samples to generate. rng : numpy.random.Generator Random number generator. Returns ------- LatentsRealization Sampled latent values. """ X = rng.normal(size=(x_dim, n_samples)) return LatentsRealization(data=X)
@dataclass class LatentsHyperPriorGP: """GP kernel hyperpriors. Stub for GPFA/mDLAG. GP hyperparameters are learnable. Parameters ---------- kernel : str, default "rbf" Kernel type. timescale : float, default 50.0 Characteristic timescale of the GP kernel. variance : float, default 1.0 Signal variance of the GP kernel. """ kernel: str = "rbf" timescale: float = 50.0 variance: float = 1.0 @dataclass class LatentsPriorGP: """GP latent prior. Stub for GPFA/mDLAG. Parameters ---------- hyperprior : LatentsHyperPriorGP, default LatentsHyperPriorGP() GP kernel hyperpriors. """ hyperprior: LatentsHyperPriorGP = field(default_factory=LatentsHyperPriorGP) def sample( self, x_dim: int, n_samples: int, n_timepoints: int, rng: np.random.Generator, ) -> LatentsRealization: """Sample from GP prior. Parameters ---------- x_dim : int Number of latent dimensions. n_samples : int Number of samples (trials). n_timepoints : int Number of time points per sample. rng : numpy.random.Generator Random number generator. Returns ------- LatentsRealization Sampled latent trajectories. Raises ------ NotImplementedError GP prior sampling not yet implemented. """ msg = "GP prior sampling not yet implemented" raise NotImplementedError(msg)