fusin001's picture
UI clarity: remove notes column and rename OpenAI section to ChatGPT Top-1
b470575 verified
Raw
History Blame Contribute Delete
8.35 kB
from __future__ import annotations
import base64
import json
import os
from functools import lru_cache
from pathlib import Path
import gradio as gr
import gradio_client.utils as gr_client_utils
from openai import OpenAI
from transformers import pipeline
MODEL_ID = "fusin001/pokemon-vit"
OPEN_SOURCE_MODEL_ID = "google/siglip-base-patch16-224"
LABELS = ["charizard", "charmander", "charmeleon", "ditto", "eevee", "ekans"]
EXAMPLE_DIR = Path(__file__).resolve().parent / "example_images"
_ORIGINAL_JSON_SCHEMA_TO_PYTHON_TYPE = gr_client_utils._json_schema_to_python_type
def _patched_json_schema_to_python_type(schema, defs=None):
# Gradio client can crash when JSON schema uses boolean additionalProperties.
if isinstance(schema, bool):
return "Any" if schema else "None"
return _ORIGINAL_JSON_SCHEMA_TO_PYTHON_TYPE(schema, defs)
gr_client_utils._json_schema_to_python_type = _patched_json_schema_to_python_type
def format_predictions(results):
return [{"label": item["label"], "score": round(float(item["score"]), 4)} for item in results]
def predictions_to_markdown(title: str, preds: list[dict]) -> str:
lines = [f"**{title}**"]
for i, item in enumerate(preds[:3], start=1):
lines.append(f"{i}. {item['label']} ({item['score']})")
return "\n".join(lines)
def comparison_to_markdown(rows: list[list]) -> str:
table = [
"| Model | Prediction | Score |",
"| --- | --- | --- |",
]
for model, prediction, score in rows:
table.append(f"| {model} | {prediction} | {score} |")
return "\n".join(table)
@lru_cache(maxsize=1)
def load_custom_model():
return pipeline("image-classification", model=MODEL_ID)
@lru_cache(maxsize=1)
def load_clip_model():
return pipeline("zero-shot-image-classification", model=OPEN_SOURCE_MODEL_ID)
@lru_cache(maxsize=1)
def load_openai_client():
api_key = os.getenv("OPENAI_API_KEY")
if not api_key:
return None
return OpenAI(api_key=api_key)
def image_to_data_url(image_path: str) -> str:
mime_type = "image/png" if image_path.lower().endswith(".png") else "image/jpeg"
with open(image_path, "rb") as image_file:
encoded = base64.b64encode(image_file.read()).decode("utf-8")
return f"data:{mime_type};base64,{encoded}"
def openai_predict(image_path: str) -> dict:
client = load_openai_client()
if client is None:
return {"label": "OPENAI_API_KEY missing", "score": 0.0, "rationale": "Add the secret in the Space settings."}
prompt = "Classify the Pokémon image into exactly one label from: charizard, charmander, charmeleon, ditto, eevee, ekans. Return JSON with keys label, confidence, rationale."
response = client.chat.completions.create(
model="gpt-4o-mini",
temperature=0,
response_format={"type": "json_object"},
messages=[
{"role": "system", "content": "You are a precise image classifier for Pokémon images."},
{
"role": "user",
"content": [
{"type": "text", "text": prompt},
{"type": "image_url", "image_url": {"url": image_to_data_url(image_path)}},
],
},
],
)
content = response.choices[0].message.content or "{}"
try:
parsed = json.loads(content)
except json.JSONDecodeError:
parsed = {"label": content.strip(), "confidence": None, "rationale": "Unparsed OpenAI response."}
label = str(parsed.get("label", "unknown")).strip()
confidence = parsed.get("confidence")
try:
confidence = float(confidence)
except (TypeError, ValueError):
confidence = 0.0
rationale = str(parsed.get("rationale", "")).strip()
return {"label": label, "score": confidence, "rationale": rationale}
def compare_models(image_path: str):
custom_model = load_custom_model()
clip_model = load_clip_model()
custom_predictions = format_predictions(custom_model(image_path, top_k=3))
clip_predictions = format_predictions(clip_model(image_path, candidate_labels=LABELS, top_k=3))
openai_prediction = openai_predict(image_path)
comparison = [
["Custom ViT", custom_predictions[0]["label"], custom_predictions[0]["score"]],
["SigLIP", clip_predictions[0]["label"], clip_predictions[0]["score"]],
["ChatGPT", openai_prediction["label"], openai_prediction["score"]],
]
summary = (
f"### Prediction Summary\n\n"
f"- **Custom ViT:** {custom_predictions[0]['label']} ({custom_predictions[0]['score']})\n"
f"- **SigLIP:** {clip_predictions[0]['label']} ({clip_predictions[0]['score']})\n"
f"- **ChatGPT:** {openai_prediction['label']} ({openai_prediction['score']})"
)
openai_text = (
"**ChatGPT Top-1**\n"
f"1. {openai_prediction['label']} ({openai_prediction['score']})\n\n"
f"**Reasoning:** {openai_prediction.get('rationale', '')}"
)
custom_text = predictions_to_markdown("Custom ViT Top-3", custom_predictions)
clip_text = predictions_to_markdown("SigLIP Top-3", clip_predictions)
comparison_md = comparison_to_markdown(comparison)
return custom_text, clip_text, openai_text, comparison_md, summary
example_images = [
str(EXAMPLE_DIR / "charizard_11.png"),
str(EXAMPLE_DIR / "charmander_1.png"),
str(EXAMPLE_DIR / "eevee_11.png"),
]
css = """
:root {
--bg: #0b0f19;
--panel: #171c28;
--panel-2: #1e2433;
--ink: #f2f4f8;
--muted: #b3bbca;
--accent: #f97316;
--accent-2: #fb923c;
--border: #2a3244;
}
body, .gradio-container {
background: var(--bg);
color: var(--ink);
}
#hero {
border-radius: 14px;
padding: 12px 2px 14px 2px;
color: var(--ink);
}
.panel {
background: var(--panel);
border: 1px solid var(--border);
border-radius: 14px;
padding: 14px;
}
.gradio-container {
max-width: 1260px !important;
}
h1, h2, h3, h4, p, li, label, .prose, .prose * {
color: var(--ink) !important;
}
.prose p {
color: var(--muted) !important;
}
button.primary {
background: linear-gradient(90deg, var(--accent), var(--accent-2)) !important;
border: none !important;
color: #0b0f19 !important;
font-weight: 700 !important;
}
.gr-form, .gr-box, .gr-panel {
background: var(--panel) !important;
border-color: var(--border) !important;
}
table {
border-collapse: collapse;
width: 100%;
}
th, td {
border: 1px solid var(--border) !important;
color: var(--ink) !important;
background: var(--panel-2) !important;
padding: 8px;
}
th {
font-weight: 700;
}
@media (max-width: 900px) {
.panel {
padding: 12px;
}
}
"""
with gr.Blocks(css=css, theme=gr.themes.Base(primary_hue="orange", neutral_hue="slate")) as demo:
gr.Markdown(
"""
<div id="hero">
<h1>Pokemon Classification Comparison</h1>
<p>Upload an image and compare results from a trained ViT model, a zero-shot SigLIP model, and OpenAI vision.</p>
</div>
"""
)
with gr.Row():
with gr.Column(scale=1, elem_classes=["panel"]):
image_input = gr.Image(type="filepath", label="Upload an image")
compare_button = gr.Button("Compare models", variant="primary")
gr.Examples(examples=example_images, inputs=image_input, label="Example images")
with gr.Column(scale=1, elem_classes=["panel"]):
summary_output = gr.Markdown()
comparison_output = gr.Markdown(label="Comparison table")
with gr.Row():
custom_output = gr.Markdown(label="Custom ViT top-3")
clip_output = gr.Markdown(label="SigLIP top-3")
openai_output = gr.Markdown(label="ChatGPT Top-1")
compare_button.click(
fn=compare_models,
inputs=image_input,
outputs=[custom_output, clip_output, openai_output, comparison_output, summary_output],
)
if __name__ == "__main__":
launch_kwargs = {
"server_name": "0.0.0.0",
"server_port": int(os.getenv("PORT", "7860")),
}
# Some Space runtimes report localhost as inaccessible; share avoids startup failure.
if os.getenv("SPACE_ID"):
demo.launch(**launch_kwargs, share=True)
else:
demo.launch(**launch_kwargs, share=False)