| import torch |
| import torch.nn as nn |
| from torchvision import models, transforms |
| from PIL import Image |
| import gradio as gr |
|
|
| |
| model = models.resnet18() |
| model.fc = nn.Linear(model.fc.in_features, 4) |
| checkpoint = torch.hub.load_state_dict_from_url( |
| 'https://huggingface.co/wandikafp/resnet18-tom-and-jerry-classifier/resolve/main/pytorch_model.bin', |
| map_location=torch.device('cpu') |
| ) |
| model.load_state_dict(checkpoint) |
| model.eval() |
|
|
| |
| transform = transforms.Compose([ |
| transforms.Resize((224, 224)), |
| transforms.ToTensor(), |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), |
| ]) |
|
|
| |
| def classify_image(image): |
| image = Image.fromarray(image) |
| image = transform(image).unsqueeze(0) |
|
|
| with torch.no_grad(): |
| outputs = model(image) |
| _, predicted = torch.max(outputs, 1) |
|
|
| labels = ['tom', 'jerry', 'tom_jerry_0', 'tom_jerry_1'] |
| return labels[predicted.item()] |
|
|
| |
| interface = gr.Interface( |
| fn=classify_image, |
| inputs="image", |
| outputs="label", |
| title="Tom and Jerry Classifier", |
| description="Classify images as 'tom', 'jerry', 'tom_jerry_0', or 'tom_jerry_1'." |
| ) |
|
|
| |
| interface.launch() |
|
|