Ibrahim-Geek/encode_and_extract_phrases
0
1import io2import os3from fastapi import FastAPI, Form4from sentence_transformers import SentenceTransformer5from keybert import KeyBERT6from pydantic import BaseModel7from typing import Union, List8 9app = FastAPI()10 11HF_TOKEN = os.getenv("HF_TOKEN")12 13class TextInput(BaseModel):14 text: Union[str, List[str]]15 16class Model:17 18 keybert_model = None19 encoding_model = None20 21 @classmethod22 def get_keybert_model(cls) -> None:23 24 if cls.keybert_model is None:25 26 cls.keybert_model = KeyBERT("sentence-transformers/all-mpnet-base-v2")27 28 # warmup model29 _ = cls.keybert_model.extract_keywords("Dummy testing to warmup model")30 31 @classmethod32 def get_encoding_model(cls) -> None:33 34 if cls.encoding_model is None:35 36 cls.encoding_model = SentenceTransformer("sentence-transformers/all-mpnet-base-v2")37 38 # warmup encoding model39 _ = cls.encoding_model.encode("Dummy testing to warm up model")40 41 @classmethod42 def load_models(cls) -> None:43 cls.get_encoding_model()44 cls.get_keybert_model()45 46Model.load_models()47 48 49@app.get("/")50def ping():51 52 Model.load_models()53 54 return {"status": "Models Warmed"}55 56 57@app.post("/get_encoding")58async def get_encoding(input_text:TextInput) -> list:59 60 embeddings = Model.encoding_model.encode(input_text.text).tolist()61 62 return embeddings63 64 65@app.post("/extract_keyword_phrases")66async def extract_keyword_phrases(input_text: TextInput) -> list:67 68 key_phrases = Model.keybert_model.extract_keywords(69 docs=input_text.text, keyphrase_ngram_range=(1, 2), top_n=-170 )71 72 result = [phrase[0] for phrase in key_phrases]73 74 return result75 