Adchay/subject-topic-predictor
0
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"}