TrashNet ResNet18

Fine-tuned ResNet18 (ImageNet-pretrained backbone, custom classification head) for 6-class waste image classification on the TrashNet dataset.

Used by the Trash Classification System — a FastAPI + Streamlit app that serves this model for real-time inference.

Classes

cardboard, glass, metal, paper, plastic, trash

Training

  • Backbone: torchvision.models.resnet18 (IMAGENET1K_V1 weights), custom Linear(64) -> ReLU -> Dropout(0.5) -> Linear(num_classes) head
  • Optimizer: AdamW + CosineAnnealingLR
  • Class-weighted cross-entropy loss (TrashNet's trash class is underrepresented ~3.6x vs paper)
  • Hyperparameters (lr, weight_decay, batch_size) selected via Optuna hyperparameter search (25 trials), then trained for 30 epochs with the winning config
  • Trained on the University of Manchester CSF3 HPC cluster (SLURM, NVIDIA L40S GPU)

Usage

from huggingface_hub import hf_hub_download
import torch
from torchvision import models
import torch.nn as nn

def get_model(hidden_size=64, num_classes=6):
    model = models.resnet18(weights=None)
    model.fc = nn.Sequential(
        nn.Linear(model.fc.in_features, hidden_size),
        nn.ReLU(),
        nn.Dropout(p=0.5),
        nn.Linear(hidden_size, num_classes),
    )
    return model

checkpoint_path = hf_hub_download(
    repo_id="tonghahaha/trashnet-resnet18",
    filename="best_resnet18_trashnet.pth",
)
model = get_model(num_classes=6)
model.load_state_dict(torch.load(checkpoint_path, map_location="cpu", weights_only=True))
model.eval()
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support