changed api structure and added object detection

This commit is contained in:
2026-05-20 23:35:31 +02:00
parent 4eb0e3217e
commit a519e77db1
19 changed files with 1012 additions and 292 deletions
+152
View File
@@ -0,0 +1,152 @@
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
class SeparableConvolution(nn.Module):
def __init__(self, c_in, c_out, kernel_size):
super().__init__()
self.depthwise = nn.Conv2d(c_in, c_in, kernel_size, groups=c_in, padding=kernel_size//2)
self.bn1 = nn.BatchNorm2d(c_in)
self.pointwise = nn.Conv2d(c_in, c_out, kernel_size=1)
self.bn2 = nn.BatchNorm2d(c_out)
def forward(self, x):
x = self.depthwise(x)
x = self.bn1(x)
x = F.relu(x)
x = self.pointwise(x)
x = self.bn2(x)
x = F.relu(x)
return x
class SkipBlock(nn.Module):
def __init__(self, c_in, c_out, kernel_size=3):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(c_in, c_out, kernel_size, padding=kernel_size//2),
nn.BatchNorm2d(c_out),
nn.ReLU(inplace=True),
nn.Conv2d(c_out, c_out, kernel_size, padding=kernel_size//2),
nn.BatchNorm2d(c_out),
nn.ReLU(inplace=True),
nn.Conv2d(c_out, c_out, kernel_size, padding=kernel_size//2),
nn.BatchNorm2d(c_out),
nn.ReLU(inplace=True)
)
self.conv_skip = nn.Sequential(
nn.Conv2d(c_in, c_out, 1),
nn.BatchNorm2d(c_out),
nn.ReLU(inplace=True)
)
def forward(self, x):
return(F.relu(self.conv_skip(x) + self.conv(x), inplace=True))
class Bird_CNN(nn.Module):
def __init__(self, c_in, c_hidden, c_out):
super().__init__()
self.model = nn.Sequential(
nn.Conv2d(c_in, c_hidden, kernel_size=3, padding=1),
nn.BatchNorm2d(c_hidden),
nn.ReLU(inplace=True),
SkipBlock(c_in=c_hidden, c_out=c_hidden),
SkipBlock(c_in=c_hidden, c_out=c_hidden),
SkipBlock(c_in=c_hidden, c_out=c_hidden),
SkipBlock(c_in=c_hidden, c_out=c_hidden),
SkipBlock(c_in=c_hidden, c_out=c_hidden*2),
SkipBlock(c_in=c_hidden*2, c_out=c_hidden*2),
SkipBlock(c_in=c_hidden*2, c_out=c_hidden*2),
SkipBlock(c_in=c_hidden*2, c_out=c_hidden*2),
nn.Conv2d(c_hidden*2, c_hidden*4, kernel_size=3, padding=1),
nn.ReLU(inplace=True),
nn.AdaptiveAvgPool2d((1, 1)),
nn.Flatten(),
nn.Linear(c_hidden*4, c_out),
nn.Dropout(0.3)
)
def forward(self, x):
return self.model(x)
def trainCNN(model, optimizer, loss_module, train_data_loader, validation_data_loader, device, num_epochs, SAVE_PATH, save=False):
best_val = 0
for epoch in range(num_epochs):
############
# Training #
############
model.train()
true_preds, count = 0, 0
for data_inputs, classes in tqdm(train_data_loader, desc=f"Train Epoch {epoch+1}", leave=False):
data_inputs = data_inputs.to(device)
classes = classes.to(device)
preds = model(data_inputs)
loss = loss_module(preds, classes)
optimizer.zero_grad()
loss.backward()
optimizer.step()
true_preds += (preds.argmax(dim=1) == classes).sum().item()
count += data_inputs.size(0)
train_acc = true_preds / count
torch.cuda.empty_cache()
##############
# Validation #
##############
model.eval()
true_preds, count = 0, 0
for data_inputs, classes in tqdm(validation_data_loader, desc=f"Validate Epoch {epoch+1}", leave=False):
with torch.no_grad():
data_inputs = data_inputs.to(device)
classes = classes.to(device)
preds = model(data_inputs)
loss = loss_module(preds, classes)
true_preds += (preds.argmax(dim=1) == classes).sum().item()
count += data_inputs.size(0)
val_acc = true_preds / count
if(save and best_val < val_acc):
best_val = val_acc
save_dir = os.path.join(SAVE_PATH, "bird_cnn")
os.makedirs(save_dir, exist_ok=True)
save_path = os.path.join(save_dir, f"bird_cnn")
torch.save(model.state_dict(), save_path)
print(f"epoch: {epoch+1} | train accuracy: {int(train_acc * 1000) / 10}% | validation accuracy: {int(val_acc * 1000) / 10}%")
torch.cuda.empty_cache()
def sample(model, img, device, SAVE_PATH, model_name="bird_cnn", folder="bird_cnn"):
with torch.no_grad():
full_path = os.path.join(SAVE_PATH, folder, model_name)
state_dict = torch.load(full_path, weights_only=False)
model.load_state_dict(state_dict)
model.eval()
img = img.to(device)
pred = model(img)
probs = F.softmax(pred, dim=1)
return(torch.max(probs, dim=1))
+74
View File
@@ -0,0 +1,74 @@
import torch
import torch.nn as nn
import torchvision
from torchvision import datasets, transforms
from torch.utils.data import DataLoader, random_split
import matplotlib.pyplot as plt
from util import TransformedSubset, visualizeData
from .bird_cnn import Bird_CNN, sample, trainCNN
from enum import Enum
from PIL import Image
SAVE_PATH = "../saved_models"
IMAGE_SIZE = 128
class bird_species(Enum):
Common_Kingfisher = 0
Common_Myna = 1
House_Crow = 2
Indian_Peacock = 3
Indian_Pitta = 4
Ruddy_Shelduck = 5
Sarus_Crane = 6
transform_augemnt = transforms.Compose([
transforms.RandomAffine(
degrees=35, # no rotation
translate=(0.2, 0.2) # shift up to 20% horizontally/vertically
),
transforms.RandomHorizontalFlip(p=0.5),
transforms.Resize(IMAGE_SIZE),
transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor()
])
transform = transforms.Compose([
transforms.Resize(IMAGE_SIZE),
transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor()
])
dataset = datasets.ImageFolder("data/CUB_200_2011/images")
train_size = int(0.8 * len(dataset))
val_size = len(dataset) - train_size
train_subset, val_subset = random_split(dataset, [train_size, val_size])
train_dataset = TransformedSubset(train_subset, transform_augemnt)
val_dataset = TransformedSubset(val_subset, transform)
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)
device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0")
print("Using device", device)
model = Bird_CNN(c_in=3, c_hidden=16, c_out=200)
model.to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
loss_module = nn.CrossEntropyLoss()
trainCNN(model, optimizer, loss_module, train_loader, val_loader, device, 500, SAVE_PATH=SAVE_PATH, save=True)
exit()
image_path = "test1.jpg"
image = Image.open(image_path).convert("RGB")
image = transform(image)
image = image.unsqueeze(0)
confidence, pred = sample(model=model, img=image, device=device, SAVE_PATH=SAVE_PATH)
print(f"Species: {bird_species(pred.item()).name} | Confidence: {int(confidence.item()*100)/100}")