医学图像分割入门:手把手教你用Python和PyTorch跑通第一个U-Net项目(附完整代码)
·
医学图像分割入门:手把手教你用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()
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)