Skip to content

Align observation process API with user's mental model #662

Description

@cdc-mitzimorris

Problem

The current observation process API requires users to manage timeline offsets between their data and the model's internal representation. This creates confusion and potential for off-by-one errors.

User's mental model

  • "I have 90 days of hospital data, days 0-89"
  • "I provide 90 observations"
  • "I get back 90 days of inferred infections"

Current implementation

The model internally requires an initialization period (n_init) to satisfy the renewal equation and delay convolutions:

  • Infections array: 133 days total (43 init + 90 user days)
  • Predicted counts: First 42 days are NaN (insufficient convolution history), giving 91 valid days
  • User's observations: 90 days, mapping to model days 43-132

Where: n_init = max(len(generation_interval), len(delay_pmf_1), len(delay_pmf_2), ...)

The issue

NaN filtering alone gives 91 valid days (42-132), but user has 90 observations (43-132). The observation process needs to know where user's data starts — not just filter NaN.

Proposed solution

Add an n_init parameter to observation sample() methods that specifies where the user's observation period begins.

Counts.sample() revised signature

def sample(
    self,
    infections: ArrayLike,
    obs: ArrayLike | None = None,
    n_init: int | None = None,
) -> ObservationSample:
    predicted_counts = self._predicted_obs(infections)
    self._deterministic("predicted", predicted_counts)

    if n_init is None:
        n_init = self.min_valid_index

    predicted_obs = predicted_counts[n_init:]

    if obs is not None and len(obs) != len(predicted_obs):
        raise ValueError(
            f"obs length {len(obs)} != predicted length {len(predicted_obs)}. "
            f"Expected {len(predicted_obs)} observations for days {n_init}+."
        )

    observed = self.noise.sample(
        name=self._sample_site_name("obs"),
        predicted=predicted_obs,
        obs=obs,
    )

    return ObservationSample(observed=observed, predicted=predicted_counts)

Usage standalone (no ModelBuilder)

  • User has 90 days of data
  • Counts uses default n_init = min_valid_index = 42
  • call: counts.sample(infections=infections_132, obs=obs_90)
  • return: predicted_counts[42:] has length 90, matches obs

Usage with ModelBuilder

  • Model computes n_init = 43 (max of all lookbacks)
  • Model passes n_init to observation
  • call: counts.sample(infections=infections_133, obs=obs_90, n_init=43)
  • return: predicted_counts[43:] has length 90, matches obs

Summary of changes

  • Counts.sample(): Add n_init parameter, slice predicted_counts[n_init:]
  • CountsBySubpop.sample(): Add n_init parameter, offset times internally
  • Measurements.sample(): Add n_init parameter, offset times internally
  • MultiSignalModel.sample(): Pass n_init to each observation's sample()
  • User API: User provides observations in natural coordinates (day 0 = first data point)

User experience after changes

# User's data
hospital_obs = jnp.array([...])  # 90 days, indexed 0-89

# With ModelBuilder (recommended)
model.fit(
    n_days=90,
    hospital={"obs": hospital_obs},  # No offset needed
)

# Standalone Counts (advanced)
counts.sample(
    infections=full_infections,
    obs=hospital_obs,
    n_init=43,
)

Related

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Labels

No labels
No labels

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions