Source code for dxpoint.viz

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