diff --git a/meridian/data/input_data.py b/meridian/data/input_data.py index 9a893aae3..7cb231744 100644 --- a/meridian/data/input_data.py +++ b/meridian/data/input_data.py @@ -292,6 +292,19 @@ def __post_init__(self): self._validate_times() self._validate_geos() self._validate_no_negative_values() + self._validate_frequencies() + + def _validate_frequencies(self) -> None: + """Validates that frequency values are at least 1.""" + for field, loggable_field in [ + (constants.FREQUENCY, "Frequency"), + (constants.ORGANIC_FREQUENCY, "Organic Frequency"), + ]: + da = getattr(self, field) + if da is not None and (da.values < 1).any(): + raise ValueError( + f"{loggable_field} values must be at least 1." + ) def _coerce_object_arrays_to_float(self): """Coerces object-typed DataArrays to float.""" @@ -619,9 +632,7 @@ def _validate_no_negative_values(self) -> None: constants.MEDIA_SPEND: "Media Spend", constants.RF_SPEND: "RF Spend", constants.REACH: "Reach", - constants.FREQUENCY: "Frequency", constants.ORGANIC_REACH: "Organic Reach", - constants.ORGANIC_FREQUENCY: "Organic Frequency", constants.REVENUE_PER_KPI: "Revenue per KPI", constants.KPI: "KPI", } diff --git a/meridian/data/input_data_test.py b/meridian/data/input_data_test.py index 015445f97..0863cc4c7 100644 --- a/meridian/data/input_data_test.py +++ b/meridian/data/input_data_test.py @@ -259,12 +259,6 @@ def test_validate_kpi_wrong_type(self): name="Reach", kpi_type=constants.REVENUE, ), - dict( - testcase_name="frequency", - field=constants.FREQUENCY, - name="Frequency", - kpi_type=constants.REVENUE, - ), dict( testcase_name="kpi", field=constants.KPI, @@ -278,6 +272,7 @@ def test_validate_kpi_wrong_type(self): kpi_type=constants.NON_REVENUE, ), ) + def test_validate_no_negative_values( self, field: str, name: str, kpi_type: str ): @@ -301,6 +296,24 @@ def test_validate_no_negative_values( frequency=maybe_flip(self.lagged_frequency, constants.FREQUENCY), ) + def test_validate_frequency_threshold(self): + with self.assertRaisesRegex( + ValueError, + expected_regex="Frequency values must be at least 1.", + ): + frequency = self.not_lagged_frequency.copy(deep=True) + frequency.values[0, 0, 0] = 0.5 + input_data.InputData( + controls=self.not_lagged_controls, + kpi=self.not_lagged_kpi, + kpi_type=constants.NON_REVENUE, + population=self.population, + revenue_per_kpi=self.revenue_per_kpi, + reach=self.not_lagged_reach, + frequency=frequency, + rf_spend=self.rf_spend, + ) + def test_validate_media_channels_duplicate_names(self): media = test_utils.random_media_da( n_geos=self.n_geos, @@ -2176,24 +2189,12 @@ def _verify_copy(self, original: input_data.InputData, deep: bool): name="Reach", kpi_type=constants.NON_REVENUE, ), - dict( - testcase_name="frequency", - field=constants.FREQUENCY, - name="Frequency", - kpi_type=constants.NON_REVENUE, - ), dict( testcase_name="organic_reach", field=constants.ORGANIC_REACH, name="Organic Reach", kpi_type=constants.NON_REVENUE, ), - dict( - testcase_name="organic_frequency", - field=constants.ORGANIC_FREQUENCY, - name="Organic Frequency", - kpi_type=constants.NON_REVENUE, - ), dict( testcase_name="kpi", field=constants.KPI, diff --git a/meridian/data/test_utils.py b/meridian/data/test_utils.py index 4bc2298ea..43d73ac3a 100644 --- a/meridian/data/test_utils.py +++ b/meridian/data/test_utils.py @@ -1142,6 +1142,7 @@ def random_frequency_da( abs(np.random.normal(3, 5, size=(n_geos, n_media_times, n_rf_channels))) + nonzero_shift ) + frequency = np.maximum(frequency, 1.0) channels = ( explicit_rf_channel_names