##############################################################
# Chapter 15 - Using Marketing Mix Models to Optimize the    #
#              Marketing Mix                                 #
# Chestnut Ridge HDTV Marketing Mix Model (Python)           #
##############################################################

# This script reproduces the full Chestnut Ridge HDTV marketing mix model example in Python:
#   Step 1. Install packages (if needed), load packages, and set seed
#   Step 2. Read in the marketing mix data
#   Step 3. Look at means of variables
#   Step 4. Create natural log, lag, and weekday variables
#   Step 5. Check for unit root (augmented Dickey-Fuller test)
#   Step 6. Check for multicollinearity
#   Step 7. Run the marketing mix regression
#   Step 8. Create predicted values for quantity and revenue
#   Step 9. Set up shared formatting for the date axis and number labels
#   Step 10. Plot marketing spend and pricing over time (Figure 15.6)
#   Step 11. Plot predicted vs. actual quantity sold (Figure 15.7)
#   Step 12. Plot predicted vs. actual revenue (Figure 15.8)
# Run it from top to bottom. Each step prints its result or shows a plot.
# The "# %%" lines split the script into cells you can run one at a time
# in VS Code, Spyder, or PyCharm.

# %%
##############################################################
# Step 1: Install packages (if needed), load packages, and   #
#         set seed                                           #
##############################################################

## Install Packages (if needed - only installs packages that are missing)
## Each entry is {name used with import: name used with pip install}
import sys, subprocess, importlib.util
packages = {"pandas": "pandas", "numpy": "numpy", "statsmodels": "statsmodels", "matplotlib": "matplotlib"}
new_packages = [pip_name for import_name, pip_name in packages.items() if importlib.util.find_spec(import_name) is None]
if new_packages: subprocess.check_call([sys.executable, "-m", "pip", "install", *new_packages])

## Load Packages
import warnings
from tkinter import Tk, filedialog

import numpy as np
import pandas as pd
import statsmodels.formula.api as smf
from statsmodels.tsa.stattools import adfuller
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
from matplotlib.ticker import FuncFormatter

np.random.seed(1)

## Function to choose a file with a dialog box
## (the Python version of file.choose() in R - no need to set the working directory)
def file_choose(title="Choose a file"):
    root = Tk()
    root.withdraw()                     # hide the empty Tk window
    root.attributes("-topmost", True)   # bring the dialog to the front
    path = filedialog.askopenfilename(title=title, filetypes=[("CSV files", "*.csv"), ("All files", "*.*")])
    root.destroy()
    if not path:
        raise SystemExit("No file was chosen.")
    return path


# %%
##############################################################
# Step 2: Read in the marketing mix data                     #
##############################################################

mmix = pd.read_csv(file_choose("Choose mmix_data.csv"))  ## Choose the file mmix_data.csv


# %%
##############################################################
# Step 3: Look at means of variables                         #
##############################################################

print(mmix.mean(numeric_only=True).round(2))


# %%
##############################################################
# Step 4: Create natural log, lag, and weekday variables     #
##############################################################

mmix["ln_quantity"] = np.log(mmix["quantity"])
mmix["ln_price"] = np.log(mmix["price"])
mmix["ln_digital_ad"] = np.log(mmix["digital_ad"])
mmix["ln_digital_search"] = np.log(mmix["digital_search"])
mmix["ln_print"] = np.log(mmix["print"] + 1)  ## Some time periods have 0 spend
mmix["ln_tv"] = np.log(mmix["tv"])

## Lag of ln_quantity (day 1 has no prior day, so it is set to 0)
mmix["lln_quantity"] = mmix["ln_quantity"].shift(1, fill_value=0)

## Day-of-week variable (used as dummy variables in the regression)
## Note: day_name() returns English day names by default, whatever your
## computer's language setting. The days sort alphabetically, so Friday
## becomes the baseline day in the regression. (Passing a locale, e.g.,
## day_name(locale="de_DE"), would change the names and the baseline day.)
mmix["date"] = pd.to_datetime(mmix["date"])
mmix["weekdays"] = mmix["date"].dt.day_name()


# %%
##############################################################
# Step 5: Check for unit root (augmented Dickey-Fuller test) #
##############################################################

## Lag order and trend match R's tseries::adf.test() defaults. The test
## statistic matches R exactly; the p-value differs slightly because R's
## adf.test() reports any p-value below 0.01 as 0.01.
warnings.simplefilter("ignore", FutureWarning)  ## Hide a statsmodels notice
adf_lag = int((len(mmix) - 1) ** (1 / 3))
adf = adfuller(mmix["ln_quantity"], maxlag=adf_lag, regression="ct", autolag=None)
print("Augmented Dickey-Fuller Test")
print(f"Dickey-Fuller = {adf[0]:.4f}, Lag order = {adf_lag}, p-value = {adf[1]:.4f}")
print("alternative hypothesis: stationary")


# %%
##############################################################
# Step 6: Check for multicollinearity                        #
##############################################################

cor_vars = ["ln_quantity", "lln_quantity", "ln_price", "ln_digital_ad",
            "ln_digital_search", "ln_print", "ln_tv"]
cor_table = mmix[cor_vars].corr()
print(cor_table.round(2).to_string())

## Combine ln_digital_ad and ln_digital_search
mmix["ln_digital"] = np.log(mmix["digital_ad"] + mmix["digital_search"])


# %%
##############################################################
# Step 7: Run the marketing mix regression                   #
##############################################################

mmix_reg = smf.ols("ln_quantity ~ lln_quantity + ln_price + ln_digital + ln_print"
                   " + ln_tv + C(weekdays)", data=mmix).fit()
print(mmix_reg.summary())


# %%
##############################################################
# Step 8: Create predicted values for quantity and revenue   #
##############################################################

mmix["pred_quantity"] = np.exp(mmix_reg.fittedvalues)
mmix["pred_revenue"] = mmix["price"] * mmix["pred_quantity"]


# %%
##############################################################
# Step 9: Set up shared formatting for the date axis and     #
#         number labels                                      #
##############################################################

def format_date_axis(ax):
    ax.xaxis.set_major_locator(mdates.MonthLocator(interval=2))
    ax.xaxis.set_major_formatter(mdates.DateFormatter("%b %Y"))
    plt.setp(ax.get_xticklabels(), rotation=45, ha="right")

comma = FuncFormatter(lambda x, pos: f"{x:,.0f}")
dollar = FuncFormatter(lambda x, pos: f"${x:,.0f}")


# %%
##############################################################
# Step 10: Plot marketing spend and pricing over time        #
#          (Figure 15.6)                                     #
##############################################################

mix_vars = ["digital_ad", "digital_search", "print", "tv", "price"]
mix_labels = ["Digital Ad ($)", "Digital Search ($)", "Print ($)", "TV ($)",
              "Price ($)"]

fig, axes = plt.subplots(len(mix_vars), 1, sharex=True, figsize=(11, 8))
for ax, var, label in zip(axes, mix_vars, mix_labels):
    ax.plot(mmix["date"], mmix[var], color="#2a78d6", linewidth=0.6)
    ax.set_ylabel(label, rotation=0, ha="right", va="center")
    ax.yaxis.set_major_formatter(comma)
    ax.grid(color="#e5e5e5", linewidth=0.5)
    ax.spines[["top", "right"]].set_visible(False)
format_date_axis(axes[-1])
axes[-1].set_xlabel("Date")
fig.suptitle("Marketing Spend and Pricing Over Time", x=0.02, ha="left")
fig.tight_layout()
plt.show()


# %%
##############################################################
# Step 11: Plot predicted vs. actual quantity sold (Figure   #
#          15.7)                                             #
##############################################################

fig, ax = plt.subplots(figsize=(11, 6))
ax.plot(mmix["date"], mmix["quantity"], color="#2a78d6", linewidth=0.7,
        label="Actual Quantity")
ax.plot(mmix["date"], mmix["pred_quantity"], color="#eb6834", linewidth=0.7,
        linestyle="--", label="Predicted Quantity")
ax.set_title("Predicted vs. Actual Quantity Sold", loc="left")
ax.set_xlabel("Date")
ax.set_ylabel("Quantity Sold")
ax.yaxis.set_major_formatter(comma)
ax.grid(color="#e5e5e5", linewidth=0.5)
ax.spines[["top", "right"]].set_visible(False)
ax.legend(loc="upper center", ncol=2, frameon=False)
format_date_axis(ax)
fig.tight_layout()
plt.show()


# %%
##############################################################
# Step 12: Plot predicted vs. actual revenue (Figure 15.8)   #
##############################################################

fig, ax = plt.subplots(figsize=(11, 6))
ax.plot(mmix["date"], mmix["revenue"], color="#2a78d6", linewidth=0.7,
        label="Actual Revenue")
ax.plot(mmix["date"], mmix["pred_revenue"], color="#eb6834", linewidth=0.7,
        linestyle="--", label="Predicted Revenue")
ax.set_title("Predicted vs. Actual Revenue", loc="left")
ax.set_xlabel("Date")
ax.set_ylabel("Revenue")
ax.yaxis.set_major_formatter(dollar)
ax.grid(color="#e5e5e5", linewidth=0.5)
ax.spines[["top", "right"]].set_visible(False)
ax.legend(loc="upper center", ncol=2, frameon=False)
format_date_axis(ax)
fig.tight_layout()
plt.show()
