baseballtoma's picture
Add Gradio app.py with dynamic id_to_label
a0cc019 verified
Raw
History Blame Contribute Delete
3.03 kB
# app.py として保存することを想定したコードです
import gradio as gr
from PIL import Image as PILImage
import torch
from transformers import ViltProcessor, ViltForQuestionAnswering
# Hugging Face Hubからファインチューニング済みのモデルとプロセッサーを読み込みます
# ここでは、先ほどアップロードしたリポジトリIDを指定します
model_id = "baseballtoma/my-vilt-image-text-classifier-cloud"
# 推論に使用するid_to_labelマッピングをここで定義します
# これはファインチューニング時に使用したマッピングと一致している必要があります
app_id_to_label = {0: 'ゲリラ豪雨', 1: '乱層雲', 2: '台風', 3: '層雲', 4: '巻層雲', 5: '巻積雲', 6: '巻雲(乱)', 7: '巻雲(直)', 8: '席層雲(多)', 9: '席層雲(少)', 10: '晴天', 11: '猛暑', 12: '積乱雲', 13: '積雲', 14: '高層雲', 15: '高積雲(大)', 16: '高積雲(小)'}
processor = ViltProcessor.from_pretrained(model_id)
model = ViltForQuestionAnswering.from_pretrained(model_id)
# 推論を実行する関数を定義します
def predict(image, text):
if image is None:
return "画像をアップロードしてください。"
if not text:
return "テキストを入力してください。"
# PIL ImageをRGBに変換
image = image.convert("RGB")
# 画像とテキストをモデルの入力形式に変換
encoding = processor(images=image, text=text, return_tensors="pt", padding=True, truncation=True)
# 推論を実行 (CPUで実行されることを想定)
model.eval()
with torch.no_grad():
outputs = model(**encoding)
# ロジットから最も確率の高いラベルを取得
logits = outputs.logits
predicted_class_id = logits.argmax(-1).item()
# 予測されたクラスIDを元のテキストラベルに変換します
if predicted_class_id in app_id_to_label:
predicted_text = app_id_to_label[predicted_class_id]
else:
predicted_text = f"未知のクラスID: {predicted_class_id}"
return predicted_text
# Gradioインターフェースを作成します
iface = gr.Interface(
fn=predict, # 推論関数
inputs=[gr.Image(type="pil"), gr.Textbox(label="テキスト(質問)を入力")], # 入力コンポーネント (画像アップロードとテキスト入力)
outputs="text", # 出力コンポーネント (テキスト)
title="画像とテキストの分類デモ", # デモのタイトル
description="画像をアップロードし、テキストを入力して、モデルの予測結果を確認します。"
)
# インターフェースを起動します
# Hugging Face Spacesでは、このスクリプトが実行されると自動的に起動されます
# ローカルでテストする場合は、launch()の引数にshare=Trueなどを指定できます
iface.launch()