File size: 1,730 Bytes
0f5fdb7
6cd7286
 
 
92d70d5
 
0f5fdb7
92d70d5
 
4dce658
92d70d5
 
 
4dce658
92d70d5
 
 
4dce658
92d70d5
 
 
 
 
5ec9c86
92d70d5
5ec9c86
92d70d5
 
 
e69509c
35502c8
 
92d70d5
35502c8
92d70d5
35502c8
92d70d5
 
e69509c
 
35502c8
 
 
 
e69509c
35502c8
 
 
 
4dce658
92d70d5
 
35502c8
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
import os

os.environ["TRANSFORMERS_CACHE"] = "/tmp/huggingface"
os.environ["HF_HOME"] = "/tmp/huggingface"
os.environ["NUMBA_CACHE_DIR"] = "/tmp/numba_cache"
os.environ["NUMBA_DISABLE_JIT"] = "1"

from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse
from bertopic import BERTopic
from sklearn.feature_extraction.text import CountVectorizer
from sklearn.decomposition import TruncatedSVD
from sklearn.cluster import KMeans

vectorizer_model = CountVectorizer()
dimensionality_model = TruncatedSVD(n_components=5)
clustering_model = KMeans(n_clusters=5, random_state=42)

topic_model = BERTopic(
    vectorizer_model=vectorizer_model,
    umap_model=dimensionality_model,
    hdbscan_model=clustering_model
)

app = FastAPI()

@app.post("/predict")
async def predict(request: Request):
    data = await request.json()

    documents = []

    if "text" in data:
        documents = [doc.strip() for doc in data["text"].split("\n") if doc.strip()]
    elif "data" in data and isinstance(data["data"], list):
        documents = [doc.strip() for doc in data["data"] if isinstance(doc, str) and doc.strip()]
    else:
        return JSONResponse({"error": "No input text provided."}, status_code=400)

    if len(documents) < 2:
        return JSONResponse({"error": "Please provide at least 2 documents for topic modeling."}, status_code=400)

    topics, probs = topic_model.fit_transform(documents)
    topic_info = topic_model.get_topic_info()

    return {
        "topics": topic_info.to_dict(orient="records"),
        "topic_assignments": topics
    }

@app.get("/")
async def root():
    return {"message": "BERTopic FastAPI is running! Use POST /predict with {'text': '...'} or {'data': [...]}."}