birdscanner-api / train.py
thesoikindustries24's picture
Add custom trained hybrid model for Hen and Peacock (using LFS)
824fa9b
Raw
History Blame Contribute Delete
10.1 kB
import argparse
import os
import sys
import urllib.request
from pathlib import Path
from PIL import Image
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from transformers import AutoImageProcessor, AutoModelForImageClassification
# List of public domain images for the 'OTHER' class (using Unsplash to avoid rate limits)
OTHER_BIRD_URLS = [
f"https://images.unsplash.com/{photo_id}?w=500&auto=format&fit=crop&q=80"
for photo_id in [
"photo-1452570053594-1b985d6ea890",
"photo-1480044965905-02098d419e96",
"photo-1516233758813-a38d024919c5",
"photo-1551085254-e96b210db58a",
"photo-1522441815192-d9f04eb0615c",
"photo-1506220926022-cc5c12abdb35",
"photo-1518998053901-5348d3961a04",
"photo-1511823794984-b87716139b88",
"photo-1470116890351-be0a9b418409",
"photo-1539664030485-a936c7d29fc0",
"photo-1444464666168-49d633b86797",
"photo-1504386106331-3e4e71712b38",
"photo-1555041469-a586c61ea9bc",
"photo-1525462519782-b55cef2f882a",
"photo-1509023467868-1c40786379f6",
"photo-1454496522488-7a8e488e8606",
"photo-1549488344-1f9b8d2bd1f3",
"photo-1528183429752-a97d0bf99b5a",
"photo-1510137600163-2729bc695ac1",
"photo-1465153690352-10c1b295ec7e",
"photo-1497250681960-ef046c08a56e",
"photo-1526336024438-db9b601fc211",
"photo-1534067783941-51c9c23eccfd",
"photo-1560015534-eca11e59876a",
"photo-1551806235-a05ff789c3c7",
]
]
def download_image(url, filepath):
"""Downloads an image from a URL with a standard browser User-Agent header."""
req = urllib.request.Request(
url,
headers={'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36'}
)
try:
with urllib.request.urlopen(req, timeout=15) as response:
with open(filepath, 'wb') as f:
f.write(response.read())
# Verify it is a valid image
with Image.open(filepath) as img:
img.verify()
return True
except Exception as e:
if os.path.exists(filepath):
os.remove(filepath)
print(f" [Warning] Failed to download or verify {url}: {e}")
return False
def setup_other_class(other_dir):
"""Sets up the OTHER class directory and downloads sample bird images."""
os.makedirs(other_dir, exist_ok=True)
existing_files = [f for f in os.listdir(other_dir) if f.lower().endswith(('.png', '.jpg', '.jpeg', '.webp'))]
if len(existing_files) >= 10:
print(f"-> 'OTHER' directory already has {len(existing_files)} images. Skipping download.")
return
print("-> Downloading background bird images for 'OTHER' class (to avoid forgetting)...")
downloaded_count = 0
for i, url in enumerate(OTHER_BIRD_URLS):
dest_path = os.path.join(other_dir, f"other_bird_{i}.jpg")
print(f" Downloading image {i+1}/{len(OTHER_BIRD_URLS)}...")
if download_image(url, dest_path):
downloaded_count += 1
print(f"-> Completed downloading. Added {downloaded_count} images to 'OTHER' directory.")
class BirdDataset(Dataset):
def __init__(self, image_paths, labels, image_processor):
self.image_paths = image_paths
self.labels = labels
self.image_processor = image_processor
def __len__(self):
return len(self.image_paths)
def __getitem__(self, idx):
img_path = self.image_paths[idx]
label = self.labels[idx]
try:
image = Image.open(img_path).convert("RGB")
inputs = self.image_processor(images=image, return_tensors="pt")
pixel_values = inputs["pixel_values"].squeeze(0)
return pixel_values, torch.tensor(label, dtype=torch.long)
except Exception as e:
# Fallback for corrupt images during training
# We return a dummy image (zeros) and the label
print(f" [Warning] Error loading image {img_path}: {e}. Using zero-filled tensor.")
dummy_pixel = torch.zeros((3, 224, 224))
return dummy_pixel, torch.tensor(label, dtype=torch.long)
def train_model(args):
dataset_path = Path(args.dataset_dir)
hen_dir = dataset_path / "HEN"
peacock_dir = dataset_path / "PEACOCK"
other_dir = dataset_path / "OTHER"
# 1. Verify dataset structure
if not hen_dir.exists() or not peacock_dir.exists():
print("Error: Dataset directory must contain 'HEN' and 'PEACOCK' subdirectories.")
print(f"Looked in: {args.dataset_dir}")
sys.exit(1)
# Count source images
hen_imgs = list(hen_dir.glob("*.[jJ][pP][gG]")) + list(hen_dir.glob("*.[jJ][pP][eE][gG]")) + list(hen_dir.glob("*.[pP][nN][gG]"))
peacock_imgs = list(peacock_dir.glob("*.[jJ][pP][gG]")) + list(peacock_dir.glob("*.[jJ][pP][eE][gG]")) + list(peacock_dir.glob("*.[pP][nN][gG]"))
print(f"Found {len(hen_imgs)} images in HEN/")
print(f"Found {len(peacock_imgs)} images in PEACOCK/")
if len(hen_imgs) == 0 or len(peacock_imgs) == 0:
print("Error: Both HEN and PEACOCK folders must contain at least a few images to train.")
sys.exit(1)
# 2. Setup OTHER class automatically
setup_other_class(str(other_dir))
other_imgs = list(other_dir.glob("*.[jJ][pP][gG]")) + list(other_dir.glob("*.[jJ][pP][eE][gG]")) + list(other_dir.glob("*.[pP][nN][gG]"))
# 3. Gather paths and labels
class_names = ["HEN", "OTHER", "PEACOCK"] # Alphabetical order
class_to_idx = {name: i for i, name in enumerate(class_names)}
image_paths = []
labels = []
for name in class_names:
dir_path = dataset_path / name
imgs = list(dir_path.glob("*.[jJ][pP][gG]")) + list(dir_path.glob("*.[jJ][pP][eE][gG]")) + list(dir_path.glob("*.[pP][nN][gG]")) + list(dir_path.glob("*.[wW][eE][bB][pP]"))
for img in imgs:
# Simple pre-check to make sure it loads
try:
with Image.open(img) as temp_img:
temp_img.draft("RGB", (32, 32))
image_paths.append(str(img))
labels.append(class_to_idx[name])
except Exception:
print(f" [Warning] Skipping corrupt file: {img}")
print(f"Total valid training samples: {len(image_paths)}")
for name in class_names:
idx = class_to_idx[name]
count = labels.count(idx)
print(f" Class '{name}': {count} images")
# 4. Load online pretrained model and processor
print("-> Loading pre-trained base model and image processor...")
image_processor = AutoImageProcessor.from_pretrained("chriamue/bird-species-classifier")
model = AutoModelForImageClassification.from_pretrained("chriamue/bird-species-classifier")
# 5. Freeze base layers to make training super fast and prevent catastrophic overfitting
print("-> Freezing feature extractor base weights (training classification head only)...")
for param in model.parameters():
param.requires_grad = False
# 6. Replace classification head
# The original was: (classifier): Linear(in_features=1408, out_features=525, bias=True)
num_features = model.classifier.in_features
model.classifier = nn.Linear(num_features, len(class_names))
# Configure model config metadata so Hugging Face saves labels correctly
model.config.id2label = {i: name for i, name in enumerate(class_names)}
model.config.label2id = {name: i for i, name in enumerate(class_names)}
model.config.num_labels = len(class_names)
# 7. Create DataLoader
dataset = BirdDataset(image_paths, labels, image_processor)
dataloader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True)
# 8. Setup optimizer and loss
optimizer = torch.optim.AdamW(model.classifier.parameters(), lr=args.lr)
criterion = nn.CrossEntropyLoss()
# 9. Training Loop
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)
model.train()
print(f"-> Starting training on device: {device}...")
for epoch in range(args.epochs):
running_loss = 0.0
correct = 0
total = 0
for batch_x, batch_y in dataloader:
batch_x, batch_y = batch_x.to(device), batch_y.to(device)
optimizer.zero_grad()
outputs = model(pixel_values=batch_x)
logits = outputs.logits
loss = criterion(logits, batch_y)
loss.backward()
optimizer.step()
running_loss += loss.item() * batch_x.size(0)
_, predicted = torch.max(logits, 1)
total += batch_y.size(0)
correct += (predicted == batch_y).sum().item()
epoch_loss = running_loss / total
epoch_acc = correct / total
print(f" Epoch {epoch+1}/{args.epochs} - Loss: {epoch_loss:.4f} - Accuracy: {epoch_acc:.4f}")
# 10. Save fine-tuned model
print(f"-> Saving fine-tuned model to {args.output_dir}...")
os.makedirs(args.output_dir, exist_ok=True)
model.save_pretrained(args.output_dir)
image_processor.save_pretrained(args.output_dir)
print("-> Done! Custom model training completed successfully.")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Fine-tune bird species classifier on custom dataset.")
parser.add_argument("--dataset_dir", type=str, default="dataset", help="Directory containing HEN and PEACOCK folders")
parser.add_argument("--output_dir", type=str, default="custom_model", help="Directory to save custom model weights")
parser.add_argument("--epochs", type=int, default=8, help="Number of training epochs")
parser.add_argument("--batch_size", type=int, default=8, help="DataLoader batch size")
parser.add_argument("--lr", type=float, default=1e-3, help="Learning rate for the classification head")
args = parser.parse_args()
train_model(args)