Una implementación de codificación en MONAI para la segmentación del bazo en 3D de extremo a extremo utilizando UNet en volúmenes de TC médicos

En este tutorial, creamos un canal de segmentación de imágenes médicas 3D de extremo a extremo utilizando MONAI para segmentar el bazo en el conjunto de datos Medical Segmentation Decathlon Task09. Trabajamos con tomografías computarizadas volumétricas, aplicamos transformaciones de imágenes médicas, como alineación de orientación, normalización de espaciado de vóxeles, ventanas de intensidad, recorte de primer plano y muestreo basado en parches, y luego entrenamos un modelo UNet 3D para la segmentación binaria de órganos. También utilizamos entrenamiento de precisión mixta, pérdida de DiceCE, inferencia de ventana deslizante, validación basada en Dice y visualización cualitativa para comprender cómo aprende el modelo y cómo se comparan sus predicciones con las máscaras de verdad sobre el terreno. Además, pasamos de volúmenes médicos sin procesar a un sistema completo de segmentación de tren, validación y visualización.

Copiar código

!pip instalar -q “monai[nibabel,tqdm,matplotlib]==1.5.2” 2>/dev/null importar sistema operativo, tiempo, glob, archivo temporal, advertencias importar numpy como np importar antorcha importar matplotlib.pyplot como plt desde torch.amp importar autocast, GradScaler desde monai.apps importar DecathlonDataset desde monai.data importar DataLoader, decollate_batch desde monai.networks.nets importar UNet desde monai.networks.layers importar Norm desde monai.losses importar DiceCELoss desde monai.metrics importar DiceMetric desde monai.inferers importar slide_window_inference desde monai.utils importar set_determinism desde monai.transforms importar ( Compose, LoadImaged, GuaranteeChannelFirstd, GuaranteeTyped, Orientationd, Spacingd, ScaleIntensityRanged, CropForegroundd, RandCropByPosNegLabeld, RandFlipd, RandRotate90d, RandShiftIntensityd, AsDiscrete,) advertencias.filterwarnings(“ignorar”)

Comenzamos instalando MONAI con las dependencias de visualización y imágenes médicas necesarias. Luego importamos PyTorch, NumPy, Matplotlib y los módulos principales de MONAI necesarios para conjuntos de datos, transformaciones, entrenamiento de modelos, métricas e inferencia. También suprimimos las advertencias para mantener limpia la salida del cuaderno mientras nos centramos en el flujo de trabajo de segmentación.

Copiar código

QUICK_RUN = True dispositivo = torch.device(“cuda” if torch.cuda.is_available() else “cpu”) root_dir = tempfile.mkdtemp() roi_size = (96, 96, 96) num_samples = 4 batch_size = 2 max_epochs = 15 if QUICK_RUN else 200 val_every = 3 train_cache = 8 if QUICK_RUN else 24 val_cache = 2 if QUICK_RUN else 6 set_determinism(seed=0) print(f”Dispositivo: {dispositivo} | épocas: {max_epochs} | directorio de datos: {root_dir}”) train_transforms = Compose(common + [
image_key=”image”, image_threshold=0),
RandFlipd(keys=[“image”, “label”]prob=0.2, eje_espacial=0), RandFlipd(claves=[“image”, “label”]prob=0.2, eje_espacial=1), RandFlipd(claves=[“image”, “label”]prob=0.2, eje_espacial=2), RandRotate90d(claves=[“image”, “label”]prob=0.2, max_k=3), RandShiftIntensityd(claves=[“image”]compensaciones = 0,10, prob = 0,5), GuaranteeTyped (claves =[“image”, “label”]), ]) val_transforms = Componer(común + [EnsureTyped(keys=[“image”, “label”])])

Definimos la configuración principal para el tutorial, incluido el dispositivo, el directorio del conjunto de datos, el tamaño del parche, el tamaño del lote, la cantidad de épocas y la configuración de la caché. Luego creamos el proceso de preprocesamiento para volúmenes de CT cargando imágenes, alineando la orientación, volviendo a muestrear el espaciado de vóxeles, escalando intensidades y recortando el primer plano. También definimos las transformaciones de entrenamiento y validación, y el proceso de entrenamiento incluye cultivos aleatorios, volteos, rotaciones y cambios de intensidad.

Copiar código

train_ds = DecathlonDataset( root_dir=root_dir, task=”Task09_Spleen”, sección=”entrenamiento”, transform=train_transforms, download=True, val_frac=0.2, cache_num=train_cache, num_workers=2, seed=0) val_ds = DecathlonDataset( root_dir=root_dir, task=”Task09_Spleen”, sección=”validación”, transform=val_transforms, download=False, val_frac=0.2, cache_num=val_cache, num_workers=2, seed=0) train_loader = DataLoader(train_ds, lote_size=batch_size, shuffle=True, num_workers=2, pin_memory=torch.cuda.is_available()) val_loader = DataLoader(val_ds, tamaño_lote=1, shuffle=False, num_workers=1, pin_memory=torch.cuda.is_available()) print(f”Volumenes del tren: {len(train_ds)} | Volúmenes Val: {len(val_ds)}”) loss_fn = DiceCELoss(to_onehot_y=True, softmax=True) optimizador = torch.optim.AdamW(model.parameters(), lr=1e-4, Weight_decay=1e-5) planificador = torch.optim.lr_scheduler.CosineAnnealingLR(optimizador, T_max=max_epochs) escalador = GradScaler(“cuda”, enable=torch.cuda.is_available()) dice_metric = DiceMetric(include_background=False, reducción=”media”) post_pred = Redactar([AsDiscrete(argmax=True, to_onehot=2)]) post_label = Redactar([AsDiscrete(to_onehot=2)])

Cargamos el conjunto de datos de bazo Decathlon Task09 de segmentación médica utilizando DecathlonDataset de MONAI. Dividimos los datos en secciones de entrenamiento y validación, aplicamos las transformaciones apropiadas y empaquetamos ambos conjuntos de datos con cargadores de datos estilo PyTorch. Luego creamos un modelo 3D UNet, definimos la pérdida de DiceCE, configuramos el optimizador AdamW, el programador de tasa de aprendizaje, el escalador de precisión mixta, la métrica Dice y los pasos de posprocesamiento.

Copiar código

best_dice, best_epoch = -1.0, -1 loss_hist, dice_hist, dice_epochs = [], [], []
best_path = os.path.join(root_dir, “best_spleen_unet.pth”) para la época en el rango (1, max_epochs + 1): model.train(); epoch_loss, t0 = 0.0, time.time() para lote en train_loader: x, y = lote[“image”].to(dispositivo), lote[“label”].to(dispositivo) optimizador.zero_grad(set_to_none=True) con autocast(“cuda”, enable=torch.cuda.is_available()): logits = model(x) pérdida = loss_fn(logits, y) scaler.scale(loss).backward() scaler.step(optimizador); scaler.update() epoch_loss += pérdida.item() planificador.step() epoch_loss /= len(train_loader); loss_hist.append(epoch_loss) print(f”[{epoch:3d}/{max_epochs}] pérdida={epoch_loss:.4f} ” f”lr={scheduler.get_last_lr()[0]:.2e} ({time.time()-t0:.0f}s)”) si época % val_every == 0 o época == max_epochs: model.eval(); dice_metric.reset() con torch.no_grad(): para vb en val_loader: vx, vy = vb[“image”].a(dispositivo), vb[“label”].to(dispositivo) con autocast(“cuda”, habilitado=torch.cuda.is_available()): vout = slide_window_inference(vx, roi_size, 4, modelo, superposición=0.5) vout = [post_pred(o) for o in decollate_batch(vout)]
vlab = [post_label(o) for o in decollate_batch(vy)]
dice_metric(y_pred=vout, y=vlab) d = dice_metric.aggregate().item() dice_hist.append(d); dice_epochs.append(época) if d > best_dice: best_dice, best_epoch = d, epoch torch.save(model.state_dict(), best_path) print(f” >> val Dice={d:.4f} (best={best_dice:.4f} @ {best_epoch})”) print(f”\nHecho. Mejor dado medio {best_dice:.4f} en la época {mejor_época}.”)

Ejecutamos el ciclo de entrenamiento completo, donde cada época entrena 3D UNet en parches volumétricos recortados del conjunto de datos del bazo. Usamos precisión mixta automática para reducir el uso de memoria y acelerar el entrenamiento cuando hay una GPU disponible. También validamos el modelo a intervalos regulares mediante inferencia de ventana deslizante, realizamos un seguimiento de la puntuación de Dice y guardamos el punto de control de mejor rendimiento.

Copiar código

higo, hacha = plt.subplots(1, 2, tamaño de higo=(12, 4)) hacha[0].plot(rango(1, len(loss_hist)+1), loss_hist, “-o”, ms=3) hacha[0].set(title=”Pérdida de entrenamiento”, xlabel=”época”, ylabel=”Pérdida de DiceCE”) hacha[1].plot(dice_epochs, dice_hist, “-o”, color=”verdemar”, ms=4) hacha[1].set(title=”Validación significa Dados”, xlabel=”época”, ylabel=”Dados”); hacha[1].set_ylim(0, 1) plt.tight_layout(); plt.show() model.load_state_dict(torch.load(best_path, map_location=dispositivo)); model.eval() con torch.no_grad(): muestra = siguiente(iter(val_loader)) img = muestra[“image”].to(dispositivo) con autocast(“cuda”, habilitado=torch.cuda.is_available()): pred = slide_window_inference(img, roi_size, 4, modelo, superposición=0.5) pred = torch.argmax(pred, dim=1).cpu().numpy()[0]
img_np, lab_np = img.cpu().numpy()[0, 0]muestra[“label”].numpy()[0, 0]
z = int(np.argmax(lab_np.sum(axis=(0, 1)))) fig, ax = plt.subplots(1, 3, figsize=(13, 5)) ax[0].imshow(img_np[:, :, z]cmap=”gris”); hacha[0].set_title(“corte CT”) hacha[1].imshow(lab_np[:, :, z]cmap=”viridis”); hacha[1].set_title(“Verdad sobre el terreno”) hacha[2].imshow(pred[:, :, z]cmap=”viridis”); hacha[2].set_title(“Predicción”) para a en ax: a.axis(“off”) plt.tight_layout(); plt.mostrar()

Primero trazamos la pérdida de entrenamiento y la puntuación de dados de validación para ver cómo mejora el modelo con el tiempo. Luego recargamos el punto de control del modelo mejor guardado y ejecutamos la inferencia en un único volumen de validación mediante predicción de ventana deslizante. Visualizamos el corte de TC, la máscara de verdad sobre el terreno y la segmentación prevista uno al lado del otro para inspeccionar el rendimiento cualitativo del modelo.

En conclusión, finalizamos un flujo de trabajo práctico basado en MONAI para la segmentación del bazo en 3D utilizando un modelo 3D UNet. Preparamos el conjunto de datos de Medical Segmentation Decathlon, transformamos y aumentamos los volúmenes de CT, entrenamos el modelo con pérdida de DiceCE, lo validamos mediante inferencia de ventana deslizante y realizamos un seguimiento tanto de la pérdida como de la puntuación de Dice a lo largo del tiempo. También inspeccionamos visualmente la predicción final comparando el corte de CT, la etiqueta de verdad fundamental y el resultado del modelo uno al lado del otro. Ahora, tenemos una comprensión clara de cómo MONAI respalda las tareas de segmentación médica, desde la carga y el preprocesamiento de datos hasta el entrenamiento, la evaluación, los puntos de control y el análisis cualitativo de modelos.

Consulte los códigos completos con Notebook. Además, no dude en seguirnos en Twitter y no olvide unirse a nuestro SubReddit de más de 150.000 ML y suscribirse a nuestro boletín. ¡Esperar! estas en telegrama? Ahora también puedes unirte a nosotros en Telegram.

¿Necesita asociarse con nosotros para promocionar su repositorio de GitHub O su página principal de Hugging O su lanzamiento de producto O seminario web, etc.? Conéctate con nosotros

La publicación Una implementación de codificación en MONAI para la segmentación del bazo 3D de extremo a extremo utilizando UNet en volúmenes de TC médicos apareció por primera vez en MarkTechPost.