Skip to content

Commit 91ac27b

Browse files
nmahadevan-googThe Meridian Authors
authored andcommitted
Add proto and serde support for per-channel media_effects_dist.
PiperOrigin-RevId: 938055551
1 parent 46145fe commit 91ac27b

5 files changed

Lines changed: 89 additions & 10 deletions

File tree

meridian/model/spec.py

Lines changed: 22 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -232,7 +232,9 @@ class ModelSpec:
232232
prior: prior_distribution.PriorDistribution = dataclasses.field(
233233
default_factory=prior_distribution.PriorDistribution,
234234
)
235-
media_effects_dist: str = constants.MEDIA_EFFECTS_LOG_NORMAL
235+
media_effects_dist: str | Mapping[str, str] = (
236+
constants.MEDIA_EFFECTS_LOG_NORMAL
237+
)
236238
hill_before_adstock: bool = False
237239
max_lag: int = 8
238240
unique_sigma_for_each_geo: bool = False
@@ -258,10 +260,26 @@ class ModelSpec:
258260

259261
def __post_init__(self):
260262
# Validate media_effects_dist.
261-
if self.media_effects_dist not in constants.MEDIA_EFFECTS_DISTRIBUTIONS:
263+
if isinstance(self.media_effects_dist, str):
264+
if self.media_effects_dist not in constants.MEDIA_EFFECTS_DISTRIBUTIONS:
265+
raise ValueError(
266+
"The `media_effects_dist` parameter"
267+
f" '{self.media_effects_dist}' must be one of"
268+
f" {sorted(constants.MEDIA_EFFECTS_DISTRIBUTIONS)}."
269+
)
270+
elif isinstance(self.media_effects_dist, Mapping):
271+
for channel, dist in self.media_effects_dist.items():
272+
if dist not in constants.MEDIA_EFFECTS_DISTRIBUTIONS:
273+
raise ValueError(
274+
f"The `media_effects_dist` for channel '{channel}'"
275+
" must be one of"
276+
f" {sorted(constants.MEDIA_EFFECTS_DISTRIBUTIONS)},"
277+
f" but got '{dist}'."
278+
)
279+
else:
262280
raise ValueError(
263-
f"The `media_effects_dist` parameter '{self.media_effects_dist}' must"
264-
f" be one of {sorted(constants.MEDIA_EFFECTS_DISTRIBUTIONS)}."
281+
"Unsupported type for `media_effects_dist` parameter:"
282+
f" {type(self.media_effects_dist)}."
265283
)
266284
# Support paid_media_prior_type for backwards compatibility.
267285
if self.paid_media_prior_type is not None:

meridian/model/spec_test.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -74,6 +74,34 @@ def test_spec_inits_invalid_media_effects_fails(self, dist, error_message):
7474
with self.assertRaisesWithLiteralMatch(ValueError, error_message):
7575
spec.ModelSpec(media_effects_dist=dist)
7676

77+
def test_spec_inits_valid_media_effects_mapping_works(self):
78+
dist_map = {"ch1": "normal", "ch2": "log_normal"}
79+
model_spec = spec.ModelSpec(media_effects_dist=dist_map)
80+
self.assertEqual(model_spec.media_effects_dist, dist_map)
81+
82+
@parameterized.named_parameters(
83+
(
84+
"invalid_mapping_value",
85+
{"ch1": "bad_dist"},
86+
ValueError,
87+
(
88+
"The `media_effects_dist` for channel 'ch1' must be one of"
89+
" ['log_normal', 'normal'], but got 'bad_dist'."
90+
),
91+
),
92+
(
93+
"invalid_type",
94+
123,
95+
ValueError,
96+
"Unsupported type for `media_effects_dist` parameter: <class 'int'>.",
97+
),
98+
)
99+
def test_spec_inits_invalid_media_effects_mapping_or_type_fails(
100+
self, dist, expected_error, error_message
101+
):
102+
with self.assertRaisesWithLiteralMatch(expected_error, error_message):
103+
spec.ModelSpec(media_effects_dist=dist)
104+
77105
@parameterized.named_parameters(
78106
("hill", constants.HILL),
79107
("none", "none"),

meridian/schema/serde/constants.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,8 @@
2525
GLOBAL_ADSTOCK_DECAY = 'global_adstock_decay'
2626
ADSTOCK_DECAY_BY_CHANNEL = 'adstock_decay_by_channel'
2727
DEFAULT_DECAY = 'geometric'
28+
MEDIA_EFFECTS_DIST_BY_CHANNEL = 'media_effects_dist_by_channel'
29+
2830

2931
SATURATION_SPEC = 'saturation_spec'
3032
GLOBAL_SATURATION = 'global_saturation'

meridian/schema/serde/hyperparameters.py

Lines changed: 26 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -87,9 +87,6 @@ class HyperparametersSerde(
8787
def serialize(self, obj: spec.ModelSpec) -> meridian_pb.Hyperparameters:
8888
"""Serializes the given ModelSpec into a `Hyperparameters` proto."""
8989
hyperparameters_proto = meridian_pb.Hyperparameters(
90-
media_effects_dist=media_effects_converter.to_proto(
91-
obj.media_effects_dist
92-
),
9390
hill_before_adstock=obj.hill_before_adstock,
9491
unique_sigma_for_each_geo=obj.unique_sigma_for_each_geo,
9592
media_prior_type=paid_media_prior_type_converter.to_proto(
@@ -112,6 +109,15 @@ def serialize(self, obj: spec.ModelSpec) -> meridian_pb.Hyperparameters:
112109
),
113110
enable_aks=obj.enable_aks,
114111
)
112+
if isinstance(obj.media_effects_dist, str):
113+
hyperparameters_proto.media_effects_dist = (
114+
media_effects_converter.to_proto(obj.media_effects_dist)
115+
)
116+
elif isinstance(obj.media_effects_dist, Mapping):
117+
for channel, dist in obj.media_effects_dist.items():
118+
hyperparameters_proto.media_effects_dist_by_channel.channel_media_effects_dists[
119+
channel
120+
] = media_effects_converter.to_proto(dist)
115121
if obj.max_lag is not None:
116122
hyperparameters_proto.max_lag = obj.max_lag
117123

@@ -290,10 +296,24 @@ def deserialize(
290296
else:
291297
saturation_spec = sc.DEFAULT_SATURATION
292298

299+
if (
300+
serialized.HasField(sc.MEDIA_EFFECTS_DIST_BY_CHANNEL)
301+
and serialized.media_effects_dist_by_channel.channel_media_effects_dists
302+
):
303+
channel_dists = (
304+
serialized.media_effects_dist_by_channel.channel_media_effects_dists
305+
)
306+
media_effects_dist = {
307+
channel: media_effects_converter.from_proto(dist)
308+
for channel, dist in channel_dists.items()
309+
}
310+
else:
311+
media_effects_dist = media_effects_converter.from_proto(
312+
serialized.media_effects_dist
313+
)
314+
293315
return spec.ModelSpec(
294-
media_effects_dist=media_effects_converter.from_proto(
295-
serialized.media_effects_dist
296-
),
316+
media_effects_dist=media_effects_dist,
297317
hill_before_adstock=serialized.hill_before_adstock,
298318
max_lag=max_lag,
299319
unique_sigma_for_each_geo=serialized.unique_sigma_for_each_geo,

proto/mmm/v1/model/meridian/meridian_model.proto

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -334,6 +334,13 @@ enum ComputationPrecision {
334334
FLOAT64 = 2;
335335
}
336336

337+
// Specifies the media effects distribution for each channel.
338+
message MediaEffectsDistByChannel {
339+
// A map where keys are channel names and values are the media effects
340+
// distribution to use for that channel.
341+
map<string, MediaEffectsDistribution> channel_media_effects_dists = 1;
342+
}
343+
337344
// Specifies the adstock decay function for each channel.
338345
message AdstockDecayByChannel {
339346
// A map where keys are channel names and values are the adstock decay
@@ -539,6 +546,10 @@ message Hyperparameters {
539546
// If `None`, the minimum value is used as baseline for each non-media
540547
// treatments channel. This attribute is used as the default value for the
541548
// corresponding argument to `Analyzer` methods.
549+
// Channel-specific media effects distributions. Defaults to 'log_normal' for
550+
// channels not specified in the map.
551+
MediaEffectsDistByChannel media_effects_dist_by_channel = 28;
552+
542553
repeated NonMediaBaselineValue non_media_baseline_values = 25;
543554
}
544555

0 commit comments

Comments
 (0)