医学图像分割入门:手把手教你用Python和PyTorch跑通第一个U-Net项目

作为一名临床医生或生物医学背景的研究者,你可能经常需要从CT、MRI等医学图像中提取特定组织结构。传统手动标注耗时费力,而深度学习技术正在彻底改变这一现状。本文将带你从零开始,用不到100行代码实现一个基础的细胞图像分割模型,让你在30分钟内获得第一个AI分割结果。

1. 环境配置与数据准备

1.1 极简开发环境搭建

对于初学者,推荐使用Anaconda创建独立的Python环境,避免依赖冲突。以下是快速上手指南:

conda create -n medseg python=3.8
conda activate medseg
pip install torch==1.12.0 torchvision==0.13.0
pip install opencv-python matplotlib tqdm

提示:PyTorch 1.x版本对新手更友好,且兼容大多数经典模型代码。如果使用GPU加速,需额外安装CUDA版本的PyTorch。

1.2 获取ISBI细胞分割数据集

我们将使用经典的ISBI挑战赛数据集,包含30张训练图像和30张测试图像,每张图像大小均为512x512像素。通过以下代码自动下载并解压:

import urllib.request
import tarfile

url = "http://www.codes.bio/wp-content/uploads/2019/10/ISBI_dataset.tar.gz"
urllib.request.urlretrieve(url, "ISBI_dataset.tar.gz")
with tarfile.open("ISBI_dataset.tar.gz") as tar:
    tar.extractall()

数据集目录结构如下:

ISBI_dataset/
├── train
│   ├── image
│   └── label
└── test
    ├── image
    └── label

2. 构建最小化U-Net模型

2.1 模型核心组件

U-Net的核心在于编码器-解码器结构和跳跃连接。下面是最简实现:

import torch
import torch.nn as nn

class MiniUNet(nn.Module):
    def __init__(self):
        super().__init__()
        # 编码器
        self.enc1 = self.conv_block(1, 64)
        self.enc2 = self.conv_block(64, 128)
        
        # 解码器 
        self.up1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)
        self.dec1 = self.conv_block(128, 64)
        
        # 输出层
        self.final = nn.Conv2d(64, 1, kernel_size=1)

    def conv_block(self, in_c, out_c):
        return nn.Sequential(
            nn.Conv2d(in_c, out_c, 3, padding=1),
            nn.ReLU(),
            nn.Conv2d(out_c, out_c, 3, padding=1),
            nn.ReLU()
        )

    def forward(self, x):
        # 编码
        e1 = self.enc1(x)
        p1 = nn.MaxPool2d(2)(e1)
        e2 = self.enc2(p1)
        p2 = nn.MaxPool2d(2)(e2)
        
        # 解码
        u1 = self.up1(p2)
        c1 = torch.cat([u1, e2], dim=1)
        d1 = self.dec1(c1)
        
        return torch.sigmoid(self.final(d1))

2.2 数据加载器实现

创建自定义Dataset类处理图像和掩码:

from torch.utils.data import Dataset
import cv2
import numpy as np

class CellDataset(Dataset):
    def __init__(self, img_dir, transform=None):
        self.img_dir = img_dir
        self.image_paths = sorted(glob(f"{img_dir}/image/*.png"))
        
    def __len__(self):
        return len(self.image_paths)
    
    def __getitem__(self, idx):
        img_path = self.image_paths[idx]
        mask_path = img_path.replace("image", "label")
        
        image = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)
        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)
        
        # 归一化处理
        image = image.astype(np.float32) / 255.0
        mask = (mask > 127).astype(np.float32)
        
        return torch.tensor(image).unsqueeze(0), torch.tensor(mask).unsqueeze(0)

3. 训练流程与可视化

3.1 训练循环实现

以下是精简版的训练脚本:

from tqdm import tqdm

def train(model, train_loader, optimizer, criterion, device):
    model.train()
    total_loss = 0
    
    for images, masks in tqdm(train_loader):
        images, masks = images.to(device), masks.to(device)
        
        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, masks)
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item()
    
    return total_loss / len(train_loader)

# 初始化
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = MiniUNet().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.BCELoss()

# 数据加载
train_dataset = CellDataset("ISBI_dataset/train")
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=4, shuffle=True)

# 训练
for epoch in range(10):
    loss = train(model, train_loader, optimizer, criterion, device)
    print(f"Epoch {epoch+1}, Loss: {loss:.4f}")

3.2 实时可视化结果

添加结果可视化函数,在训练过程中观察分割效果:

import matplotlib.pyplot as plt

def visualize(model, loader, device, num_samples=3):
    model.eval()
    with torch.no_grad():
        for i, (images, masks) in enumerate(loader):
            if i >= num_samples: break
            images, masks = images.to(device), masks.to(device)
            preds = model(images)
            
            plt.figure(figsize=(12,4))
            plt.subplot(1,3,1)
            plt.imshow(images[0].cpu().squeeze(), cmap='gray')
            plt.title("Input")
            
            plt.subplot(1,3,2)
            plt.imshow(masks[0].cpu().squeeze(), cmap='gray')
            plt.title("Ground Truth")
            
            plt.subplot(1,3,3)
            plt.imshow(preds[0].cpu().squeeze() > 0.5, cmap='gray')
            plt.title("Prediction")
            plt.show()

# 在每个epoch后调用
visualize(model, train_loader, device)

4. 进阶技巧与优化建议

4.1 数据增强策略

在CellDataset类中添加随机增强:

from torchvision import transforms

class CellDataset(Dataset):
    def __init__(self, img_dir):
        # ...
        self.transform = transforms.Compose([
            transforms.RandomHorizontalFlip(p=0.5),
            transforms.RandomVerticalFlip(p=0.5),
            transforms.RandomRotation(15)
        ])
    
    def __getitem__(self, idx):
        # ...
        if self.transform:
            # 将numpy数组转为PIL图像进行变换
            image = Image.fromarray((image*255).astype(np.uint8))
            mask = Image.fromarray((mask*255).astype(np.uint8))
            
            seed = np.random.randint(2147483647)
            torch.manual_seed(seed)
            image = self.transform(image)
            torch.manual_seed(seed)
            mask = self.transform(mask)
            
            image = np.array(image) / 255.0
            mask = (np.array(mask) > 127).astype(np.float32)
        # ...

4.2 评估指标实现

添加Dice系数计算函数:

def dice_coeff(pred, target):
    smooth = 1.
    pred_flat = pred.view(-1)
    target_flat = target.view(-1)
    intersection = (pred_flat * target_flat).sum()
    return (2. * intersection + smooth) / (pred_flat.sum() + target_flat.sum() + smooth)

4.3 学习率调度

使用PyTorch内置的学习率调度器:

from torch.optim.lr_scheduler import ReduceLROnPlateau

scheduler = ReduceLROnPlateau(optimizer, 'min', patience=2, factor=0.5)

# 在训练循环后调用
scheduler.step(loss)

在Jupyter Notebook中运行这些代码块时,我发现一个常见问题是内存泄漏。解决方法是在每个epoch结束后手动清除缓存:

torch.cuda.empty_cache()
Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐