| --- |
| license: mit |
| tags: |
| - image-classification |
| - pytorch |
| - resnet18 |
| - waste-classification |
| datasets: |
| - trashnet |
| --- |
| |
| # TrashNet ResNet18 |
|
|
| Fine-tuned **ResNet18** (ImageNet-pretrained backbone, custom classification head) for 6-class waste image classification on the [TrashNet](https://github.com/garythung/trashnet) dataset. |
|
|
| Used by the [Trash Classification System](https://github.com/yutongyu-ai/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 |
|
|
| ```python |
| 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() |
| ``` |
|
|