diff --git a/meridian/model/context.py b/meridian/model/context.py index 0af6bb63b..f1adede5b 100644 --- a/meridian/model/context.py +++ b/meridian/model/context.py @@ -17,8 +17,9 @@ from collections.abc import Mapping, Sequence import dataclasses +import datetime import functools -from typing import Any +from typing import Any, cast import warnings from meridian import backend @@ -61,6 +62,267 @@ def _get_decay(decay_spec: str | Sequence[str], index: int) -> str: return decay_spec[index] if not isinstance(decay_spec, str) else decay_spec +def _nearest_coordinates_hint( + target: datetime.date, dates: Sequence[datetime.date] +) -> str: + """Describes the time coordinates bracketing `target`, for error messages.""" + earlier = [date for date in dates if date < target] + later = [date for date in dates if date > target] + nearest = [str(d) for d in earlier[-1:] + later[:1]] + if not nearest: + return "" + return f" The nearest are {' and '.join(nearest)}." + + +def _compile_date_range_mask( + date_ranges: Sequence[spec.DateRange], + dates: Sequence[datetime.date], + *, + spec_name: str, +) -> np.ndarray: + """Compiles a sequence of `DateRange`s into a boolean mask over `dates`. + + Each `DateRange` is the closed interval `[start_date, end_date]`, matching + the semantics documented on `spec.DateRange`: both bounds are inclusive. An + omitted bound leaves that side open, so `DateRange()` selects every date. The + union of all the given ranges is taken. + + Note: Every bound that is present must be one of `dates`. + + Args: + date_ranges: The date ranges to compile. + dates: The date coordinates to compile against, in order. + spec_name: The `ModelSpec` attribute being compiled, used in error messages. + + Returns: + A boolean array of shape `(len(dates),)`, `True` wherever the date falls + inside at least one of `date_ranges`. + + Raises: + ValueError: If a bound is not one of `dates`. + """ + known_dates = set(dates) + mask = np.zeros(len(dates), dtype=bool) + for date_range in date_ranges: + # `DateRange.__post_init__` already normalizes these, but the declared + # attribute type stays polymorphic (`Date | None`), so normalize again to + # compare `date` against `date` rather than against `str`. + start = ( + tc.normalize_date(date_range.start_date) + if date_range.start_date is not None + else None + ) + end = ( + tc.normalize_date(date_range.end_date) + if date_range.end_date is not None + else None + ) + for bound_name, bound in (("start_date", start), ("end_date", end)): + if bound is not None and bound not in known_dates: + raise ValueError( + f"`{spec_name}` has a `DateRange` whose `{bound_name}` ({bound}) is" + " not one of the input data's time coordinates. Date range bounds" + " must name an exact time coordinate." + + _nearest_coordinates_hint(bound, dates) + ) + mask |= np.array( + [ + (start is None or date >= start) and (end is None or date <= end) + for date in dates + ], + dtype=bool, + ) + return mask + + +def _resolve_name_indices( + names: Sequence[str], + universe: Sequence[str], + *, + spec_name: str, + dim_name: str, +) -> list[int]: + """Resolves coordinate names to their positional indices. + + Args: + names: The names to resolve. + universe: The ordered coordinate values to resolve against. + spec_name: The `ModelSpec` attribute being compiled, used in error messages. + dim_name: The input data dimension being resolved against, used in error + messages. + + Returns: + The index of each name in `universe`, in the order given. + + Raises: + ValueError: If any name is absent from `universe`. + """ + index_of = {name: index for index, name in enumerate(universe)} + unknown = [name for name in names if name not in index_of] + if unknown: + raise ValueError( + f"`{spec_name}` refers to {dim_name} that are not in the input data:" + f" {sorted(unknown)}. Available {dim_name}: {sorted(universe)}." + ) + return [index_of[name] for name in names] + + +def _compile_name_mask( + names: Sequence[str], + universe: Sequence[str], + *, + spec_name: str, + dim_name: str, +) -> np.ndarray: + """Compiles a selection of coordinate names into a boolean mask. + + Args: + names: The selected names. + universe: The ordered coordinate values to compile against. + spec_name: The `ModelSpec` attribute being compiled, used in error messages. + dim_name: The input data dimension being resolved against, used in error + messages. + + Returns: + A boolean array of shape `(len(universe),)`, `True` at the selected names. + + Raises: + ValueError: If any name is absent from `universe`. + """ + mask = np.zeros(len(universe), dtype=bool) + mask[ + _resolve_name_indices( + names, universe, spec_name=spec_name, dim_name=dim_name + ) + ] = True + return mask + + +def _compile_calibration_spec( + calibration: spec.CalibrationSpec, + dates: Sequence[datetime.date], + channels: Sequence[str], + *, + spec_name: str, + dim_name: str, +) -> np.ndarray: + """Compiles a `CalibrationSpec` into a boolean calibration period array. + + Args: + calibration: The declarative calibration specification. + dates: The media time coordinates to compile against, in order. + channels: The channel coordinates to compile against, in order. + spec_name: The `ModelSpec` attribute being compiled, used in error messages. + dim_name: The channel dimension being resolved against, used in error + messages. + + Returns: + A boolean array of shape `(len(dates), len(channels))`. + + Raises: + ValueError: If a channel name is absent from `channels`, or a date range + bound is not one of `dates`. + """ + entries = calibration.spec + n_channels = len(channels) + + # `CalibrationSpec.__post_init__` guarantees the sequence is homogeneous and + # non-empty, so the first element determines the scope of the whole spec. + # The type checker cannot carry that guarantee across the sequence, hence the + # casts. + if isinstance(entries[0], spec.DateRange): + global_mask = _compile_date_range_mask( + cast(Sequence[spec.DateRange], entries), dates, spec_name=spec_name + ) + return np.tile(global_mask[:, np.newaxis], (1, n_channels)) + + # Per-channel scope. A channel that no entry mentions is left *unrestricted* + # (all `True`), matching the documented meaning of an unset + # `roi_calibration_period`: "If `None`, all times are used." Defaulting an + # unmentioned channel to all `False` would instead zero out its aggregated + # spend, making the denominator of its ROI prior zero. + compiled = np.zeros((len(dates), n_channels), dtype=bool) + is_mentioned = np.zeros(n_channels, dtype=bool) + for entry in cast(Sequence[spec.ChannelCalibrationSpec], entries): + mask = _compile_date_range_mask( + entry.date_ranges, dates, spec_name=spec_name + ) + for index in _resolve_name_indices( + entry.channels, channels, spec_name=spec_name, dim_name=dim_name + ): + compiled[:, index] |= mask + is_mentioned[index] = True + compiled[:, ~is_mentioned] = True + return compiled + + +def _compile_geo_holdout_specs( + geo_specs: Sequence[spec.GeoHoldoutSpec], + dates: Sequence[datetime.date], + geos: Sequence[str], +) -> np.ndarray: + """Compiles per-geo holdout specs into a boolean holdout mask. + + Unlike calibration, a geo that no entry mentions is simply not held out, so + it compiles to all `False`. + + Args: + geo_specs: The per-geo holdout specifications. + dates: The time coordinates to compile against, in order. + geos: The geo coordinates to compile against, in order. + + Returns: + A boolean array of shape `(len(geos), len(dates))`. + + Raises: + ValueError: If a geo name is absent from `geos`, or a date range bound is + not one of `dates`. + """ + compiled = np.zeros((len(geos), len(dates)), dtype=bool) + for geo_spec in geo_specs: + mask = _compile_date_range_mask( + geo_spec.date_ranges, dates, spec_name="holdout" + ) + for index in _resolve_name_indices( + geo_spec.geos, geos, spec_name="holdout", dim_name="geos" + ): + compiled[index, :] |= mask + return compiled + + +def _draw_random_holdout( + random_spec: spec.RandomHoldoutSpec, + n_geos: int, + n_times: int, +) -> np.ndarray: + """Draws a random holdout mask, stratified by geo. + + Each geo independently holds out exactly `round(ratio * n_times)` time + periods, sampled without replacement. Stratifying by geo -- rather than + drawing over the flattened `n_geos * n_times` cell space -- guarantees that + every geo contributes both training and test rows. + + Note that random samples are drawn sequentially in the order geos appear in + the input data coordinates; permuting geo order will produce different samples + for the same seed. + + Args: + random_spec: The random holdout specification. + n_geos: The number of geos. + n_times: The number of time periods. + + Returns: + A boolean array of shape `(n_geos, n_times)`. + """ + rng = np.random.default_rng(random_spec.seed) + n_holdout = min(int(round(random_spec.ratio * n_times)), n_times) + mask = np.zeros((n_geos, n_times), dtype=bool) + if n_holdout > 0: + for geo_index in range(n_geos): + mask[geo_index, rng.choice(n_times, size=n_holdout, replace=False)] = True + return mask + + @dataclasses.dataclass(frozen=True) class SaturationSpec: """Specification for each channel's saturation function. @@ -110,6 +372,11 @@ def __init__( self._validate_media_spend_for_paid_channels() self._validate_rf_spend_for_paid_channels() + # TODO: Deduplicate with `_validate_model_spec_shapes`. Both + # methods run from `__init__` and validate the same legacy `ModelSpec` + # arrays (`roi_calibration_period`, `rf_roi_calibration_period`, + # `holdout_id`, `control_population_scaling_id`) against the same shapes, + # with near-identical error messages. def _validate_data_dependent_model_spec(self): """Validates that the data dependent model specs have correct shapes.""" @@ -183,6 +450,7 @@ def _validate_data_dependent_model_spec(self): f" ({self.n_non_media_channels},)`." ) + # TODO: Deduplicate with `_validate_data_dependent_model_spec`. def _validate_model_spec_shapes(self): """Validate shapes of model_spec attributes.""" if self._model_spec.roi_calibration_period is not None: @@ -696,6 +964,238 @@ def holdout_id(self) -> backend.Tensor | None: tensor = backend.to_tensor(self._model_spec.holdout_id, dtype=backend.bool_) return tensor[backend.newaxis, ...] if self.is_national else tensor + # -------------------------------------------------------------------------- + # Compiled model spec properties. + # + # These resolve `ModelSpec`'s declarative attributes against this context's + # `InputData` coordinates, producing the positional boolean arrays the model + # engine consumes. They are the single place where a channel name becomes a + # column index and a date range becomes a row mask. + # + # Precedence is *legacy first*: when both a declarative attribute and its + # deprecated array counterpart are set, the deprecated array wins. This + # matches the contract `ModelSpec.__post_init__` currently advertises in its + # conflict warning (" takes precedence for backward compatibility"). + # -------------------------------------------------------------------------- + + def _coordinate_names(self, coordinate: Any) -> list[str]: + """Returns an input data coordinate's values as a list of strings.""" + if coordinate is None: + return [] + return [str(value) for value in coordinate.values] + + @functools.cached_property + def compiled_roi_calibration_period(self) -> np.ndarray | None: + """The effective ROI calibration period for media channels. + + Resolved from the declarative `ModelSpec.roi_calibration` against the input + data's media time and media channel coordinates, or taken as-is from the + deprecated `ModelSpec.roi_calibration_period`. + + Returns: + A boolean array of shape `(n_media_times, n_media_channels)`, or `None` + if neither attribute is set. + + Raises: + ValueError: If the spec names a media channel not in the input data, or + a date range bound that is not a time coordinate. + """ + if self._model_spec.roi_calibration_period is not None: + return self._model_spec.roi_calibration_period + if self._model_spec.roi_calibration is None: + return None + return _compile_calibration_spec( + self._model_spec.roi_calibration, + self._input_data.media_time_coordinates.all_dates, + self._coordinate_names(self._input_data.media_channel), + spec_name="roi_calibration", + dim_name="media channels", + ) + + @functools.cached_property + def compiled_rf_roi_calibration_period(self) -> np.ndarray | None: + """The effective ROI calibration period for reach & frequency channels. + + Resolved from the declarative `ModelSpec.rf_roi_calibration` against the + input data's media time and RF channel coordinates, or taken as-is from the + deprecated `ModelSpec.rf_roi_calibration_period`. + + Returns: + A boolean array of shape `(n_media_times, n_rf_channels)`, or `None` if + neither attribute is set. + + Raises: + ValueError: If the spec names an RF channel not in the input data, or a + date range bound that is not a time coordinate. + """ + if self._model_spec.rf_roi_calibration_period is not None: + return self._model_spec.rf_roi_calibration_period + if self._model_spec.rf_roi_calibration is None: + return None + return _compile_calibration_spec( + self._model_spec.rf_roi_calibration, + self._input_data.media_time_coordinates.all_dates, + self._coordinate_names(self._input_data.rf_channel), + spec_name="rf_roi_calibration", + dim_name="RF channels", + ) + + @functools.cached_property + def compiled_holdout_id(self) -> np.ndarray | None: + """The effective holdout mask. + + Resolved from the declarative `ModelSpec.holdout` against the input data's + time and geo coordinates, or taken as-is from the deprecated + `ModelSpec.holdout_id`. + + For a declarative holdout, a `resolved` draw always wins and is never + re-drawn; see `spec.RandomHoldoutSpec` for why a seed alone cannot + reproduce a draw. Only when a `RandomHoldoutSpec` carries no `resolved` + draw is one made here, once, and memoized for the lifetime of this context. + + Returns: + A boolean array of shape `(n_times,)` for a national model or + `(n_geos, n_times)` otherwise -- the same convention as the deprecated + `ModelSpec.holdout_id` -- or `None` if neither attribute is set. + + Raises: + ValueError: If the spec names a geo not in the input data, or a date + range bound that is not a time coordinate. + """ + if self._model_spec.holdout_id is not None: + return self._model_spec.holdout_id + holdout = self._model_spec.holdout + if holdout is None: + return None + + dates = self._input_data.time_coordinates.all_dates + geos = self._coordinate_names(self._input_data.geo) + + if holdout.resolved is not None: + compiled = _compile_geo_holdout_specs(holdout.resolved, dates, geos) + elif isinstance(holdout.spec, spec.RandomHoldoutSpec): + compiled = _draw_random_holdout(holdout.spec, len(geos), len(dates)) + elif isinstance(holdout.spec[0], spec.DateRange): + # A global holdout applies the same date mask to every geo. + global_mask = _compile_date_range_mask( + cast(Sequence[spec.DateRange], holdout.spec), + dates, + spec_name="holdout", + ) + compiled = np.tile(global_mask[np.newaxis, :], (len(geos), 1)) + else: + compiled = _compile_geo_holdout_specs( + cast(Sequence[spec.GeoHoldoutSpec], holdout.spec), dates, geos + ) + + # National models carry a 1-D holdout, matching the legacy convention that + # `_validate_model_spec_shapes` enforces. + return compiled[0] if self.is_national else compiled + + @functools.cached_property + def compiled_control_population_scaling_id(self) -> np.ndarray | None: + """The effective population-scaling selection for control variables. + + Resolved from the declarative `ModelSpec.population_scaled_controls` + against the input data's control variable coordinates, or taken as-is from + the deprecated `ModelSpec.control_population_scaling_id`. + + Returns: + A boolean array of shape `(n_controls,)`, or `None` if neither attribute + is set. + + Raises: + ValueError: If the spec names a control variable not in the input data. + """ + if self._model_spec.control_population_scaling_id is not None: + return self._model_spec.control_population_scaling_id + if self._model_spec.population_scaled_controls is None: + return None + return _compile_name_mask( + self._model_spec.population_scaled_controls, + self._coordinate_names(self._input_data.control_variable), + spec_name="population_scaled_controls", + dim_name="control variables", + ) + + @functools.cached_property + def compiled_non_media_population_scaling_id(self) -> np.ndarray | None: + """The effective population-scaling selection for non-media channels. + + Resolved from the declarative + `ModelSpec.population_scaled_non_media_channels` against the input data's + non-media channel coordinates, or taken as-is from the deprecated + `ModelSpec.non_media_population_scaling_id`. + + Returns: + A boolean array of shape `(n_non_media_channels,)`, or `None` if neither + attribute is set. + + Raises: + ValueError: If the spec names a non-media channel not in the input data. + """ + if self._model_spec.non_media_population_scaling_id is not None: + return self._model_spec.non_media_population_scaling_id + if self._model_spec.population_scaled_non_media_channels is None: + return None + return _compile_name_mask( + self._model_spec.population_scaled_non_media_channels, + self._coordinate_names(self._input_data.non_media_channel), + spec_name="population_scaled_non_media_channels", + dim_name="non-media channels", + ) + + def resolve_non_media_baseline_values( + self, + values: Mapping[str, float | str] | Sequence[float | str] | None, + ) -> list[float | str] | None: + """Resolves non-media baseline values into positional channel order. + + A mapping only needs to name the channels whose baseline differs from the + default; any channel it omits falls back to `'min'`. + + Args: + values: A mapping from non-media channel name to baseline value, a + sequence already in channel order, or `None`. + + Returns: + A list of length `n_non_media_channels` in channel order, or `None` if + `values` is `None`. + + Raises: + ValueError: If a mapping key is not a known non-media channel. + """ + if values is None: + return None + if not isinstance(values, Mapping): + return list(values) + channels = self._coordinate_names(self._input_data.non_media_channel) + _resolve_name_indices( + list(values.keys()), + channels, + spec_name="non_media_baseline_values", + dim_name="non-media channels", + ) + return [ + values.get(channel, constants.NON_MEDIA_BASELINE_MIN) + for channel in channels + ] + + @functools.cached_property + def compiled_non_media_baseline_values(self) -> list[float | str] | None: + """`ModelSpec.non_media_baseline_values`, in positional channel order. + + Returns: + A list of length `n_non_media_channels`, or `None` if the attribute is + unset. + + Raises: + ValueError: If the attribute is a mapping naming an unknown channel. + """ + return self.resolve_non_media_baseline_values( + self._model_spec.non_media_baseline_values + ) + def _warn_setting_ignored_priors(self): """Raises a warning if ignored priors are set.""" default_distribution = prior_distribution.PriorDistribution() diff --git a/meridian/model/context_test.py b/meridian/model/context_test.py index cc48a3800..adc2a9d78 100644 --- a/meridian/model/context_test.py +++ b/meridian/model/context_test.py @@ -13,6 +13,7 @@ # limitations under the License. from collections.abc import Collection, Mapping, Sequence +import datetime import types from typing import Any from unittest import mock @@ -2173,5 +2174,673 @@ def test_get_channel_parameter_tensor_beta_g_success(self): test_utils.assert_allclose(tensor, expected) +class CompiledModelSpecTest( + test_utils.MeridianTestCase, + model_test_data.WithInputDataSamples, +): + """Tests for `ModelContext`'s compiled model spec properties.""" + + input_data_samples = model_test_data.WithInputDataSamples + + @classmethod + def setUpClass(cls): + super().setUpClass() + model_test_data.WithInputDataSamples.setup() + + def _context( + self, + data: input_data.InputData, + model_spec: spec.ModelSpec, + ) -> context.ModelContext: + return context.ModelContext(input_data=data, model_spec=model_spec) + + # --- Invariants the compilation logic depends on ---------------------------- + + # `_compile_calibration_spec` and `compiled_holdout_id` decide how to + # interpret an entire sequence by inspecting only its first element. That is + # sound only because the spec dataclasses reject mixed-scope and empty + # sequences at construction. These tests pin those guarantees next to the + # code that relies on them, so relaxing the validation cannot quietly become + # a mis-compilation. + # + # Static typing already rejects a *literal* mixed list, so each case below + # builds the sequence as `list[Any]`. That is deliberate: it reproduces the + # dynamically-built sequence the runtime guard actually exists to catch. + + def test_calibration_spec_rejects_mixed_scopes(self): + mixed: list[Any] = [ + spec.DateRange("2021-01-25", "2021-02-01"), + spec.ChannelCalibrationSpec( + channels=["ch_1"], + date_ranges=[spec.DateRange("2021-01-25", "2021-02-01")], + ), + ] + with self.assertRaisesRegex(ValueError, "the two cannot be mixed"): + spec.CalibrationSpec(spec=mixed) + + def test_holdout_spec_rejects_mixed_scopes(self): + mixed: list[Any] = [ + spec.DateRange("2021-01-25", "2021-02-01"), + spec.GeoHoldoutSpec( + geos=["geo_0"], + date_ranges=[spec.DateRange("2021-01-25", "2021-02-01")], + ), + ] + with self.assertRaisesRegex(ValueError, "the two cannot be mixed"): + spec.HoldoutSpec(spec=mixed) + + def test_calibration_spec_rejects_empty_sequence(self): + """An empty sequence would make the `entries[0]` scope probe raise.""" + with self.assertRaisesRegex(ValueError, "cannot be empty"): + spec.CalibrationSpec(spec=[]) + + def test_holdout_spec_rejects_empty_sequence(self): + with self.assertRaisesRegex(ValueError, "cannot be empty"): + spec.HoldoutSpec(spec=[]) + + # --- Date range bound validation ------------------------------------------- + + # A `DateRange` bound must name an exact time coordinate. A bound landing + # between two coordinates has no well-defined period boundary, so it has no + # faithful half-open `DateInterval` representation. + + def test_compiled_roi_calibration_rejects_start_date_off_coordinate(self): + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + off_coordinate = dates[10] + datetime.timedelta(days=1) + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[spec.DateRange(off_coordinate, dates[20])] + ), + ) + with self.assertRaisesRegex( + ValueError, + "`roi_calibration` has a `DateRange` whose `start_date`" + f" \\({off_coordinate}\\) is not one of the input data's time" + " coordinates", + ): + _ = self._context(data, model_spec).compiled_roi_calibration_period + + def test_compiled_roi_calibration_rejects_end_date_off_coordinate(self): + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + off_coordinate = dates[20] + datetime.timedelta(days=1) + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[spec.DateRange(dates[10], off_coordinate)] + ), + ) + with self.assertRaisesRegex(ValueError, "`end_date`"): + _ = self._context(data, model_spec).compiled_roi_calibration_period + + def test_off_coordinate_bound_error_names_the_bracketing_coordinates(self): + """The message points at the two coordinates the bad bound falls between.""" + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + off_coordinate = dates[10] + datetime.timedelta(days=1) + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[spec.DateRange(off_coordinate, dates[20])] + ), + ) + with self.assertRaisesRegex( + ValueError, f"nearest are {dates[10]} and {dates[11]}" + ): + _ = self._context(data, model_spec).compiled_roi_calibration_period + + def test_compiled_holdout_rejects_off_coordinate_bound(self): + """The holdout path validates bounds too, not just calibration.""" + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + off_coordinate = dates[10] + datetime.timedelta(days=1) + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=[spec.DateRange(off_coordinate, dates[20])] + ) + ) + with self.assertRaisesRegex(ValueError, "`holdout` has a `DateRange`"): + _ = self._context(data, model_spec).compiled_holdout_id + + def test_compiled_per_geo_holdout_rejects_off_coordinate_bound(self): + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + off_coordinate = dates[10] + datetime.timedelta(days=1) + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=[ + spec.GeoHoldoutSpec( + geos=["geo_0"], + date_ranges=[spec.DateRange(off_coordinate, dates[20])], + ) + ] + ) + ) + with self.assertRaisesRegex(ValueError, "`holdout` has a `DateRange`"): + _ = self._context(data, model_spec).compiled_holdout_id + + def test_omitted_date_range_bounds_are_not_validated(self): + """An open bound is not a coordinate, and must stay legal.""" + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec(spec=[spec.DateRange()]), + ) + compiled = self._context(data, model_spec).compiled_roi_calibration_period + + assert compiled is not None + np.testing.assert_array_equal( + compiled, np.ones((len(dates), 3), dtype=bool) + ) + + # --- ROI calibration ------------------------------------------------------ + + def test_compiled_roi_calibration_period_unset_is_none(self): + model_context = self._context( + self.input_data_with_media_and_rf, spec.ModelSpec() + ) + self.assertIsNone(model_context.compiled_roi_calibration_period) + + def test_compiled_roi_calibration_period_passes_through_legacy_array(self): + data = self.input_data_with_media_and_rf + legacy = np.zeros((len(data.media_time), 3), dtype=bool) + legacy[5:10, 1] = True + model_context = self._context( + data, spec.ModelSpec(roi_calibration_period=legacy) + ) + np.testing.assert_array_equal( + model_context.compiled_roi_calibration_period, legacy + ) + + def test_compiled_roi_calibration_period_prefers_legacy_over_declarative( + self, + ): + """The deprecated array wins, matching what `ModelSpec` warns it will do.""" + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + legacy = np.zeros((len(dates), 3), dtype=bool) + legacy[0, 0] = True + with self.assertWarns(UserWarning): + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[spec.DateRange(dates[10], dates[20])] + ), + roi_calibration_period=legacy, + ) + model_context = self._context(data, model_spec) + np.testing.assert_array_equal( + model_context.compiled_roi_calibration_period, legacy + ) + + def test_compiled_roi_calibration_period_global_date_ranges(self): + """A global spec applies one date mask to every channel.""" + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[spec.DateRange(dates[10], dates[20])] + ), + ) + compiled = self._context(data, model_spec).compiled_roi_calibration_period + + assert compiled is not None + self.assertEqual(compiled.shape, (len(dates), 3)) + expected = np.zeros(len(dates), dtype=bool) + # `[start, end]`: both index 10 and index 20 are included. + expected[10:21] = True + for channel in range(3): + np.testing.assert_array_equal(compiled[:, channel], expected) + + def test_compiled_roi_calibration_period_unions_multiple_date_ranges(self): + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[ + spec.DateRange(dates[10], dates[20]), + spec.DateRange(dates[30], dates[35]), + ] + ), + ) + compiled = self._context(data, model_spec).compiled_roi_calibration_period + + assert compiled is not None + expected = np.zeros(len(dates), dtype=bool) + expected[10:21] = True + expected[30:36] = True + np.testing.assert_array_equal(compiled[:, 0], expected) + + @parameterized.named_parameters( + dict(testcase_name="both_bounds_open", start=None, end=None), + dict(testcase_name="open_start", start=None, end=10), + dict(testcase_name="open_end", start=10, end=None), + ) + def test_compiled_roi_calibration_period_open_bounds( + self, start: int | None, end: int | None + ): + """An omitted bound is resolved against the data, not left empty.""" + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[ + spec.DateRange( + dates[start] if start is not None else None, + dates[end] if end is not None else None, + ) + ] + ), + ) + compiled = self._context(data, model_spec).compiled_roi_calibration_period + + assert compiled is not None + expected = np.zeros(len(dates), dtype=bool) + # `end` is inclusive, but Python slicing is not, hence the `+ 1`. + expected[slice(start, end + 1 if end is not None else None)] = True + np.testing.assert_array_equal(compiled[:, 0], expected) + + def test_compiled_roi_calibration_period_per_channel(self): + """Channels the spec does not name are left unrestricted, not zeroed.""" + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + assert data.media_channel is not None + channels = list(data.media_channel.values) + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[ + spec.ChannelCalibrationSpec( + channels=[channels[1]], + date_ranges=[spec.DateRange(dates[10], dates[20])], + ) + ] + ), + ) + compiled = self._context(data, model_spec).compiled_roi_calibration_period + + assert compiled is not None + expected_named = np.zeros(len(dates), dtype=bool) + expected_named[10:21] = True + np.testing.assert_array_equal(compiled[:, 1], expected_named) + + # An all-`False` column here would zero the channel's aggregated spend and + # make its ROI prior denominator zero, so unnamed channels must be `True`. + unrestricted = np.ones(len(dates), dtype=bool) + np.testing.assert_array_equal(compiled[:, 0], unrestricted) + np.testing.assert_array_equal(compiled[:, 2], unrestricted) + + def test_compiled_roi_calibration_period_unknown_channel_fails(self): + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + media_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + roi_calibration=spec.CalibrationSpec( + spec=[ + spec.ChannelCalibrationSpec( + channels=["not_a_channel"], + date_ranges=[spec.DateRange(dates[0], dates[5])], + ) + ] + ), + ) + model_context = self._context(data, model_spec) + with self.assertRaisesRegex( + ValueError, + r"`roi_calibration` refers to media channels that are not in the input" + r" data: \['not_a_channel'\]", + ): + _ = model_context.compiled_roi_calibration_period + + # --- RF ROI calibration --------------------------------------------------- + + def test_compiled_rf_roi_calibration_period_global_date_ranges(self): + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + rf_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + rf_roi_calibration=spec.CalibrationSpec( + spec=[spec.DateRange(dates[10], dates[20])] + ), + ) + compiled = self._context( + data, model_spec + ).compiled_rf_roi_calibration_period + + assert compiled is not None + self.assertEqual(compiled.shape, (len(dates), 2)) + expected = np.zeros(len(dates), dtype=bool) + expected[10:21] = True + np.testing.assert_array_equal(compiled[:, 0], expected) + + def test_compiled_rf_roi_calibration_period_unknown_channel_fails(self): + data = self.input_data_with_media_and_rf + dates = data.media_time_coordinates.all_dates + model_spec = spec.ModelSpec( + rf_prior_type=constants.TREATMENT_PRIOR_TYPE_ROI, + rf_roi_calibration=spec.CalibrationSpec( + spec=[ + spec.ChannelCalibrationSpec( + channels=["not_an_rf_channel"], + date_ranges=[spec.DateRange(dates[0], dates[5])], + ) + ] + ), + ) + model_context = self._context(data, model_spec) + with self.assertRaisesRegex( + ValueError, + r"`rf_roi_calibration` refers to RF channels that are not in the input" + r" data: \['not_an_rf_channel'\]", + ): + _ = model_context.compiled_rf_roi_calibration_period + + # --- Holdout -------------------------------------------------------------- + + def test_compiled_holdout_id_unset_is_none(self): + model_context = self._context( + self.input_data_with_media_and_rf, spec.ModelSpec() + ) + self.assertIsNone(model_context.compiled_holdout_id) + + def test_compiled_holdout_id_passes_through_legacy_array(self): + data = self.input_data_with_media_and_rf + legacy = np.zeros((len(data.geo), len(data.time)), dtype=bool) + legacy[0, :5] = True + model_context = self._context(data, spec.ModelSpec(holdout_id=legacy)) + np.testing.assert_array_equal(model_context.compiled_holdout_id, legacy) + + def test_compiled_holdout_id_prefers_legacy_over_declarative(self): + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + legacy = np.zeros((len(data.geo), len(dates)), dtype=bool) + legacy[0, :5] = True + with self.assertWarns(UserWarning): + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec(spec=[spec.DateRange(dates[10], dates[20])]), + holdout_id=legacy, + ) + model_context = self._context(data, model_spec) + np.testing.assert_array_equal(model_context.compiled_holdout_id, legacy) + + def test_compiled_holdout_id_global_date_ranges(self): + """A global holdout applies the same date mask to every geo.""" + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + n_geos = len(data.geo) + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec(spec=[spec.DateRange(dates[10], dates[20])]) + ) + compiled = self._context(data, model_spec).compiled_holdout_id + + assert compiled is not None + self.assertEqual(compiled.shape, (n_geos, len(dates))) + expected_row = np.zeros(len(dates), dtype=bool) + expected_row[10:21] = True + for geo in range(n_geos): + np.testing.assert_array_equal(compiled[geo], expected_row) + + def test_compiled_holdout_id_per_geo(self): + """Geos the spec does not name are simply not held out.""" + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + geos = [str(geo) for geo in data.geo.values] + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=[ + spec.GeoHoldoutSpec( + geos=[geos[1]], + date_ranges=[spec.DateRange(dates[10], dates[20])], + ) + ] + ) + ) + compiled = self._context(data, model_spec).compiled_holdout_id + + assert compiled is not None + expected = np.zeros((len(geos), len(dates)), dtype=bool) + expected[1, 10:21] = True + np.testing.assert_array_equal(compiled, expected) + + def test_compiled_holdout_id_unknown_geo_fails(self): + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=[ + spec.GeoHoldoutSpec( + geos=["not_a_geo"], + date_ranges=[spec.DateRange(dates[0], dates[5])], + ) + ] + ) + ) + model_context = self._context(data, model_spec) + with self.assertRaisesRegex( + ValueError, + r"`holdout` refers to geos that are not in the input data:" + r" \['not_a_geo'\]", + ): + _ = model_context.compiled_holdout_id + + def test_compiled_holdout_id_random_holds_out_exact_ratio_per_geo(self): + """The draw is stratified: every geo holds out the same exact count.""" + data = self.input_data_with_media_and_rf + n_times = len(data.time) + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=spec.RandomHoldoutSpec(ratio=0.2, seed=17) + ) + ) + compiled = self._context(data, model_spec).compiled_holdout_id + + assert compiled is not None + self.assertEqual(compiled.shape, (len(data.geo), n_times)) + expected_per_geo = round(0.2 * n_times) + np.testing.assert_array_equal( + compiled.sum(axis=1), + np.full(len(data.geo), expected_per_geo), + ) + self.assertEqual(compiled.sum(), expected_per_geo * len(data.geo)) + + def test_compiled_holdout_id_random_draws_differ_across_geos(self): + """Stratification must not degenerate into the same draw for every geo.""" + data = self.input_data_with_media_and_rf + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=spec.RandomHoldoutSpec(ratio=0.2, seed=17) + ) + ) + compiled = self._context(data, model_spec).compiled_holdout_id + + assert compiled is not None + distinct_rows = {row.tobytes() for row in compiled} + self.assertLen(distinct_rows, len(data.geo)) + + def test_compiled_holdout_id_random_is_reproducible_for_a_seed(self): + data = self.input_data_with_media_and_rf + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=spec.RandomHoldoutSpec(ratio=0.3, seed=99) + ) + ) + first = self._context(data, model_spec).compiled_holdout_id + second = self._context(data, model_spec).compiled_holdout_id + np.testing.assert_array_equal(first, second) + + def test_compiled_holdout_id_random_is_drawn_only_once(self): + """Without a seed the draw is random, so memoization must hold it fixed.""" + data = self.input_data_with_media_and_rf + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec(spec=spec.RandomHoldoutSpec(ratio=0.5)) + ) + model_context = self._context(data, model_spec) + np.testing.assert_array_equal( + model_context.compiled_holdout_id, model_context.compiled_holdout_id + ) + + def test_compiled_holdout_id_resolved_wins_over_random_spec(self): + """A materialized draw is authoritative and is never re-drawn.""" + data = self.input_data_with_media_and_rf + dates = data.time_coordinates.all_dates + geos = [str(geo) for geo in data.geo.values] + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec( + spec=spec.RandomHoldoutSpec(ratio=0.5, seed=3), + resolved=[ + spec.GeoHoldoutSpec( + geos=[geos[0]], + date_ranges=[spec.DateRange(dates[0], dates[5])], + ) + ], + ) + ) + compiled = self._context(data, model_spec).compiled_holdout_id + + assert compiled is not None + expected = np.zeros((len(geos), len(dates)), dtype=bool) + expected[0, 0:6] = True + np.testing.assert_array_equal(compiled, expected) + # A 0.5 draw would have held out roughly half the grid, not 6 cells. + self.assertEqual(compiled.sum(), 6) + + def test_compiled_holdout_id_national_is_one_dimensional(self): + """National models keep the 1-D convention of the legacy attribute.""" + data = self.national_input_data_media_and_rf + dates = data.time_coordinates.all_dates + model_spec = spec.ModelSpec( + holdout=spec.HoldoutSpec(spec=[spec.DateRange(dates[10], dates[20])]) + ) + compiled = self._context(data, model_spec).compiled_holdout_id + + assert compiled is not None + self.assertEqual(compiled.shape, (len(dates),)) + expected = np.zeros(len(dates), dtype=bool) + expected[10:21] = True + np.testing.assert_array_equal(compiled, expected) + + # --- Population scaling --------------------------------------------------- + + def test_compiled_control_population_scaling_id_unset_is_none(self): + model_context = self._context( + self.input_data_with_media_and_rf, spec.ModelSpec() + ) + self.assertIsNone(model_context.compiled_control_population_scaling_id) + + def test_compiled_control_population_scaling_id_from_names(self): + data = self.input_data_with_media_and_rf + assert data.control_variable is not None + controls = [str(control) for control in data.control_variable.values] + model_spec = spec.ModelSpec(population_scaled_controls=[controls[1]]) + compiled = self._context( + data, model_spec + ).compiled_control_population_scaling_id + + expected = np.zeros(len(controls), dtype=bool) + expected[1] = True + np.testing.assert_array_equal(compiled, expected) + + def test_compiled_control_population_scaling_id_prefers_legacy(self): + data = self.input_data_with_media_and_rf + assert data.control_variable is not None + controls = [str(control) for control in data.control_variable.values] + legacy = np.zeros(len(controls), dtype=bool) + legacy[0] = True + with self.assertWarns(UserWarning): + model_spec = spec.ModelSpec( + population_scaled_controls=[controls[1]], + control_population_scaling_id=legacy, + ) + compiled = self._context( + data, model_spec + ).compiled_control_population_scaling_id + np.testing.assert_array_equal(compiled, legacy) + + def test_compiled_control_population_scaling_id_unknown_name_fails(self): + data = self.input_data_with_media_and_rf + model_spec = spec.ModelSpec(population_scaled_controls=["not_a_control"]) + model_context = self._context(data, model_spec) + with self.assertRaisesRegex( + ValueError, + r"`population_scaled_controls` refers to control variables that are not" + r" in the input data: \['not_a_control'\]", + ): + _ = model_context.compiled_control_population_scaling_id + + def test_compiled_non_media_population_scaling_id_from_names(self): + data = self.input_data_non_media_and_organic + assert data.non_media_channel is not None + channels = [str(channel) for channel in data.non_media_channel.values] + model_spec = spec.ModelSpec( + population_scaled_non_media_channels=[channels[0]] + ) + compiled = self._context( + data, model_spec + ).compiled_non_media_population_scaling_id + + expected = np.zeros(len(channels), dtype=bool) + expected[0] = True + np.testing.assert_array_equal(compiled, expected) + + def test_compiled_non_media_population_scaling_id_unknown_name_fails(self): + data = self.input_data_non_media_and_organic + model_spec = spec.ModelSpec( + population_scaled_non_media_channels=["not_a_channel"] + ) + model_context = self._context(data, model_spec) + with self.assertRaisesRegex( + ValueError, + r"`population_scaled_non_media_channels` refers to non-media channels" + r" that are not in the input data: \['not_a_channel'\]", + ): + _ = model_context.compiled_non_media_population_scaling_id + + # --- Non-media baseline values -------------------------------------------- + + def test_compiled_non_media_baseline_values_unset_is_none(self): + model_context = self._context( + self.input_data_non_media_and_organic, spec.ModelSpec() + ) + self.assertIsNone(model_context.compiled_non_media_baseline_values) + + def test_compiled_non_media_baseline_values_sequence_is_passed_through(self): + data = self.input_data_non_media_and_organic + model_spec = spec.ModelSpec(non_media_baseline_values=["max", 1.5]) + self.assertEqual( + self._context(data, model_spec).compiled_non_media_baseline_values, + ["max", 1.5], + ) + + def test_compiled_non_media_baseline_values_mapping_defaults_to_min(self): + """A mapping only names the channels that differ from the default.""" + data = self.input_data_non_media_and_organic + assert data.non_media_channel is not None + channels = [str(channel) for channel in data.non_media_channel.values] + model_spec = spec.ModelSpec(non_media_baseline_values={channels[1]: "max"}) + self.assertEqual( + self._context(data, model_spec).compiled_non_media_baseline_values, + [constants.NON_MEDIA_BASELINE_MIN, "max"], + ) + + def test_compiled_non_media_baseline_values_unknown_channel_fails(self): + data = self.input_data_non_media_and_organic + model_spec = spec.ModelSpec( + non_media_baseline_values={"not_a_channel": "max"} + ) + model_context = self._context(data, model_spec) + with self.assertRaisesRegex( + ValueError, + r"`non_media_baseline_values` refers to non-media channels that are not" + r" in the input data: \['not_a_channel'\]", + ): + _ = model_context.compiled_non_media_baseline_values + + if __name__ == "__main__": absltest.main() diff --git a/meridian/model/spec.py b/meridian/model/spec.py index 523359c55..1d2510c7b 100644 --- a/meridian/model/spec.py +++ b/meridian/model/spec.py @@ -64,6 +64,11 @@ class DateRange: A range whose bounds are equal is valid and selects that single date. + A bound that is present must name one of the input data's time coordinates + exactly; a bound falling between two coordinates is rejected. Because it + depends on the data, the check happens when the range is compiled against an + `InputData`, not at construction time. + Note: The `mmm.v1.common.DateInterval` proto that these specs serialize to uses the opposite, *half-open* `[start_date, end_date)` convention. The