MahdiHasan/Image_Classifier
0
1import torch
2import torch.nn as nn
3import torch.optim as optim
4
5class ResidualBlock(nn.Module):
6 def __init__(self, in_channels, out_channels, stride=1):
7 super(ResidualBlock, self).__init__()
8 self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)
9 self.bn1 = nn.BatchNorm2d(out_channels)
10 self.relu = nn.ReLU(inplace=True)
11 self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)
12 self.bn2 = nn.BatchNorm2d(out_channels)
13
14 self.shortcut = nn.Sequential()
15 if stride != 1 or in_channels != out_channels:
16 self.shortcut = nn.Sequential(
17 nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),
18 nn.BatchNorm2d(out_channels)
19 )
20
21 def forward(self, x):
22 residual = x
23 out = self.conv1(x)
24 out = self.bn1(out)
25 out = self.relu(out)
26 out = self.conv2(out)
27 out = self.bn2(out)
28 out += self.shortcut(residual)
29 out = self.relu(out)
30 return out
31
32# Define the ResNet model (same as before)
33
34class ResNet(nn.Module):
35 def __init__(self, num_classes=4):
36 super(ResNet, self).__init__()
37 self.in_channels = 64
38 self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, bias=False)
39 self.bn1 = nn.BatchNorm2d(64)
40 self.relu = nn.ReLU(inplace=True)
41 self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
42
43 self.layer1 = self._make_layer(64, 2)
44 self.layer2 = self._make_layer(128, 2, stride=2)
45 self.layer3 = self._make_layer(256, 2, stride=2)
46 self.layer4 = self._make_layer(512, 2, stride=2)
47
48 self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
49 self.fc = nn.Linear(512, num_classes)
50
51 def _make_layer(self, out_channels, num_blocks, stride=1):
52 layers = []
53 layers.append(ResidualBlock(self.in_channels, out_channels, stride))
54 self.in_channels = out_channels
55 for _ in range(1, num_blocks):
56 layers.append(ResidualBlock(out_channels, out_channels))
57 return nn.Sequential(*layers)
58
59 def forward(self, x):
60 out = self.conv1(x)
61 out = self.bn1(out)
62 out = self.relu(out)
63 out = self.maxpool(out)
64
65 out = self.layer1(out)
66 out = self.layer2(out)
67 out = self.layer3(out)
68 out = self.layer4(out)
69
70 out = self.avgpool(out)
71 out = out.view(out.size(0), -1) # Flatten before FC
72 out = self.fc(out)
73 return out