MindSense / app.py
Yashdesai07's picture
Update app.py
a2890ed verified
Raw
History Blame Contribute Delete
6.08 kB
"""
FastAPI backend + frontend server mapped for Hugging Face Spaces.
Models are loaded ONCE at startup and reused for every request.
"""
import os
import sys
import shutil
import tempfile
import subprocess
import gradio as gr
from fastapi import FastAPI, File, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from contextlib import asynccontextmanager
# Dynamically compute paths relative to this file's location
_here = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, os.path.join(_here, "scripts"))
import torch
from feature_extractors import (
_load_audio_model, _load_text_model,
_load_whisper, _load_face_model,
extract_all
)
from fusion_model import AttentionFusion, ConcatFusion, EMBED_DIM
# Safe cross-platform path resolution
WEIGHTS_PATH = os.environ.get("FUSION_WEIGHTS", os.path.join(_here, "models", "fusion.pt"))
LABELS = ["low risk", "moderate risk", "high risk"]
# Global model handle -- set at startup
_fusion_model = None
def build_model(kind="attention", **kwargs):
if kind == "concat":
return ConcatFusion(**kwargs)
return AttentionFusion(**kwargs)
def get_device():
if torch.cuda.is_available():
return "cuda"
return "cpu"
@asynccontextmanager
async def lifespan(app):
"""Load ALL models once when the server starts."""
global _fusion_model
device = get_device()
print(f"\n[startup] Loading models on {device} ...")
print("[startup] 1/4 wav2vec2 audio model...")
_load_audio_model()
print("[startup] 2/4 Whisper ASR...")
_load_whisper()
print("[startup] 3/4 Text model (MentalBERT / RoBERTa)...")
_load_text_model()
print("[startup] 4/4 ViT face model...")
_load_face_model()
print("[startup] 5/5 Fusion head...")
checkpoint = torch.load(WEIGHTS_PATH, map_location=device)
_fusion_model = build_model(checkpoint["model_type"], embed_dim=EMBED_DIM).to(device)
_fusion_model.load_state_dict(checkpoint["state_dict"])
_fusion_model.eval()
print("[startup] All models loaded β€” ready to serve requests.\n")
yield
# cleanup on shutdown (nothing needed for torch models)
app = FastAPI(title="Mental Health Multimodal Fusion API", lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
def extract_audio_from_video(video_path: str, audio_path: str):
subprocess.run([
"ffmpeg", "-y", "-i", video_path,
"-vn", "-acodec", "pcm_s16le",
"-ar", "16000", "-ac", "1", audio_path
], check=True)
# try:
# import noisereduce as nr
# import soundfile as sf
# import librosa
# y, sr = librosa.load(raw_path, sr=16000)
# y_clean = nr.reduce_noise(y=y, sr=sr, stationary=True)
# sf.write(audio_path, y_clean, sr)
# os.remove(raw_path)
# except Exception:
# os.rename(raw_path, audio_path)
def predict_with_loaded_models(audio_path: str, video_path: str) -> dict:
"""Run inference using already-loaded global models -- no reload."""
device = get_device()
audio_emb, text_emb, face_emb, transcript = extract_all(audio_path, video_path)
a = torch.from_numpy(audio_emb).unsqueeze(0).to(device)
t = torch.from_numpy(text_emb).unsqueeze(0).to(device)
f = torch.from_numpy(face_emb).unsqueeze(0).to(device)
with torch.no_grad():
if hasattr(_fusion_model, "pool_weights"):
logits, modality_weights = _fusion_model(a, t, f, return_weights=True)
modality_weights = modality_weights.squeeze(0).cpu().tolist()
else:
logits = _fusion_model(a, t, f)
modality_weights = None
probs = torch.softmax(logits, dim=-1).squeeze(0).cpu().tolist()
pred_idx = int(torch.argmax(logits, dim=-1).item())
return {
"transcript": transcript,
"prediction": LABELS[pred_idx],
"probabilities": dict(zip(LABELS, probs)),
"modality_contribution": dict(zip(["audio", "text", "face"], modality_weights))
if modality_weights else None,
}
# ── API routes ──────────────────────────────────────────────────────────────
@app.get("/health")
def health():
return {
"status": "ok",
"weights_loaded": os.path.exists(WEIGHTS_PATH),
"model_ready": _fusion_model is not None,
"device": get_device(),
}
@app.post("/predict")
async def predict_endpoint(video: UploadFile = File(...)):
with tempfile.TemporaryDirectory() as tmp:
video_path = os.path.join(tmp, video.filename)
with open(video_path, "wb") as f:
shutil.copyfileobj(video.file, f)
audio_path = os.path.join(tmp, "extracted_audio.wav")
extract_audio_from_video(video_path, audio_path)
result = predict_with_loaded_models(audio_path, video_path)
return result
# ── Static files mount ──────────────────────────────────────────────────────
STATIC_DIR = os.path.join(_here, "scripts")
app.mount("/", StaticFiles(directory=STATIC_DIR, html=True), name="frontend")
# ── Gradio Mount Gateway Definition ─────────────────────────────────────────
# DO NOT reassign 'app = gr.mount_gradio_app(...)'.
# Instead, keep 'app' and 'demo' as separate global objects.
with gr.Blocks() as demo:
gr.Markdown("# MindSense API Gateway Core Active")
gr.Markdown("Your custom HTML frontend is running at the root `/` URL pathway.")
# Mount your FastAPI app *into* the Gradio interface structure securely
gr.mount_gradio_app(app, demo, path="/gradio_backend")
if __name__ == "__main__":
import uvicorn
uvicorn.run("app:app", host="0.0.0.0", port=7860)