# cluster_situations_kmeans_mpnet.py
import pandas as pd
import numpy as np
from sklearn.preprocessing import normalize
from sklearn.cluster import KMeans
from sklearn.metrics import silhouette_score
from sentence_transformers import SentenceTransformer

INPUT_FILE = "jtbd_openai_results.csv"
OUTPUT_FILE = "jtbd_situations_with_clusters.csv"

# 1. Load results
df = pd.read_csv(INPUT_FILE)

# Keep only rows where GPT said jtbd=YES and situation is not empty
df_yes = df[df["jtbd"] == "YES"].dropna(subset=["situation"]).reset_index(drop=True)

# Remove duplicate comments
df_yes = df_yes.drop_duplicates(subset=["comment"]).reset_index(drop=True)

print(f"👉 Loaded {len(df_yes)} unique comments with situations for clustering")

# 2. Extract situations
situations = df_yes["situation"].tolist()

# 3. Embed situations with mpnet
print("Embedding situations with all-mpnet-base-v2...")
model = SentenceTransformer("sentence-transformers/all-mpnet-base-v2")
X = model.encode(situations, show_progress_bar=True)

# 4. Normalize
X_norm = normalize(X)

# 5. Find best k using silhouette score
print("Finding best k...")
best_k = None
best_score = -1
for k in range(2, 21):  # test k = 2...20
    kmeans = KMeans(n_clusters=k, random_state=42, n_init=10)
    labels = kmeans.fit_predict(X_norm)
    if len(set(labels)) > 1:  # silhouette needs >1 cluster
        score = silhouette_score(X_norm, labels)
        print(f"k={k}, silhouette={score:.3f}")
        if score > best_score:
            best_score = score
            best_k = k

print(f"\n✅ Best k = {best_k} with silhouette score {best_score:.3f}")

# 6. Run final clustering
kmeans = KMeans(n_clusters=best_k, random_state=42, n_init=10)
labels = kmeans.fit_predict(X_norm)

df_yes["situation_cluster"] = labels

# 7. Save
df_yes.to_csv(OUTPUT_FILE, index=False)
print(f"\n✅ Clustered situations saved to {OUTPUT_FILE}")

# 8. Show cluster sizes + examples
print("\nCluster distribution:")
for cluster_id in sorted(set(labels)):
    cluster_df = df_yes[df_yes["situation_cluster"] == cluster_id]
    print(f"Cluster {cluster_id}: {len(cluster_df)} situations")
    for ex in cluster_df["situation"].head(3).tolist():
        print("  -", ex[:120])
