##########################################################
# Factor Analysis in Python - Chestnut Ridge Example     #
# Retailer Brand Audit from Survey Data                  #
##########################################################

# This script reproduces the full Chestnut Ridge retailer brand audit example in Python:
#   Step 1. Look at the average survey responses by retailer
#   Step 2. Determine the number of factors (eigenvalues and scree plot)
#   Step 3. Run a factor analysis with four factors (default varimax rotation)
#   Step 4. Rerun the factor analysis with an oblique rotation (heat map of loadings)
#   Step 5. Name the factors and score each retailer on the factors
#   Step 6. Heat map of factor scores by retailer
# 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", "matplotlib": "matplotlib", "scipy": "scipy", "statsmodels": "statsmodels"}
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 pandas as pd                                      # data manipulation
import numpy as np                                       # matrix calculations
import matplotlib.pyplot as plt                          # plots and heat maps
from scipy.stats import chi2                             # chi-square test of model fit
from statsmodels.multivariate.factor import Factor       # factor analysis
from statsmodels.multivariate.factor_rotation import rotate_factors  # factor rotation
from tkinter import Tk, filedialog                       # file-choosing dialog box

## 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

## Set Seed
# Setting the seed ensures your results are the same as the results in this example.
np.random.seed(1)

# Import Data
retailer_survey = pd.read_csv(file_choose("Choose retailer_survey.csv"))  ## Choose retailer_survey.csv file

# Take a first look at the data
# 600 rows: 100 respondents rating each of 6 retailers (Chestnut Ridge = CR, and A-E)
# on 12 statements, each measured on a 1-7 scale
print(retailer_survey.head())
print(retailer_survey["Retailer"].value_counts().sort_index())


# %%
#####################################################
# Step 1: Average Survey Responses by Retailer      #
#####################################################

# Average response to each of the 12 statements for each retailer
retailer_means = retailer_survey.groupby("Retailer").mean()

print("\nAverage Survey Responses by Retailer")
print(retailer_means.round(2).to_string())


# %%
#####################################################
# Step 2: Determine the Number of Factors           #
#####################################################

# Factor analysis uses only the 12 statements, so remove the retailer column
retailer_factors = retailer_survey.drop(columns="Retailer")
statements = retailer_factors.columns

# Eigenvalues of the correlation matrix (sorted from largest to smallest)
eigenvalues = np.sort(np.linalg.eigvalsh(retailer_factors.corr()))[::-1]
print(eigenvalues.round(8))

# Latent root (eigenvalue) criterion: keep factors with an eigenvalue of at least 1
print((eigenvalues >= 1).sum())

# Percentage of variance criterion: share of the variance explained by each factor
variance_explained = pd.DataFrame({"factor": range(1, len(eigenvalues) + 1),
                                   "eigenvalue": eigenvalues,
                                   "pct_variance": eigenvalues / eigenvalues.sum(),
                                   "cumulative_pct": eigenvalues.cumsum() / eigenvalues.sum()})
print(variance_explained.round(3).to_string(index=False))

# Scree plot of the eigenvalues
fig, ax = plt.subplots(figsize=(9, 6))
ax.plot(variance_explained["factor"], variance_explained["eigenvalue"],
        color="blue", marker="o", markerfacecolor="black")
ax.axhline(1, linestyle="--", color="red")
ax.set_xticks(range(1, len(eigenvalues) + 1))
ax.set_title("Scree Plot of Eigenvalues\nDashed line marks an eigenvalue of 1")
ax.set_xlabel("Factor")
ax.set_ylabel("Eigenvalue")
plt.show()

# Four eigenvalues are greater than 1 (3.27, 2.90, 2.67, and 2.05), and together
# they explain about 91% of the variance, so we use a four-factor solution.


# %%
#####################################################
# Step 3: Factor Analysis with Four Factors         #
#####################################################

# Factor analysis software numbers the factors arbitrarily and may flip the sign of a
# factor. To match the R version of this example, this function flips each factor so its
# loadings are mostly positive and orders the factors from largest to smallest
# sum of squared loadings.
# It also returns the new order and the signs so the factor scores can be matched later.
def order_factors(loadings):
    signs = np.sign(loadings.sum(axis=0))
    order = np.argsort(-(loadings ** 2).sum(axis=0))
    return (loadings * signs)[:, order], order, signs[order]

# Run factor analysis with 4 factors using maximum likelihood (the same method as R)
fa_fit = Factor(retailer_factors, n_factor=4, method="ml").fit()
unrotated = np.real(fa_fit.loadings)
uniquenesses = pd.Series(fa_fit.uniqueness, index=statements)

# Varimax rotation (R's default). Each statement's loadings are rescaled to length 1
# before rotating and scaled back afterward (Kaiser normalization, as in R).
row_length = np.sqrt((unrotated ** 2).sum(axis=1, keepdims=True))
varimax, _ = rotate_factors(unrotated / row_length, "varimax")
varimax, _, _ = order_factors(varimax * row_length)
factor_labels = ["Factor1", "Factor2", "Factor3", "Factor4"]
varimax = pd.DataFrame(varimax, index=statements, columns=factor_labels)

print("\nUniquenesses:")
print(uniquenesses.round(3).to_string())

# Loadings close to 0 (below 0.1) are left blank, as in R's output
print("\nLoadings:")
print(varimax.map(lambda v: f"{v:.3f}" if abs(v) >= 0.1 else "").to_string())

ss_loadings = (varimax ** 2).sum()
print(pd.DataFrame({"SS loadings": ss_loadings,
                    "Proportion Var": ss_loadings / len(statements),
                    "Cumulative Var": ss_loadings.cumsum() / len(statements)}).T.round(3).to_string())

# Test of the hypothesis that 4 factors are sufficient (same test R reports)
n, p, m = len(retailer_factors), len(statements), 4
corr = retailer_factors.corr().values
model_corr = unrotated @ unrotated.T + np.diag(fa_fit.uniqueness)
objective = (np.linalg.slogdet(model_corr)[1] - np.linalg.slogdet(corr)[1]
             + np.trace(corr @ np.linalg.inv(model_corr)) - p)
chi_square = (n - 1 - (2 * p + 5) / 6 - 2 * m / 3) * objective
dof = ((p - m) ** 2 - p - m) / 2
print(f"\nThe chi square statistic is {chi_square:.2f} on {dof:.0f} degrees of freedom.")
print(f"The p-value is {chi2.sf(chi_square, dof):.3f}")


# %%
#####################################################
# Step 4: Factor Analysis with Oblique Rotation     #
#####################################################

# Rerun the factor analysis with an oblique (oblimin) rotation
retailer_fa = Factor(retailer_factors, n_factor=4, method="ml").fit()
retailer_fa.rotate("oblimin")
oblimin, order, signs = order_factors(np.real(retailer_fa.loadings))
oblimin = pd.DataFrame(oblimin, index=statements, columns=factor_labels)

# Factor loadings (all loadings are shown, including those close to 0)
print(oblimin.round(2).to_string())

# List the statements so that statements on the same factor are next to each other
statement_order = ["Latest", "Trends", "Stylish", "Quality", "Last", "Fit",
                   "Satisfied", "Purchase", "Recommend", "Value", "Bargain", "Worth"]
loadings_plot = oblimin.loc[statement_order]

# Heat map of factor loadings (darker red means a stronger factor loading)
fig, ax = plt.subplots(figsize=(8, 7))
image = ax.imshow(loadings_plot, cmap="Reds", aspect="auto")
for i in range(loadings_plot.shape[0]):
    for j in range(loadings_plot.shape[1]):
        ax.text(j, i, f"{loadings_plot.iloc[i, j]:.2f}", ha="center", va="center")
ax.set_xticks(range(len(factor_labels)), factor_labels)
ax.set_yticks(range(len(statement_order)), statement_order)
fig.colorbar(image, ax=ax, label="Loading")
ax.set_title("Factor Loadings from Survey")
plt.show()


# %%
#####################################################
# Step 5: Name the Factors and Score the Retailers  #
#####################################################

# The three statements with the highest loadings on each factor (Table 11.4)
for factor in factor_labels:
    print(factor, ":", ", ".join(oblimin[factor].nlargest(3).index))

# Name each factor to capture the essence of the statements that load on it
factor_names = ["Innovative", "High Quality", "Loyalty", "Good Value"]

# Bartlett factor scores for each response (put in the same order and sign as the loadings),
# with the retailer that was rated
scores = np.real(retailer_fa.factor_scoring(method="bartlett"))
retailer_scores = pd.DataFrame(scores[:, order] * signs, columns=factor_names)
retailer_scores["Retailer"] = retailer_survey["Retailer"]

# Average factor score for each retailer
# Factor scores are standardized (mean 0, standard deviation 1), so a positive score
# means the retailer is perceived above average on that factor
retailer_fa_mean = retailer_scores.groupby("Retailer").mean()

print("\nAverage Factor Scores by Retailer")
print(retailer_fa_mean.round(3).to_string())


# %%
#####################################################
# Step 6: Heat Map of Factor Scores by Retailer     #
#####################################################

# Heat map of factor scores by retailer (darker blue means a higher factor score)
fig, ax = plt.subplots(figsize=(8, 6))
image = ax.imshow(retailer_fa_mean, cmap="Blues", aspect="auto")
for i in range(retailer_fa_mean.shape[0]):
    for j in range(retailer_fa_mean.shape[1]):
        ax.text(j, i, f"{retailer_fa_mean.iloc[i, j]:.2f}", ha="center", va="center")
ax.set_xticks(range(len(factor_names)), factor_names)
ax.set_yticks(range(len(retailer_fa_mean)), retailer_fa_mean.index)
ax.set_ylabel("Retailer")
fig.colorbar(image, ax=ax, label="Factor Score")
ax.set_title("Factor Score by Retailer")
plt.show()
