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
30 changes: 18 additions & 12 deletions meridian/model/knots.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,6 @@
from scipy import interpolate
from statsmodels.regression import linear_model


__all__ = [
'KnotInfo',
'get_knot_info',
Expand Down Expand Up @@ -260,7 +259,7 @@ def __init__(self, data: input_data.InputData):
def automatic_knot_selection(
self,
base_penalty: np.ndarray | None = None,
min_internal_knots: int = 1,
min_internal_knots: int | None = None,
max_internal_knots: int | None = None,
required_knots: Collection[int] | None = None,
excluded_knots: Collection[int] | None = None,
Expand All @@ -271,7 +270,9 @@ def automatic_knot_selection(
Args:
base_penalty: A vector of positive penalty values. The adaptive spline
regression is performed for every value of penalty.
min_internal_knots: The minimum number of internal knots. Defaults to 1.
min_internal_knots: The minimum number of internal knots. If None,
defaults to 0 for knot selection while ensuring initial degrees of
freedom are sufficient.
max_internal_knots: The maximum number of internal knots. If None, this
value is calculated as the number of initial knots minus the total count
of all treatment and control variables. Otherwise, the user-provided
Expand Down Expand Up @@ -334,8 +335,14 @@ def automatic_knot_selection(
raise ValueError('The same knot cannot be both required and excluded.')

knots = self._calculate_initial_knots(x, excluded_knots_arr)
min_internal_knots_for_validation = (
1 if min_internal_knots is None else min_internal_knots
)
max_internal_knots = self._calculate_and_validate_max_internal_knots(
knots, min_internal_knots, max_internal_knots
knots, min_internal_knots_for_validation, max_internal_knots
)
min_internal_knots_for_selection = (
0 if min_internal_knots is None else min_internal_knots
)
geo_scaling_factor = 1 / np.sqrt(len(self._data.geo))
penalty = geo_scaling_factor * base_penalty
Expand Down Expand Up @@ -376,18 +383,19 @@ def automatic_knot_selection(
)
).tolist()
if not any(
min_internal_knots <= k <= max_internal_knots
min_internal_knots_for_selection <= k <= max_internal_knots
for k in available_knots_lengths
):
raise ValueError(
f'The range [{min_internal_knots}, {max_internal_knots}] does not'
' contain any of the available knot lengths:'
f' {pprint.pformat(available_knots_lengths)}'
f'The range [{min_internal_knots_for_selection},'
f' {max_internal_knots}] does not contain any of the available knot'
f' lengths: {pprint.pformat(available_knots_lengths)}'
)

n_knots = np.array([len(x) for x in aspline[constants.KNOTS_SELECTED]])
feasible_idx = np.where(
(n_knots >= min_internal_knots) & (n_knots <= max_internal_knots)
(n_knots >= min_internal_knots_for_selection)
& (n_knots <= max_internal_knots)
)[0]
information_criterion = aspline[constants.AIC][feasible_idx]
knots_sel = [aspline[constants.KNOTS_SELECTED][i] for i in feasible_idx]
Expand Down Expand Up @@ -449,9 +457,7 @@ def _calculate_initial_knots(
)

if not np.all(np.isin(excluded_knots, knots)):
raise ValueError(
'The excluded knots are not legitimate knot locations.'
)
raise ValueError('The excluded knots are not legitimate knot locations.')
is_included = ~np.isin(knots, excluded_knots)
knots = knots[is_included]
return knots
Expand Down
17 changes: 17 additions & 0 deletions meridian/model/knots_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -786,6 +786,23 @@ def test_aks_returns_correct_knot_info(self):
actual_knot_info.weights, expected_knot_info.weights
)

def test_aks_no_internal_knots_selected_returns_empty(self):
data = test_utils.sample_input_data_from_dataset(
test_utils.random_dataset(
n_geos=3,
n_times=50,
n_media_times=50,
n_controls=2,
n_media_channels=2,
),
"non_revenue",
)
# A large base penalty forces all internal knots to be dropped.
huge_penalty = np.array([1e6, 1e7, 1e8])
aks_obj = knots.AKS(data)
result = aks_obj.automatic_knot_selection(base_penalty=huge_penalty)
self.assertEmpty(result.knots)

def test_aspline(self):
x = np.array([0, 1, 2, 3, 4, 5, 7, 8, 9, 10, 11, 12, 13, 14, 15])
y = np.array([10, 2, 3, 4, 50, 6, 30, 4, 5, 6, 3, 4, 5, 6, 6])
Expand Down
Loading