bone_2026/src/model/unet.py

122 lines
4.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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()