agrobot_base/Python/OAK/datasets/fast_scnn.py

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