Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions dynamax/hidden_markov_model/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@
from dynamax.hidden_markov_model.models.logreg_hmm import LogisticRegressionHMM
from dynamax.hidden_markov_model.models.multinomial_hmm import MultinomialHMM
from dynamax.hidden_markov_model.models.poisson_hmm import PoissonHMM
from dynamax.hidden_markov_model.models.inputdriven_linreg_hmm import InputDrivenLinearRegressionHMM
from dynamax.hidden_markov_model.models.inputdriven_categoricalreg_hmm import InputDrivenCategoricalRegressionHMM

from dynamax.hidden_markov_model.inference import HMMPosterior
from dynamax.hidden_markov_model.inference import HMMPosteriorFiltered
Expand Down
3 changes: 2 additions & 1 deletion dynamax/hidden_markov_model/models/abstractions.py
Original file line number Diff line number Diff line change
Expand Up @@ -602,7 +602,8 @@ def _inference_args(self, params: HMMParameterSet,
emissions: Array,
inputs: Optional[Array]) -> Tuple:
"""Return the arguments needed for inference."""
return (self.initial_component._compute_initial_probs(params.initial, inputs),
initial_inputs = pytree_slice(inputs, 0) if inputs is not None else None
return (self.initial_component._compute_initial_probs(params.initial, initial_inputs),
self.transition_component._compute_transition_matrices(params.transitions, inputs),
self.emission_component._compute_conditional_logliks(params.emissions, emissions, inputs))

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
"""
Categorial Regression hidden Markov model (HMM) with state-dependent weights and input-driven state transitions.
"""

from typing import Any, Dict, NamedTuple, Optional, Tuple, Union
import jax.random as jr
from jaxtyping import Array, Float, Int, PyTree
import optax

from dynamax.hidden_markov_model.models.abstractions import HMM, HMMParameterSet, HMMPropertySet
from dynamax.hidden_markov_model.models.categorical_glm_hmm import ParamsCategoricalRegressionHMMEmissions, CategoricalRegressionHMMEmissions

from dynamax.hidden_markov_model.models.inputdriven_linreg_hmm import InputDrivenHMMInitialState, ParamsInputDrivenHMMInitialState
from dynamax.hidden_markov_model.models.inputdriven_linreg_hmm import InputDrivenHMMTransitions, ParamsInputDrivenHMMTransitions


class ParamsInputDrivenCategoricalRegressionHMM(NamedTuple):
"""Parameters for an input-driven categorical regression HMM."""
initial: ParamsInputDrivenHMMInitialState
transitions: ParamsInputDrivenHMMTransitions
emissions: ParamsCategoricalRegressionHMMEmissions


class InputDrivenCategoricalRegressionHMM(HMM):
r"""An HMM whose emissions come from a categorical regression with state-dependent weights and
initial-state and transition distributions are driven by an external input, rather than fixed.

The initial distribution (see `InputDrivenHMMInitialState`) and the transition
distribution (see `InputDrivenHMMTransitions`) are both multinomial logistic
regressions on the input,

$$p(z_1 \mid u_1, \theta) = \mathrm{Cat}(z_1 \mid \mathrm{softmax}(W^{\mathsf{init}} u_1 + b^{\mathsf{init}}))$$

$$p(z_t \mid z_{t-1}, u_t, \theta) = \mathrm{Cat}(z_t \mid \mathrm{softmax}(W^{\mathsf{trans}}_{z_{t-1}} u_t + b^{\mathsf{trans}}_{z_{t-1}}))$$

The emission distribution is a state-dependent linear regression, as in `CategoricalRegressionHMM`,

$$p(y_t \mid z_t, u_t, \theta) = \mathrm{Cat}(y_{t} \mid \mathrm{softmax}(W^{\mathsf{emis}}_{z_t} u_t + b^{\mathsf{emis}}_{z_t}))$$

:param num_states: number of discrete states $K$
:param input_dim: input dimension $M$
:param emission_dim: emission dimension $N$
:param m_step_optimizer: ``optax`` optimizer, like Adam, used for the transition M-step.
:param m_step_num_iters: number of optimizer steps per M-step.
"""

def __init__(self,
num_states: int,
num_classes: int,
input_dim: int,
m_step_optimizer: optax.GradientTransformation = optax.adam(1e-2),
m_step_num_iters: int = 50):
self.num_classes = num_classes
self.input_dim = input_dim
initial_component = InputDrivenHMMInitialState(num_states, input_dim)
transition_component = InputDrivenHMMTransitions(num_states, input_dim, m_step_optimizer=m_step_optimizer, m_step_num_iters=m_step_num_iters)
emission_component = CategoricalRegressionHMMEmissions(num_states, num_classes, input_dim, m_step_optimizer=m_step_optimizer, m_step_num_iters=m_step_num_iters)
super().__init__(num_states, initial_component, transition_component, emission_component)

@property
def inputs_shape(self):
"""Return the shape of the input."""
return (self.input_dim,)

def initialize(self,
key: Array=jr.PRNGKey(0),
method: str="prior",
transition_weights: Optional[Float[Array, "num_states num_states input_dim"]] = None,
transition_biases: Optional[Float[Array, "num_states num_states"]] = None,
emission_weights: Optional[Float[Array, "num_states emission_dim input_dim"]]=None,
emission_biases: Optional[Float[Array, "num_states emission_dim"]]=None,
) -> Tuple[HMMParameterSet, HMMPropertySet]:
"""Initialize the model parameters and their corresponding properties.

You can either specify parameters manually via the keyword arguments, or you can have
them set automatically. If any parameters are not specified, you must supply a PRNGKey.

Args:
key: random number generator for unspecified parameters. Must not be None if there are any unspecified parameters.
method: method for initializing unspecified parameters. Currently only "prior" is supported.
transition_weights: manually specified transition weights.
transition_biases: manually specified transition biases.
emission_weights: manually specified emission weights.
emission_biases: manually specified emission biases.
Returns:
Model parameters and their properties.
"""
key1, key2, key3 = jr.split(key , 3)
params, props = dict(), dict()
params["initial"], props["initial"] = self.initial_component.initialize(key1, method=method)
params["transitions"], props["transitions"] = self.transition_component.initialize(key2, method=method, weights=transition_weights, biases=transition_biases)
params["emissions"], props["emissions"] = self.emission_component.initialize(key=key3, method=method, emission_weights=emission_weights, emission_biases=emission_biases)
return ParamsInputDrivenCategoricalRegressionHMM(**params), ParamsInputDrivenCategoricalRegressionHMM(**props)
240 changes: 240 additions & 0 deletions dynamax/hidden_markov_model/models/inputdriven_linreg_hmm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,240 @@
"""
Linear regression hidden Markov model (HMM) with state-dependent emission weights and
input-driven initial-state and transition distributions.
"""
import jax.random as jr
import tensorflow_probability.substrates.jax.distributions as tfd
import optax

from dynamax.parameters import ParameterProperties, ParameterSet
from dynamax.hidden_markov_model.models.abstractions import HMM, HMMInitialState, HMMTransitions
from dynamax.hidden_markov_model.models.abstractions import HMMParameterSet, HMMPropertySet
from dynamax.hidden_markov_model.inference import HMMPosterior
from dynamax.hidden_markov_model.models.linreg_hmm import LinearRegressionHMMEmissions
from dynamax.hidden_markov_model.models.linreg_hmm import ParamsLinearRegressionHMMEmissions
from dynamax.types import Scalar

from typing import NamedTuple, Optional, Tuple, Union
from jaxtyping import Array, Float, Int


class ParamsInputDrivenHMMInitialState(NamedTuple):
"""Parameters for the initial distribution of an input-driven HMM."""
weights: Union[Float[Array, "num_states input_dim"], ParameterProperties]
biases: Union[Float[Array, "num_states"], ParameterProperties]


class ParamsInputDrivenHMMTransitions(NamedTuple):
"""Parameters for the transitions of an input-driven HMM."""
weights: Union[Float[Array, "num_states num_states input_dim"], ParameterProperties]
biases: Union[Float[Array, "num_states num_states"], ParameterProperties]


class ParamsInputDrivenLinearRegressionHMM(NamedTuple):
"""Parameters for an input-driven linear regression HMM."""
initial: ParamsInputDrivenHMMInitialState
transitions: ParamsInputDrivenHMMTransitions
emissions: ParamsLinearRegressionHMMEmissions


class InputDrivenHMMInitialState(HMMInitialState):
"""
HMM initial-state distribution for an input-driven HMM.
The initial-state probabilities depend on the input at the first timestep:
P(z_1 = k | u_1) = softmax(W_k @ u_1 + b_k)
"""
def __init__(self,
num_states: int,
input_dim: int,
m_step_optimizer: optax.GradientTransformation = optax.adam(1e-2),
m_step_num_iters: int = 50):
super().__init__(m_step_optimizer=m_step_optimizer, m_step_num_iters=m_step_num_iters)
self.num_states = num_states
self.input_dim = input_dim

def distribution(self, params: ParamsInputDrivenHMMInitialState, inputs=Float[Array, " input_dim"]) -> tfd.Distribution:
"""Return the distribution object of the initial distribution."""
logits = params.weights @ inputs + params.biases
return tfd.Categorical(logits=logits)

def initialize(
self,
key: Optional[Array] = None,
method: str = "prior",
**kwargs
) -> Tuple[ParamsInputDrivenHMMInitialState, ParamsInputDrivenHMMInitialState]:
"""Initialize the model parameters and their corresponding properties."""

if key is None:
raise ValueError("key must be provided.")

# Initialize the initial probabilities with small random weights
key_w, key_b = jr.split(key)
weights = jr.normal(key_w, (self.num_states, self.input_dim)) * 0.01
biases = jr.normal(key_b, (self.num_states,)) * 0.01

# Package the results into dictionaries
params = ParamsInputDrivenHMMInitialState(weights=weights, biases=biases)
props = ParamsInputDrivenHMMInitialState(weights=ParameterProperties(), biases=ParameterProperties())
return params, props

def log_prior(self, params: ParamsInputDrivenHMMInitialState) -> Scalar:
"""Compute the log prior of the parameters."""
return 0.0


class InputDrivenHMMTransitions(HMMTransitions):
"""
HMM transitions for an input-driven HMM.
The transition probabilities depend on external inputs/covariates:
P(z_t | z_{t-1}, u_t) where u_t are inputs at time t

For each previous state j, we use multinomial logistic regression:
P(z_t = k | z_{t-1} = j, u_t) = softmax(W_j @ u_t + b_j)[k]
"""

def __init__(
self,
num_states: int,
input_dim: int,
m_step_optimizer: optax.GradientTransformation = optax.adam(1e-2),
m_step_num_iters: int = 50
):
super().__init__(m_step_optimizer=m_step_optimizer, m_step_num_iters=m_step_num_iters)
self.num_states = num_states
self.input_dim = input_dim

def distribution(
self,
params: ParamsInputDrivenHMMTransitions,
state: Union[int, Int[Array, ""]],
inputs: Float[Array, " input_dim"]) -> tfd.Distribution:
"""
Return the distribution over the next state given the current state and input.
"""
if inputs is None:
raise ValueError("Inputs must be provided for input-driven transitions")

logits = params.weights[state] @ inputs + params.biases[state]
return tfd.Categorical(logits=logits)

def initialize(
self,
key: Optional[Array] = None,
method: str = "prior",
weights=None,
biases=None,
**kwargs
) -> Tuple[ParamsInputDrivenHMMTransitions, ParamsInputDrivenHMMTransitions]:
"""Initialize the model parameters and their corresponding properties."""

if key is None:
raise ValueError("key must be provided.")

# Initialize with small random weights so transitions start near uniform
key_w, key_b = jr.split(key)
_weights = jr.normal(key_w, (self.num_states, self.num_states, self.input_dim)) * 0.01
_biases = jr.normal(key_b, (self.num_states, self.num_states)) * 0.01

# Only use the values above if the user hasn't specified their own
default = lambda x, x0: x if x is not None else x0
params = ParamsInputDrivenHMMTransitions(
weights=default(weights, _weights),
biases=default(biases, _biases))
props = ParamsInputDrivenHMMTransitions(
weights=ParameterProperties(),
biases=ParameterProperties())
return params, props

def log_prior(self, params: ParamsInputDrivenHMMTransitions) -> Scalar:
"""Return the log-prior probability of the emission parameters.

Currently, there is no prior so this function returns 0.
"""
return 0.0


class InputDrivenLinearRegressionHMM(HMM):
r"""An HMM with linear-regression emissions whose initial-state and transition
distributions are driven by an external input, rather than fixed.

Let $y_t \in \mathbb{R}^N$ and $u_t \in \mathbb{R}^M$ denote vector-valued emissions
and inputs at time $t$, respectively. The initial distribution (see `InputDrivenHMMInitialState`)
and the transition distribution (see `InputDrivenHMMTransitions`) are both multinomial logistic
regressions on the input,

$$p(z_1 \mid u_1, \theta) = \mathrm{Cat}(z_1 \mid \mathrm{softmax}(W^{\mathsf{init}} u_1 + b^{\mathsf{init}}))$$

with *initial weights* $W^{\mathsf{init}} \in \mathbb{R}^{K \times M}$ and
*initial biases* $b^{\mathsf{init}} \in \mathbb{R}^K$, and

$$p(z_t \mid z_{t-1}, u_t, \theta) = \mathrm{Cat}(z_t \mid \mathrm{softmax}(W^{\mathsf{trans}}_{z_{t-1}} u_t + b^{\mathsf{trans}}_{z_{t-1}}))$$

with *transition weights* $W_j^{\mathsf{trans}} \in \mathbb{R}^{K \times M}$ and *transition biases* $b_j^{\mathsf{trans}} \in \mathbb{R}^{K}$.

The emission distribution is a state-dependent linear regression, as in `LinearRegressionHMM`,

$$p(y_t \mid z_t, u_t, \theta) = \mathcal{N}(y_t \mid W^{\mathsf{emis}}_{z_t} u_t + b^{\mathsf{emis}}_{z_t}, \Sigma_{z_t})$$

with *emission weights* $W_k \in \mathbb{R}^{N \times M}$, *emission biases* $b_k \in \mathbb{R}^N$,
and *emission covariances* $\Sigma_k \in \mathbb{R}_{\succeq 0}^{N \times N}$.

:param num_states: number of discrete states $K$
:param input_dim: input dimension $M$
:param emission_dim: emission dimension $N$
:param m_step_optimizer: ``optax`` optimizer, like Adam, used for the transition M-step.
:param m_step_num_iters: number of optimizer steps per M-step.
"""

def __init__(self,
num_states: int,
input_dim: int,
emission_dim: int,
m_step_optimizer: optax.GradientTransformation = optax.adam(1e-2),
m_step_num_iters: int = 50):
self.emission_dim = emission_dim
self.input_dim = input_dim
initial_component = InputDrivenHMMInitialState(num_states, input_dim)
transition_component = InputDrivenHMMTransitions(num_states, input_dim, m_step_optimizer=m_step_optimizer, m_step_num_iters=m_step_num_iters)
emission_component = LinearRegressionHMMEmissions(num_states, input_dim, emission_dim)
super().__init__(num_states, initial_component, transition_component, emission_component)

@property
def inputs_shape(self):
"""Return the shape of the input."""
return (self.input_dim,)

def initialize(self,
key: Array=jr.PRNGKey(0),
method: str="prior",
transition_weights: Optional[Float[Array, "num_states num_states input_dim"]]=None,
transition_biases: Optional[Float[Array, "num_states num_states"]]=None,
emission_weights: Optional[Float[Array, "num_states emission_dim input_dim"]]=None,
emission_biases: Optional[Float[Array, "num_states emission_dim"]]=None,
emission_covariances: Optional[Float[Array, "num_states emission_dim emission_dim"]]=None,
emissions: Optional[Float[Array, "num_timesteps emission_dim"]]=None
) -> Tuple[HMMParameterSet, HMMPropertySet]:
"""Initialize the model parameters and their corresponding properties.

You can either specify parameters manually via the keyword arguments, or you can have
them set automatically. If any parameters are not specified, you must supply a PRNGKey.

Args:
key: random number generator for unspecified parameters. Must not be None if there are any unspecified parameters.
method: method for initializing unspecified parameters. Currently only "prior" is supported.
transition_weights: manually specified transition weights.
transition_biases: manually specified transition biases.
emission_weights: manually specified emission weights.
emission_biases: manually specified emission biases.
emission_covariances: manually specified emission covariances.
emissions: emissions for initializing the parameters with kmeans.

Returns:
Model parameters and their properties.
"""
key1, key2, key3 = jr.split(key , 3)
params, props = dict(), dict()
params["initial"], props["initial"] = self.initial_component.initialize(key1, method=method)
params["transitions"], props["transitions"] = self.transition_component.initialize(key2, method=method, weights=transition_weights, biases=transition_biases)
params["emissions"], props["emissions"] = self.emission_component.initialize(key3, method=method, emission_weights=emission_weights, emission_biases=emission_biases, emission_covariances=emission_covariances, emissions=emissions)
return ParamsInputDrivenLinearRegressionHMM(**params), ParamsInputDrivenLinearRegressionHMM(**props)
19 changes: 19 additions & 0 deletions dynamax/hidden_markov_model/models/test_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,8 @@
(models.LogisticRegressionHMM, dict(num_states=4, input_dim=5), jnp.ones((NUM_TIMESTEPS, 5))),
(models.MultinomialHMM, dict(num_states=4, emission_dim=3, num_classes=5, num_trials=10), None),
(models.PoissonHMM, dict(num_states=4, emission_dim=3), None),
(models.InputDrivenLinearRegressionHMM, dict(num_states=3, emission_dim=3, input_dim=2), jr.normal(jr.PRNGKey(0),(NUM_TIMESTEPS, 2))),
(models.InputDrivenCategoricalRegressionHMM, dict(num_states=3, num_classes=3, input_dim=2), jnp.ones((NUM_TIMESTEPS, 2))),
]


Expand Down Expand Up @@ -104,6 +106,23 @@ def test_sample_and_fit_arhmm():
fitted_params, lps = arhmm.fit_sgd(params, param_props, emissions, inputs=inputs, num_epochs=10)


def test_transitions_and_initial_state_depend_on_inputs():
"""Test that the initial-state and transition distributions change with the input."""
hmm = models.InputDrivenLinearRegressionHMM(num_states=3, input_dim=5, emission_dim=3)
params, _ = hmm.initialize(jr.PRNGKey(0))

u_a = -3.0 * jnp.ones(5)
u_b = 3.0 * jnp.ones(5)

init_a = hmm.initial_distribution(params, u_a).probs_parameter()
init_b = hmm.initial_distribution(params, u_b).probs_parameter()
assert not jnp.allclose(init_a, init_b, atol=1e-3)

trans_a = hmm.transition_distribution(params, 0, u_a).probs_parameter()
trans_b = hmm.transition_distribution(params, 0, u_b).probs_parameter()
assert not jnp.allclose(trans_a, trans_b, atol=1e-3)


# @pytest.mark.skip(reason="this would introduce a torch dependency")
# def test_hmm_fit_stochastic_em(num_iters=100):
# """Evaluate stochastic em fit with respect to exact em fit."""
Expand Down