Spaces:
Sleeping
Sleeping
| """ | |
| FastAPI backend β RoBERTa Multi-Label Emotion Classifier | |
| SemEval-2018 Task 1, Subtask E-c | |
| Model is loaded from Hugging Face Hub at startup. | |
| Set environment variable HF_MODEL_REPO to your HF repo path, | |
| e.g. krishanbhati/roberta-semeval2018-emotions | |
| """ | |
| import logging | |
| import requests | |
| from fastapi import FastAPI, HTTPException | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from pydantic import BaseModel | |
| from typing import List, Dict | |
| import torch | |
| import torch.nn as nn | |
| from transformers import RobertaTokenizer, RobertaModel | |
| from huggingface_hub import hf_hub_download | |
| import numpy as np | |
| import re | |
| import os | |
| import json | |
| logging.basicConfig(level=logging.INFO) | |
| log = logging.getLogger(__name__) | |
| # ββ Try to import demoji / wordninja gracefully ββββββββββββββββββββ | |
| try: | |
| import demoji | |
| demoji.download_codes() | |
| DEMOJI_AVAILABLE = True | |
| except Exception: | |
| DEMOJI_AVAILABLE = False | |
| log.warning("demoji not available β emoji conversion disabled") | |
| try: | |
| import wordninja | |
| WORDNINJA_AVAILABLE = True | |
| except Exception: | |
| WORDNINJA_AVAILABLE = False | |
| log.warning("wordninja not available β hashtag segmentation disabled") | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Configuration | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| HF_MODEL_REPO = "oyykrishna/roberta-semeval2018-emotions" | |
| HF_MODEL_FILENAME = "best_roberta_model.pt" | |
| MAX_LEN = 128 | |
| EMOTION_LABELS = [ | |
| "anger", "anticipation", "disgust", "fear", | |
| "joy", "love", "optimism", "pessimism", | |
| "sadness", "surprise", "trust" | |
| ] | |
| # Adaptive thresholds tuned on SemEval-2018 dev set | |
| # Replace these with your actual values from the notebook output | |
| # (Stage 7 β final_results.json β adaptive_thresholds) | |
| ADAPTIVE_THRESHOLDS = { | |
| "anger": 0.53, | |
| "anticipation": 0.49, | |
| "disgust": 0.42, | |
| "fear": 0.93, | |
| "joy": 0.33, | |
| "love": 0.89, | |
| "optimism": 0.66, | |
| "pessimism": 0.73, | |
| "sadness": 0.62, | |
| "surprise": 0.82, | |
| "trust": 0.58, | |
| } | |
| # Emotion metadata for UI display | |
| EMOTION_META = { | |
| "anger": {"emoji": "π ", "color": "#E24B4A", "description": "Feeling angry or irritated"}, | |
| "anticipation": {"emoji": "π€©", "color": "#BA7517", "description": "Looking forward to something"}, | |
| "disgust": {"emoji": "π€’", "color": "#3B6D11", "description": "Strong dislike or revulsion"}, | |
| "fear": {"emoji": "π¨", "color": "#534AB7", "description": "Feeling scared or anxious"}, | |
| "joy": {"emoji": "π", "color": "#1D9E75", "description": "Feeling happy or elated"}, | |
| "love": {"emoji": "β€οΈ", "color": "#D4537E", "description": "Feeling affection or deep care"}, | |
| "optimism": {"emoji": "π", "color": "#185FA5", "description": "Hopeful about the future"}, | |
| "pessimism": {"emoji": "π", "color": "#5F5E5A", "description": "Expecting the worst outcome"}, | |
| "sadness": {"emoji": "π’", "color": "#378ADD", "description": "Feeling sad or sorrowful"}, | |
| "surprise": {"emoji": "π²", "color": "#993C1D", "description": "Feeling unexpected shock"}, | |
| "trust": {"emoji": "π€", "color": "#0F6E56", "description": "Feeling safe or confident"}, | |
| } | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Model definition (must match training architecture exactly) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class RoBERTaEmotionClassifier(nn.Module): | |
| def __init__(self, num_labels=11, dropout_p=0.1): | |
| super().__init__() | |
| self.roberta = RobertaModel.from_pretrained("roberta-base") | |
| self.dropout = nn.Dropout(p=dropout_p) | |
| self.classifier = nn.Linear(self.roberta.config.hidden_size, num_labels) | |
| def forward(self, input_ids, attention_mask): | |
| outputs = self.roberta(input_ids=input_ids, attention_mask=attention_mask) | |
| cls_output = outputs.last_hidden_state[:, 0, :] | |
| cls_output = self.dropout(cls_output) | |
| return self.classifier(cls_output) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Preprocessing (mirrors the Colab notebook Stage 3) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def segment_hashtag(match): | |
| tag = match.group(1) | |
| if WORDNINJA_AVAILABLE: | |
| return " ".join(wordninja.split(tag)).lower() | |
| return tag.lower() | |
| def clean_tweet(text: str) -> str: | |
| if not isinstance(text, str): | |
| return "" | |
| if DEMOJI_AVAILABLE: | |
| text = demoji.replace_with_desc(text, sep=" ") | |
| text = re.sub(r":", " ", text) | |
| text = re.sub(r"_", " ", text) | |
| text = re.sub(r"http\S+|www\.\S+", "[URL]", text) | |
| text = re.sub(r"@\w+", "[USER]", text) | |
| text = re.sub(r"#(\w+)", segment_hashtag, text) | |
| text = re.sub(r"<[^>]+>", "", text) | |
| text = re.sub(r"([!?.])\\1+", r"\1", text) | |
| text = re.sub(r"(.)\1{3,}", r"\1\1\1", text) | |
| text = re.sub(r"[^\x00-\x7F]+", " ", text) | |
| text = re.sub(r"\s+", " ", text).strip() | |
| return text.lower() | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Load model at startup | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| log.info(f"Device: {DEVICE}") | |
| log.info("Loading tokenizer...") | |
| tokenizer = RobertaTokenizer.from_pretrained("roberta-base") | |
| log.info(f"Downloading model from HF Hub: {HF_MODEL_REPO}") | |
| model_path = hf_hub_download( | |
| repo_id=HF_MODEL_REPO, | |
| filename=HF_MODEL_FILENAME, | |
| repo_type="model", | |
| cache_dir="/tmp/hf_cache", | |
| ) | |
| log.info(f"Model cached at: {model_path}") | |
| model = RoBERTaEmotionClassifier(num_labels=len(EMOTION_LABELS)) | |
| model.load_state_dict(torch.load(model_path, map_location=DEVICE)) | |
| model.to(DEVICE) | |
| model.eval() | |
| log.info("Model ready β ") | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # FastAPI App | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| app = FastAPI( | |
| title="Emotion Detection API", | |
| description="Multi-label emotion classification using RoBERTa fine-tuned on SemEval-2018", | |
| version="1.0.0", | |
| ) | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], # tighten in production | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Schemas | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class PredictRequest(BaseModel): | |
| text: str | |
| threshold_mode: str = "adaptive" # "adaptive" | "fixed" | |
| fixed_threshold: float = 0.5 | |
| class EmotionResult(BaseModel): | |
| label: str | |
| probability: float | |
| detected: bool | |
| emoji: str | |
| color: str | |
| description: str | |
| class PredictResponse(BaseModel): | |
| original_text: str | |
| cleaned_text: str | |
| emotions: List[EmotionResult] | |
| detected_emotions: List[str] | |
| dominant_emotion: str | None | |
| confidence_score: float | |
| class BatchRequest(BaseModel): | |
| texts: List[str] | |
| threshold_mode: str = "adaptive" | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Inference helper | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def run_inference(text: str, threshold_mode: str = "adaptive", fixed_threshold: float = 0.5) -> dict: | |
| cleaned = clean_tweet(text) | |
| if not cleaned: | |
| raise ValueError("Text is empty after preprocessing") | |
| enc = tokenizer( | |
| cleaned, | |
| max_length=MAX_LEN, | |
| padding="max_length", | |
| truncation=True, | |
| return_tensors="pt", | |
| ) | |
| with torch.no_grad(): | |
| logits = model( | |
| enc["input_ids"].to(DEVICE), | |
| enc["attention_mask"].to(DEVICE), | |
| ) | |
| probs = torch.sigmoid(logits).cpu().numpy()[0] | |
| emotions = [] | |
| for i, label in enumerate(EMOTION_LABELS): | |
| if threshold_mode == "adaptive": | |
| threshold = ADAPTIVE_THRESHOLDS[label] | |
| else: | |
| threshold = fixed_threshold | |
| detected = bool(probs[i] > threshold) | |
| meta = EMOTION_META[label] | |
| emotions.append({ | |
| "label": label, | |
| "probability": round(float(probs[i]), 4), | |
| "detected": detected, | |
| "emoji": meta["emoji"], | |
| "color": meta["color"], | |
| "description": meta["description"], | |
| }) | |
| emotions.sort(key=lambda x: x["probability"], reverse=True) | |
| detected_emotions = [e["label"] for e in emotions if e["detected"]] | |
| dominant = emotions[0]["label"] if emotions else None | |
| confidence = float(max(probs)) if len(probs) > 0 else 0.0 | |
| return { | |
| "original_text": text, | |
| "cleaned_text": cleaned, | |
| "emotions": emotions, | |
| "detected_emotions": detected_emotions, | |
| "dominant_emotion": dominant, | |
| "confidence_score": round(confidence, 4), | |
| } | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Routes | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def root(): | |
| return { | |
| "service": "Emotion Detection API", | |
| "model": "RoBERTa-base (SemEval-2018)", | |
| "labels": EMOTION_LABELS, | |
| "endpoints": ["/predict", "/predict/batch", "/thresholds", "/health"], | |
| } | |
| def health(): | |
| return {"status": "ok", "device": str(DEVICE), "model_loaded": True} | |
| def get_thresholds(): | |
| return {"adaptive_thresholds": ADAPTIVE_THRESHOLDS} | |
| def predict(req: PredictRequest): | |
| if not req.text.strip(): | |
| raise HTTPException(status_code=400, detail="Text cannot be empty") | |
| if len(req.text) > 2000: | |
| raise HTTPException(status_code=400, detail="Text too long (max 2000 chars)") | |
| try: | |
| result = run_inference(req.text, req.threshold_mode, req.fixed_threshold) | |
| return result | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=str(e)) | |
| def predict_batch(req: BatchRequest): | |
| if not req.texts: | |
| raise HTTPException(status_code=400, detail="texts list is empty") | |
| if len(req.texts) > 50: | |
| raise HTTPException(status_code=400, detail="Max 50 texts per batch") | |
| results = [] | |
| for text in req.texts: | |
| try: | |
| result = run_inference(text, req.threshold_mode) | |
| results.append({"text": text, "result": result, "error": None}) | |
| except Exception as e: | |
| results.append({"text": text, "result": None, "error": str(e)}) | |
| return {"results": results, "total": len(results)} | |