CoolFace
Modelpublic

Bektur756/vit-cats-dogs-pandas

sourceHugging Faceapache-2.0updated 8d agoView on Hugging Face
0likes25downloads
Model Card

ViT Cats, Dogs and Pandas Classifier

Vision Transformer image classifier for three classes:

  • —0: cats
  • —1: dogs
  • —2: panda

Base model

google/vit-base-patch16-224

The classifier uses the output representation of the first [CLS] token and a Linear(768, 3) classification head.

Dataset

Melisa13/Animals_dataset

SplitImagesPer class
Train408136
Validation7224
Test12040

Exact image duplicates between splits: 0.

Image preprocessing

  • —RGB input
  • —Resize and crop to 224×224
  • —Patch size: 16×16
  • —196 image patches
  • —197 tokens including [CLS]
  • —Normalization mean: (0.5, 0.5, 0.5)
  • —Normalization standard deviation: (0.5, 0.5, 0.5)

Fine-tuning

The following components were trainable:

  • —Transformer blocks 10 and 11
  • —Final LayerNorm
  • —Linear(768, 3) classifier
ConfigurationValue
Total parameters85,800,963
Trainable parameters14,179,587
Trainable share16.53%
Epochs completed7
Batch size16
Learning rate2e-5
Weight decay0.01

Results

SplitAccuracyMacro F1
Validation98.61%98.61%
Test100.00%100.00%

Test classification report:

ClassPrecisionRecallF1Support
Cats1.00001.00001.000040
Dogs1.00001.00001.000040
Panda1.00001.00001.000040

End-to-end latency

Latency includes image preprocessing and model inference with batch size 1.

DeviceMeanMedianP95Throughput
Tesla T421.86 ms18.37 ms39.92 ms45.75 images/s
Colab CPU599.39 ms371.03 ms1550.67 ms1.67 images/s

Usage

python
from PIL import Image
import torch

from transformers import (
    AutoImageProcessor,
    AutoModelForImageClassification,
)

MODEL_NAME = "Bektur756/vit-cats-dogs-pandas"

processor = AutoImageProcessor.from_pretrained(
    MODEL_NAME
)

model = AutoModelForImageClassification.from_pretrained(
    MODEL_NAME
)

image = Image.open(
    "animal.jpg"
).convert(
    "RGB"
)

inputs = processor(
    images=image,
    return_tensors="pt",
)

with torch.inference_mode():
    logits = model(
        **inputs
    ).logits

probabilities = torch.softmax(
    logits,
    dim=-1,
)[0]

predicted_id = int(
    probabilities.argmax()
)

print(
    model.config.id2label[
        predicted_id
    ]
)

print(
    probabilities.tolist()
)

## Limitations

- The test split contains only 120 images.
- The dataset is small and relatively simple.
- Perfect test accuracy does not imply perfect performance on arbitrary images.
- CPU inference is substantially slower than GPU inference.