pokemon_classifier / README.md
skshmjn's picture
Update README.md
35c2b13 verified
|
raw
history blame
907 Bytes

Pokémon Classifier

Intro

A fine-tuned version of ViT-base on a collected set of Pokémon images. You can read more about the model here.

Using the model

from transformers import ViTForImageClassification, ViTImageProcessor
from PIL import Image
import torch

# Loading in Model
device = "cuda" if torch.cuda.is_available() else "cpu"
model = ViTForImageClassification.from_pretrained( "imjeffhi/pokemon_classifier").to(device)
image_processor = ViTImageProcessor.from_pretrained('imjeffhi/pokemon_classifier')

# Caling the model on a test image
img = Image.open('test.jpg')
extracted = image_processor(images=img, return_tensors='pt').to(device)
predicted_id = model(**extracted).logits.argmax(-1).item()
predicted_pokemon = model.config.id2label[predicted_id]