import numpy as np
import matplotlib.pyplot as plt
import matplotlib.ticker as ticker
[docs]
class CurveVisualizer:
"""Handles visualization for media response curves."""
# Google Brand Colors
G_BLUE = '#4285F4'
G_RED = '#EA4335'
G_YELLOW = '#FBBC04'
G_GREEN = '#34A853'
G_GRAY = '#5F6368'
G_LIGHT_GRAY = '#F8F9FA'
[docs]
@classmethod
def plot_response_curve(cls, model, target_mroas=1.0, current_spend=None, show_intervals=True, scatter=None, include_baseline=False):
"""Generates a visualization of the media response and marginal return curves."""
min_spend = model.get_minimal_marginal_cost_point()
max_spend = model.get_diminishing_returns_point(target_mroas)
# Determine plot limits
max_x = max_spend * 1.5 if max_spend else min_spend * 4
if current_spend: max_x = max(max_x, current_spend * 1.2)
if scatter is not None: max_x = max(max_x, np.max(scatter[0]) * 1.1)
if max_spend and max_x > 100 * max_spend:
max_x = max_spend * 3.0
x_vals = np.linspace(0, max_x, 500)
has_intervals = False
interval_label = "Uncertainty Interval"
if show_intervals and model.posterior_samples:
y_return, y_return_low, y_return_high = model.predict_incremental_return(
x_vals, return_interval=True, confidence_level=0.90, include_baseline=include_baseline
)
y_mroas, y_mroas_low, y_mroas_high = model.predict_marginal_return(x_vals, return_interval=True, confidence_level=0.90)
has_intervals = True
interval_label = "90% Credible Interval"
elif show_intervals and model.covariance_matrix is not None:
y_return, y_return_low, y_return_high = model.predict_incremental_return(
x_vals, return_interval=True, confidence_level=0.95, include_baseline=include_baseline
)
y_mroas, y_mroas_low, y_mroas_high = model.predict_marginal_return(x_vals, return_interval=True, confidence_level=0.95)
has_intervals = True
interval_label = "95% Confidence Interval"
else:
y_return = model.predict_incremental_return(x_vals, include_baseline=include_baseline)
y_mroas = model.predict_marginal_return(x_vals)
plt.rcParams['font.family'] = 'sans-serif'
plt.rcParams['font.sans-serif'] = ['Roboto', 'Open Sans', 'Arial', 'DejaVu Sans']
fig, ax1 = plt.subplots(figsize=(12, 7), facecolor='white')
ax1.set_facecolor('white')
# Primary Axis: Response Curve
curve_label = "Total Return" if include_baseline else "Incremental Return"
y_axis_label = "Total Return ($)" if include_baseline else "Incremental Return ($)"
ax1.plot(x_vals, y_return, color=cls.G_BLUE, linewidth=3.5, label=curve_label, zorder=3)
if has_intervals:
ax1.fill_between(x_vals, y_return_low, y_return_high, color=cls.G_BLUE, alpha=0.15, label=interval_label, zorder=2)
ax1.set_xlabel('Spend', fontsize=11, color=cls.G_GRAY, fontweight='500', labelpad=10)
ax1.set_ylabel(y_axis_label, color=cls.G_BLUE, fontsize=11, fontweight='500', labelpad=10)
ax1.tick_params(axis='both', which='major', labelsize=10, colors=cls.G_GRAY)
# Secondary Axis: Marginal Return
ax2 = ax1.twinx()
ax2.plot(x_vals, y_mroas, color=cls.G_GRAY, linestyle=(0, (5, 2)), linewidth=1.5, label="Marginal ROAS", alpha=0.6, zorder=1)
if has_intervals:
ax2.fill_between(x_vals, y_mroas_low, y_mroas_high, color=cls.G_GRAY, alpha=0.05, zorder=0)
ax2.set_ylabel('Marginal ROAS (mROAS)', color=cls.G_GRAY, fontsize=11, fontweight='500', labelpad=10)
ax2.tick_params(axis='y', labelcolor=cls.G_GRAY, labelsize=10)
ax2.axhline(target_mroas, color=cls.G_RED, linestyle=':', linewidth=1, alpha=0.5, label=f"Target mROAS ({target_mroas})")
# Optimal Scaling Zone
if max_spend and max_spend > min_spend:
ax1.axvspan(min_spend, max_spend, color=cls.G_GREEN, alpha=0.08, label='Optimal Scaling Zone', zorder=0)
# Use blended transform (x in data coords, y in axes fraction) for robust placement
ax1.text((min_spend + max_spend) / 2.0, 0.03, 'OPTIMAL ZONE',
transform=ax1.get_xaxis_transform(),
horizontalalignment='center', verticalalignment='bottom',
fontsize=9, color=cls.G_GREEN, fontweight='bold', alpha=0.7)
# Current Spend marker
if current_spend:
ax1.axvline(current_spend, color=cls.G_RED, linestyle='--', linewidth=1.5, alpha=0.8, label=f"Current Spend (${current_spend:,.0f})", zorder=4)
curr_ret = model.predict_incremental_return(current_spend, include_baseline=include_baseline)
ax1.scatter(current_spend, curr_ret, color=cls.G_RED, s=60, edgecolors='white', linewidth=1.5, zorder=5)
# Scatter data
if scatter is not None:
scatter_spend, scatter_return = scatter
scatter_spend_adstocked = model.adstock_spend(scatter_spend)
has_adstock = (model.theta > 0) or (model.adstock_type and model.adstock_type != "none")
ax1.scatter(scatter_spend_adstocked, scatter_return, color=cls.G_BLUE, alpha=0.3, s=40, edgecolors='white', linewidth=0.8,
label="Historical Data (Adstocked)" if has_adstock else "Historical Data", zorder=1)
# Markers for key points
if min_spend > 0:
ax2.scatter(min_spend, model.predict_marginal_return(min_spend), marker='o', color=cls.G_YELLOW, s=100, edgecolors=cls.G_GRAY, linewidth=1, label="Peak Efficiency", zorder=6)
# Formatting with smart scaling ($M, $k, $)
def format_spend(x, p):
if abs(x) >= 1e6:
return f'${x*1e-6:g}M'
elif abs(x) >= 1e3:
return f'${x*1e-3:g}k'
else:
return f'${x:g}'
def format_return(x, p):
if abs(x) >= 1e6:
return f'{x*1e-6:g}M'
elif abs(x) >= 1e3:
return f'{x*1e-3:g}k'
else:
return f'{x:g}'
ax1.xaxis.set_major_formatter(ticker.FuncFormatter(format_spend))
ax1.yaxis.set_major_formatter(ticker.FuncFormatter(format_return))
ax1.set_ylim(bottom=0)
ax2.set_ylim(bottom=0)
# Hide spines
ax1.spines['top'].set_visible(False)
ax1.spines['right'].set_visible(False)
ax1.spines['left'].set_color(cls.G_LIGHT_GRAY)
ax1.spines['bottom'].set_color(cls.G_LIGHT_GRAY)
ax2.spines['top'].set_visible(False)
ax2.spines['right'].set_visible(False)
ax2.spines['left'].set_visible(False)
ax1.grid(True, linestyle='-', alpha=0.1, color=cls.G_GRAY)
# Legends
lines1, labels1 = ax1.get_legend_handles_labels()
lines2, labels2 = ax2.get_legend_handles_labels()
ax1.legend(lines1 + lines2, labels1 + labels2, loc='center right', frameon=True, facecolor='white', framealpha=1.0, fontsize=10)
# Title
plt.title(f'Media Response Analysis: {model.channel_name}', loc='left', fontsize=16, fontweight='bold', pad=25, color='#202124')
# Subtitle with parameters
fig.text(0.125, 0.91, f'Hill Curve Parameters: α={model.alpha:.2f}, K={model.K:,.0f}, β={model.beta:,.0f}',
fontsize=10, color=cls.G_GRAY)
plt.tight_layout()
return fig