38 lines
1.2 KiB
Python
38 lines
1.2 KiB
Python
import torch.nn as nn
|
|
|
|
class FastSCNN(nn.Module):
|
|
def __init__(self, num_classes):
|
|
super().__init__()
|
|
self.down1 = nn.Sequential(
|
|
nn.Conv2d(3, 32, 3, stride=2, padding=1),
|
|
nn.BatchNorm2d(32),
|
|
nn.ReLU(inplace=True)
|
|
)
|
|
self.down2 = nn.Sequential(
|
|
nn.Conv2d(32, 48, 3, stride=2, padding=1),
|
|
nn.BatchNorm2d(48),
|
|
nn.ReLU(inplace=True)
|
|
)
|
|
self.down3 = nn.Sequential(
|
|
nn.Conv2d(48, 64, 3, stride=2, padding=1),
|
|
nn.BatchNorm2d(64),
|
|
nn.ReLU(inplace=True)
|
|
)
|
|
self.classifier = nn.Sequential(
|
|
nn.Conv2d(64, num_classes, 1)
|
|
)
|
|
# Upsampling substituído por ConvTranspose2d
|
|
self.up1 = nn.ConvTranspose2d(num_classes, num_classes, kernel_size=2, stride=2)
|
|
self.up2 = nn.ConvTranspose2d(num_classes, num_classes, kernel_size=2, stride=2)
|
|
self.up3 = nn.ConvTranspose2d(num_classes, num_classes, kernel_size=2, stride=2)
|
|
|
|
def forward(self, x):
|
|
x = self.down1(x)
|
|
x = self.down2(x)
|
|
x = self.down3(x)
|
|
x = self.classifier(x)
|
|
x = self.up1(x)
|
|
x = self.up2(x)
|
|
x = self.up3(x)
|
|
return x
|