quazarcom commited on
Commit
0982826
·
verified ·
1 Parent(s): 77dc3bc

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +53 -0
app.py ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torchvision.transforms as transforms
3
+ from PIL import Image
4
+ import gradio as gr
5
+ from timm import create_model
6
+ import torch.nn as nn
7
+ import os
8
+
9
+ # ==== Your custom model ====
10
+ class VisionTransformer(nn.Module):
11
+ def __init__(self, num_classes, model_name):
12
+ super(VisionTransformer, self).__init__()
13
+ self.model = create_model(model_name, pretrained=False, num_classes=num_classes)
14
+ self.model.head = nn.Sequential(
15
+ nn.Linear(self.model.num_features, 512),
16
+ nn.ReLU(),
17
+ nn.Dropout(0.5),
18
+ nn.Linear(512, num_classes)
19
+ )
20
+
21
+ def forward(self, x):
22
+ return self.model(x)
23
+
24
+ # ==== Load the trained weights ====
25
+ model_path = "./models/your_model.pth"
26
+ device = torch.device("cpu") # MPS not available on HF Spaces
27
+
28
+ model = VisionTransformer(num_classes=2, model_name="vit_small_patch16_224")
29
+ model.load_state_dict(torch.load(model_path, map_location=device))
30
+ model.eval()
31
+
32
+ # ==== Image transforms ====
33
+ transform = transforms.Compose([
34
+ transforms.Resize((224, 224)),
35
+ transforms.ToTensor(),
36
+ transforms.Normalize(mean=[0.5]*3, std=[0.5]*3)
37
+ ])
38
+
39
+ def predict(img: Image.Image):
40
+ img_tensor = transform(img).unsqueeze(0)
41
+ with torch.no_grad():
42
+ outputs = model(img_tensor)
43
+ _, pred = torch.max(outputs, 1)
44
+ label = "Fake" if pred.item() == 0 else "Real"
45
+ return label
46
+
47
+ gr.Interface(
48
+ fn=predict,
49
+ inputs=gr.Image(type="pil"),
50
+ outputs="label",
51
+ title="Luxury Item Checker",
52
+ description="Upload an image to check if it's a Real or Fake item."
53
+ ).launch()