Team Ai
Apppublic

codeby-hp/finetuneTinyBERT-SentimentClassification

sourceHugging Faceupdated 10mo agoView on Hugging Face
1likes
app.py111 linesDownload Raw Back to fastapi_app
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