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)