waseem11's picture
Update app.py
974290d verified
Raw
History Blame Contribute Delete
2.73 kB
import os
import re
from contextlib import asynccontextmanager
from typing import Literal
import torch
from fastapi import FastAPI
from pydantic import BaseModel, Field
from transformers import AutoConfig, AutoModelForSequenceClassification, AutoTokenizer
MODEL_NAME = os.getenv("MODEL_NAME", "distilbert-base-multilingual-cased")
MODEL_WEIGHTS_PATH = os.getenv("MODEL_WEIGHTS_PATH", "depression_model.pt")
TOKENIZER_PATH = os.getenv("TOKENIZER_PATH", "tokenizer")
_device: torch.device | None = None
_tokenizer = None
_model = None
def clean_text(text: str) -> str:
text = str(text).lower()
text = re.sub(r"http\S+", "", text) # remove urls
text = re.sub(r"@\w+", "", text) # remove mentions
text = re.sub(r"[^\w\s]", "", text) # remove emojis + punctuation
text = re.sub(r"\s+", " ", text) # remove extra spaces
return text.strip()
def _load_model() -> None:
global _device, _tokenizer, _model
_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
tokenizer_source = TOKENIZER_PATH if os.path.isdir(TOKENIZER_PATH) else MODEL_NAME
_tokenizer = AutoTokenizer.from_pretrained(tokenizer_source)
# Avoid downloading large base-model weights. We only need the config to
# construct the architecture, then we load your local trained weights.
config = AutoConfig.from_pretrained(MODEL_NAME, num_labels=2)
_model = AutoModelForSequenceClassification.from_config(config)
state_dict = torch.load(MODEL_WEIGHTS_PATH, map_location="cpu")
_model.load_state_dict(state_dict)
_model.to(_device)
_model.eval()
@asynccontextmanager
async def lifespan(_: FastAPI):
_load_model()
yield
app = FastAPI(title="Depression Text Classifier", version="1.0.0", lifespan=lifespan)
class PredictRequest(BaseModel):
text: str = Field(..., min_length=1, description="Input text to classify")
class PredictResponse(BaseModel):
label: Literal["Depressed", "Non-Depressed"]
class_id: int
@app.get("/health")
def health() -> dict:
return {"status": "ok"}
@app.post("/predict", response_model=PredictResponse)
def predict(req: PredictRequest) -> PredictResponse:
if _model is None or _tokenizer is None or _device is None:
_load_model()
text = clean_text(req.text)
inputs = _tokenizer(
text,
return_tensors="pt",
truncation=True,
padding=True,
max_length=128,
)
inputs = {k: v.to(_device) for k, v in inputs.items()}
with torch.no_grad():
outputs = _model(**inputs)
pred = int(torch.argmax(outputs.logits, dim=1).item())
label = "Depressed" if pred == 1 else "Non-Depressed"
return PredictResponse(label=label, class_id=pred)