rasha102004's picture
Update app.py (#7)
4b93afd
Raw
History Blame Contribute Delete
6.83 kB
import gradio as gr
import torch
import torch.nn as nn
import sys
import os
import uvicorn
from fastapi import FastAPI, UploadFile, File
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from models.audio_model import FrozenExtractor, LinearProjection, AGRU, PersonalityRegressor
from models.visual_model import VisualPersonalityModel
from models.fusion_model import FusionMLP
from inference.audio_inference import get_audio_embedding
from inference.visual_inference import get_visual_embedding, _load_face_detector
from inference.fusion_inference import get_fusion_prediction, format_predictions, TRAITS
# ── Device ─────────────────────────────────────────────────────────────────────
DEVICE = torch.device("cpu")
# ── Architecture config ────────────────────────────────────────────────────────
VISUAL_CFG = {
"backbone": "efficientnet_b0",
"embed_dim": 512,
"dropout": 0.3,
"traits": TRAITS,
}
N_SHALLOW = 9
GRU_DIM = 256
GRU_HEADS = 8
REG_FFN = 1024
REG_MLP = 256
REG_HEADS = 8
REG_DROPOUT = 0.1
PROJ_DROPOUT = 0.1
def _unwrap_ckpt(ckpt):
if not isinstance(ckpt, dict):
return ckpt
for key in ("model_state", "model_state_dict", "state_dict", "model"):
if key in ckpt:
return ckpt[key]
return ckpt
def _load_ckpt(model, path):
ckpt = torch.load(path, map_location=DEVICE, weights_only=False)
state = _unwrap_ckpt(ckpt)
model.load_state_dict(state)
return model
print("⏳ Loading visual model...")
visual_model = VisualPersonalityModel(VISUAL_CFG).to(DEVICE)
visual_model = _load_ckpt(visual_model, "checkpoints/best_model.pt")
visual_model.eval()
print("βœ“ Visual model ready")
print("⏳ Loading fusion model...")
fusion_model = FusionMLP(
in_dim=1024, hidden_dims=[512, 256], num_traits=5, dropout=0.3
).to(DEVICE)
fusion_model = _load_ckpt(fusion_model, "checkpoints/best_fusion_model.pt")
fusion_model.eval()
print("βœ“ Fusion model ready")
print("⏳ Loading audio model...")
frozen_ext = FrozenExtractor("facebook/wav2vec2-base", n_shallow=N_SHALLOW).to(DEVICE)
frozen_ext.eval()
from transformers import Wav2Vec2Model
_w2v = Wav2Vec2Model.from_pretrained("facebook/wav2vec2-base")
deep_transformer = nn.ModuleList(_w2v.encoder.layers[N_SHALLOW:]).to(DEVICE)
del _w2v
proj = LinearProjection(768, GRU_DIM, dropout=PROJ_DROPOUT).to(DEVICE)
agru = AGRU(dim=GRU_DIM, num_heads=GRU_HEADS).to(DEVICE)
regressor = PersonalityRegressor(
in_dim=GRU_DIM * 2,
ffn_dim=REG_FFN,
mlp_hidden=REG_MLP,
num_traits=5,
num_heads=REG_HEADS,
dropout=REG_DROPOUT,
).to(DEVICE)
audio_ckpt = torch.load("checkpoints/best_audio_model.pt", map_location=DEVICE, weights_only=False)
audio_state = _unwrap_ckpt(audio_ckpt)
if "deep_transformer" in audio_state:
deep_transformer.load_state_dict(audio_state["deep_transformer"])
proj.load_state_dict(audio_state["proj"])
agru.load_state_dict(audio_state["agru"])
regressor.load_state_dict(audio_state["regressor"])
else:
dt_state = {k[len("deep_transformer."):]: v for k, v in audio_state.items() if k.startswith("deep_transformer.")}
proj_state = {k[len("proj."):]: v for k, v in audio_state.items() if k.startswith("proj.")}
agru_state = {k[len("agru."):]: v for k, v in audio_state.items() if k.startswith("agru.")}
reg_state = {k[len("regressor."):]: v for k, v in audio_state.items() if k.startswith("regressor.")}
if dt_state:
deep_transformer.load_state_dict(dt_state)
proj.load_state_dict(proj_state)
agru.load_state_dict(agru_state)
regressor.load_state_dict(reg_state)
for m in [deep_transformer, proj, agru, regressor]:
m.eval()
print("βœ“ Audio model ready")
print("⏳ Loading face detector...")
face_detector = _load_face_detector()
print("βœ“ Face detector ready")
def process_video_bytes(video_bytes):
visual_preds, visual_emb = get_visual_embedding(video_bytes, visual_model, DEVICE, face_detector)
audio_preds, audio_emb = get_audio_embedding(video_bytes, frozen_ext, deep_transformer, proj, agru, regressor, DEVICE)
fusion_preds = get_fusion_prediction(audio_emb, visual_emb, fusion_model, DEVICE)
results = format_predictions(fusion_preds, visual_preds, audio_preds)
return results
app = FastAPI()
# Wide open CORS to completely eliminate React local testing errors
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=False,
allow_methods=["*"],
allow_headers=["*"],
)
@app.post("/api/predict")
async def api_predict(video: UploadFile = File(...)):
"""This route is exclusively for your React app to call via Axios"""
try:
video_bytes = await video.read()
results = process_video_bytes(video_bytes)
return JSONResponse(content=results)
except Exception as e:
import traceback
return JSONResponse(content={"error": str(e), "trace": traceback.format_exc()}, status_code=500)
def gradio_predict(video_path):
if video_path is None:
return None, "⚠️ Please upload a video first."
try:
with open(video_path, "rb") as f:
video_bytes = f.read()
results = process_video_bytes(video_bytes)
fusion = results["fusion"]
md = "## 🧠 Fusion Results (Main)\n\n"
for trait, score in fusion.items():
bar = "β–ˆ" * int(score / 5) + "β–‘" * (20 - int(score / 5))
md += f"**{trait}** `{score:.1f}%` {bar}\n\n"
return results, md
except Exception as e:
import traceback
return None, f"❌ Error: {str(e)}\n\n```{traceback.format_exc()}```"
with gr.Blocks(title="Personality Prediction API") as demo:
gr.Markdown("# 🧠 Personality Prediction")
with gr.Row():
with gr.Column(scale=1):
video_input = gr.Video(label="πŸ“Ή Upload Video")
predict_btn = gr.Button("πŸ” Predict Personality", variant="primary", size="lg")
with gr.Column(scale=1):
summary_out = gr.Markdown(label="Results Summary")
with gr.Accordion("πŸ“Š Full Results", open=False):
json_out = gr.JSON(label="Raw Scores (%)")
predict_btn.click(
fn=gradio_predict,
inputs=video_input,
outputs=[json_out, summary_out]
)
# Mount Gradio over the FastAPI app
app = gr.mount_gradio_app(app, demo, path="/")
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=7860)