Desarrollo personalizado de UNet++: Cómo expandir nuevos codificadores y decodificadores
UNet++ es un modelo avanzado de segmentación de imágenes que mejora la precisión de segmentación mediante conexiones anidadas y densas. Este artículo explica cómo extender UNet++ con nuevos codificadores y decodificadores para adaptarlo a diferentes tareas.
Arquitectura de UNet++: Un vistazo general
La principal ventaja de UNet++ radica en su diseño de red innovador, que utiliza conexiones anidadas para mejorar la propagación y reutilización de características. La figura siguiente muestra la estructura general de UNet++, donde los nodos en amarillo representan capas de convolución y las flechas en虚线 indican las conexiones anidadas.
Imagen: Estructura general de UNet++, mostrando conexiones anidadas y mecanismos de fusión de características multiescala
Papel fundamental de los codificadores y decodificadores
- Codificadores: Extraen características desde la imagen de entrada mediante operaciones de downsampling, reduciendo gradualmente la resolución espacial y aumentando el número de canales de características.
- Decodificadores: Upsamplian las características de alta dimensionalidad hasta el tamaño original de la imagen, fusionando características de diferentes escalas para una segmentación precisa.
Pasos para expandir nuevos codificadores
Implementación básica de codificadores
El codificador típico de UNet++ está compuesto por capas de convolución y capas de pooling. La implementación base se encuentra en el archivo pytorch/nnunet/network_architecture/generic_modular_UNet.py, donde la clase PlainConvUNetEncoder define el codificador estándar.
Guía para desarrollar codificadores personalizados
- Crear una nueva clase de codificador
Hereda de nn.Module y define los siguientes métodos clave:
class CustomEncoder(nn.Module):
def __init__(self, input_channels, base_num_features, num_blocks_per_stage, ...):
super().__init__()
# Inicializa las capas de convolución y pooling
def forward(self, x, return_skips=True):
# Implementa la propagación adelante, retornando mapeos de características o conexiones anidadas
skips = []
for stage in self.stages:
x = stage(x)
skips.append(x)
return skips
- Definir la estructura de las capas convolucionales
Utiliza la clase StackedConvLayers del proyecto para construir bloques convolucionales, personalizando el tamaño del kernel, la función de activación, etc.:
current_stage = StackedConvLayers(
input_features,
output_features,
kernel_size=3,
props=network_props,
num_blocks=2
)
- Integrar en la red neuronal
Modifica la clase PlainConvUNet para reemplazar el codificador predeterminado por el codificador personalizado:
self.encoder = CustomEncoder(...)
Ejemplo de codificador basado en ResNet
El proyecto ya incluye un codificador basado en ResNet, cuya estructura se define en keras/segmentation_models/backbones/classification_models/imgs/graphs/resnet50.png. ResNet aborda el problema de entrenar redes profundas mediante conexiones residuals, lo que lo hace ideal para usarse como codificador en UNet++.
Imagen: Estructura de ResNet50, mostrando bloques residual y el proceso de extracción de características
Pasos para expandir nuevos decodificadores
Implementación básica de decodificadores
El decodificador de UNet++ utiliza upsampling y conexiones anidadas para fusionar características de diferentes escalas. La implementación base se encuentra en la clase PlainConvUNetDecoder del archivo pytorch/nnunet/network_architecture/generic_modular_UNet.py.
Guía para desarrollar decodificadores personalizados
- Crear una nueva clase de decodificador
class CustomDecoder(nn.Module):
def __init__(self, encoder, num_classes, deep_supervision=False):
super().__init__()
self.encoder = encoder
self.num_classes = num_classes
# Inicializa las capas de upsampling y convolución
def forward(self, skips):
# Implementa la lógica de decodificación, upsampling y fusión de características
x = skips[0] # Inicia desde la capa más profunda
for i in range(len(self.tus)):
x = self.tusi # Upsampling
x = torch.cat((x, skips[i+1]), dim=1) # Conexión anidada para fusionar características
x = self.stagesi # Procesamiento convolucional
return self.segmentation_output(x)
- Implementar el upsampling
Puedes elegir entre convoluciones transpuestas o interpolación para el upsampling:
# Upsampling con convolución transpuesta
self.tus.append(nn.ConvTranspose2d(
in_channels,
out_channels,
kernel_size=2,
stride=2
))
# O interpolación
self.upsample = nn.Upsample(scale_factor=2, mode='bilinear')
- Agregar supervisión profunda
Para mejorar la precisión de segmentación, implementa supervisión profunda:
if deep_supervision:
self.deep_supervision_outputs.append(
nn.Conv2d(features_skip, num_classes, 1)
)
Integración y pruebas de codificadores y decodificadores
Armado de la red neuronal
Después de desarrollar los codificadores y decodificadores personalizados, integra ambos en un modelo UNet++ completo:
class CustomUNet(SegmentationNetwork):
def __init__(self, input_channels, num_classes):
super().__init__()
self.encoder = CustomEncoder(input_channels, ...)
self.decoder = CustomDecoder(self.encoder, num_classes, ...)
def forward(self, x):
skips = self.encoder(x)
return self.decoder(skips)
Evaluación del desempeño
Realiza pruebas visuales para evaluar el desempeño del modelo:
Imagen: Comparación de desempeño en tareas de segmentación de pólizas, hígado y núcleos celulares, donde UNet++ muestra una detección de bordes más precisa
Conclusión y mejores prácticas
- Mentener coherencia en las interfaces: Asegúrate de que los nuevos codificadores/decodificadores cumplan con las interfaces existentes para una integración sin problemas.
- Diseño modular: Separa diferentes funcionalidades en módulos independientes para mejorar la reutilizabilidad.
- Pruebas progresivas: Prueba las componentes de manera independiente antes de realizar pruebas end-to-end.
- Referirse a implementaciones existentes: Puedes consultar
pytorch/nnunet/network_architecture/generic_modular_residual_UNet.pypara ver ejemplos de codificadores residual.
A través de los métodos descritos, los desarrolladores pueden adaptar UNet++ de manera flexible para diferentes aplicaciones y requisitos de desempeño. La arquitectura modular de UNet++ permite integrar avances tecnológicos continuamente, mejorando constantemente su capacidad de segmentación.