共计 1968 个字符,预计需要花费 5 分钟才能阅读完成。
1. 为什么需要 3D CNN?
在医学影像分析中,CT 扫描数据本质上是三维体数据(长×宽×切片数)。2D CNN 只能逐片处理,会丢失切片间的空间关联。通过 3D 卷积核(如 3×3×3)可同时捕捉:

- 单切片的局部特征(X/ Y 轴)
- 相邻切片的解剖结构连续性(Z 轴)
2. 核心原理对比
2.1 参数量计算
对于输入通道 $C_{in}$ 和输出通道 $C_{out}$:
-
2D 卷积 参数量:
$$K_{2D} = C_{in} \times C_{out} \times k_h \times k_w$$ -
3D 卷积 参数量:
$$K_{3D} = C_{in} \times C_{out} \times k_h \times k_w \times k_d$$
当 $k_h=k_w=k_d=3$ 时,3D 卷积参数量是 2D 的 3 倍
2.2 感受野变化
经过 $L$ 层卷积后:
$$RF_{3D} = 1 + \sum_{l=1}^L (k_l – 1) \times \prod_{i=1}^{l-1} s_i$$
其中 $s_i$ 为第 $i$ 层的 stride 值
3. PyTorch 实战
3.1 数据加载器
class MedicalDataset(Dataset):
def __init__(self, dicom_dir):
"""
dicom_dir: DICOM 文件目录
每个病例包含多个.dcm 文件
"""
self.samples = []
for case_id in os.listdir(dicom_dir):
# 读取 DICOM 序列并排序
slices = [pydicom.dcmread(f) for f in
sorted(glob(f"{dicom_dir}/{case_id}/*.dcm"))]
# 转换为 HU 单位
volume = np.stack([s.pixel_array*s.RescaleSlope + s.RescaleIntercept
for s in slices])
self.samples.append(torch.FloatTensor(volume))
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
# 添加通道维度 (C×D×H×W)
return self.samples[idx].unsqueeze(0)
3.2 网络架构
class CNN3D(nn.Module):
def __init__(self, in_channels=1):
super().__init__()
self.net = nn.Sequential(
# Block 1
nn.Conv3d(in_channels, 32, kernel_size=3, padding=1),
nn.BatchNorm3d(32),
nn.ReLU(),
nn.MaxPool3d(2),
# Block 2
nn.Conv3d(32, 64, 3, padding=1),
nn.BatchNorm3d(64),
nn.ReLU(),
nn.Dropout3d(0.3),
nn.MaxPool3d(2),
# 全局平均池化替代全连接层
nn.AdaptiveAvgPool3d(1),
nn.Flatten(),
nn.Linear(64, 2)
)
def forward(self, x):
return self.net(x)
3.3 显存优化
# 梯度检查点技术(需 PyTorch>=1.8)from torch.utils.checkpoint import checkpoint
class MemoryEfficientBlock(nn.Module):
def forward(self, x):
# 仅在反向传播时重新计算中间结果
return checkpoint(self._forward, x)
def _forward(self, x):
# 实际计算逻辑
return x * 2
4. 性能分析
4.1 卷积核尺寸影响
| kernel_size | 参数量 | 推理时间(ms) |
|---|---|---|
| 3×3×3 | 1× | 15.2 |
| 5×5×5 | 4.63× | 28.7 |
| 7×7×7 | 12.7× | 49.1 |
4.2 Padding 策略对比
- Valid 卷积:输出尺寸减小,丢失边缘信息
- Same 卷积:通过 padding 保持尺寸,但可能引入无效填充
对于医学影像推荐使用:
# 动态计算 padding 值
padding = (kernel_size - 1) // 2
5. 常见问题
5.1 视频时序对齐
- 使用光流法补偿帧间运动
- 在 DataLoader 中实现帧采样策略:
# 等间隔采样 16 帧
frame_indices = np.linspace(0, total_frames-1, 16, dtype=int)
5.2 DICOM 预处理
- 窗宽窗位调整(WW/WL)
- 体素值标准化到[-1,1]
- 处理缺失切片(插值补偿)
6. 延伸思考
- 时空不对齐数据:可尝试 3D ConvLSTM 或 Transformer 结构
- 点云数据局限:3D CNN 需要规则网格,点云更适合 PointNet++ 等网络
正文完
发表至: 未分类
近三天内
