Team Ai
Apppublic

Spaceboy23/topic_modelling

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
tools.py424 linesDownload Raw Back to root
1"""
2tools.py - BERTopic Agentic AI Tools
3Stateless, functional tools with zero try/except blocks.
4All errors propagate to LLM via @tool.
5"""
6
7import pandas as pd
8import numpy as np
9import re
10import json
11from pathlib import Path
12from typing import Dict, Any, List
13
14import nltk
15from nltk.tokenize import sent_tokenize
16from sentence_transformers import SentenceTransformer
17from sklearn.cluster import AgglomerativeClustering
18from sklearn.metrics.pairwise import cosine_similarity
19from bertopic import BERTopic
20from bertopic.vectorizers import ClassTfidfTransformer
21from bertopic.representation import KeyBERTInspired
22import plotly.graph_objects as go
23from plotly.subplots import make_subplots
24import plotly.express as px
25
26from langchain_core.tools import tool
27from langchain_mistralai import ChatMistralAI
28from langchain_core.prompts import PromptTemplate
29from langchain_core.output_parsers import JsonOutputParser
30
31# Download NLTK data
32nltk.download('punkt', quiet=True)
33nltk.download('punkt_tab', quiet=True)
34
35# Global paths
36DATA_DIR = Path("./data")
37DATA_DIR.mkdir(exist_ok=True)
38
39def clean_abstract_text(text: str) -> str:
40    """Remove publisher boilerplate noise from abstracts."""
41    patterns = [
42        r'©\s*\d{4}.*?rights reserved',
43        r'All rights reserved',
44        r'Elsevier.*?Ltd',
45        r'Springer.*?Verlag',
46        r'Published by.*?$',
47        r'\[ABSTRACT FROM AUTHOR\]',
48        r'\[ABSTRACT FROM PUBLISHER\]',
49    ]
50    cleaned = text
51    cleaned = re.sub('|'.join(patterns), '', cleaned, flags=re.IGNORECASE)
52    cleaned = re.sub(r'\s+', ' ', cleaned).strip()
53    return cleaned
54
55@tool
56def load_scopus_csv(filepath: str) -> str:
57    """
58    Load Scopus CSV, count papers, split abstracts into sentences, remove boilerplate.
59    Returns stats string with paper count and sentence count.
60    """
61    # FIX: Use encoding_errors='replace' to prevent Scopus invisible character crashes
62    df = pd.read_csv(filepath, encoding='utf-8-sig', encoding_errors='replace')
63    
64    paper_count = len(df)
65    
66    df['Abstract'] = df['Abstract'].fillna('').astype(str)
67    df['cleaned_abstract'] = df['Abstract'].apply(clean_abstract_text)
68    
69    sentences = []
70    paper_ids = []
71    
72    for idx, row in df.iterrows():
73        sents = sent_tokenize(row['cleaned_abstract'])
74        sentences.extend(sents)
75        paper_ids.extend([idx] * len(sents))
76        
77    sentence_df = pd.DataFrame({
78        'sentence': sentences,
79        'paper_id': paper_ids
80    })
81    
82    sentence_df.to_csv(DATA_DIR / 'sentences.csv', index=False)
83    df.to_csv(DATA_DIR / 'papers.csv', index=False)
84    
85    return f"Loaded {paper_count} papers, extracted {len(sentences)} sentences. Saved to data/sentences.csv and data/papers.csv"
86
87@tool
88def run_bertopic_discovery(run_key: str, threshold: float = 0.7) -> str:
89    """
90    Embed sentences with all-MiniLM-L6-v2, cluster in 384d space using AgglomerativeClustering.
91    NO UMAP dimensionality reduction. Generate Plotly charts and save outputs.
92    """
93    sentences_df = pd.read_csv(DATA_DIR / 'sentences.csv')
94    sentences = sentences_df['sentence'].tolist()
95    
96    # Embed sentences
97    embedder = SentenceTransformer('all-MiniLM-L6-v2')
98    embeddings = embedder.encode(sentences, normalize_embeddings=True, show_progress_bar=True)
99    
100    # Cluster directly in 384d space
101    clusterer = AgglomerativeClustering(
102        n_clusters=None,
103        metric='cosine',
104        linkage='average',
105        distance_threshold=threshold
106    )
107    cluster_labels = clusterer.fit_predict(embeddings)
108    
109    # Build topic summaries
110    unique_labels = sorted(set(cluster_labels))
111    topic_summaries = []
112    
113    for label in unique_labels:
114        cluster_mask = cluster_labels == label
115        cluster_embeddings = embeddings[cluster_mask]
116        cluster_sentences_idx = np.where(cluster_mask)[0]
117        
118        centroid = cluster_embeddings.mean(axis=0).reshape(1, -1)
119        similarities = cosine_similarity(cluster_embeddings, centroid).flatten()
120        
121        top_5_idx = similarities.argsort()[-5:][::-1]
122        top_sentences = [sentences[cluster_sentences_idx[i]] for i in top_5_idx]
123        
124        paper_ids = sentences_df.iloc[cluster_sentences_idx]['paper_id'].unique()
125        
126        topic_summaries.append({
127            'topic_id': int(label),
128            'sentence_count': int(cluster_mask.sum()),
129            'paper_count': int(len(paper_ids)),
130            'top_sentences': top_sentences,
131            'centroid': centroid.tolist()
132        })
133        
134    # Save outputs
135    run_dir = DATA_DIR / run_key
136    run_dir.mkdir(parents=True, exist_ok=True)
137    
138    with open(run_dir / 'summaries.json', 'w') as f:
139        json.dump(topic_summaries, f, indent=2)
140        
141    np.save(run_dir / 'emb.npy', embeddings)
142    np.save(run_dir / 'labels.npy', cluster_labels)
143    
144    # Generate Plotly charts
145    embeddings_2d = embeddings[:, :2]  # Use first 2 dimensions for visualization
146    
147    # Chart 1: Intertopic Distance Map
148    fig1 = px.scatter(
149        x=embeddings_2d[:, 0],
150        y=embeddings_2d[:, 1],
151        color=cluster_labels,
152        title='Intertopic Distance Map',
153        labels={'x': 'Dimension 1', 'y': 'Dimension 2', 'color': 'Topic'}
154    )
155    fig1.write_html(run_dir / 'intertopic_map.html', include_plotlyjs='cdn')
156    
157    # Chart 2: Topic Bars
158    topic_sizes = pd.Series(cluster_labels).value_counts().sort_index()
159    fig2 = go.Figure(data=[go.Bar(x=topic_sizes.index, y=topic_sizes.values)])
160    fig2.update_layout(title='Topic Sizes', xaxis_title='Topic ID', yaxis_title='Sentence Count')
161    fig2.write_html(run_dir / 'topic_bars.html', include_plotlyjs='cdn')
162    
163    # Chart 3: Hierarchy (dendrogram simulation)
164    fig3 = go.Figure(data=[go.Scatter(
165        x=list(range(len(unique_labels))),
166        y=[len([c for c in cluster_labels if c == l]) for l in unique_labels],
167        mode='markers+lines',
168        marker=dict(size=10)
169    )])
170    fig3.update_layout(title='Topic Hierarchy', xaxis_title='Topic ID', yaxis_title='Size')
171    fig3.write_html(run_dir / 'hierarchy.html', include_plotlyjs='cdn')
172    
173    # Chart 4: Heatmap (topic similarity)
174    topic_centroids = np.array([np.mean(embeddings[cluster_labels == l], axis=0) for l in unique_labels])
175    similarity_matrix = cosine_similarity(topic_centroids)
176    fig4 = go.Figure(data=go.Heatmap(z=similarity_matrix, x=unique_labels, y=unique_labels))
177    fig4.update_layout(title='Topic Similarity Heatmap')
178    fig4.write_html(run_dir / 'heatmap.html', include_plotlyjs='cdn')
179    
180    return f"Discovered {len(unique_labels)} topics. Saved summaries, embeddings, and 4 charts to {run_dir}"
181
182@tool
183def label_topics_with_llm(run_key: str) -> str:
184    """
185    Send top 100 topics to Mistral API for labeling.
186    Parse JSON output: label, category, confidence, reasoning, niche.
187    """
188    run_dir = DATA_DIR / run_key
189    
190    with open(run_dir / 'summaries.json', 'r') as f:
191        summaries = json.load(f)
192        
193    top_100 = sorted(summaries, key=lambda x: x['sentence_count'], reverse=True)[:100]
194    
195    llm = ChatMistralAI(model='mistral-small-latest', temperature=0)
196    parser = JsonOutputParser()
197    
198    prompt = PromptTemplate(
199        template="""You are a research domain expert. Analyze these topic clusters from academic papers and provide structured labels.
200For each topic, return JSON with:
201- label: Short research area name (2-5 words)
202- category: Broader domain (e.g., "Machine Learning", "Healthcare", "Finance")
203- confidence: 0-1 score
204- reasoning: Brief explanation
205- niche: true if highly specialized, false if mainstream
206Topics:
207{topics}
208Return a JSON array with one object per topic. {format_instructions}""",
209        input_variables=['topics'],
210        partial_variables={'format_instructions': parser.get_format_instructions()}
211    )
212    
213    topics_text = json.dumps([{
214        'topic_id': t['topic_id'],
215        'top_sentences': t['top_sentences'][:3]
216    } for t in top_100], indent=2)
217    
218    chain = prompt | llm | parser
219    labels = chain.invoke({'topics': topics_text})
220    
221    with open(run_dir / 'labels.json', 'w') as f:
222        json.dump(labels, f, indent=2)
223        
224    return f"Labeled {len(labels)} topics. Saved to {run_dir / 'labels.json'}"
225
226@tool
227def consolidate_into_themes(run_key: str, theme_map: Dict[str, List[int]]) -> str:
228    """
229    Group topics into researcher-approved themes.
230    Recompute centroids, recount sentences/papers.
231    """
232    run_dir = DATA_DIR / run_key
233    
234    with open(run_dir / 'summaries.json', 'r') as f:
235        summaries = json.load(f)
236        
237    embeddings = np.load(run_dir / 'emb.npy')
238    cluster_labels = np.load(run_dir / 'labels.npy')
239    
240    summaries_dict = {s['topic_id']: s for s in summaries}
241    themes = []
242    
243    for theme_name, topic_ids in theme_map.items():
244        topic_ids = [int(tid) for tid in topic_ids]
245        theme_mask = np.isin(cluster_labels, topic_ids)
246        theme_embeddings = embeddings[theme_mask]
247        
248        centroid = theme_embeddings.mean(axis=0)
249        sentence_count = int(theme_mask.sum())
250        
251        paper_ids = set()
252        for tid in topic_ids:
253            sentences_df = pd.read_csv(DATA_DIR / 'sentences.csv')
254            topic_sentences = sentences_df[cluster_labels == tid]
255            paper_ids.update(topic_sentences['paper_id'].unique())
256            
257        themes.append({
258            'theme_name': theme_name,
259            'topic_ids': topic_ids,
260            'sentence_count': sentence_count,
261            'paper_count': len(paper_ids),
262            'centroid': centroid.tolist()
263        })
264        
265    with open(run_dir / 'themes.json', 'w') as f:
266        json.dump(themes, f, indent=2)
267        
268    return f"Consolidated {len(themes)} themes. Saved to {run_dir / 'themes.json'}"
269
270@tool
271def compare_with_taxonomy(run_key: str) -> str:
272    """
273    Map themes to PAJAIS 25-category taxonomy using Mistral.
274    Return pajais_match, confidence, reasoning, is_novel.
275    """
276    run_dir = DATA_DIR / run_key
277    
278    with open(run_dir / 'themes.json', 'r') as f:
279        themes = json.load(f)
280        
281    pajais_categories = [
282        "Artificial Intelligence", "Blockchain", "Business Analytics", "Cloud Computing",
283        "Cybersecurity", "Data Science", "Digital Innovation", "E-Commerce", "Enterprise Systems",
284        "Ethics and Privacy", "Healthcare IT", "Human-Computer Interaction", "Internet of Things",
285        "IT Governance", "Knowledge Management", "Mobile Computing", "Social Media",
286        "Software Development", "Strategic IS", "Supply Chain", "Technology Adoption",
287        "Telecommunications", "Virtual Reality", "Web Technologies", "NOVEL"
288    ]
289    
290    llm = ChatMistralAI(model='mistral-small-latest', temperature=0)
291    parser = JsonOutputParser()
292    
293    prompt = PromptTemplate(
294        template="""Map each research theme to the closest PAJAIS category.
295PAJAIS Categories: {categories}
296Themes:
297{themes}
298For each theme, return JSON with:
299- theme_name: Original theme name
300- pajais_match: Best matching category (or NOVEL if none fit)
301- confidence: 0-1 score
302- reasoning: Explanation
303- is_novel: true if NOVEL, false otherwise
304Return a JSON array. {format_instructions}""",
305        input_variables=['themes', 'categories'],
306        partial_variables={'format_instructions': parser.get_format_instructions()}
307    )
308    
309    themes_text = json.dumps([{
310        'theme_name': t['theme_name'],
311        'sentence_count': t['sentence_count'],
312        'paper_count': t['paper_count']
313    } for t in themes], indent=2)
314    
315    chain = prompt | llm | parser
316    mappings = chain.invoke({
317        'themes': themes_text,
318        'categories': ', '.join(pajais_categories)
319    })
320    
321    with open(run_dir / 'taxonomy_map.json', 'w') as f:
322        json.dump(mappings, f, indent=2)
323        
324    return f"Mapped {len(mappings)} themes to taxonomy. Saved to {run_dir / 'taxonomy_map.json'}"
325
326@tool
327def generate_comparison_csv() -> str:
328    """
329    Load themes from abstract and title runs.
330    Create side-by-side comparison DataFrame.
331    """
332    abstract_dir = DATA_DIR / 'abstract_run'
333    title_dir = DATA_DIR / 'title_run'
334    
335    with open(abstract_dir / 'themes.json', 'r') as f:
336        abstract_themes = json.load(f)
337        
338    with open(title_dir / 'themes.json', 'r') as f:
339        title_themes = json.load(f)
340        
341    rows = []
342    max_len = max(len(abstract_themes), len(title_themes))
343    
344    for i in range(max_len):
345        row = {}
346        
347        if i < len(abstract_themes):
348            at = abstract_themes[i]
349            row['abstract_theme'] = at['theme_name']
350            row['abstract_sentences'] = at['sentence_count']
351            row['abstract_papers'] = at['paper_count']
352        else:
353            row['abstract_theme'] = ''
354            row['abstract_sentences'] = 0
355            row['abstract_papers'] = 0
356            
357        if i < len(title_themes):
358            tt = title_themes[i]
359            row['title_theme'] = tt['theme_name']
360            row['title_sentences'] = tt['sentence_count']
361            row['title_papers'] = tt['paper_count']
362        else:
363            row['title_theme'] = ''
364            row['title_sentences'] = 0
365            row['title_papers'] = 0
366            
367        rows.append(row)
368        
369    comparison_df = pd.DataFrame(rows)
370    comparison_df.to_csv(DATA_DIR / 'comparison.csv', index=False)
371    
372    return f"Generated comparison CSV with {len(rows)} rows. Saved to {DATA_DIR / 'comparison.csv'}"
373
374@tool
375def export_narrative(run_key: str) -> str:
376    """
377    Send themes and taxonomy to Mistral to draft a 500-word Section 7 literature review.
378    References B&C phases.
379    """
380    run_dir = DATA_DIR / run_key
381    
382    with open(run_dir / 'themes.json', 'r') as f:
383        themes = json.load(f)
384        
385    with open(run_dir / 'taxonomy_map.json', 'r') as f:
386        taxonomy = json.load(f)
387        
388    llm = ChatMistralAI(model='mistral-small-latest', temperature=0.7)
389    
390    prompt = PromptTemplate(
391        template="""You are writing Section 7 (Literature Review) for an academic paper on IS research trends.
392Themes discovered:
393{themes}
394Taxonomy mapping:
395{taxonomy}
396Write a 500-word narrative that:
3971. References the Bandara & Cresswell (B&C) framework phases
3982. Highlights dominant research themes
3993. Notes novel or emerging areas
4004. Connects to PAJAIS categories
4015. Uses academic tone
402Output only the prose, no preamble.""",
403        input_variables=['themes', 'taxonomy']
404    )
405    
406    themes_text = json.dumps([{
407        'theme_name': t['theme_name'],
408        'paper_count': t['paper_count']
409    } for t in themes], indent=2)
410    
411    taxonomy_text = json.dumps([{
412        'theme_name': m['theme_name'],
413        'pajais_match': m['pajais_match'],
414        'is_novel': m['is_novel']
415    } for m in taxonomy], indent=2)
416    
417    chain = prompt | llm
418    narrative = chain.invoke({'themes': themes_text, 'taxonomy': taxonomy_text})
419    narrative_text = narrative.content
420    
421    with open(run_dir / 'narrative.txt', 'w') as f:
422        f.write(narrative_text)
423        
424    return f"Generated 500-word narrative. Saved to {run_dir / 'narrative.txt'}"