Desarrollo personalizado de UNet++: Cómo expandir nuevos codificadores y decodificadores

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

  1. 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

  1. 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
)

  1. 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

  1. 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)

  1. 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')

  1. 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

  1. Mentener coherencia en las interfaces: Asegúrate de que los nuevos codificadores/decodificadores cumplan con las interfaces existentes para una integración sin problemas.
  2. Diseño modular: Separa diferentes funcionalidades en módulos independientes para mejorar la reutilizabilidad.
  3. Pruebas progresivas: Prueba las componentes de manera independiente antes de realizar pruebas end-to-end.
  4. Referirse a implementaciones existentes: Puedes consultar pytorch/nnunet/network_architecture/generic_modular_residual_UNet.py para 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.

Etiquetas: deep_learning neural_networks image_segmentation PyTorch computer_vision

Publicado el 9-14 16:57