##################################################
## Discriminant Analysis and Classification in R #
##################################################

# This script reproduces the full Chestnut Ridge discriminant analysis and targeting example in R:
#   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 to the console or Plots pane.

## Install Packages (if needed - only installs packages that are missing)
packages <- c("MASS", "dplyr", "tidyr", "ggplot2", "sf", "tigris")
new_packages <- packages[!(packages %in% installed.packages()[, "Package"])]
if (length(new_packages) > 0) install.packages(new_packages)

## Load Packages and Set Seed
library(MASS)
library(dplyr)
library(tidyr)
library(ggplot2)
library(sf)
library(tigris)
select <- dplyr::select ## MASS also has a select() function; use the dplyr version
set.seed(1)

## Read in Segment Data and Classification Data
seg <- read.csv(file.choose()) ## Choose segmentation_result.csv file 
class <- read.csv(file.choose()) ## Choose retail_classification.csv file


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

## Run Discriminant Analysis
fit <- lda(segment ~ married + own_home + household_size + income + age, data = seg)
fit ## print the summary statistics of your discriminant analysis

## Check which Discriminant Functions are Significant
## Segment is treated as a factor (6 groups), so each ANOVA has 5 degrees of freedom
ldaPred <- predict(fit, seg)
ld <- ldaPred$x
for (i in 1:ncol(ld)) {
  cat("\n---", colnames(ld)[i], "---\n")
  print(anova(lm(ld[, i] ~ factor(seg$segment))))
}

## Check Discriminant Model Fit (Confusion Table)
pred.seg <- ldaPred$class
tseg <- table(seg$segment, pred.seg)
tseg # print table
sum(diag(tseg))/nrow(seg) # print percent correct (hit rate)

## Compare Hit Rate to Benchmarks
1/length(unique(seg$segment)) # random chance: 1 in 6
sum(fit$prior^2) # proportional chance criterion (accounts for unequal segment sizes)


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

## Run Classification Using Discriminant Function
pred.class <- predict(fit, class)$class
tclass <- table(pred.class)
tclass # print table

## Add Predicted Segment to Classification Data
class.seg <- cbind(class, pred.class)


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

## Zip codes are stored as numbers, so restore the leading zero (e.g., 7726 -> "07726")
class.seg <- class.seg %>%
  mutate(zip_code = formatC(zip_code, width = 5, flag = "0"))

## Number of Prospects in each Zip from each Segment (Figure 4.3)
zip.seg <- class.seg %>%
  count(zip_code, pred.class) %>%
  pivot_wider(names_from = pred.class, values_from = n, values_fill = 0, names_sort = TRUE)
print(zip.seg, n = 20) # print the first 20 zip codes

## 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)
options(tigris_use_cache = TRUE)
zip.count <- class.seg %>% count(zip_code, name = "prospects")
zip.shapes <- zctas(cb = TRUE, year = 2020, starts_with = zip.count$zip_code)
zip.map <- zip.shapes %>%
  inner_join(zip.count, by = c("ZCTA5CE20" = "zip_code"))
state.map <- states(cb = TRUE, year = 2020) %>%
  filter(STUSPS %in% c("NJ", "PA", "DE", "MD", "DC", "VA", "WV"))

ggplot() +
  geom_sf(data = state.map, fill = "grey95", color = "grey60") +
  geom_sf(data = zip.map, aes(fill = prospects), color = NA) +
  scale_fill_distiller(palette = "Blues", direction = 1, name = "Number of\nProspects") +
  coord_sf(xlim = st_bbox(zip.map)[c(1, 3)], ylim = st_bbox(zip.map)[c(2, 4)]) +
  labs(title = "Heat Map of Prospects by Zip Code") +
  theme_void()

## 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 <- class.seg %>% count(zip_code, pred.class, name = "prospects")
zip.seg.map <- zip.shapes %>%
  inner_join(zip.seg.count, by = c("ZCTA5CE20" = "zip_code"))

for (s in sort(unique(as.character(zip.seg.map$pred.class)))) {
  print(
    ggplot() +
      geom_sf(data = state.map, fill = "grey95", color = "grey60") +
      geom_sf(data = filter(zip.seg.map, pred.class == s), aes(fill = prospects), color = NA) +
      scale_fill_distiller(palette = "Blues", direction = 1, name = "Number of\nProspects") +
      coord_sf(xlim = st_bbox(zip.map)[c(1, 3)], ylim = st_bbox(zip.map)[c(2, 4)]) +
      labs(title = paste("Heat Map of Segment", s, "Prospects by Zip Code")) +
      theme_void()
  )
}

## 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.bases <- seg %>%
  mutate(sales_per_year = avg_order_size * avg_order_freq * (1 - return_rate),
         profit_per_year = sales_per_year * 0.52 - avg_mktg_cnt * 0.75) %>%
  group_by(segment) %>%
  summarize(avg_mktg_cnt = mean(avg_mktg_cnt),
            avg_order_freq = mean(avg_order_freq),
            avg_order_size = mean(avg_order_size),
            crossbuy = mean(crossbuy),
            multichannel = mean(multichannel),
            per_sale = mean(per_sale),
            return_rate = mean(return_rate),
            count = n(),
            sales_per_year = sum(sales_per_year),
            profit_per_year = sum(profit_per_year))
print(as.data.frame(seg.bases), digits = 3) # print table
