"""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()