ResNet / app.py
JangTaeng's picture
Upload 8 files
b1fe7e3 verified
Raw
History Blame Contribute Delete
3.71 kB
"""Gradio ์›น ๋ฐ๋ชจ: Food-101 101๊ฐœ ์Œ์‹ ํด๋ž˜์Šค ๋ถ„๋ฅ˜๊ธฐ.
์ด ํŒŒ์ผ์€ Hugging Face Spaces์—์„œ ์ž๋™ ์‹คํ–‰๋ฉ๋‹ˆ๋‹ค (app_file: app.py).
์ž‘๋™ ๋ฐฉ์‹:
1) ๊ธฐ๋ณธ๊ฐ’: ํ—ˆ๋ธŒ์˜ 'nateraw/food' ๊ณต๊ฐœ ๋ชจ๋ธ ์‚ฌ์šฉ
(Pretrained ResNet + Food-101 fine-tuning์œผ๋กœ ~90% ์ •ํ™•๋„)
2) ์ง์ ‘ ํ•™์Šตํ•œ MyResNet์ด ์žˆ๋‹ค๋ฉด ์•„๋ž˜ USE_CUSTOM_RESNET=True๋กœ ๋ณ€๊ฒฝ
"""
import torch
import torch.nn.functional as F
import gradio as gr
# ============================================================
# ์„ค์ •
# ============================================================
USE_CUSTOM_RESNET = False # True: ์ž์ฒด MyResNet ์‚ฌ์šฉ
MODEL_ID = "nateraw/food" # ๋˜๋Š” "your-username/my-resnet18-food101"
# ============================================================
# ๋ชจ๋ธ ๋กœ๋”ฉ
# ============================================================
print(f"๋ชจ๋ธ ๋กœ๋”ฉ ์ค‘: {MODEL_ID}")
if USE_CUSTOM_RESNET:
# ์ž์ฒด ํ•™์Šตํ•œ MyResNet ๋ถˆ๋Ÿฌ์˜ค๊ธฐ
from configuration_myresnet import MyResNetConfig
from modeling_myresnet import MyResNetForImageClassification
from torchvision.transforms import Compose, Resize, ToTensor, Normalize
model = MyResNetForImageClassification.from_pretrained(MODEL_ID)
_transform = Compose([
Resize((224, 224)),
ToTensor(),
Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
def preprocess(image):
return _transform(image.convert("RGB")).unsqueeze(0)
else:
# ๊ณต๊ฐœ ๋ชจ๋ธ ๋ถˆ๋Ÿฌ์˜ค๊ธฐ (AutoImageProcessor๊ฐ€ ์ „์ฒ˜๋ฆฌ ์ž๋™ ์ฒ˜๋ฆฌ)
from transformers import AutoImageProcessor, AutoModelForImageClassification
processor = AutoImageProcessor.from_pretrained(MODEL_ID)
model = AutoModelForImageClassification.from_pretrained(MODEL_ID)
def preprocess(image):
inputs = processor(images=image.convert("RGB"), return_tensors="pt")
return inputs["pixel_values"]
model.eval()
device = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(device)
id2label = model.config.id2label
print(f"๋””๋ฐ”์ด์Šค: {device}")
print(f"ํด๋ž˜์Šค ์ˆ˜: {len(id2label)}")
# ============================================================
# ์˜ˆ์ธก ํ•จ์ˆ˜
# ============================================================
def classify(image):
"""์ด๋ฏธ์ง€๋ฅผ Top-5 ์Œ์‹ ํด๋ž˜์Šค๋กœ ๋ถ„๋ฅ˜ํ•ฉ๋‹ˆ๋‹ค."""
if image is None:
return {}
pixel_values = preprocess(image).to(device)
with torch.no_grad():
logits = model(pixel_values=pixel_values).logits
probs = F.softmax(logits, dim=-1)[0].cpu()
top5_probs, top5_idx = torch.topk(probs, k=5)
return {
id2label[idx.item()].replace("_", " ").title(): float(prob)
for prob, idx in zip(top5_probs, top5_idx)
}
# ============================================================
# Gradio UI
# ============================================================
TITLE = "๐Ÿฝ๏ธ Food Image Classifier"
DESCRIPTION = """
์Œ์‹ ์‚ฌ์ง„์„ ์—…๋กœ๋“œํ•˜๋ฉด **Food-101** ๋ฐ์ดํ„ฐ์…‹์˜ 101๊ฐœ ์Œ์‹ ์ค‘ ๊ฐ€์žฅ ์œ ์‚ฌํ•œ ๊ฒƒ์„ ์ฐพ์•„ **Top-5** ๊ฒฐ๊ณผ๋กœ ๋ณด์—ฌ์ค๋‹ˆ๋‹ค.
**์ง€์› ์Œ์‹ ์˜ˆ์‹œ:** ๐Ÿ• Pizza ยท ๐Ÿฃ Sushi ยท ๐Ÿ” Hamburger ยท ๐Ÿฅฉ Steak ยท ๐Ÿฅž Pancakes ยท ๐Ÿœ Ramen ยท ๐Ÿฆ Ice Cream ยท ๐Ÿฅ˜ Bibimbap ยท ๐ŸŒฎ Tacos ยท ๐ŸฅŸ Gyoza ยท ...
**๋ชจ๋ธ:** ResNet-18 (Pretrained on ImageNet, Fine-tuned on Food-101)
"""
demo = gr.Interface(
fn=classify,
inputs=gr.Image(type="pil", label="์Œ์‹ ์ด๋ฏธ์ง€ ์—…๋กœ๋“œ"),
outputs=gr.Label(num_top_classes=5, label="์˜ˆ์ธก ๊ฒฐ๊ณผ"),
title=TITLE,
description=DESCRIPTION,
flagging_mode="never",
)
if __name__ == "__main__":
demo.launch()