import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): """Два сверточных слоя с GELU активацией""" def __init__(self, in_channels, out_channels, dropout=0.0): super().__init__() self.double_conv = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), nn.BatchNorm2d(out_channels), nn.GELU(), nn.Dropout2d(dropout), # добавляем dropout для регуляризации nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), nn.BatchNorm2d(out_channels), nn.GELU(), ) def forward(self, x): return self.double_conv(x) class DownBlock(nn.Module): """Блок энкодера: DoubleConv + MaxPool""" def __init__(self, in_channels, out_channels, dropout=0.0): super().__init__() self.double_conv = DoubleConv(in_channels, out_channels, dropout) self.maxpool = nn.MaxPool2d(2) def forward(self, x): # Сохраняем результат DoubleConv для skip connection conv_out = self.double_conv(x) pooled_out = self.maxpool(conv_out) return conv_out, pooled_out class UpBlock(nn.Module): """Блок декодера: Upsample + Concatenate + DoubleConv""" def __init__(self, in_channels, out_channels, dropout=0.0): super().__init__() # Транспонированная свертка для апсемплинга self.up_conv = nn.ConvTranspose2d( in_channels, out_channels, kernel_size=2, stride=2 ) # После конкатенации будет (out_channels * 2) self.double_conv = DoubleConv(out_channels * 2, out_channels, dropout) def forward(self, x, skip_connection): x = self.up_conv(x) # Обрезаем skip_connection до размера x (если размеры не совпадают) diffY = skip_connection.size()[2] - x.size()[2] diffX = skip_connection.size()[3] - x.size()[3] x = F.pad(x, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2]) # Конкатенируем x = torch.cat([skip_connection, x], dim=1) return self.double_conv(x) class UNet(nn.Module): """Полная архитектура U-Net""" def __init__(self, in_channels=1, out_classes=2, features=[64, 128, 256, 512], dropout=0.1): super().__init__() self.encoder = nn.ModuleList() self.decoder = nn.ModuleList() # Encoder path prev_channels = in_channels for f in features: self.encoder.append(DownBlock(prev_channels, f, dropout)) prev_channels = f # Bottleneck self.bottleneck = DoubleConv(prev_channels, features[-1] * 2, dropout) # Decoder path (идем в обратном порядке) for f in reversed(features): self.decoder.append(UpBlock(f * 2, f, dropout)) # Final classification layer self.final_conv = nn.Conv2d(features[0], out_classes, kernel_size=1) def forward(self, x): skip_connections = [] # Encoder for down_block in self.encoder: skip, x = down_block(x) skip_connections.append(skip) # Bottleneck x = self.bottleneck(x) # Decoder for up_block in self.decoder: # Берем последний skip connection skip = skip_connections.pop() x = up_block(x, skip) # Final layer x = self.final_conv(x) return x def test_unet(): """Тестируем модель на случайном тензоре""" model = UNet(in_channels=1, out_classes=2) x = torch.randn(1, 1, 512, 512) # batch_size, channels, height, width with torch.no_grad(): y = model(x) print(f"Input shape: {x.shape}") print(f"Output shape: {y.shape}") print(f"Number of parameters: {sum(p.numel() for p in model.parameters()):,}") if __name__ == "__main__": test_unet()