CoolFace
Apppublic

Adchay/subject-topic-predictor

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py95 linesDownload Raw Back to root
1"""2FastAPI server inside Hugging Face Space3POST /predict  ->  zero-shot subject prediction + save to TiDB4"""5import os6import time7from contextlib import asynccontextmanager8 9import mysql.connector10import torch11from transformers import AutoTokenizer, AutoModelForSequenceClassification12from fastapi import FastAPI, HTTPException13from pydantic import BaseModel14 15# ---------- load model ONCE ----------16MODEL_NAME = "MoritzLaurer/deberta-v3-large-zeroshot-v1.1-all-33"17LABELS = [18    "Mathematics", "Physics", "Chemistry", "Biology",19    "History", "Geography", "Literature", "Computer-Science"20]21 22ml_models = {}23 24@asynccontextmanager25async def lifespan(app: FastAPI):26    # load at startup27    tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)28    model = AutoModelForSequenceClassification.from_pretrained(MODEL_NAME)29    model.eval()30    if torch.cuda.is_available():31        model.cuda()32    ml_models["tokenizer"] = tokenizer33    ml_models["model"] = model34    yield35    # shutdown36    ml_models.clear()37 38app = FastAPI(lifespan=lifespan)39 40# ---------- DB helper ----------41def get_conn():42    return mysql.connector.connect(43        host=os.getenv("DB_HOST"),44        port=int(os.getenv("DB_PORT", 4000)),45        user=os.getenv("DB_USER"),46        password=os.getenv("DB_PASS"),47        database=os.getenv("DB_NAME"),48        ssl_ca=os.getenv("DB_SSL_CA_PATH") or None49    )50 51# ---------- request schema ----------52class PredictRequest(BaseModel):53    student_id: str54    text: str55 56# ---------- API endpoint ----------57@app.post("/predict")58def predict(req: PredictRequest):59    if not req.text.strip():60        raise HTTPException(400, "Empty text")61    tok = ml_models["tokenizer"](62        req.text,63        padding=True,64        truncation=True,65        return_tensors="pt"66    )67    if torch.cuda.is_available():68        tok = {k: v.cuda() for k, v in tok.items()}69    with torch.no_grad():70        logits = ml_models["model"](**tok).logits71        probs = torch.softmax(logits, dim=-1)[0]72        idx = int(torch.argmax(probs))73        subject = LABELS[idx]74 75    # save to DB76    try:77        conn = get_conn()78        cur = conn.cursor()79        cur.execute(80            "INSERT INTO log_table (student_id, input_sample, subject, prediction_time) "81            "VALUES (%s, %s, %s, %s)",82            (req.student_id, req.text, subject, time.strftime('%Y-%m-%d %H:%M:%S'))83        )84        conn.commit()85        cur.close()86        conn.close()87    except Exception as e:88        print("DB error:", e)89        raise HTTPException(500, "DB write failed")90 91    return {"subject": subject}92 93@app.get("/")94def root():95    return {"message": "Subject predictor is running"}