From 24567b1e47e20dcf5e3534875ee917408562ee38 Mon Sep 17 00:00:00 2001 From: Umesh Singla Date: Thu, 16 Jul 2026 15:23:41 -0700 Subject: [PATCH 1/3] Add linear and categorical regression HMM classes with input-driven state transitions --- dynamax/hidden_markov_model/__init__.py | 2 + .../models/inputdriven_categoricalreg_hmm.py | 93 +++++++ .../models/inputdriven_linreg_hmm.py | 240 ++++++++++++++++++ 3 files changed, 335 insertions(+) create mode 100644 dynamax/hidden_markov_model/models/inputdriven_categoricalreg_hmm.py create mode 100644 dynamax/hidden_markov_model/models/inputdriven_linreg_hmm.py diff --git a/dynamax/hidden_markov_model/__init__.py b/dynamax/hidden_markov_model/__init__.py index 3932a1ee8..61bf5a680 100644 --- a/dynamax/hidden_markov_model/__init__.py +++ b/dynamax/hidden_markov_model/__init__.py @@ -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 diff --git a/dynamax/hidden_markov_model/models/inputdriven_categoricalreg_hmm.py b/dynamax/hidden_markov_model/models/inputdriven_categoricalreg_hmm.py new file mode 100644 index 000000000..93f849137 --- /dev/null +++ b/dynamax/hidden_markov_model/models/inputdriven_categoricalreg_hmm.py @@ -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) diff --git a/dynamax/hidden_markov_model/models/inputdriven_linreg_hmm.py b/dynamax/hidden_markov_model/models/inputdriven_linreg_hmm.py new file mode 100644 index 000000000..00e0953db --- /dev/null +++ b/dynamax/hidden_markov_model/models/inputdriven_linreg_hmm.py @@ -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) From 23ee3837398080f3d233efe1fc1bc66e1eb77fc3 Mon Sep 17 00:00:00 2001 From: Umesh Singla Date: Sat, 18 Jul 2026 16:51:12 -0700 Subject: [PATCH 2/3] add tests for inputdriven hmms --- .../hidden_markov_model/models/test_models.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/dynamax/hidden_markov_model/models/test_models.py b/dynamax/hidden_markov_model/models/test_models.py index 73a2969a4..0cb822d5d 100644 --- a/dynamax/hidden_markov_model/models/test_models.py +++ b/dynamax/hidden_markov_model/models/test_models.py @@ -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))), ] @@ -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.""" From 19243aeb187c8b5a14064c0058eac5c3a29882b2 Mon Sep 17 00:00:00 2001 From: Umesh Singla Date: Sat, 18 Jul 2026 16:52:14 -0700 Subject: [PATCH 3/3] fix inconsistent arguments to initial_distribution _compute_initial_probs is called from two places: _inference_args() and _single_expected_log_like() in m_step. But _inference_args() passes the full input sequence to it, while _single_expected_log_like() passes only the first timestep's input (as retrieved from collect_suff_stats). StandardHMMInitialState never surfaces this, since its distribution() ignores inputs entirely. Fix by slicing inputs to the first timestep specifically for the initial component, leaving the transition/emission components' full-sequence inputs unchanged. --- dynamax/hidden_markov_model/models/abstractions.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/dynamax/hidden_markov_model/models/abstractions.py b/dynamax/hidden_markov_model/models/abstractions.py index 1dd6b946b..fcbf565a5 100644 --- a/dynamax/hidden_markov_model/models/abstractions.py +++ b/dynamax/hidden_markov_model/models/abstractions.py @@ -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))