CoolFace
Apppublic

lyfesan/Personality-Interpretation-Testing

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes
app.py319 linesDownload Raw Back to root
1import gradio as gr2import requests3import base644import json5import tempfile6import os7import concurrent.futures8from io import BytesIO9from PIL import Image10 11INFERENCE_API_URL = os.getenv("INFERENCE_API_URL", "http://127.0.0.1:8000")12INTERPRETATION_API_URL = os.getenv("INTERPRETATION_API_URL", "http://127.0.0.1:8080")13 14 15def get_available_models():16    """Fetch available models from the FastAPI server."""17    try:18        response = requests.get(f"{INFERENCE_API_URL}/models", timeout=(2, 2))19        if response.status_code == 200:20            models_data = response.json().get("available_models", [])21            return [(f"{m.get('name', m.get('id'))}", m.get("id")) for m in models_data]22    except Exception as e:23        print(f"Warning: Could not fetch models from API ({e}). Using defaults.")24    return [("SwinV2 (swinv2)", "swinv2"), ("ViT (vit)", "vit"), ("PVTv2 (pvtv2)", "pvtv2")]25 26def predict(image, model_type):27    if image is None:28        return {"error": "Please upload an image."}, None29    if not model_type:30        return {"error": "Please select a model."}, None31    32    # Convert PIL Image to Base6433    if image.mode in ("RGBA", "P"):34        image = image.convert("RGB")35    buffered = BytesIO()36    image.save(buffered, format="JPEG")37    img_str = base64.b64encode(buffered.getvalue()).decode("utf-8")38    39    payload = {40        "model_type": model_type,41        "image_base64": img_str42    }43    44    try:45        #gr.Info("Analyzing image... (Note: If the backend Space was asleep, waking it up may take up to 2-3 minutes!)")46        response = requests.post(f"{INFERENCE_API_URL}/predict", json=payload, timeout=180)47        if response.status_code == 200:48            data = response.json()49            predictions = data.get("predictions", {})50            cropped_b64 = data.get("cropped_face_base64")51            52            cropped_img = None53            if cropped_b64:54                try:55                    img_data = base64.b64decode(cropped_b64)56                    cropped_img = Image.open(BytesIO(img_data)).convert("RGB")57                except Exception:58                    pass59                    60            return predictions, cropped_img61        else:62            return {"error": f"HTTP {response.status_code}", "details": response.text}, None63    except Exception as e:64        return {"error": "Connection failed. Is the API running?", "details": str(e)}, None65 66def get_inference_models():67    """Fetch inference models from the interpretation API."""68    try:69        response = requests.get(f"{INTERPRETATION_API_URL}/inference-models", timeout=(2, 2))70        if response.status_code == 200:71            data = response.json()72            if isinstance(data, dict):73                return data.get("available_models", [])74            return data75    except Exception as e:76        print(f"Warning: Could not fetch inference models ({e}).")77    return ["swinv2", "vit", "pvtv2"]78 79def get_llm_models():80    """Fetch allowed LLM models from the interpretation API."""81    try:82        response = requests.get(f"{INTERPRETATION_API_URL}/llm-models", timeout=(2, 2))83        if response.status_code == 200:84            models = response.json()85            return [(m["name"], m["id"]) for m in models]86    except Exception as e:87        print(f"Warning: Could not fetch LLM models ({e}).")88    return [("Gemma 4 31B (free)", "google/gemma-4-31b-it:free")]89 90def get_response_styles():91    """Fetch allowed response styles from the interpretation API."""92    try:93        response = requests.get(f"{INTERPRETATION_API_URL}/response-styles", timeout=(2, 2))94        if response.status_code == 200:95            styles = response.json()96            return [(s["name"], s["id"]) for s in styles]97    except Exception as e:98        print(f"Warning: Could not fetch response styles ({e}).")99    return [("Comprehensive (ID)", "comprehensive_id")]100 101def interpret(image, inference_model, llm_model, style_id):102    """Send image to the interpretation API via multipart/form-data."""103    if image is None:104        return {}, "Please upload an image."105    if not inference_model:106        return {}, "Please select an inference model."107    if not llm_model:108        return {}, "Please select an LLM model."109 110    # Convert PIL image to bytes for multipart upload111    if image.mode in ("RGBA", "P"):112        image = image.convert("RGB")113    buffered = BytesIO()114    image.save(buffered, format="JPEG")115    buffered.seek(0)116 117    try:118        files = {"image": ("image.jpg", buffered, "image/jpeg")}119        data = {120            "inference_model": inference_model,121            "llm_model": llm_model,122            "style_id": style_id,123        }124        response = requests.post(125            f"{INTERPRETATION_API_URL}/interpret",126            files=files,127            data=data,128            timeout=120,129        )130        if response.status_code == 200:131            result = response.json()132            traits = result.get("predictions", {})133            interpretation = result.get("interpretation", "No interpretation returned.")134            return traits, interpretation135        else:136            err = response.json().get("error", response.text)137            return {}, f"Error {response.status_code}: {err}"138    except Exception as e:139        return {}, f"Connection failed. Is the interpretation API running?\n{e}"140 141def export_result(image, inf_model, llm_id, style_id, traits, interpretation):142    """Exports the results to a JSON file and returns the temp file path."""143    if not traits and not interpretation:144        return None145        146    img_b64 = None147    if image is not None:148        if image.mode in ("RGBA", "P"):149            image = image.convert("RGB")150        buffered = BytesIO()151        image.save(buffered, format="JPEG")152        img_b64 = base64.b64encode(buffered.getvalue()).decode("utf-8")153        154    data = {155        "parameters": {156            "inference_model": inf_model,157            "llm_model": llm_id,158            "response_style": style_id159        },160        "results": {161            "predictions": traits,162            "interpretation": interpretation163        },164        "image_base64": img_b64165    }166    167    fd, path = tempfile.mkstemp(suffix=".json", prefix="personality_export_")168    with os.fdopen(fd, 'w', encoding='utf-8') as f:169        json.dump(data, f, indent=4)170        171    return path172 173 174def load_all_dropdowns():175    """Fetch all dynamic models, LLMs, and styles concurrently to populate dropdowns on page load/refresh."""176    with concurrent.futures.ThreadPoolExecutor() as executor:177        f_models = executor.submit(get_available_models)178        f_inf = executor.submit(get_inference_models)179        f_llm = executor.submit(get_llm_models)180        f_styles = executor.submit(get_response_styles)181        182        try:183            models = f_models.result(timeout=2.5)184        except concurrent.futures.TimeoutError:185            models = [("SwinV2 (swinv2)", "swinv2"), ("ViT (vit)", "vit"), ("PVTv2 (pvtv2)", "pvtv2")]186            187        try:188            inf_models_raw = f_inf.result(timeout=2.5)189        except concurrent.futures.TimeoutError:190            inf_models_raw = ["swinv2", "vit", "pvtv2"]191            192        try:193            llm_models = f_llm.result(timeout=2.5)194        except concurrent.futures.TimeoutError:195            llm_models = [("Gemma 4 31B (free)", "google/gemma-4-31b-it:free")]196            197        try:198            response_styles = f_styles.result(timeout=2.5)199        except concurrent.futures.TimeoutError:200            response_styles = [("Comprehensive (ID)", "comprehensive_id")]201    202    id_to_name = {m_id: m_name for m_name, m_id in models}203    204    inf_models = []205    for m in inf_models_raw:206        if isinstance(m, dict):207            inf_models.append((m.get("name", m.get("id")), m.get("id")))208        else:209            inf_models.append((id_to_name.get(m, m), m))210 211    return (212        gr.update(choices=models, value=models[0][1] if models else None),213        gr.update(choices=inf_models, value=inf_models[0][1] if inf_models else None),214        gr.update(choices=llm_models, value=llm_models[0][1] if llm_models else None),215        gr.update(choices=response_styles, value=response_styles[0][1] if response_styles else None)216    )217 218 219def build_app():220    models = get_available_models()221    inf_models_raw = get_inference_models()222    223    id_to_name = {m_id: m_name for m_name, m_id in models}224    225    inf_models = []226    for m in inf_models_raw:227        if isinstance(m, dict):228            inf_models.append((m.get("name", m.get("id")), m.get("id")))229        else:230            inf_models.append((id_to_name.get(m, m), m))231 232    llm_models = get_llm_models()233    response_styles = get_response_styles()234 235    with gr.Blocks(title="Personality Interpretation") as demo:236        gr.Markdown("# Personality Analysis")237 238        with gr.Tabs():239            with gr.TabItem("🔬 Inference"):240                with gr.Row():241                    with gr.Column():242                        image_input = gr.Image(type="pil", label="Face Image")243                        model_dropdown = gr.Dropdown(244                            choices=models, 245                            value=models[0][1] if models else None, 246                            label="Inference Model"247                        )248 249                        submit_btn = gr.Button("Predict Personality", variant="primary")250                        251                    with gr.Column():252                        output_json = gr.JSON(label="Personality Traits (Big Five)")253                        cropped_output = gr.Image(type="pil", label="Extracted Face (Model Input)")254        255                # Action mappings256                submit_btn.click(257                    fn=predict,258                    inputs=[image_input, model_dropdown],259                    outputs=[output_json, cropped_output]260                )261                262            with gr.TabItem("✨ Interpretation"):263                with gr.Row():264                    with gr.Column():265                        interp_image = gr.Image(type="pil", label="Face Image")266                        with gr.Row():267                            interp_inf_dropdown = gr.Dropdown(268                                choices=inf_models,269                                value=inf_models[0][1] if inf_models else None,270                                label="Inference Model",271                            )272                            interp_llm_dropdown = gr.Dropdown(273                                choices=llm_models,274                                value=llm_models[0][1] if llm_models else None,275                                label="LLM Model",276                            )277                        style_dropdown = gr.Dropdown(278                            choices=response_styles,279                            value=response_styles[0][1] if response_styles else None,280                            label="Response Style"281                        )282                        interp_btn = gr.Button("Interpret Personality", variant="primary")283                    with gr.Column():284                        interp_traits = gr.JSON(label="Predicted Traits (Big Five)")285                        interp_text = gr.Markdown(label="LLM Interpretation", value="*Interpretation will appear here...*")286                        287                        export_btn = gr.DownloadButton("Export Result as JSON", variant="secondary")288 289                def on_interpret(image, inf_model, llm_id, style_id):290                    return interpret(image, inf_model, llm_id, style_id)291 292                interp_btn.click(293                    fn=on_interpret,294                    inputs=[interp_image, interp_inf_dropdown, interp_llm_dropdown, style_dropdown],295                    outputs=[interp_traits, interp_text],296                )297                298                export_btn.click(299                    fn=export_result,300                    inputs=[interp_image, interp_inf_dropdown, interp_llm_dropdown, style_dropdown, interp_traits, interp_text],301                    outputs=[export_btn]302                )303 304        demo.load(305            fn=load_all_dropdowns,306            inputs=[],307            outputs=[model_dropdown, interp_inf_dropdown, interp_llm_dropdown, style_dropdown]308        )309 310    return demo311 312 313if __name__ == "__main__":314    app = build_app()315    server_name = os.getenv("GRADIO_SERVER_NAME", "0.0.0.0")316    server_port = int(os.getenv("GRADIO_SERVER_PORT", 7860))317    app.launch(server_name=server_name, server_port=server_port, share=False, theme=gr.themes.Soft())318 319