122 lines
4.1 KiB
Python
122 lines
4.1 KiB
Python
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() |