changed cnn architecture
This commit is contained in:
|
Before Width: | Height: | Size: 218 KiB After Width: | Height: | Size: 218 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 430 KiB |
Binary file not shown.
+26
-24
@@ -32,42 +32,44 @@ class Bird_CNN(nn.Module):
|
|||||||
self.conv_init = nn.Sequential(
|
self.conv_init = nn.Sequential(
|
||||||
nn.Conv2d(c_in, c_hidden, kernel_size=3, padding=1),
|
nn.Conv2d(c_in, c_hidden, kernel_size=3, padding=1),
|
||||||
nn.BatchNorm2d(c_hidden),
|
nn.BatchNorm2d(c_hidden),
|
||||||
nn.ReLU(),
|
|
||||||
nn.Conv2d(c_hidden, c_hidden, 3, stride=2, padding=1)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 1x1 conv branch
|
|
||||||
self.branch1 = SeparableConvolution(c_in=c_hidden, c_out=64, kernel_size=1)
|
|
||||||
|
|
||||||
# 1x1 -> 3x3 conv branch
|
|
||||||
self.branch2 = SeparableConvolution(c_in=c_hidden, c_out=128, kernel_size=3)
|
|
||||||
|
|
||||||
# 1x1 -> 5x5 conv branch
|
|
||||||
self.branch3 = SeparableConvolution(c_in=c_hidden, c_out=32, kernel_size=5)
|
|
||||||
|
|
||||||
# 3x3 max pooling -> 1x1 conv branch
|
|
||||||
self.branch4 = nn.Sequential(
|
|
||||||
nn.MaxPool2d(kernel_size=3, stride=1, padding=1),
|
|
||||||
nn.Conv2d(c_hidden, 32, kernel_size=1),
|
|
||||||
nn.ReLU()
|
nn.ReLU()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.conv_1 = SeparableConvolution(c_in=c_hidden, c_out=c_hidden*2, kernel_size=3)
|
||||||
|
self.conv_skip_1 = nn.Sequential(
|
||||||
|
nn.Conv2d(c_hidden, c_hidden*2, 1),
|
||||||
|
nn.BatchNorm2d(c_hidden*2)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.conv_2 = SeparableConvolution(c_in=c_hidden*2, c_out=c_hidden*4, kernel_size=3)
|
||||||
|
self.conv_skip_2 = nn.Sequential(
|
||||||
|
nn.Conv2d(c_hidden*2, c_hidden*4, 1),
|
||||||
|
nn.BatchNorm2d(c_hidden*4)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.conv_3 = SeparableConvolution(c_in=c_hidden*4, c_out=c_hidden*8, kernel_size=3)
|
||||||
|
self.conv_skip_3 = nn.Sequential(
|
||||||
|
nn.Conv2d(c_hidden*4, c_hidden*8, 1),
|
||||||
|
nn.BatchNorm2d(c_hidden*8)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.conv_4 = SeparableConvolution(c_in=c_hidden*8, c_out=c_hidden*16, kernel_size=3)
|
||||||
|
|
||||||
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
||||||
self.flatten = nn.Flatten()
|
self.flatten = nn.Flatten()
|
||||||
self.linear = nn.Linear(256, c_out)
|
self.linear = nn.Linear(c_hidden*16, c_out)
|
||||||
|
|
||||||
self.dropout = nn.Dropout(0.3)
|
self.dropout = nn.Dropout(0.3)
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
x = self.conv_init(x)
|
x = self.conv_init(x)
|
||||||
|
|
||||||
b1 = self.branch1(x)
|
x = F.relu(self.conv_skip_1(x) + self.conv_1(x))
|
||||||
b2 = self.branch2(x)
|
x = F.relu(self.conv_skip_2(x) + self.conv_2(x))
|
||||||
b3 = self.branch3(x)
|
x = F.relu(self.conv_skip_3(x) + self.conv_3(x))
|
||||||
b4 = self.branch4(x)
|
|
||||||
|
x = self.conv_4(x)
|
||||||
|
|
||||||
x = torch.cat([b1, b2, b3, b4], dim=1)
|
|
||||||
x = F.relu(x)
|
|
||||||
x = self.avgpool(x)
|
x = self.avgpool(x)
|
||||||
x = torch.flatten(x, 1)
|
x = torch.flatten(x, 1)
|
||||||
|
|
||||||
|
|||||||
+5
-5
@@ -10,8 +10,7 @@ from enum import Enum
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
SAVE_PATH = "./saved_models"
|
SAVE_PATH = "./saved_models"
|
||||||
#IMAGE_SIZE = (1141, 850)
|
IMAGE_SIZE = 256
|
||||||
IMAGE_SIZE = (300, 300)
|
|
||||||
|
|
||||||
class bird_species(Enum):
|
class bird_species(Enum):
|
||||||
Common_Kingfisher = 0
|
Common_Kingfisher = 0
|
||||||
@@ -24,6 +23,7 @@ class bird_species(Enum):
|
|||||||
|
|
||||||
transform = transforms.Compose([
|
transform = transforms.Compose([
|
||||||
transforms.Resize(IMAGE_SIZE),
|
transforms.Resize(IMAGE_SIZE),
|
||||||
|
transforms.CenterCrop(IMAGE_SIZE),
|
||||||
transforms.ToTensor()
|
transforms.ToTensor()
|
||||||
])
|
])
|
||||||
|
|
||||||
@@ -34,13 +34,13 @@ val_size = len(dataset) - train_size
|
|||||||
|
|
||||||
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
|
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
|
||||||
|
|
||||||
train_loader = DataLoader(train_dataset, batch_size=26, shuffle=True)
|
train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True)
|
||||||
val_loader = DataLoader(val_dataset, batch_size=26, shuffle=False)
|
val_loader = DataLoader(val_dataset, batch_size=8, shuffle=False)
|
||||||
|
|
||||||
device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0")
|
device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0")
|
||||||
print("Using device", device)
|
print("Using device", device)
|
||||||
|
|
||||||
model = Bird_CNN(c_in=3, c_hidden=15, c_out=7)
|
model = Bird_CNN(c_in=3, c_hidden=32, c_out=7)
|
||||||
model.to(device)
|
model.to(device)
|
||||||
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
|
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
|
||||||
loss_module = nn.CrossEntropyLoss()
|
loss_module = nn.CrossEntropyLoss()
|
||||||
|
|||||||
Binary file not shown.
+2
-7
@@ -12,13 +12,7 @@ import torch.nn.functional as F
|
|||||||
from bird_cnn import Bird_CNN
|
from bird_cnn import Bird_CNN
|
||||||
|
|
||||||
BUILD_PATH = "./build_models"
|
BUILD_PATH = "./build_models"
|
||||||
#IMAGE_SIZE = (1141, 850)
|
IMAGE_SIZE = 256
|
||||||
IMAGE_SIZE = (300, 300)
|
|
||||||
|
|
||||||
transform = transforms.Compose([
|
|
||||||
transforms.Resize(IMAGE_SIZE),
|
|
||||||
transforms.ToTensor()
|
|
||||||
])
|
|
||||||
|
|
||||||
class bird_species(Enum):
|
class bird_species(Enum):
|
||||||
Common_Kingfisher = 0
|
Common_Kingfisher = 0
|
||||||
@@ -31,6 +25,7 @@ class bird_species(Enum):
|
|||||||
|
|
||||||
transform = transforms.Compose([
|
transform = transforms.Compose([
|
||||||
transforms.Resize(IMAGE_SIZE),
|
transforms.Resize(IMAGE_SIZE),
|
||||||
|
transforms.CenterCrop(IMAGE_SIZE),
|
||||||
transforms.ToTensor()
|
transforms.ToTensor()
|
||||||
])
|
])
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user