Team Ai
Apppublic

codeby-hp/sentiment-classification

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
app.py193 linesDownload Raw Back to root
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