#####################################################
# Bass Diffusion Model                              #
# Chapter 13 - Chestnut Ridge Smart Watch Example   #
#####################################################

# This script reproduces the full Chestnut Ridge smart watch Bass diffusion example in Python:
#   Step 1. Data setup (cumulative and lagged cumulative sales)
#   Step 2. Estimate the Bass diffusion model (N, p, and q)
#   Step 3. Forecast sales for 50 quarters
#   Step 4. Visualize actual vs. predicted sales
#   Step 5. Use the forecasts for sales planning
# 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.

# %%
## 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 and Set Seed
from tkinter import Tk, filedialog

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import statsmodels.formula.api as smf

rng = np.random.default_rng(1)

## Show all columns and rows when printing tables
pd.set_option("display.max_columns", None)
pd.set_option("display.max_rows", None)
pd.set_option("display.width", 200)


## 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 1: Data Setup                                #
#####################################################

## Read in the sales data (see Table 13.4)
sales = pd.read_csv(file_choose("Choose sales_data.csv"))  ## Choose the file sales_data.csv
sales["date"] = pd.to_datetime(sales["date"], format="%m/%d/%Y")
print(sales.head())

## Create cumulative sales and lag cumulative sales variables
## N(t-1) is cumulative sales up to the prior quarter, and N(t-1)^2 is its square
## The lag for the first quarter is 0 since there are no sales before launch
sales["cumsales"] = sales["sales"].cumsum()
sales["cumsales_1"] = sales["cumsales"].shift(1, fill_value=0)
sales["cumsales2"] = sales["cumsales"] ** 2
sales["cumsales2_1"] = sales["cumsales2"].shift(1, fill_value=0)
print(sales)


# %%
#####################################################
# Step 2: Estimate the Bass Diffusion Model         #
#####################################################

## Run the Bass Diffusion Model regression (equation 13.5)
## n(t) = a + b*N(t-1) + c*N(t-1)^2
bass = smf.ols("sales ~ cumsales_1 + cumsales2_1", data=sales).fit()
print(bass.summary())  ## R-squared = 0.7752

## Determine N, p, and q from the regression coefficients
a = bass.params["Intercept"]    ## Intercept                   ## 2.3721
b = bass.params["cumsales_1"]   ## Coefficient on cumsales_1   ## 0.1179
c = bass.params["cumsales2_1"]  ## Coefficient on cumsales2_1  ## -0.0002692
N1 = (-b + np.sqrt(b**2 - 4*a*c)) / (2*c)
N2 = (-b - np.sqrt(b**2 - 4*a*c)) / (2*c)
N = max(N1, N2)  ## Market potential (millions of units)
p = a / N        ## Coefficient of innovation
q = b + p        ## Coefficient of imitation
print(f"N = {N:.2f}, p = {p:.4f}, q = {q:.4f}")  ## N = 457.34, p = 0.0052, q = 0.1231


# %%
#####################################################
# Step 3: Forecast Sales                            #
#####################################################

## Forecast sales for 50 quarters using the Bass diffusion model
## n(t) = p*N + (q - p)*N(t-1) - (q/N)*N(t-1)^2
periods = 50
forecasts = pd.DataFrame({"t": range(1, periods + 1),
                          "date": pd.date_range(sales["date"].min(), periods=periods, freq="QS")})
psales = np.zeros(periods)
pcumsales = np.zeros(periods + 1)
for i in range(periods):
    psales[i] = p*N + (q - p)*pcumsales[i] - (q/N)*pcumsales[i]**2
    pcumsales[i+1] = pcumsales[i] + psales[i]
forecasts["psales"] = psales
forecasts["pcumsales"] = pcumsales[1:]  ## Remove the first value which is 0

## Add the actual sales to the forecasts (NaN for quarters with no data yet)
forecasts = forecasts.merge(sales[["period", "sales", "cumsales"]],
                            left_on="t", right_on="period", how="left").drop(columns="period")
print(forecasts)


# %%
#####################################################
# Step 4: Visualize the Sales Forecasts             #
#####################################################

## Function to plot actual vs. predicted sales by quarter
## Actual values only exist for the first 24 quarters
def plot_forecast(actual, predicted, title, ylab):
    fig, ax = plt.subplots(figsize=(8, 5))
    ax.plot(forecasts["t"], actual, color="darkorange", linewidth=2, label="Actual")
    ax.plot(forecasts["t"], predicted, color="steelblue", linewidth=2, label="Predicted")
    ax.set_xticks(range(2, periods + 1, 2))
    ax.grid(color="0.92")
    ax.spines[:].set_visible(False)
    ax.set_title(title, loc="left")
    ax.set_xlabel("Quarter")
    ax.set_ylabel(ylab)
    ax.legend(loc="upper center", bbox_to_anchor=(0.5, -0.12), ncol=2, frameon=False)
    plt.tight_layout()
    plt.show()


## Figure 13.4: Actual vs. Predicted Sales by Quarter
plot_forecast(forecasts["sales"], forecasts["psales"],
              "Actual vs. Predicted Sales Per Quarter", "Sales (millions of units)")

## Figure 13.5: Actual vs. Predicted Cumulative Sales by Quarter
plot_forecast(forecasts["cumsales"], forecasts["pcumsales"],
              "Actual vs. Predicted Cumulative Sales Per Quarter", "Cumulative Sales (millions of units)")

## Quarter in which predicted sales peak
print(forecasts.loc[forecasts["psales"].idxmax(), ["t", "date", "psales"]])  ## Quarter 27 (10/1/2020), 15.28 million


# %%
#####################################################
# Step 5: Using Forecasts for Sales Planning        #
#####################################################

## Predicted market share of the Chestnut Ridge smart watch (first choice rule, Chapter 12)
cr_share = 0.30

## Expected Chestnut Ridge sales in the next quarter (quarter 25)
next_qtr = forecasts[forecasts["t"] == sales["period"].max() + 1].iloc[0]
print(next_qtr["psales"])             ## 15.13 million smart watches in the market
print(next_qtr["psales"] * cr_share)  ## 4.54 million Chestnut Ridge smart watches

## Expected Chestnut Ridge sales for each remaining forecast quarter
future = forecasts.loc[forecasts["t"] > sales["period"].max(), ["t", "date", "psales"]].copy()
future["cr_sales"] = future["psales"] * cr_share
print(future)
print(future["cr_sales"].sum())  ## 73.51 million over quarters 25 to 50
