Team Ai
Apppublic

vt12w/Topic_Modelling

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
tools.py562 linesDownload Raw Back to root
1"""2tools.py — 7 LangChain @tool functions for BERTopic thematic analysis.3Zero if/else, zero for/while, zero try/except.4"""5 6import json7import os8import numpy as np9import pandas as pd10import re11import plotly.express as px12import plotly.graph_objects as go13 14from langchain_core.tools import tool15from sentence_transformers import SentenceTransformer16from sklearn.cluster import AgglomerativeClustering17from sklearn.metrics.pairwise import cosine_similarity18from langchain_google_genai import ChatGoogleGenerativeAI19from langchain_core.prompts import PromptTemplate20from langchain_core.output_parsers import JsonOutputParser21 22# ---------------------------------------------------------------------------23# Constants24# ---------------------------------------------------------------------------25 26BOILERPLATE_PATTERNS = [27    r"^(abstract|introduction|conclusion|references|acknowledgements?)[\s:]*$",28    r"^\s*\d+\s*$",29    r"^(et al|ibid|op cit)[\.,]",30    r"^\s*$",31]32 33RUN_CONFIGS = {34    "abstract": ["Abstract"],35    "title": ["Title"],36}37 38_model_cache: dict = {}39 40 41def _get_model() -> SentenceTransformer:42    key = "all-MiniLM-L6-v2"43    _model_cache.setdefault(key, SentenceTransformer(key))44    return _model_cache[key]45 46 47def _is_boilerplate(text: str) -> bool:48    return any(bool(re.search(p, text.strip(), re.IGNORECASE)) for p in BOILERPLATE_PATTERNS)49 50 51def _load_state() -> dict:52    return json.loads(open("state.json").read()) if os.path.exists("state.json") else {}53 54 55def _save_state(state: dict) -> None:56    with open("state.json", "w") as f:57        json.dump(state, f, indent=2, default=str)58 59 60# ---------------------------------------------------------------------------61# Tool 1 — load_scopus_csv62# ---------------------------------------------------------------------------63 64@tool(handle_tool_error=True)65def load_scopus_csv(file_path: str, run_config: str = "abstract") -> str:66    """Load a Scopus CSV export, count papers and sentences, and apply boilerplate filtering.67 68    Args:69        file_path: Path to the Scopus CSV file.70        run_config: Either 'abstract' or 'title' to select which columns to analyse.71 72    Returns:73        JSON string with paper count, sentence count, column names, and a sample of sentences.74    """75    df = pd.read_csv(file_path)76    columns = RUN_CONFIGS.get(run_config, RUN_CONFIGS["abstract"])77    available = list(filter(lambda c: c in df.columns, columns))78 79    raw_sentences = list(80        map(81            lambda row: str(row).strip(),82            pd.concat([df[c].dropna() for c in available]).tolist(),83        )84    )85 86    clean_sentences = list(filter(lambda s: not _is_boilerplate(s) and len(s) > 20, raw_sentences))87 88    state = _load_state()89    state.update(90        {91            "file_path": file_path,92            "run_config": run_config,93            "columns": available,94            "paper_count": len(df),95            "sentence_count": len(clean_sentences),96            "sentences": clean_sentences,97            "column_names": list(df.columns),98        }99    )100    _save_state(state)101 102    return json.dumps(103        {104            "paper_count": len(df),105            "sentence_count": len(clean_sentences),106            "boilerplate_removed": len(raw_sentences) - len(clean_sentences),107            "columns_used": available,108            "sample_sentences": clean_sentences[:5],109            "all_columns": list(df.columns),110        }111    )112 113 114# ---------------------------------------------------------------------------115# Tool 2 — run_bertopic_discovery116# ---------------------------------------------------------------------------117 118@tool(handle_tool_error=True)119def run_bertopic_discovery(n_topics_hint: int = 50) -> str:120    """Embed sentences, cluster with AgglomerativeClustering, find top topics, generate Plotly charts.121 122    Args:123        n_topics_hint: Approximate number of topics to discover (used as max).124 125    Returns:126        JSON string with topic summaries and chart HTML paths.127    """128    state = _load_state()129    sentences = state["sentences"]130 131    model = _get_model()132    embeddings = model.encode(sentences, normalize_embeddings=True, show_progress_bar=False)133    np.save("emb.npy", embeddings)134 135    clustering = AgglomerativeClustering(136        n_clusters=None,137        metric="cosine",138        linkage="average",139        distance_threshold=0.3,140    )141    labels = clustering.fit_predict(embeddings)142 143    unique_labels = list(set(labels.tolist()))144    cluster_data = list(145        map(146            lambda lbl: {147                "topic_id": int(lbl),148                "sentence_indices": list(np.where(labels == lbl)[0]),149                "size": int(np.sum(labels == lbl)),150            },151            unique_labels,152        )153    )154    cluster_data.sort(key=lambda x: x["size"], reverse=True)155    top_clusters = cluster_data[:min(n_topics_hint, len(cluster_data))]156 157    def _centroid_and_evidence(cd: dict) -> dict:158        idxs = cd["sentence_indices"]159        vecs = embeddings[idxs]160        centroid = vecs.mean(axis=0)161        sims = cosine_similarity([centroid], vecs)[0]162        top5_local = list(np.argsort(sims)[-5:][::-1])163        top5_global = list(map(lambda i: int(idxs[i]), top5_local))164        return {165            **cd,166            "centroid": centroid.tolist(),167            "top_sentences": list(map(lambda i: sentences[i], top5_global)),168            "top_sentence_indices": top5_global,169        }170 171    enriched = list(map(_centroid_and_evidence, top_clusters))172 173    summaries = list(174        map(175            lambda i_cd: {176                "topic_id": i_cd[1]["topic_id"],177                "rank": i_cd[0] + 1,178                "size": i_cd[1]["size"],179                "top_sentences": i_cd[1]["top_sentences"],180                "label": f"Topic {i_cd[0] + 1}",181                "approved": False,182                "rename_to": "",183                "reasoning": "",184            },185            enumerate(enriched),186        )187    )188 189    with open("summaries.json", "w") as f:190        json.dump(summaries, f, indent=2)191 192    # ---- Charts ----193    sizes = list(map(lambda s: s["size"], summaries[:30]))194    labels_top = list(map(lambda s: s["label"], summaries[:30]))195 196    fig1 = px.bar(x=labels_top, y=sizes, title="Top 30 Topic Sizes", labels={"x": "Topic", "y": "Count"})197    fig1.write_html("chart_topic_sizes.html")198 199    fig2 = px.pie(values=sizes[:10], names=labels_top[:10], title="Top 10 Topics Share")200    fig2.write_html("chart_topic_pie.html")201 202    all_sizes = list(map(lambda c: c["size"], cluster_data))203    fig3 = px.histogram(x=all_sizes, nbins=20, title="Topic Size Distribution", labels={"x": "Size", "y": "Count"})204    fig3.write_html("chart_size_dist.html")205 206    cumulative = list(np.cumsum(sorted(all_sizes, reverse=True)) / sum(all_sizes) * 100)207    fig4 = go.Figure(go.Scatter(y=cumulative, mode="lines", name="Cumulative Coverage"))208    fig4.update_layout(title="Cumulative Topic Coverage", xaxis_title="Topics", yaxis_title="Coverage %")209    fig4.write_html("chart_coverage.html")210 211    state.update(212        {213            "summaries": summaries,214            "n_clusters": len(unique_labels),215            "charts": ["chart_topic_sizes.html", "chart_topic_pie.html", "chart_size_dist.html", "chart_coverage.html"],216        }217    )218    _save_state(state)219 220    return json.dumps(221        {222            "n_clusters_found": len(unique_labels),223            "top_topics_saved": len(summaries),224            "charts_generated": ["chart_topic_sizes.html", "chart_topic_pie.html", "chart_size_dist.html", "chart_coverage.html"],225            "sample_topics": summaries[:3],226        }227    )228 229 230# ---------------------------------------------------------------------------231# Tool 3 — label_topics_with_llm232# ---------------------------------------------------------------------------233 234@tool(handle_tool_error=True)235def label_topics_with_llm(batch_size: int = 100) -> str:236    """Send top topics to Mistral via PromptTemplate + JsonOutputParser to generate concise labels.237 238    Args:239        batch_size: Number of top topics to label (max 100).240 241    Returns:242        JSON string confirming how many topics were labelled.243    """244    state = _load_state()245    summaries = state.get("summaries", [])246    top = summaries[:min(batch_size, 100)]247 248    llm = ChatGoogleGenerativeAI(model="gemini-1.5-flash", temperature=0.2)249    parser = JsonOutputParser()250 251    prompt = PromptTemplate(252        template=(253            "You are an expert qualitative researcher performing thematic analysis.\n"254            "Given the following topics (each with top sentences), generate a concise 3-6 word label "255            "for each topic that captures its core theme.\n\n"256            "Topics:\n{topics_json}\n\n"257            "Respond ONLY with a JSON array of objects with keys 'topic_id' and 'label'.\n"258            "No markdown, no explanation.\n"259        ),260        input_variables=["topics_json"],261    )262 263    topics_input = json.dumps(264        list(265            map(266                lambda s: {"topic_id": s["topic_id"], "top_sentences": s["top_sentences"][:3]},267                top,268            )269        )270    )271 272    chain = prompt | llm | parser273    result = chain.invoke({"topics_json": topics_input})274 275    label_map = dict(map(lambda r: (r["topic_id"], r["label"]), result))276 277    updated = list(278        map(279            lambda s: {**s, "label": label_map.get(s["topic_id"], s["label"])},280            summaries,281        )282    )283 284    state["summaries"] = updated285    _save_state(state)286 287    with open("summaries.json", "w") as f:288        json.dump(updated, f, indent=2)289 290    return json.dumps({"labelled_count": len(result), "sample_labels": result[:5]})291 292 293# ---------------------------------------------------------------------------294# Tool 4 — consolidate_into_themes295# ---------------------------------------------------------------------------296 297@tool(handle_tool_error=True)298def consolidate_into_themes(approved_groups: str) -> str:299    """Merge approved topic groups into themes and recompute centroids.300 301    Args:302        approved_groups: JSON string of list of groups, each group being a list of topic_ids.303                         E.g. '[[1,2,3],[4,5]]'304 305    Returns:306        JSON string with consolidated theme summaries.307    """308    groups = json.loads(approved_groups)309    state = _load_state()310    summaries = state.get("summaries", [])311    sentences = state.get("sentences", [])312    embeddings = np.load("emb.npy")313 314    summary_map = dict(map(lambda s: (s["topic_id"], s), summaries))315 316    def _merge_group(i_group: tuple) -> dict:317        idx, group = i_group318        members = list(map(lambda tid: summary_map.get(tid, {}), group))319        all_sentences = list(320            set(321                sum(map(lambda m: m.get("top_sentences", []), members), [])322            )323        )324        name_parts = list(map(lambda m: m.get("label", ""), members))325        merged_label = " / ".join(filter(None, name_parts[:3]))326        total_size = sum(map(lambda m: m.get("size", 0), members))327        return {328            "theme_id": idx + 1,329            "topic_ids": group,330            "label": merged_label,331            "size": total_size,332            "top_sentences": all_sentences[:10],333            "approved": True,334        }335 336    themes = list(map(_merge_group, enumerate(groups)))337 338    state["themes"] = themes339    _save_state(state)340 341    with open("themes.json", "w") as f:342        json.dump(themes, f, indent=2)343 344    return json.dumps({"themes_created": len(themes), "themes": themes})345 346 347# ---------------------------------------------------------------------------348# Tool 5 — compare_with_taxonomy349# ---------------------------------------------------------------------------350 351@tool(handle_tool_error=True)352def compare_with_taxonomy(taxonomy_name: str = "PAJAIS") -> str:353    """Map consolidated themes to PAJAIS 25 categories via Mistral.354 355    Args:356        taxonomy_name: Name of the taxonomy to compare against (default: PAJAIS).357 358    Returns:359        JSON string with theme-to-category mappings.360    """361    state = _load_state()362    themes = state.get("themes", [])363 364    pajais_categories = [365        "Intelligent Systems & AI Applications",366        "Machine Learning & Deep Learning",367        "Natural Language Processing",368        "Computer Vision & Image Processing",369        "Decision Support Systems",370        "Knowledge Representation & Reasoning",371        "Human-Computer Interaction",372        "Information Retrieval & Search",373        "Data Mining & Analytics",374        "Recommender Systems",375        "Autonomous Agents & Multi-Agent Systems",376        "Robotics & Automation",377        "Healthcare & Medical AI",378        "Business Intelligence & Analytics",379        "Ethics, Fairness & Explainability",380        "Neural Networks & Architectures",381        "Optimisation & Evolutionary Computing",382        "Cybersecurity & Adversarial AI",383        "Social Media & Network Analysis",384        "Education Technology & E-Learning",385        "Supply Chain & Operations Management",386        "Financial Technology & FinTech",387        "Smart Cities & IoT",388        "Environmental & Sustainability AI",389        "Legal & Regulatory AI",390    ]391 392    llm = ChatGoogleGenerativeAI(model="gemini-1.5-flash", temperature=0.2)393    parser = JsonOutputParser()394 395    prompt = PromptTemplate(396        template=(397            "You are a research classification expert.\n"398            "Map each theme to the most appropriate PAJAIS category.\n\n"399            "PAJAIS Categories:\n{categories}\n\n"400            "Themes:\n{themes_json}\n\n"401            "Respond ONLY with a JSON array of objects with keys: "402            "'theme_id', 'theme_label', 'pajais_category', 'justification'.\n"403        ),404        input_variables=["categories", "themes_json"],405    )406 407    chain = prompt | llm | parser408    result = chain.invoke(409        {410            "categories": "\n".join(map(lambda c: f"- {c}", pajais_categories)),411            "themes_json": json.dumps(412                list(map(lambda t: {"theme_id": t["theme_id"], "label": t["label"]}, themes))413            ),414        }415    )416 417    state["taxonomy_mapping"] = result418    _save_state(state)419 420    with open("taxonomy_mapping.json", "w") as f:421        json.dump(result, f, indent=2)422 423    return json.dumps({"mappings": result})424 425 426# ---------------------------------------------------------------------------427# Tool 6 — generate_comparison_csv428# ---------------------------------------------------------------------------429 430@tool(handle_tool_error=True)431def generate_comparison_csv() -> str:432    """Generate a side-by-side abstract vs title topic comparison CSV.433 434    Returns:435        JSON string confirming the CSV path and row count.436    """437    state = _load_state()438    file_path = state.get("file_path", "")439    df = pd.read_csv(file_path)440 441    abstract_col = next(iter(filter(lambda c: "abstract" in c.lower(), df.columns)), None)442    title_col = next(iter(filter(lambda c: "title" in c.lower(), df.columns)), None)443 444    themes = state.get("themes", [])445    taxonomy = state.get("taxonomy_mapping", [])446    tax_map = dict(map(lambda t: (t["theme_id"], t.get("pajais_category", "")), taxonomy))447 448    def _row_to_record(row: pd.Series) -> dict:449        abstract_text = str(row.get(abstract_col, "")) if abstract_col else ""450        title_text = str(row.get(title_col, "")) if title_col else ""451        return {452            "Title": title_text,453            "Abstract": abstract_text,454            "Theme (Abstract)": next(455                iter(456                    map(457                        lambda t: t["label"],458                        filter(459                            lambda t: any(abstract_text[:50] in s for s in t.get("top_sentences", [])),460                            themes,461                        ),462                    )463                ),464                "Unclassified",465            ),466            "PAJAIS Category": next(467                iter(468                    map(469                        lambda t: tax_map.get(t["theme_id"], ""),470                        filter(471                            lambda t: any(abstract_text[:50] in s for s in t.get("top_sentences", [])),472                            themes,473                        ),474                    )475                ),476                "",477            ),478        }479 480    records = list(map(_row_to_record, [row for _, row in df.iterrows()]))481    comparison_df = pd.DataFrame(records)482    comparison_df.to_csv("comparison_abstract_vs_title.csv", index=False)483 484    return json.dumps(485        {486            "csv_path": "comparison_abstract_vs_title.csv",487            "row_count": len(comparison_df),488            "columns": list(comparison_df.columns),489        }490    )491 492 493# ---------------------------------------------------------------------------494# Tool 7 — export_narrative495# ---------------------------------------------------------------------------496 497@tool(handle_tool_error=True)498def export_narrative() -> str:499    """Generate a 500-word Section 7 narrative report via Mistral and export it.500 501    Returns:502        JSON string with the narrative text and output file path.503    """504    state = _load_state()505    themes = state.get("themes", [])506    taxonomy = state.get("taxonomy_mapping", [])507 508    llm = ChatGoogleGenerativeAI(model="gemini-1.5-flash", temperature=0.2)509 510    prompt = PromptTemplate(511        template=(512            "You are an academic researcher writing Section 7 of a systematic literature review.\n"513            "Following Braun & Clarke (2006) thematic analysis methodology, write a cohesive ~500-word "514            "narrative that:\n"515            "1. Introduces the overarching themes discovered\n"516            "2. Describes each theme with supporting evidence\n"517            "3. Discusses relationships between themes\n"518            "4. Connects themes to the PAJAIS taxonomy categories\n"519            "5. Concludes with implications for the field\n\n"520            "Themes:\n{themes_json}\n\n"521            "Taxonomy Mappings:\n{taxonomy_json}\n\n"522            "Write the section now. Use academic tone. No bullet points.\n"523        ),524        input_variables=["themes_json", "taxonomy_json"],525    )526 527    chain = prompt | llm528    response = chain.invoke(529        {530            "themes_json": json.dumps(531                list(532                    map(533                        lambda t: {534                            "theme_id": t["theme_id"],535                            "label": t["label"],536                            "size": t["size"],537                            "sample_sentences": t.get("top_sentences", [])[:3],538                        },539                        themes,540                    )541                )542            ),543            "taxonomy_json": json.dumps(taxonomy),544        }545    )546 547    narrative_text = response.content548 549    with open("section7_narrative.txt", "w") as f:550        f.write(narrative_text)551 552    state["narrative"] = narrative_text553    _save_state(state)554 555    return json.dumps(556        {557            "narrative_path": "section7_narrative.txt",558            "word_count": len(narrative_text.split()),559            "preview": narrative_text[:300],560        }561    )562