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
15 changes: 13 additions & 2 deletions meridian/data/input_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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",
}
Expand Down
37 changes: 19 additions & 18 deletions meridian/data/input_data_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
):
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
1 change: 1 addition & 0 deletions meridian/data/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading