olive-gemma / app.py
multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
ec8284c verified
Raw
History Blame Contribute Delete
4.01 kB
import spaces # MUST come before any CUDA-touching import
import torch
import gradio as gr
from transformers import AutoProcessor, PaliGemmaForConditionalGeneration
MODEL_ID = "JamesZar/OliveGemma-3B"
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = PaliGemmaForConditionalGeneration.from_pretrained(
MODEL_ID, torch_dtype=torch.bfloat16
).to("cuda").eval()
QUESTIONS = [
"What is the name of this dish?",
"What are the likely ingredients of this dish?",
"What visible ingredients can you see?",
"What visual evidence supports this dish?",
"How is this dish different from a visually similar one?",
]
@spaces.GPU(duration=60)
def recognize(image, question: str, max_new_tokens: int = 64) -> str:
"""Recognise a Mediterranean or European dish from an image and answer a question about it.
Args:
image: A food photograph.
question: What to ask the model about the food (dish name, ingredients, etc.).
max_new_tokens: Maximum number of new tokens to generate.
"""
from PIL import Image
if image is None:
return "Please upload an image."
if not isinstance(image, Image.Image):
image = Image.open(image)
image = image.convert("RGB")
# PaliGemma prompt format used during training
prompt = f"<image>answer en {question}\n"
inputs = processor(text=prompt, images=image, return_tensors="pt").to(model.device)
in_len = inputs["input_ids"].shape[-1]
with torch.inference_mode():
out = model.generate(
**inputs,
max_new_tokens=int(max_new_tokens),
do_sample=False,
)
answer = processor.decode(out[0][in_len:], skip_special_tokens=True).strip()
return answer
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks() as demo:
gr.Markdown(
"# OliveGemma 馃珤\n"
"A 3B visual-language model for fine-grained Mediterranean & European food recognition. "
"Upload a food photo and ask about the dish name, ingredients, or visual evidence.\n\n"
"Model: [`JamesZar/OliveGemma-3B`](https://huggingface.co/JamesZar/OliveGemma-3B) 路 "
"Paper: [2608.03428](https://huggingface.co/papers/2608.03428) 路 "
"Code: [GitHub](https://github.com/tsiokris/OliveGemma)"
)
with gr.Column(elem_id="col-container"):
with gr.Row():
image_input = gr.Image(type="pil", label="Food image", scale=1)
with gr.Column(scale=1):
question_input = gr.Dropdown(
choices=QUESTIONS,
value=QUESTIONS[0],
label="Question",
interactive=True,
)
recognize_btn = gr.Button("Recognise", variant="primary")
output_text = gr.Textbox(label="Answer", lines=4, interactive=False)
with gr.Accordion("Advanced settings", open=False):
max_tokens = gr.Slider(
minimum=16, maximum=256, value=64, step=16,
label="Max new tokens",
)
gr.Examples(
examples=[
["examples/pizza_board.jpg", QUESTIONS[0]],
["examples/sushi_nigiri.jpg", QUESTIONS[0]],
["examples/pancakes_berries.jpg", QUESTIONS[0]],
["examples/gourmet_burger.jpg", QUESTIONS[1]],
["examples/macarons.jpg", QUESTIONS[0]],
["examples/cake.jpg", QUESTIONS[2]],
],
inputs=[image_input, question_input],
outputs=output_text,
fn=recognize,
cache_examples=True,
cache_mode="lazy",
)
recognize_btn.click(
fn=recognize,
inputs=[image_input, question_input, max_tokens],
outputs=output_text,
api_name="recognize",
)
if __name__ == "__main__":
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)