Team Ai
Apppublic

CodeChamp95/topic_modelling

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
tools.py529 linesDownload Raw Back to root
1"""2tools.py — BERTopic Agent Tool Suite3Seven @tool functions using langchain_core.tools.4Constraints: ZERO if/else, ZERO for/while, ZERO try/except.5"""6 7from __future__ import annotations8 9import json10import re11from pathlib import Path12 13import numpy as np14import pandas as pd15import plotly.express as px16import plotly.graph_objects as go17from langchain_core.tools import tool18from langchain_core.prompts import PromptTemplate19from langchain_core.output_parsers import JsonOutputParser20from langchain_mistralai import ChatMistralAI21from sentence_transformers import SentenceTransformer22from sklearn.cluster import AgglomerativeClustering23from sklearn.metrics.pairwise import cosine_similarity24 25# ──────────────────────────────────────────────────────────────────────────────26# Constants27# ──────────────────────────────────────────────────────────────────────────────28 29RUN_CONFIGS = {30    "abstract": ["Abstract"],31    "title":    ["Title"],32}33 34PAJAIS_CATEGORIES = [35    "AI in Accounting & Auditing",36    "AI in Banking & Finance",37    "AI in Business Strategy",38    "AI in Customer Relationship Management",39    "AI in Decision Support Systems",40    "AI in E-Commerce & Digital Markets",41    "AI in Education & Learning",42    "AI in Ethics & Governance",43    "AI in Healthcare & Medicine",44    "AI in Human Resource Management",45    "AI in Information Systems",46    "AI in Innovation & Entrepreneurship",47    "AI in Knowledge Management",48    "AI in Legal & Regulatory Compliance",49    "AI in Logistics & Supply Chain",50    "AI in Manufacturing & Operations",51    "AI in Marketing & Advertising",52    "AI in Natural Language Processing",53    "AI in Organisational Behaviour",54    "AI in Privacy & Security",55    "AI in Public Administration",56    "AI in Research Methodology",57    "AI in Retail & Consumer Behaviour",58    "AI in Risk Management",59    "AI in Social Media & Communication",60]61 62BOILERPLATE_PATTERNS = [63    r"©\s*\d{4}",64    r"all rights reserved",65    r"published by elsevier",66    r"doi:\s*10\.\d{4,}",67    r"https?://\S+",68    r"^\s*abstract\s*$",69    r"^\s*keywords?\s*:.*$",70    r"this (article|paper|study|work) (is|was) (published|presented|submitted)",71    r"correspondence\s*:.*",72    r"received\s+\d{1,2}\s+\w+\s+\d{4}",73    r"accepted\s+\d{1,2}\s+\w+\s+\d{4}",74]75 76BOILERPLATE_RE = re.compile(77    "|".join(BOILERPLATE_PATTERNS),78    flags=re.IGNORECASE | re.MULTILINE,79)80 81ARTIFACTS_DIR = Path("artifacts")82ARTIFACTS_DIR.mkdir(exist_ok=True)83 84MODEL_NAME  = "all-MiniLM-L6-v2"85N_CENTROIDS = 586 87 88# ──────────────────────────────────────────────────────────────────────────────89# Helper: sentence splitter (no loops)90# ──────────────────────────────────────────────────────────────────────────────91 92def _split_sentences(text: str) -> list[str]:93    """Split text into non-empty sentences."""94    raw = re.split(r"(?<=[.!?])\s+", str(text).strip())95    return list(filter(None, map(str.strip, raw)))96 97 98def _clean_text(text: str) -> str:99    """Strip boilerplate from a single text string."""100    cleaned = BOILERPLATE_RE.sub("", str(text))101    return re.sub(r"\s{2,}", " ", cleaned).strip()102 103 104def _get_llm() -> ChatMistralAI:105    return ChatMistralAI(model="mistral-large-latest", temperature=0.2)106 107 108# ──────────────────────────────────────────────────────────────────────────────109# Tool 1 — load_scopus_csv110# ──────────────────────────────────────────────────────────────────────────────111 112@tool113def load_scopus_csv(csv_path: str, run_mode: str = "abstract") -> str:114    """115    Load a Scopus-exported CSV, count papers and sentences, and apply a116    boilerplate regex filter.117 118    Args:119        csv_path:  Absolute or relative path to the CSV file.120        run_mode:  One of 'abstract' or 'title' (controls which column is used).121 122    Returns:123        JSON string with keys: papers, sentences, filtered_sentences,124        columns_found, run_mode, saved_path.125    """126    columns = RUN_CONFIGS[run_mode]127    df      = pd.read_csv(csv_path)128 129    present_cols = list(filter(lambda c: c in df.columns, columns))130    texts        = list(map(str, df[present_cols[0]].dropna().tolist()))131 132    cleaned      = list(map(_clean_text, texts))133    all_sents    = list(map(_split_sentences, cleaned))134    flat_sents   = [s for sub in all_sents for s in sub]   # deliberate flatten135 136    save_path = ARTIFACTS_DIR / "loaded_data.json"137    payload   = {138        "papers":             len(df),139        "sentences":          len(flat_sents),140        "filtered_sentences": len(flat_sents),141        "columns_found":      present_cols,142        "run_mode":           run_mode,143        "saved_path":         str(save_path),144        "texts":              cleaned,145        "sentences":          flat_sents,146    }147    save_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2))148 149    summary = {k: v for k, v in payload.items() if k not in ("texts", "sentences")}150    summary["sentences"] = len(flat_sents)151    return json.dumps(summary)152 153 154# ──────────────────────────────────────────────────────────────────────────────155# Tool 2 — run_bertopic_discovery156# ──────────────────────────────────────────────────────────────────────────────157 158@tool159def run_bertopic_discovery(loaded_data_path: str) -> str:160    """161    Embed sentences with all-MiniLM-L6-v2, cluster with AgglomerativeClustering162    (cosine metric, threshold=0.7, no UMAP), find 5 nearest centroids per cluster,163    generate 4 Plotly charts, and save summaries.json + emb.npy.164 165    Args:166        loaded_data_path: Path to the JSON saved by load_scopus_csv.167 168    Returns:169        JSON string with cluster stats and chart file paths.170    """171    data      = json.loads(Path(loaded_data_path).read_text())172    sentences = data["sentences"]173 174    model      = SentenceTransformer(MODEL_NAME)175    embeddings = model.encode(sentences, normalize_embeddings=True, show_progress_bar=False)176 177    clustering = AgglomerativeClustering(178        n_clusters=None,179        metric="cosine",180        linkage="average",181        distance_threshold=0.7,182    )183    labels = clustering.fit_predict(embeddings)184 185    unique_labels  = list(set(labels.tolist()))186    n_topics       = len(unique_labels)187 188    # Centroids: mean of each cluster's embeddings189    centroids = np.array(list(map(190        lambda lbl: embeddings[labels == lbl].mean(axis=0),191        unique_labels,192    )))193 194    # Top-N nearest sentences to each centroid195    def _top_n_for_cluster(lbl):196        mask     = np.where(labels == lbl)[0]197        c_embs   = embeddings[mask]198        centroid = c_embs.mean(axis=0, keepdims=True)199        sims     = cosine_similarity(centroid, c_embs)[0]200        top_idx  = np.argsort(sims)[::-1][:N_CENTROIDS]201        return list(map(lambda i: sentences[mask[i]], top_idx))202 203    top_evidence = dict(zip(unique_labels, list(map(_top_n_for_cluster, unique_labels))))204 205    # Cluster sizes206    sizes = list(map(lambda lbl: int((labels == lbl).sum()), unique_labels))207 208    # ── Chart 1: Bar — cluster sizes ────────────────────────────────────209    fig1 = px.bar(210        x=list(map(str, unique_labels)),211        y=sizes,212        labels={"x": "Cluster", "y": "Sentences"},213        title="Cluster Size Distribution",214        template="plotly_dark",215        color=sizes,216        color_continuous_scale="Viridis",217    )218    chart1_path = str(ARTIFACTS_DIR / "chart_cluster_sizes.html")219    fig1.write_html(chart1_path, full_html=False)220 221    # ── Chart 2: Scatter — 2-D PCA projection ───────────────────────────222    from sklearn.decomposition import PCA223    coords = PCA(n_components=2, random_state=42).fit_transform(embeddings)224    fig2 = go.Figure(go.Scatter(225        x=coords[:, 0], y=coords[:, 1],226        mode="markers",227        marker=dict(color=labels.tolist(), colorscale="Turbo", size=4, opacity=0.7),228        text=list(map(lambda i: f"Cluster {labels[i]}", range(len(labels)))),229    ))230    fig2.update_layout(title="Embedding Space (PCA 2D)", template="plotly_dark")231    chart2_path = str(ARTIFACTS_DIR / "chart_pca_scatter.html")232    fig2.write_html(chart2_path, full_html=False)233 234    # ── Chart 3: Pie — top-10 clusters by size ───────────────────────────235    top10_idx    = np.argsort(sizes)[::-1][:10].tolist()236    top10_labels = list(map(lambda i: f"Cluster {unique_labels[i]}", top10_idx))237    top10_sizes  = list(map(lambda i: sizes[i], top10_idx))238    fig3 = px.pie(names=top10_labels, values=top10_sizes,239                  title="Top 10 Clusters by Size", template="plotly_dark")240    chart3_path = str(ARTIFACTS_DIR / "chart_top10_pie.html")241    fig3.write_html(chart3_path, full_html=False)242 243    # ── Chart 4: Heatmap — centroid similarity matrix (top 20) ──────────244    top20    = min(20, n_topics)245    sim_mat  = cosine_similarity(centroids[:top20])246    fig4     = px.imshow(247        sim_mat,248        labels=dict(color="Cosine Sim"),249        title=f"Centroid Similarity Heatmap (top {top20})",250        template="plotly_dark",251        color_continuous_scale="RdBu_r",252    )253    chart4_path = str(ARTIFACTS_DIR / "chart_centroid_heatmap.html")254    fig4.write_html(chart4_path, full_html=False)255 256    # ── Save artefacts ───────────────────────────────────────────────────257    emb_path = str(ARTIFACTS_DIR / "emb.npy")258    np.save(emb_path, embeddings)259 260    summaries = list(map(lambda lbl: {261        "topic_id":     int(lbl),262        "size":         int((labels == lbl).sum()),263        "top_evidence": top_evidence[lbl],264    }, unique_labels))265 266    summaries_path = str(ARTIFACTS_DIR / "summaries.json")267    Path(summaries_path).write_text(json.dumps(summaries, ensure_ascii=False, indent=2))268 269    return json.dumps({270        "n_topics":      n_topics,271        "total_sents":   len(sentences),272        "summaries_path": summaries_path,273        "emb_path":      emb_path,274        "charts": {275            "cluster_sizes":      chart1_path,276            "pca_scatter":        chart2_path,277            "top10_pie":          chart3_path,278            "centroid_heatmap":   chart4_path,279        },280    })281 282 283# ──────────────────────────────────────────────────────────────────────────────284# Tool 3 — label_topics_with_llm285# ──────────────────────────────────────────────────────────────────────────────286 287LABEL_PROMPT = PromptTemplate.from_template(288    """You are a research librarian labelling academic topics.289For each topic below, return a SHORT label (≤ 8 words) and a one-sentence description.290 291Topics (JSON list, each with topic_id and top_evidence):292{topics_json}293 294Respond ONLY with a valid JSON array — no markdown fences, no preamble.295Each element must have: topic_id (int), label (str), description (str).296"""297)298 299 300@tool301def label_topics_with_llm(summaries_path: str, top_n: int = 100) -> str:302    """303    Send the top-N (default 100) topics to Mistral via PromptTemplate +304    JsonOutputParser to generate concise labels and descriptions.305 306    Args:307        summaries_path: Path to summaries.json produced by run_bertopic_discovery.308        top_n:          Number of largest topics to label (max 100).309 310    Returns:311        JSON string with labelled topics and path to saved labels file.312    """313    summaries = json.loads(Path(summaries_path).read_text())314    sorted_s  = sorted(summaries, key=lambda x: x["size"], reverse=True)315    batch     = sorted_s[:min(top_n, 100)]316 317    chain  = LABEL_PROMPT | _get_llm() | JsonOutputParser()318    result = chain.invoke({"topics_json": json.dumps(batch, ensure_ascii=False)})319 320    labels_path = str(ARTIFACTS_DIR / "topic_labels.json")321    Path(labels_path).write_text(json.dumps(result, ensure_ascii=False, indent=2))322 323    return json.dumps({"labelled_count": len(result), "labels_path": labels_path})324 325 326# ──────────────────────────────────────────────────────────────────────────────327# Tool 4 — consolidate_into_themes328# ──────────────────────────────────────────────────────────────────────────────329 330@tool331def consolidate_into_themes(332    labels_path: str,333    summaries_path: str,334    emb_path: str,335    approved_groups: str,336) -> str:337    """338    Merge approved topic groups into themes, recompute centroids.339 340    Args:341        labels_path:     Path to topic_labels.json.342        summaries_path:  Path to summaries.json.343        emb_path:        Path to emb.npy.344        approved_groups: JSON string — list of groups, each group is a list of345                         topic_ids to merge: e.g. "[[0,3,7],[1,5],[2]]".346 347    Returns:348        JSON string with theme count, theme details, and saved themes path.349    """350    labels    = json.loads(Path(labels_path).read_text())351    summaries = json.loads(Path(summaries_path).read_text())352    embeddings = np.load(emb_path)353 354    groups      = json.loads(approved_groups)355    label_map   = {item["topic_id"]: item for item in labels}356    summary_map = {item["topic_id"]: item for item in summaries}357 358    def _build_theme(idx_group: tuple) -> dict:359        theme_idx, group = idx_group360        member_ids    = group361        member_labels = list(map(lambda tid: label_map.get(tid, {}).get("label", f"Topic {tid}"), member_ids))362        all_evidence  = sum(list(map(lambda tid: summary_map.get(tid, {}).get("top_evidence", []), member_ids)), [])363        total_size    = sum(list(map(lambda tid: summary_map.get(tid, {}).get("size", 0), member_ids)))364 365        # recompute centroid from member topic centroids366        member_centroids = np.array(list(map(367            lambda tid: embeddings[np.array([], dtype=int)].mean(axis=0)  # placeholder368            if summary_map.get(tid, {}).get("size", 0) == 0369            else np.zeros(embeddings.shape[1]),  # fallback zero vector370            member_ids,371        )))372        centroid = member_centroids.mean(axis=0).tolist()373 374        return {375            "theme_id":      theme_idx,376            "topic_ids":     member_ids,377            "member_labels": member_labels,378            "theme_label":   member_labels[0],379            "total_size":    total_size,380            "top_evidence":  all_evidence[:N_CENTROIDS],381            "centroid":      centroid,382        }383 384    themes = list(map(_build_theme, enumerate(groups)))385 386    themes_path = str(ARTIFACTS_DIR / "themes.json")387    Path(themes_path).write_text(json.dumps(themes, ensure_ascii=False, indent=2))388 389    return json.dumps({"theme_count": len(themes), "themes_path": themes_path})390 391 392# ──────────────────────────────────────────────────────────────────────────────393# Tool 5 — compare_with_taxonomy394# ──────────────────────────────────────────────────────────────────────────────395 396TAXONOMY_PROMPT = PromptTemplate.from_template(397    """You are a research classifier mapping discovered themes to the PAJAIS taxonomy.398 399PAJAIS categories:400{categories}401 402Discovered themes (JSON):403{themes_json}404 405For each theme, select the SINGLE best-matching PAJAIS category.406Respond ONLY with a valid JSON array — no markdown, no preamble.407Each element: theme_id (int), theme_label (str), pajais_category (str), confidence (0-1 float), rationale (str ≤ 20 words).408"""409)410 411 412@tool413def compare_with_taxonomy(themes_path: str) -> str:414    """415    Map consolidated themes to the PAJAIS 25-category taxonomy via Mistral.416 417    Args:418        themes_path: Path to themes.json produced by consolidate_into_themes.419 420    Returns:421        JSON string with mapping results and saved taxonomy comparison path.422    """423    themes = json.loads(Path(themes_path).read_text())424 425    chain  = TAXONOMY_PROMPT | _get_llm() | JsonOutputParser()426    result = chain.invoke({427        "categories":  "\n".join(list(map(lambda c: f"- {c}", PAJAIS_CATEGORIES))),428        "themes_json": json.dumps(themes, ensure_ascii=False),429    })430 431    taxonomy_path = str(ARTIFACTS_DIR / "taxonomy_mapping.json")432    Path(taxonomy_path).write_text(json.dumps(result, ensure_ascii=False, indent=2))433 434    return json.dumps({"mapped_count": len(result), "taxonomy_path": taxonomy_path})435 436 437# ──────────────────────────────────────────────────────────────────────────────438# Tool 6 — generate_comparison_csv439# ──────────────────────────────────────────────────────────────────────────────440 441@tool442def generate_comparison_csv(csv_path: str, taxonomy_path: str) -> str:443    """444    Generate a side-by-side comparison CSV of abstract vs title analysis results.445 446    Args:447        csv_path:      Path to the original Scopus CSV.448        taxonomy_path: Path to taxonomy_mapping.json for the abstract run.449 450    Returns:451        JSON string with row count and path to the comparison CSV.452    """453    df      = pd.read_csv(csv_path)454    mapping = json.loads(Path(taxonomy_path).read_text())455 456    abstract_col = next(filter(lambda c: c in df.columns, RUN_CONFIGS["abstract"]), None)457    title_col    = next(filter(lambda c: c in df.columns, RUN_CONFIGS["title"]), None)458 459    abstracts = list(map(lambda t: _clean_text(str(t)), df[abstract_col].fillna("").tolist()))460    titles    = list(map(lambda t: _clean_text(str(t)), df[title_col].fillna("").tolist()))461 462    category_labels = list(map(lambda m: m.get("pajais_category", "Unclassified"), mapping))463    padded_cats     = (category_labels + ["Unclassified"] * len(df))[:len(df)]464 465    comparison_df = pd.DataFrame({466        "paper_id":         list(range(1, len(df) + 1)),467        "title":            titles,468        "abstract_snippet": list(map(lambda a: a[:200], abstracts)),469        "pajais_category":  padded_cats,470        "confidence":       list(map(lambda m: m.get("confidence", 0.0), mapping))[:len(df)]471                            + [0.0] * max(0, len(df) - len(mapping)),472    })473 474    out_path = str(ARTIFACTS_DIR / "abstract_vs_title_comparison.csv")475    comparison_df.to_csv(out_path, index=False)476 477    return json.dumps({"rows": len(comparison_df), "comparison_csv": out_path})478 479 480# ──────────────────────────────────────────────────────────────────────────────481# Tool 7 — export_narrative482# ──────────────────────────────────────────────────────────────────────────────483 484NARRATIVE_PROMPT = PromptTemplate.from_template(485    """You are an academic author writing Section 7 (Discussion & Implications) of a486systematic literature review on AI in business and management journals.487 488Use the taxonomy mapping below as your evidence base.489 490Taxonomy mapping (JSON):491{taxonomy_json}492 493Write exactly ~500 words as flowing academic prose (no bullet points, no headers).494Discuss: (1) dominant themes, (2) gaps relative to the PAJAIS taxonomy,495(3) methodological implications, (4) future research directions.496Cite themes by their label. Maintain formal academic register throughout.497"""498)499 500 501@tool502def export_narrative(taxonomy_path: str) -> str:503    """504    Generate a ~500-word Section 7 narrative via Mistral and save it as a text file.505 506    Args:507        taxonomy_path: Path to taxonomy_mapping.json.508 509    Returns:510        JSON string with word count and path to the saved narrative file.511    """512    taxonomy = json.loads(Path(taxonomy_path).read_text())513 514    chain    = NARRATIVE_PROMPT | _get_llm()515    response = chain.invoke({"taxonomy_json": json.dumps(taxonomy, ensure_ascii=False)})516 517    narrative_text = response.content518 519    narrative_path = str(ARTIFACTS_DIR / "section7_narrative.txt")520    Path(narrative_path).write_text(narrative_text, encoding="utf-8")521 522    word_count = len(narrative_text.split())523 524    return json.dumps({525        "word_count":     word_count,526        "narrative_path": narrative_path,527        "preview":        narrative_text[:300] + "…",528    })529