Spaces:
Sleeping
Sleeping
| import torch | |
| import torchvision.transforms as transforms | |
| from PIL import Image | |
| import gradio as gr | |
| from timm import create_model | |
| import torch.nn as nn | |
| import os | |
| class VisionTransformer(nn.Module): | |
| def __init__(self, num_classes, model_name): | |
| super(VisionTransformer, self).__init__() | |
| self.model = create_model(model_name, pretrained=False, num_classes=num_classes) | |
| self.model.head = nn.Sequential( | |
| nn.Linear(self.model.num_features, 512), | |
| nn.ReLU(), | |
| nn.Dropout(0.5), | |
| nn.Linear(512, num_classes) | |
| ) | |
| def forward(self, x): | |
| return self.model(x) | |
| model_path = "./models/vit_small_patch16_224_final.pth" | |
| device = torch.device("cpu") | |
| model = VisionTransformer(num_classes=2, model_name="vit_small_patch16_224") | |
| model.load_state_dict(torch.load(model_path, map_location=device)) | |
| model.eval() | |
| transform = transforms.Compose([ | |
| transforms.Resize((224, 224)), | |
| transforms.ToTensor(), | |
| transforms.Normalize(mean=[0.5]*3, std=[0.5]*3) | |
| ]) | |
| # ✅ Define classify_image first | |
| def classify_image(img: Image.Image): | |
| img_tensor = transform(img).unsqueeze(0) | |
| with torch.no_grad(): | |
| outputs = model(img_tensor) | |
| _, predicted = torch.max(outputs, 1) | |
| label = 'Fake' if predicted.item() == 0 else 'Real' | |
| return label | |
| # ✅ Then wrap it in predict() | |
| def predict(img: Image.Image): | |
| return classify_image(img) | |
| # ✅ All set to go | |
| gr.Interface( | |
| fn=predict, | |
| inputs=gr.Image(type="pil"), | |
| outputs="label", | |
| title="Luxury Item Authenticity Detector", | |
| description="Upload an image to check if it's a real or fake item." | |
| ).launch() | |