codeby-hp/sentiment-classification
0
1from fastapi import FastAPI, Request, Form2from fastapi.responses import HTMLResponse3from fastapi.templating import Jinja2Templates4from fastapi.staticfiles import StaticFiles5import mlflow6import pickle7import os8import pandas as pd9import numpy as np10from nltk.stem import WordNetLemmatizer11from nltk.corpus import stopwords12import string13import re14import dagshub15import nltk16 17import warnings18warnings.simplefilter("ignore", UserWarning)19warnings.filterwarnings("ignore")20 21from dotenv import load_dotenv22 23load_dotenv()24 25# Download required NLTK data26try:27 nltk.download('stopwords', quiet=True)28 nltk.download('wordnet', quiet=True)29 nltk.download('omw-1.4', quiet=True)30except:31 pass32 33def lemmatization(text):34 """Lemmatize the text."""35 lemmatizer = WordNetLemmatizer()36 text = text.split()37 text = [lemmatizer.lemmatize(word) for word in text]38 return " ".join(text)39 40def remove_stop_words(text):41 """Remove stop words from the text."""42 stop_words = set(stopwords.words("english"))43 text = [word for word in str(text).split() if word not in stop_words]44 return " ".join(text)45 46def removing_numbers(text):47 """Remove numbers from the text."""48 text = ''.join([char for char in text if not char.isdigit()])49 return text50 51def lower_case(text):52 """Convert text to lower case."""53 text = text.split()54 text = [word.lower() for word in text]55 return " ".join(text)56 57def removing_punctuations(text):58 """Remove punctuations from the text."""59 text = re.sub('[%s]' % re.escape(string.punctuation), ' ', text)60 text = text.replace('؛', "")61 text = re.sub('\s+', ' ', text).strip()62 return text63 64def removing_urls(text):65 """Remove URLs from the text."""66 url_pattern = re.compile(r'https?://\S+|www\.\S+')67 return url_pattern.sub(r'', text)68 69def remove_small_sentences(df):70 """Remove sentences with less than 3 words."""71 for i in range(len(df)):72 if len(df.text.iloc[i].split()) < 3:73 df.text.iloc[i] = np.nan74 75def normalize_text(text):76 text = lower_case(text)77 text = remove_stop_words(text)78 text = removing_numbers(text)79 text = removing_punctuations(text)80 text = removing_urls(text)81 text = lemmatization(text)82 83 return text84 85# Below code block is for local use86# -------------------------------------------------------------------------------------87# mlflow.set_tracking_uri('https://dagshub.com/CodeBy-HP/Sentiment-Classification-Mlflow-DVC.mlflow')88# dagshub.init(repo_owner='CodeBy-HP', repo_name='Sentiment-Classification-Mlflow-DVC', mlflow=True)89# -------------------------------------------------------------------------------------90 91# Below code block is for production use92# -------------------------------------------------------------------------------------93# Set up DagsHub credentials for MLflow tracking94dagshub_token = os.getenv("CAPSTONE_TEST")95if not dagshub_token:96 raise EnvironmentError("CAPSTONE_TEST environment variable is not set")97 98os.environ["MLFLOW_TRACKING_USERNAME"] = dagshub_token99os.environ["MLFLOW_TRACKING_PASSWORD"] = dagshub_token100 101dagshub_url = "https://dagshub.com"102repo_owner = "CodeBy-HP"103repo_name = "Sentiment-Classification-Mlflow-DVC"104# Set up MLflow tracking URI105mlflow.set_tracking_uri(f'{dagshub_url}/{repo_owner}/{repo_name}.mlflow')106# -------------------------------------------------------------------------------------107 108 109# Initialize FastAPI app110app = FastAPI(title="Sentiment Analysis API", version="1.0.0")111 112# Set up Jinja2 templates113current_file_dir = os.path.dirname(os.path.abspath(__file__))114templates_dir = os.path.join(current_file_dir, "templates")115templates = Jinja2Templates(directory=templates_dir)116 117# ------------------------------------------------------------------------------------------118# Model and vectorizer setup119model_name = "my_model"120 121# Get the path to the vectorizer file122current_dir = os.path.dirname(os.path.abspath(__file__))123vectorizer_path = os.path.join(current_dir, 'models', 'vectorizer.pkl')124if not os.path.exists(vectorizer_path):125 # Try alternative paths126 alt_paths = [127 os.path.join(os.getcwd(), 'models', 'vectorizer.pkl'),128 os.path.join(current_dir, '..', 'models', 'vectorizer.pkl'),129 '/app/models/vectorizer.pkl' # Docker path130 ]131 for path in alt_paths:132 if os.path.exists(path):133 vectorizer_path = path134 break135 136def get_latest_model_version(model_name):137 client = mlflow.MlflowClient()138 latest_version = client.get_latest_versions(model_name, stages=["Production"])139 if not latest_version:140 latest_version = client.get_latest_versions(model_name, stages=["None"])141 return latest_version[0].version if latest_version else None142 143model_version = get_latest_model_version(model_name)144model_uri = f'models:/{model_name}/{model_version}'145print(f"Fetching model from: {model_uri}")146model = mlflow.sklearn.load_model(model_uri)147vectorizer = pickle.load(open(vectorizer_path, 'rb'))148 149# Routes150@app.get("/", response_class=HTMLResponse)151async def home(request: Request):152 """Render the home page."""153 return templates.TemplateResponse(154 request=request,155 name="index.html",156 context={"result": None}157 )158 159@app.post("/predict", response_class=HTMLResponse)160async def predict(request: Request, text: str = Form(...)):161 """Handle sentiment prediction."""162 # Clean text163 cleaned_text = normalize_text(text)164 165 # Convert to features166 features = vectorizer.transform([cleaned_text])167 # Convert to array without column names to avoid sklearn warning168 features_array = features.toarray()169 170 # Predict171 result = model.predict(features_array)172 prediction = int(result[0])173 174 # Get probability scores for confidence175 # Note: predict_proba returns [prob_negative, prob_positive]176 probabilities = model.predict_proba(features_array)[0]177 confidence = float(probabilities[prediction]) * 100 # Convert to percentage178 179 return templates.TemplateResponse(180 request=request,181 name="index.html",182 context={"result": prediction, "confidence": confidence}183 )184 185@app.get("/health")186async def health_check():187 """Health check endpoint for monitoring."""188 return {"status": "healthy", "model_version": model_version}189 190if __name__ == "__main__":191 import uvicorn192 uvicorn.run(app, host="0.0.0.0", port=8000)193 