merge_models_and_idata#

pymc_marketing.mmm.budget_optimizer.merge_models_and_idata(models, idatas, *, prefixes=None, merge_on='channel_data', use_every_n_draw=1)[source]#

Merge multiple PyMC models and their DataTree objects in one call.

Convenience wrapper that calls merge_models() and merge_inference_data() with the same prefixes and merge_on arguments, returning both results as a (model, idata) tuple ready for BudgetOptimizer.

Parameters:
modelslist of pm.Model

Optimization models, one per fitted MMM (e.g. from mmm.create_optimization_model(...)).

idataslist of xarray.DataTree

Posterior samples corresponding to each model (e.g. mmm.idata). Must be the same length as models.

prefixeslist of str or None, optional

Per-model prefix applied to every variable and dimension name that is not in merge_on. If None (default), prefixes are auto-generated as ["model1", "model2", ...].

merge_onstr or None, optional

Variable name that is shared across all models and therefore not prefixed in either the PyMC graph or the posterior. Defaults to "channel_data" which is the standard budget input node used by BudgetOptimizer. Pass None to prefix every variable.

use_every_n_drawint, optional

Thinning factor applied to each DataTree before merging. Keeps every n-th posterior draw. Defaults to 1 (no thinning).

Returns:
tuple of (pm.Model, xarray.DataTree)

(merged_model, merged_idata) ready to be passed directly to BudgetOptimizer.

Raises:
ValueError

If len(models) != len(idatas), or if models has fewer than 2 elements, or if prefixes length does not match.

Examples

Merge two fitted MMMs for a multi-region optimization:

from pymc_marketing.mmm import merge_models_and_idata, BudgetOptimizer

m1 = mmm_north.create_optimization_model("2025-01-01", "2025-03-31")
m2 = mmm_south.create_optimization_model("2025-01-01", "2025-03-31")

merged_model, merged_idata = merge_models_and_idata(
    models=[m1, m2],
    idatas=[mmm_north.idata, mmm_south.idata],
    prefixes=["north", "south"],
    merge_on="channel_data",
    use_every_n_draw=2,
)

optimizer = BudgetOptimizer(
    model=merged_model,
    idata=merged_idata,
    num_periods=13,
    # The model's date axis is carry-in + decisions + carry-over, each
    # flank effective_carryover_lags() wide.
    carry_in_periods=mmm_north.effective_carryover_lags(),
    adstock_periods=mmm_north.effective_carryover_lags(),
    response_variable="north_total_media_contribution_original_scale",
)
optimal, result = optimizer.allocate_budget(total_budget=100_000)