Source code for causalpy.experiments.up_down_geolift

#   Copyright 2026 - 2026 The PyMC Labs Developers
#
#   Licensed under the Apache License, Version 2.0 (the "License");
#   you may not use this file except in compliance with the License.
#   You may obtain a copy of the License at
#
#       http://www.apache.org/licenses/LICENSE-2.0
#
#   Unless required by applicable law or agreed to in writing, software
#   distributed under the License is distributed on an "AS IS" BASIS,
#   WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
#   See the License for the specific language governing permissions and
#   limitations under the License.
"""Arm-aware geo lift analysis built on synthetic control."""

from collections.abc import Mapping, Sequence
from typing import Any, Literal

import numpy as np
import pandas as pd
import xarray as xr
from matplotlib import pyplot as plt
from matplotlib.ticker import StrMethodFormatter
from sklearn.base import RegressorMixin

from causalpy.constants import HDI_PROB
from causalpy.pymc_models import PyMCModel, SoftmaxWeightedSumFitter

from .synthetic_control import SyntheticControl


[docs] class UpDownGeoLift(SyntheticControl): """Estimate signed up and down geo effects against unchanged controls. ``arms`` maps each of ``"up"``, ``"down"``, and ``"control"`` to a nonempty sequence of geo columns in a wide revenue panel. All geos share one intervention date. Fitting uses the same synthetic-control model as :class:`SyntheticControl`; the arm methods aggregate its joint posterior draws without fitting separate models or treating down geos as donors. Attribution requires adequate pre-period donor fit, no spillovers into controls, and no concurrent geo-specific changes correlated with arms. Arm labels describe assignment, not the amount of media actually delivered. Parameters ---------- data : pd.DataFrame Wide revenue panel with a unique, sorted time index and geo columns. treatment_time : int, float, or pd.Timestamp First intervention period, shared by all up and down geos. arms : mapping of str to sequence of str Nonempty and disjoint ``up``, ``down``, and ``control`` geo groups. model : PyMCModel or RegressorMixin, optional Counterfactual model. Defaults to :class:`SoftmaxWeightedSumFitter`. min_donor_correlation : float, default 0.0 Minimum pre-period correlation before warning about a donor. auto_scale_sigma : bool, default True Whether to scale the stock observation-noise prior by pre-period data. Notes ----- The inherited three-panel :meth:`plot` labels revenue and impact axes from ``data.attrs["unit"]`` when present. A per-period unit such as ``USD/week`` becomes ``USD`` on the cumulative-impact axis. """ supports_ols = False supports_bayes = True _default_model_class = SoftmaxWeightedSumFitter
[docs] def __init__( self, data: pd.DataFrame, treatment_time: int | float | pd.Timestamp, arms: Mapping[str, Sequence[str]], model: PyMCModel | RegressorMixin | None = None, *, min_donor_correlation: float = 0.0, auto_scale_sigma: bool = True, ) -> None: if set(arms) != {"up", "down", "control"}: raise ValueError("arms must contain exactly 'up', 'down', and 'control'.") normalized = {arm: list(arms[arm]) for arm in ("up", "down", "control")} if any(not geos for geos in normalized.values()): raise ValueError("Each arm must contain at least one geo.") all_geos = sum(normalized.values(), []) if len(all_geos) != len(set(all_geos)): raise ValueError("Geo assignments overlap or contain duplicates.") if not data.index.is_unique or not data.index.is_monotonic_increasing: raise ValueError("Revenue panel time index must be unique and sorted.") if not data.columns.is_unique: raise ValueError("Revenue panel geo columns must be unique.") missing = set(all_geos) - set(data.columns) if missing: raise ValueError( f"Assigned geos are missing from revenue: {sorted(missing)}" ) if data[all_geos].isna().any().any(): raise ValueError("Revenue panel contains missing values in assigned geos.") self.outcome_unit = data.attrs.get("unit") self.arms = normalized self.geo_arm = {geo: arm for arm, geos in normalized.items() for geo in geos} super().__init__( data=data, treatment_time=treatment_time, control_units=normalized["control"], treated_units=normalized["up"] + normalized["down"], model=model, min_donor_correlation=min_donor_correlation, auto_scale_sigma=auto_scale_sigma, )
def _plot( self, *, group: Literal["prior", "posterior"] = "posterior", round_to: int | None = None, treated_unit: str | None = None, ci_prob: float = HDI_PROB, kind: Literal["ribbon", "histogram", "spaghetti"] = "ribbon", ci_kind: Literal["hdi", "eti"] = "hdi", num_samples: int = 50, plot_predictors: bool = False, figsize: tuple[float, float] = (7, 8), **kwargs: Any, ) -> tuple[plt.Figure, list[plt.Axes]]: """Label the inherited synthetic-control diagnostics in revenue units.""" fig, axes = super()._plot( group=group, round_to=round_to, treated_unit=treated_unit, ci_prob=ci_prob, kind=kind, ci_kind=ci_kind, num_samples=num_samples, plot_predictors=plot_predictors, figsize=figsize, **kwargs, ) period_unit = f" ({self.outcome_unit})" if self.outcome_unit else "" axes[0].set_ylabel(f"Revenue{period_unit}") if group == "posterior": cumulative_unit = ( self.outcome_unit.rsplit("/", 1)[0] if self.outcome_unit else None ) axes[1].set_ylabel(f"Impact{period_unit}") axes[2].set_ylabel( f"Cumulative impact ({cumulative_unit})" if cumulative_unit else "Cumulative impact" ) if self.outcome_unit and self.outcome_unit.startswith("USD"): for axis in axes: axis.yaxis.set_major_formatter(StrMethodFormatter("${x:,.1f}")) geo = self.treated_units[0] if treated_unit is None else treated_unit axes[0].set_title(f"{geo}: {axes[0].get_title()}") return fig, axes @property def impact_draws(self) -> xr.DataArray: """Signed per-period impact with aligned joint geo posterior draws.""" return self.result.impact_post.assign_coords( arm=("treated_units", [self.geo_arm[geo] for geo in self.treated_units]) ) @property def arm_impact_draws(self) -> xr.DataArray: """Draw-wise mean impact across geos in each intervention arm.""" impact = self.impact_draws return xr.concat( [ impact.sel(treated_units=self.arms[arm]) .mean(dim="treated_units") .drop_vars("arm", errors="ignore") for arm in ("up", "down") ], dim=pd.Index(["up", "down"], name="arm"), ) def _selected_window( self, start: int | float | pd.Timestamp | None, end: int | float | pd.Timestamp | None, ) -> pd.Index: index = self.datapost.index window_start = index[0] if start is None else start window_end = index[-1] if end is None else end selected = index[(index >= window_start) & (index <= window_end)] if ( len(selected) == 0 or selected[0] != window_start or selected[-1] != window_end ): raise ValueError( "Window boundaries must be observed post-intervention periods." ) return selected
[docs] def aggregate_draws( self, *, start: int | float | pd.Timestamp | None = None, end: int | float | pd.Timestamp | None = None, aggregate: Literal["mean", "sum"] = "mean", ) -> xr.Dataset: """Aggregate a common inclusive post-intervention window within draws. The ``geo`` variable retains a ``treated_units`` dimension and ``arm_mean`` retains separate up and down values. ``mean`` yields effect per outcome period; ``sum`` yields cumulative effect over the selected periods. Parameters ---------- start : int, float, or pd.Timestamp, optional First included post-intervention period. Defaults to the first post period. end : int, float, or pd.Timestamp, optional Last included post-intervention period. Defaults to the last post period. aggregate : {"mean", "sum"}, default "mean" Draw-wise time aggregation over the inclusive window. Returns ------- xr.Dataset Joint geo and arm effect draws with window metadata. """ if aggregate not in ("mean", "sum"): raise ValueError("aggregate must be 'mean' or 'sum'.") selected = self._selected_window(start, end) geo = self.result.impact_post.sel(obs_ind=selected) arm = self.arm_impact_draws.sel(obs_ind=selected) return xr.Dataset( { "geo": getattr(geo, aggregate)(dim="obs_ind"), "arm_mean": getattr(arm, aggregate)(dim="obs_ind"), }, attrs={ "window_start": str(selected[0]), "window_end": str(selected[-1]), "n_periods": len(selected), "aggregate": aggregate, }, )
[docs] def arm_effect_table( self, *, start: int | float | pd.Timestamp | None = None, end: int | float | pd.Timestamp | None = None, aggregate: Literal["mean", "sum"] = "mean", ) -> pd.DataFrame: """Summarize geo and arm effects with 94% equal-tailed intervals. Parameters ---------- start : int, float, or pd.Timestamp, optional First included post-intervention period. end : int, float, or pd.Timestamp, optional Last included post-intervention period. aggregate : {"mean", "sum"}, default "mean" Draw-wise time aggregation over the inclusive window. Returns ------- pd.DataFrame One row per treated geo and one row per intervention arm, with posterior mean, standard deviation, interval, and window metadata. """ draws = self.aggregate_draws(start=start, end=end, aggregate=aggregate) rows: list[tuple[str, str | None, str, np.ndarray]] = [] for geo in self.treated_units: samples = draws["geo"].sel(treated_units=geo).to_numpy().ravel() rows.append(("geo", geo, self.geo_arm[geo], samples)) for arm in ("up", "down"): samples = draws["arm_mean"].sel(arm=arm).to_numpy().ravel() rows.append(("arm", None, arm, samples)) return pd.DataFrame( [ { "level": level, "geo": geo, "arm": arm, "mean": float(np.mean(samples)), "sd": float(np.std(samples)), "lower_94": float(np.quantile(samples, 0.03)), "upper_94": float(np.quantile(samples, 0.97)), **draws.attrs, } for level, geo, arm, samples in rows ] )
[docs] def to_mmm_lift( self, *, spend_baseline: pd.DataFrame, spend_realized: pd.DataFrame, channel: str, spend_unit: str, outcome_unit: str, start: int | float | pd.Timestamp | None = None, end: int | float | pd.Timestamp | None = None, ) -> pd.DataFrame: """Export per-period geo lift rows for scalar MMM saturation calibration. Both spend frames require ``(geo, channel)`` MultiIndex columns. Their values must be in ``spend_unit`` per outcome period, and revenue must be in ``outcome_unit`` per the same period. ``x`` is mean baseline spend; ``delta_x`` is mean realized minus baseline spend; ``delta_y`` and ``sigma`` summarize posterior draws of mean signed revenue impact. The selected window is inclusive and identical for spend and outcome. Baseline spend and the delivered change must each be stable across the window because a nonlinear saturation curve evaluated at mean spend generally differs from the mean of period-level responses. The six scalar likelihood columns are accompanied by arm, window, and unit metadata. The joint posterior remains in :attr:`impact_draws`. Parameters ---------- spend_baseline : pd.DataFrame Business-as-usual spend panel with ``(geo, channel)`` columns and ``attrs['unit']`` matching ``spend_unit``. spend_realized : pd.DataFrame Delivered spend panel on the same period grid and in the same units. channel : str Tested channel name used in the MMM. spend_unit : str Spend unit per outcome period, matching both spend frame attributes. outcome_unit : str Revenue unit per outcome period, matching the fitted revenue attribute. start : int, float, or pd.Timestamp, optional First included post-intervention period. end : int, float, or pd.Timestamp, optional Last included post-intervention period. Returns ------- pd.DataFrame One scalar lift row per treated geo with signed effect and spend change, posterior standard deviation, and unit and window metadata. """ if not channel or not spend_unit or not outcome_unit: raise ValueError("channel, spend_unit, and outcome_unit are required.") if self.outcome_unit != outcome_unit: raise ValueError("outcome_unit must match revenue.attrs['unit'].") selected = self._selected_window(start, end) for label, frame in ( ("spend_baseline", spend_baseline), ("spend_realized", spend_realized), ): if frame.attrs.get("unit") != spend_unit: raise ValueError(f"spend_unit must match {label}.attrs['unit'].") if ( not isinstance(frame.columns, pd.MultiIndex) or frame.columns.nlevels != 2 ): raise ValueError( f"{label} must have (geo, channel) MultiIndex columns." ) if not frame.columns.is_unique: raise ValueError(f"{label} geo/channel columns must be unique.") if not frame.index.is_unique or not selected.isin(frame.index).all(): raise ValueError(f"{label} does not align with the outcome window.") missing = [ (geo, channel) for geo in self.treated_units if (geo, channel) not in frame ] if missing: raise ValueError(f"{label} is missing geo/channel values: {missing}") geo_draws = self.aggregate_draws( start=selected[0], end=selected[-1], aggregate="mean" )["geo"] rows = [] for geo in self.treated_units: baseline = spend_baseline.loc[selected, (geo, channel)].to_numpy( dtype=float ) realized = spend_realized.loc[selected, (geo, channel)].to_numpy( dtype=float ) if not np.isfinite(baseline).all() or not np.isfinite(realized).all(): raise ValueError( f"Spend values are missing or nonfinite for geo '{geo}'." ) if (baseline < 0).any() or (realized < 0).any(): raise ValueError(f"Spend cannot fall below zero for geo '{geo}'.") change = realized - baseline if not np.allclose(baseline, baseline[0]) or not np.allclose( change, change[0] ): raise ValueError( f"Baseline spend and delivered change must be stable across " f"the scalar MMM window for geo '{geo}'." ) arm = self.geo_arm[geo] if (arm == "up" and (change < 0).any()) or ( arm == "down" and (change > 0).any() ): raise ValueError( f"Delivered spend direction conflicts with '{arm}' assignment for '{geo}'." ) if change.mean() == 0: raise ValueError( f"Delivered spend change must be nonzero for geo '{geo}' " "in the scalar MMM lift likelihood." ) samples = geo_draws.sel(treated_units=geo).to_numpy().ravel() delta_y = float(samples.mean()) sigma = float(samples.std()) if delta_y == 0: raise ValueError( f"Estimated lift must be nonzero for geo '{geo}' " "in the scalar MMM lift likelihood." ) if change.mean() * delta_y < 0: raise ValueError( f"Estimated lift for geo '{geo}' has the opposite sign from its " "delivered spend change, which the scalar MMM lift likelihood " "does not accept. Inspect the estimate before calibration." ) if sigma <= 0: raise ValueError( f"Posterior lift uncertainty must be positive for geo '{geo}'." ) rows.append( { "channel": channel, "geo": geo, "x": float(baseline.mean()), "delta_x": float(change.mean()), "delta_y": delta_y, "sigma": sigma, "arm": arm, "window_start": selected[0], "window_end": selected[-1], "n_periods": len(selected), "spend_unit": spend_unit, "outcome_unit": outcome_unit, } ) return pd.DataFrame(rows)