codeby-hp/finetuneTinyBERT-SentimentClassification
1
1import torch2import logging3from contextlib import asynccontextmanager4from fastapi import FastAPI, Request, Form5from fastapi.responses import HTMLResponse6from fastapi.templating import Jinja2Templates7from fastapi.staticfiles import StaticFiles8from transformers import AutoModelForSequenceClassification, AutoTokenizer9 10logging.basicConfig(level=logging.INFO)11logger = logging.getLogger(__name__)12 13model = None14tokenizer = None15device = torch.device("cuda" if torch.cuda.is_available() else "cpu")16 17 18@asynccontextmanager19async def lifespan(app: FastAPI):20 """Load model on startup and cleanup on shutdown"""21 global model, tokenizer22 23 try:24 model_id = "codeby-hp/FinetuneTinybert-SentimentClassification"25 26 logger.info(f"Loading tokenizer from {model_id}...")27 tokenizer = AutoTokenizer.from_pretrained(model_id)28 29 logger.info(f"Loading model from {model_id}...")30 model = AutoModelForSequenceClassification.from_pretrained(model_id)31 model.to(device)32 model.eval()33 34 logger.info(f"Model loaded successfully on {device}")35 except Exception as e:36 logger.error(f"Error loading model: {e}")37 raise38 39 yield40 41 logger.info("Shutting down...")42 43 44app = FastAPI(title="Sentiment Analysis API", lifespan=lifespan)45 46templates = Jinja2Templates(directory="templates")47 48 49@app.get("/", response_class=HTMLResponse)50async def home(request: Request):51 """Render the home page"""52 return templates.TemplateResponse("index.html", {"request": request})53 54 55@app.post("/predict")56async def predict(request: Request, text: str = Form(...)):57 """Predict sentiment for the given text"""58 if not text.strip():59 return templates.TemplateResponse(60 "index.html",61 {"request": request, "error": "Please enter some text to analyze"},62 )63 64 try:65 inputs = tokenizer(66 text, return_tensors="pt", truncation=True, max_length=512, padding=True67 )68 inputs = {k: v.to(device) for k, v in inputs.items()}69 70 with torch.no_grad():71 outputs = model(**inputs)72 logits = outputs.logits73 probabilities = torch.nn.functional.softmax(logits, dim=-1)74 predicted_class = torch.argmax(probabilities, dim=-1).item()75 confidence = probabilities[0][predicted_class].item()76 77 sentiment_map = {0: "Negative", 1: "Positive"}78 sentiment = sentiment_map.get(predicted_class, "Unknown")79 80 return templates.TemplateResponse(81 "index.html",82 {83 "request": request,84 "text": text,85 "sentiment": sentiment,86 "confidence": round(confidence * 100, 2),87 },88 )89 90 except Exception as e:91 logger.error(f"Prediction error: {e}")92 return templates.TemplateResponse(93 "index.html", {"request": request, "error": f"An error occurred: {str(e)}"}94 )95 96 97@app.get("/health")98async def health_check():99 """Health check endpoint"""100 return {101 "status": "healthy",102 "model_loaded": model is not None,103 "device": str(device),104 }105 106 107if __name__ == "__main__":108 import uvicorn109 110 uvicorn.run(app, host="0.0.0.0", port=7860)111 