lyfesan/Personality-Interpretation-Testing
0
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 