| 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 = torch.device("cpu") |
|
|
| |
| 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() |
|
|
| |
| 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] |
| ) |
|
|
| |
| app = gr.mount_gradio_app(app, demo, path="/") |
|
|
| if __name__ == "__main__": |
| uvicorn.run(app, host="0.0.0.0", port=7860) |