##################################################
## Discriminant Analysis and Classification in Python
##################################################

# This script reproduces the full Chestnut Ridge discriminant analysis and targeting example in Python:
#   Step 1. Discriminant analysis of the Chapter 3 segments (significance and hit rate)
#   Step 2. Classify prospects into segments using the discriminant function
#   Step 3. Marketing strategy for targeting (prospects by zip code, heat maps, segment bases)
# 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", "sklearn": "scikit-learn", "statsmodels": "statsmodels", "matplotlib": "matplotlib", "geopandas": "geopandas", "pygris": "pygris"}
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
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import geopandas as gpd
import statsmodels.formula.api as smf
from statsmodels.stats.anova import anova_lm
from sklearn.discriminant_analysis import LinearDiscriminantAnalysis
from pygris import zctas, states
from tkinter import Tk, filedialog
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

## Read in Segment Data and Classification Data
seg = pd.read_csv(file_choose("Choose segmentation_result.csv"))  ## Choose segmentation_result.csv file
pros = pd.read_csv(file_choose("Choose retail_classification.csv"))  ## Choose retail_classification.csv file


# %%
##################################################
# Step 1: Discriminant Analysis                  #
##################################################

## Run Discriminant Analysis
descriptors = ["married", "own_home", "household_size", "income", "age"]
fit = LinearDiscriminantAnalysis()
fit.fit(seg[descriptors], seg["segment"])

## scikit-learn scales the discriminant functions slightly differently than R;
## rescale them so the coefficients match R's lda() output (classification is unaffected)
n, g = len(seg), seg["segment"].nunique()
rescale = np.sqrt((n - g) / n)
scalings = fit.scalings_ * rescale

## Print the summary statistics of your discriminant analysis
ld_names = ["LD" + str(i + 1) for i in range(scalings.shape[1])]
print("Prior probabilities of groups:")
print(pd.Series(fit.priors_, index=fit.classes_).round(3), "\n")
print("Group means:")
print(pd.DataFrame(fit.means_, index=fit.classes_, columns=descriptors), "\n")
print("Coefficients of linear discriminants:")
print(pd.DataFrame(scalings, index=descriptors, columns=ld_names), "\n")
print("Proportion of trace:")
print(pd.Series(fit.explained_variance_ratio_, index=ld_names).round(4), "\n")

## Check which Discriminant Functions are Significant
## Segment is treated as a categorical variable (6 groups), so each ANOVA has 5 degrees of freedom
ld = pd.DataFrame(fit.transform(seg[descriptors]) * rescale, columns=ld_names)
ld["segment"] = seg["segment"]
for name in ld_names:
    print("---", name, "---")
    print(anova_lm(smf.ols(name + " ~ C(segment)", data=ld).fit()), "\n")

## Check Discriminant Model Fit (Confusion Table)
pred_seg = fit.predict(seg[descriptors])
tseg = pd.crosstab(seg["segment"], pd.Categorical(pred_seg, categories=fit.classes_),
                   rownames=["segment"], colnames=["pred_seg"], dropna=False)
print(tseg, "\n")  # print table
print(np.trace(tseg) / len(seg))  # print percent correct (hit rate)

## Compare Hit Rate to Benchmarks
print(1 / seg["segment"].nunique())  # random chance: 1 in 6
print((fit.priors_ ** 2).sum())  # proportional chance criterion (accounts for unequal segment sizes)


# %%
##################################################
# Step 2: Classification                         #
##################################################

## Run Classification Using Discriminant Function
pros["pred_class"] = fit.predict(pros[descriptors])
tclass = pros["pred_class"].value_counts().reindex(fit.classes_, fill_value=0)
print(tclass, "\n")  # print table


# %%
##################################################
# Step 3: Marketing Strategy for Targeting       #
##################################################

## Zip codes are stored as numbers, so restore the leading zero (e.g., 7726 -> "07726")
pros["zip_code"] = pros["zip_code"].astype(str).str.zfill(5)

## Number of Prospects in each Zip from each Segment (Figure 4.3)
zip_seg = pd.crosstab(pros["zip_code"], pros["pred_class"])
print(zip_seg.head(20), "\n")  # print the first 20 zip codes
# print(zip_seg.to_string())  # print the full table

## Heat Map of Prospects by Zip Code (Figure 4.4)
## Download zip code (ZCTA) and state boundaries from the US Census Bureau
## (the first download may take a few minutes; files are cached for later use)
zip_count = pros["zip_code"].value_counts().rename("prospects").reset_index()
zip_shapes = zctas(cb=True, year=2020, starts_with=list(zip_count["zip_code"]), cache=True)
zip_map = zip_shapes.merge(zip_count, left_on="ZCTA5CE20", right_on="zip_code")
state_map = states(cb=True, year=2020, cache=True)
state_map = state_map[state_map["STUSPS"].isin(["NJ", "PA", "DE", "MD", "DC", "VA", "WV"])].to_crs(zip_map.crs)
xmin, ymin, xmax, ymax = zip_map.total_bounds

fig, ax = plt.subplots(figsize=(10, 7))
state_map.plot(ax=ax, color="0.95", edgecolor="0.6")
zip_map.plot(ax=ax, column="prospects", cmap="Blues", legend=True,
             legend_kwds={"label": "Number of Prospects", "shrink": 0.6})
ax.set_xlim(xmin, xmax)
ax.set_ylim(ymin, ymax)
ax.set_title("Heat Map of Prospects by Zip Code")
ax.set_axis_off()
plt.show()

## Heat Map of Prospects by Zip Code for each Predicted Segment
## (one map per segment so each has its own color scale)
zip_seg_count = pros.groupby(["zip_code", "pred_class"]).size().rename("prospects").reset_index()
zip_seg_map = zip_shapes.merge(zip_seg_count, left_on="ZCTA5CE20", right_on="zip_code")

for s in sorted(zip_seg_map["pred_class"].unique()):
    fig, ax = plt.subplots(figsize=(10, 7))
    state_map.plot(ax=ax, color="0.95", edgecolor="0.6")
    zip_seg_map[zip_seg_map["pred_class"] == s].plot(
        ax=ax, column="prospects", cmap="Blues", legend=True,
        legend_kwds={"label": "Number of Prospects", "shrink": 0.6})
    ax.set_xlim(xmin, xmax)
    ax.set_ylim(ymin, ymax)
    ax.set_title("Heat Map of Segment " + str(s) + " Prospects by Zip Code")
    ax.set_axis_off()
    plt.show()

## Average Values for the Bases of each Segment (Figure 4.5; see Figure 3.6)
## Sales Per Year = Avg Order Size * Avg Order Freq * (1 - Return Rate)
## Profit Per Year = Sales Per Year * 0.52 margin - Avg Mktg Cnt * $0.75 per catalog
seg["sales_per_year"] = seg["avg_order_size"] * seg["avg_order_freq"] * (1 - seg["return_rate"])
seg["profit_per_year"] = seg["sales_per_year"] * 0.52 - seg["avg_mktg_cnt"] * 0.75
seg_bases = seg.groupby("segment").agg(
    avg_mktg_cnt=("avg_mktg_cnt", "mean"),
    avg_order_freq=("avg_order_freq", "mean"),
    avg_order_size=("avg_order_size", "mean"),
    crossbuy=("crossbuy", "mean"),
    multichannel=("multichannel", "mean"),
    per_sale=("per_sale", "mean"),
    return_rate=("return_rate", "mean"),
    count=("segment", "size"),
    sales_per_year=("sales_per_year", "sum"),
    profit_per_year=("profit_per_year", "sum"))
print(seg_bases.round(2).to_string())  # print table
