##################################################################
# Chapter 17: Using Topic Models to Glean Customer Insights      #
# Topic Model of Apple Watch Product Reviews (Python version)    #
##################################################################

# This script reproduces the full Apple Watch product review topic modeling example in Python:
#   Step 1. Look at a sample of the data
#   Step 2. Process the documents (clean and stem the review text)
#   Step 3. Determine the number of topics and estimate the topic model
#   Step 4. See which topics relate to high vs. low ratings
#   Step 5. Identify the top 3 (positive) and bottom 3 (negative) topics
#   Step 6. Visualize the topics (word clouds)
# 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.

# %%
## Python does not have a Structural Topic Model (STM) package, so this
## version mirrors the R example as closely as possible:
##   - The number of topics is chosen automatically using the same
##     approach as stm(K = 0): anchor words found with t-SNE + convex hull
##   - The topic model is a Dirichlet-Multinomial Regression (DMR) model,
##     where (like STM's prevalence =~ review_star) topic prevalence
##     depends on the star rating
##   - Topics are then related to ratings with a regression, as in
##     estimateEffect()

## 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", "scipy": "scipy", "sklearn": "scikit-learn",
            "statsmodels": "statsmodels", "matplotlib": "matplotlib", "wordcloud": "wordcloud",
            "nltk": "nltk", "tomotopy": "tomotopy"}
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 re
import string
from tkinter import Tk, filedialog

import numpy as np
import pandas as pd
import scipy.sparse as sp
from scipy.spatial import ConvexHull
from sklearn.decomposition import PCA
from sklearn.manifold import TSNE
import statsmodels.formula.api as smf
import matplotlib.pyplot as plt
from wordcloud import WordCloud
import nltk
from nltk.corpus import stopwords
from nltk.stem.snowball import SnowballStemmer
import tomotopy as tp

nltk.download("stopwords", quiet=True)
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 the Product Reviews Data
reviews = pd.read_csv(file_choose("Choose reviews_data.csv"),  ## Choose the file reviews_data.csv
                      encoding="utf-8", encoding_errors="replace")

## Make Sure the Text is Encoded Properly (works on both Mac and PC)
reviews["review_text"] = reviews["review_text"].str.replace("�", " ")


# %%
##################################################################
# Step 1: Look at a Sample of the Data                           #
##################################################################

print(reviews.head())
print(reviews["review_star"].value_counts().sort_index())


# %%
##################################################################
# Step 2: Process Documents                                      #
##################################################################

## (lowercase, remove punctuation, stop words, custom words, and numbers,
##  stem words, and keep words with at least 3 characters)
customwords = ["watch", "iphone", "apple"]
stop_words = set(stopwords.words("english"))
stemmer = SnowballStemmer("english")

def process_text(text):
    text = str(text).lower()
    text = text.translate(str.maketrans("", "", string.punctuation))
    words = [w for w in text.split() if w not in stop_words and w not in customwords]
    words = [re.sub(r"[0-9]", "", w) for w in words]
    words = [stemmer.stem(w) for w in words if w != ""]
    return [w for w in words if len(w) >= 3]

reviews["words"] = reviews["review_text"].apply(process_text)

## Remove words that appear in only one review, then remove empty reviews
doc_freq = pd.Series([w for words in reviews["words"] for w in set(words)]).value_counts()
keep_words = set(doc_freq[doc_freq > 1].index)
reviews["words"] = reviews["words"].apply(lambda words: [w for w in words if w in keep_words])
meta = reviews[reviews["words"].str.len() > 0].reset_index(drop=True)
vocab = sorted(keep_words)
print(f"Your corpus now has {len(meta)} documents, {len(vocab)} terms and "
      f"{meta['words'].str.len().sum()} tokens.")


# %%
##################################################################
# Step 3: Determine Number of Topics and Estimate the Topic      #
#         Model                                                  #
##################################################################

## Determine Number of Topics (same approach as stm with K = 0)
## First, build a word co-occurrence matrix (normalized so each row sums to 1)
word_index = {w: i for i, w in enumerate(vocab)}
rows, cols = [], []
for d, words in enumerate(meta["words"]):
    rows += [d] * len(words)
    cols += [word_index[w] for w in words]
dtm = sp.csr_matrix((np.ones(len(rows)), (rows, cols)), shape=(len(meta), len(vocab)))
nd = np.asarray(dtm.sum(axis=1)).ravel()
dtm, nd = dtm[nd >= 2], nd[nd >= 2]
divisor = nd * (nd - 1)
Htilde = sp.diags(1 / np.sqrt(divisor)) @ dtm
Hhat = np.asarray((sp.diags(1 / divisor) @ dtm).sum(axis=0)).ravel()
Q = (Htilde.T @ Htilde).toarray() - np.diag(Hhat)
Qsums = Q.sum(axis=1)
Q = Q[Qsums > 0][:, Qsums > 0]  ## drop words that never co-occur with others
Qbar = Q / Q.sum(axis=1, keepdims=True)

## Then project words into 3 dimensions and find the anchor words
## (the corners of the convex hull); each anchor word defines one topic
xpca = PCA(n_components=50, random_state=1).fit_transform(Qbar)
proj = TSNE(n_components=3, perplexity=30, init="random", random_state=1).fit_transform(xpca)
anchors = ConvexHull(proj).vertices

## See How Many Topics
num_topics = len(anchors)
print(num_topics)

## Estimate the Topic Model (topic prevalence depends on the star rating)
reviewsFit = tp.DMRModel(k=num_topics, min_cf=0, seed=1)
for words, star in zip(meta["words"], meta["review_star"]):
    reviewsFit.add_doc(words, metadata=str(star))
reviewsFit.train(1000, workers=1)

## Topic proportions for each review (columns: Topic 1, Topic 2, ...)
theta = np.array([doc.get_topic_dist() for doc in reviewsFit.docs])


# %%
##################################################################
# Step 4: See Which Topics Relate to High vs. Low Ratings        #
##################################################################

## (difference in topic prevalence between 5-star and 1-star reviews)
meta["rating"] = meta["review_star"].astype("category")
results = []
for k in range(num_topics):
    meta["topic"] = theta[:, k]
    fit = smf.ols("topic ~ C(rating)", data=meta).fit()
    ci = fit.conf_int().loc["C(rating)[T.5]"]
    results.append({"topic": k + 1, "difference": fit.params["C(rating)[T.5]"],
                    "lower": ci[0], "upper": ci[1]})
topic_diff = pd.DataFrame(results)

fig, ax = plt.subplots(figsize=(8, 12))
y = np.arange(num_topics, 0, -1)
ax.hlines(y, topic_diff["lower"], topic_diff["upper"], color="black")
ax.plot(topic_diff["difference"], y, "o", color="black")
for yi, row in zip(y, topic_diff.itertuples()):
    ax.text(row.difference, yi + 0.3, str(row.topic), ha="center", fontsize=8)
ax.axvline(0, linestyle="--", color="grey")
ax.set_yticks([])
ax.set_xlabel("Lower Rating ... Higher Rating")
ax.set_title("Relationship between Topic and Rating")
plt.show()


# %%
##################################################################
# Step 5: Identify the Top 3 (Positive) and Bottom 3 (Negative)  #
#         Topics                                                 #
##################################################################

topic_diff = topic_diff.sort_values("difference", ascending=False)
positive_topics = topic_diff["topic"].head(3).tolist()
negative_topics = topic_diff["topic"].tail(3).tolist()[::-1]  ## most negative first
print(positive_topics)
print(negative_topics)

## Top Words in the Positive and Negative Topics
for k in positive_topics + negative_topics:
    top_words = [w for w, p in reviewsFit.get_topic_words(k - 1, top_n=7)]
    print(f"Topic {k} Top Words: {', '.join(top_words)}")


# %%
##################################################################
# Step 6: Visualize Topics                                       #
##################################################################

## Each word cloud is drawn as its own plot so it can be larger
def topic_cloud(k, title):
    word_probs = dict(zip(reviewsFit.used_vocabs, reviewsFit.get_topic_word_dist(k - 1)))
    cloud = WordCloud(background_color="white", max_words=100, width=1200, height=800, random_state=1)
    fig, ax = plt.subplots(figsize=(10, 7))
    ax.imshow(cloud.generate_from_frequencies(word_probs), interpolation="bilinear")
    ax.set_title(title, fontsize=16)
    ax.axis("off")
    plt.show()

## Positive Topics - Top 3
for k in positive_topics:
    topic_cloud(k, f"Positive Topic {k}")

## Negative Topics - Bottom 3
for k in negative_topics:
    topic_cloud(k, f"Negative Topic {k}")


