| 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_best.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) |
| ]) |
|
|
| def predict(img: Image.Image): |
| img_tensor = transform(img).unsqueeze(0) |
| with torch.no_grad(): |
| outputs = model(img_tensor) |
| _, pred = torch.max(outputs, 1) |
| label = "Fake" if pred.item() == 0 else "Real" |
| return label |
|
|
| gr.Interface( |
| fn=predict, |
| inputs=gr.Image(type="pil"), |
| outputs="label", |
| title="Luxury Item Checker", |
| description="Upload an image to check if it's a Real or Fake item." |
| ).launch() |
|
|