尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

HexMIL:层次注意力MIL实现CT影像篡改检测与可解释定位

HexMIL:层次注意力MIL实现CT影像篡改检测与可解释定位 医学影像AI这几年快速进入临床辅助诊断流程但与此同时一个更隐蔽的风险也在被放大既然AI能读片自然也就能被用来“改片”。如果一张CT卷里的肺结节被算法抹掉或者被伪造出原本不存在的病灶再由医生基于这张片子做诊断后果将非常严重。HexMIL这个名字看起来像一篇论文标题但它本质上是在回答一个问题如何让模型既能检测出AI篡改的CT影像又能解释它凭什么这么判断。这篇文章我会把HexMIL涉及的核心概念拆开讲清楚包括Hierarchical Attention层次化注意力、MILMultiple Instance Learning多实例学习、前摄可解释性Ante-Hoc Explainable并给出一套基于PyTorch的简化实现思路。适合对医学影像AI、弱监督学习、模型可解释性感兴趣的开发者阅读。读完你会理解层次注意力MIL的建模逻辑也知道如何把CT卷组织成“包-实例”结构来训练检测模型。1. 背景与核心概念1.1 为什么需要检测AI篡改的CT影像先假设一个场景患者做了一次胸部CT图像通过网络传到诊断系统。如果数据链路中有人截获了图像用生成模型在某个肺叶区域叠加了一个仿真结节或者把原本的结节清理掉最终诊断报告就可能完全不同。这类篡改不是简单的像素擦除或PS涂改而是AI辅助的语义级修改。生成模型的表达能力足够强可以保留组织纹理、噪声分布让人眼很难察觉。传统图像取证方法对这类篡改往往失效因为伪造痕迹不再表现为网格畸变或颜色异常而是“纹理过于真实”。这就带来一个需求检测系统需要在图像语义层面识别出不自然的局部区域同时给出证据而不是只输出一个“有问题/没问题”的二分类结果。临床上一个没有解释的二分类结果很难被信任医生需要知道模型根据哪些切片、哪些区域做出了判断。1.2 MIL多实例学习与Simulink MIL的区别搜索引擎里搜“MIL”很容易搜到Simulink中的Model-in-the-Loop模型在环测试。这跟我们讲的MIL完全是两个概念先做区分Simulink MIL指在Simulink环境中把控制器模型放到环路里做仿真测试验证模型逻辑是否满足需求属于汽车电子、控制系统的开发流程。Multiple Instance Learning多实例学习一种弱监督学习范式把数据组织成“包”和“实例”。一个包包含多个实例只有包级标签不提供实例级标签。在HexMIL中MIL指多实例学习。CT卷天然适合用MIL建模整卷是一个包每个2D切片或3D小块是一个实例。我们只知道“这卷CT被篡改了”却不知道具体哪一帧被改这就构成了典型的弱监督学习问题。1.3 Ante-Hoc可解释与Post-Hoc可解释的区别模型可解释性通常分两种路线Post-Hoc事后解释训练一个黑盒模型再用SHAP、LIME、Grad-CAM等工具解释预测结果。这类方法的优点是模型结构不受限制缺点是解释可能不忠实于模型真实决策逻辑尤其在医学影像这种高维非结构化数据上热图可能指向纹理区域而不是真实病灶。Ante-Hoc前摄可解释在模型设计阶段就把解释机制内嵌进去让模型本身具备产出解释的能力。模型不只是输出一个分数还会明确输出“哪些切片、哪些区域对判定贡献更大”。HexMIL强调的是Ante-Hoc可解释。也就是说模型最后的输出不仅包含“篡改概率”还包含一个可以映射回原始切片的注意力热图。解释不是事后补上的而是模型结构的一部分。1.4 HexMIL整体思路HexMIL整体流程可以拆成四步将CT卷按轴向通常是横断位切成2D切片每张切片是一个实例。用卷积神经网络CNN提取每张切片的特征表示。用两层注意力机制对切片特征进行聚合第一层在实例间分配注意力权重第二层在不同特征子空间或类别分支间分配权重形成层次化聚合。基于聚合后的包级特征做分类同时热图可以反投影到原始切片完成前摄可解释输出。接下来我们把这套思路逐步展开。2. 环境准备与版本说明2.1 软硬件环境本文示例代码以常见深度学习环境为例操作系统Ubuntu 20.04 / Windows 10 / Windows 11Python3.8 以上深度学习框架PyTorch视觉库torchvision数值计算NumPyCUDA根据GPU驱动自行适配如果没有GPU代码也能跑但训练会慢很多具体版本建议结合你本机环境调整。PyTorch在1.10到2.x之间对本示例影响不大核心API是保持稳定的。如果你的CUDA版本较新直接安装对应cuda版本即可。安装命令参考pip install torch torchvision numpy matplotlib scikit-learn如果你的环境是用Anaconda管理也可以先创建虚拟环境conda create -n hexmil python3.9 conda activate hexmil pip install torch torchvision numpy matplotlib scikit-learn2.2 项目结构建议按下面的目录组织代码方便后面扩展hexmil_example/ ├── data/ │ └── preprocess.py # CT卷预处理切片、归一化 ├── models/ │ ├── encoder.py # 切片特征提取器 │ └── hexmil.py # 层次注意力MIL模型 ├── utils/ │ └── visualize.py # 热图可视化 ├── train.py # 训练入口 └── config.py # 配置参数这个结构不是必须的但保持清晰划分后后续换数据集、换特征提取器时不需要大规模改动代码。2.3 关于医学影像处理库实际项目中CT影像通常是DICOM格式建议安装pydicom、SimpleITK来读取pip install pydicom SimpleITKMonai也是医学影像项目中常用的库封装了很多预处理算子。本文为了减少依赖先直接用NumPy数组演示核心逻辑。真实项目中你可以用Monai来实现重采样、窗宽窗位调整、数据增强等操作。3. 核心原理拆解3.1 为什么用“包-实例”结构而不是直接分类最直接的方法是把CT卷所有切片拼接成一个3D输入用3D CNN做分类。但这里有一个现实问题篡改区域往往很小可能只出现在连续几片切片上甚至只出现在某一个局部区域。3D CNN虽然能建模空间结构但对“局部异常”的定位能力较弱而且3D卷积的开销较大对数据量和显存要求高。MIL的思路更加简洁CT卷是一个包里面的切片是实例。因为只有包级标签模型需要在训练过程中自动学会对实例的重要性加权正常的切片拿到低权重包含篡改痕迹的切片拿到高权重。这样训练结束后我们自然可以用权重来判断哪些切片更可疑。3.2 实例级注意力假设一个CT卷被切成N张切片每张切片经过特征提取器后得到特征向量 h_ii1..N。实例级注意力要做的事是计算每个实例的重要性权重a_i exp(w^T tanh(V h_i^T)) / sum_j exp(w^T tanh(V h_j^T))这里的 V 是一个线性变换矩阵w 是一个可学习的查询向量。整个过程类似于注意力打分再经过softmax归一化。然后包表示就是加权求和z sum_i a_i * h_i这个公式最早在“Attention-based Deep Multiple Instance Learning”这篇经典论文中提出。优点是简单、可微、能够端到端训练。缺点也很明显如果所有实例的权重都趋于均匀分布模型就退化成简单的平均池化失去定位能力。这个问题在训练时需要注意后面会讲。3.3 类别级注意力层次化的第二层HexMIL进一步引入第二层注意力即类别级注意力。为什么需要第二层CT影像中包含多种组织结构和伪影一张切片的异常可能是细微的纹理变化也可能是结构性的形态改变。单一注意力池化只能得到一组全局权重但很难同时捕捉“多尺度的异常信号”。类别级注意力可以理解为模型先学习若干个注意力头或特征分组每个分组关注一种异常模式再通过注意力机制聚合这些分组的结果。简化理解流程如下实例级注意力生成切片权重得到若干不同的包表示每个包表示可以理解为一个“视角”关注不同类型的异常信号类别级注意力在这几个包表示之间做加权融合形成最终包级表示分类头基于最终包级表示输出篡改概率。这种层次化设计让模型在训练时不需要知道篡改区域的具体位置却能在推理时通过注意力图指示可疑区域。整个过程是内嵌的不需要额外训练解释模型。3.4 前摄可解释输出的形态模型的可解释输出通常包含两个部分全局解释包被分类为“篡改”的一维概率。局部解释切片级注意力权重和特征图的组合热图。将高注意力切片上的特征激活值反投影回原始像素空间就能得到与原始CT切片分辨率对齐的热图。推理阶段医生看到的不再只是一个孤立的“阳性”告警而是一张CT卷上哪里可疑、置信度多少、涉及哪些切片。这就是前摄可解释性在临床场景中的价值。3.5 损失函数与标签设计在MIL框架下训练标签是图像级别甚至患者级别的二分类标签0表示正常1表示篡改。对于单包分类最常用的是交叉熵损失或BCEWithLogitsLoss。如果某个数据集只有患者级标签可能需要先按患者合并多个CT序列再在患者级别组织包。这会增加包内实例一致性的假设难度是工程上容易被忽略的环节。4. 完整实战案例下面我们用PyTorch实现一个简化版HexMIL。需要提前说明这不是论文的完整复现而是用于理解核心思路的代码骨架。真实实验还需要根据数据规模和任务类型做很多细节调整。4.1 模拟CT数据准备实战中我们先用模拟数据验证流程。假设一个CT卷被表示成NumPy数组形状是(D, H, W)D是切片数H和W是高度和宽度。为了模拟篡改可以在连续几片切片的某个区域加入局部纹理扰动。下面的代码生成模拟数据# 文件路径data/preprocess.py import numpy as np def make_synthetic_ct(shape(32, 128, 128), seed0): 生成一个模拟CT卷。 这里使用随机噪声模拟组织纹理真实项目中应替换为DICOM序列读取。 rng np.random.default_rng(seed) ct rng.normal(0.1, 0.02, sizeshape).astype(np.float32) return ct def add_local_manipulation(ct, start_slice, size(8, 32, 32), strength0.2, seed1): 在连续切片上加入局部篡改扰动。 真实场景中的篡改更隐蔽这里用局部亮度偏移模拟异常区域。 rng np.random.default_rng(seed) manipulated ct.copy() d, h, w size manipulated[start_slice:start_slice d, h // 4:h // 4 h, w // 4:w // 4 w] strength * rng.normal(sizesize) return manipulated这个步骤的目的是生成“有包级标签”的数据集正常卷和篡改卷。训练时模型只能看到包级标签看不到篡改区域的切片编号。4.2 读取与切片训练前要把整个CT卷切成2D切片# 文件路径data/preprocess.py def ct_to_slices(ct): 将CT卷转换为切片列表。 返回形状为 (D, H, W) 的数组后续逐张提取特征。 return ct实际中还需要做归一化、缩放等操作。为了减少代码复杂度本文直接用原始值作为输入但这只是演示。真实项目中建议根据CT值范围做窗宽窗位调整再转为0-1范围内的浮点数组。4.3 特征提取器Encoder用torchvision中预训练的ResNet18作为特征提取器去掉最后的全连接层输出512维特征向量。预训练权重在ImageNet上训练虽然与CT影像分布有差异但在迁移学习场景下仍能提供基础视觉特征。# 文件路径models/encoder.py import torch import torch.nn as nn from torchvision import models import torchvision.transforms as transforms class SliceEncoder(nn.Module): def __init__(self, out_dim512, pretrainedTrue): super().__init__() resnet models.resnet18(pretrainedpretrained) # 去掉最后的全局池化和全连接层 self.features nn.Sequential(*list(resnet.children())[:-2]) self.gap nn.AdaptiveAvgPool2d((1, 1)) self.project nn.Linear(512, out_dim) def forward(self, x): # x: (B, 3, H, W) x self.features(x) x self.gap(x).flatten(1) x self.project(x) return x原始CT切片是单通道而ResNet输入是三通道。一个常用做法是在输入维度上重复三次等效于三通道。也可以把模型第一层卷积的输入通道改为1但这样无法加载预训练权重得不偿失。4.4 层次注意力MIL模型下面是最核心的模型。先实现实例级注意力池化再扩展为层次化结构。为了兼顾简洁和可读性这里把层次化实现为“多头注意力池化类别级注意力融合”# 文件路径models/hexmil.py import torch import torch.nn as nn import torch.nn.functional as F class InstanceAttentionPooling(nn.Module): 实例级注意力池化。 输入 (B, N, D) 的实例特征输出包特征和注意力权重。 def __init__(self, in_dim, hidden_dim128): super().__init__() self.attention nn.Sequential( nn.Linear(in_dim, hidden_dim), nn.Tanh(), nn.Linear(hidden_dim, 1) ) def forward(self, x): # x: (B, N, D) scores self.attention(x) # (B, N, 1) weights F.softmax(scores, dim1) bag_feat torch.sum(x * weights, dim1) # (B, D) return bag_feat, weights.squeeze(-1) class HierarchicalAttentionMIL(nn.Module): 简化版HexMIL 1. 多组实例级注意力产生多个包表示 2. 类别级注意力对多个包表示加权融合 3. 分类头输出二分类概率。 def __init__(self, in_dim, n_heads4, hidden_dim128, num_classes2): super().__init__() self.n_heads n_heads self.heads nn.ModuleList( [InstanceAttentionPooling(in_dim, hidden_dim) for _ in range(n_heads)] ) # 类别级注意力 self.class_attention nn.Sequential( nn.Linear(in_dim, hidden_dim), nn.Tanh(), nn.Linear(hidden_dim, 1) ) self.classifier nn.Sequential( nn.Linear(in_dim, hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, num_classes) ) def forward(self, x): # x: (B, N, D) bag_feats [] attn_weights [] for head in self.heads: bag_feat, weights head(x) bag_feats.append(bag_feat) attn_weights.append(weights) bag_feats torch.stack(bag_feats, dim1) # (B, num_heads, D) # 类别级注意力在num_heads维度上加权 scores self.class_attention(bag_feats) # (B, num_heads, 1) head_weights F.softmax(scores, dim1) final_bag_feat torch.sum(bag_feats * head_weights, dim1) logits self.classifier(final_bag_feat) return logits, attn_weights, head_weights这段代码的思路是多个注意力头各自从不同角度找出可疑切片产生多个包表示类别级注意力再融合这些包表示。这里的n_heads能起到类似“多专家”的作用避免单一注意力头只专注于某一种特征。4.5 组装训练流程训练流程包括几个环节把CT卷切分成切片将每个切片从(H, W)扩展为(3, H, W)送入Encoder得到(N, 512)的实例特征包作为MIL模型输入计算损失反向传播。下面给出一段简洁的训练代码# 文件路径train.py import numpy as np import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader from data.preprocess import make_synthetic_ct, add_local_manipulation from models.encoder import SliceEncoder from models.hexmil import HierarchicalAttentionMIL class CTBagDataset(Dataset): 构造包级数据。每个样本是一个CT卷标签是0或1。 这里用模拟数据构造真实使用时应替换为实际数据读取逻辑。 def __init__(self, num_samples40, seq_len24, img_size64): self.samples [] for i in range(num_samples): base_ct make_synthetic_ct(shape(seq_len, img_size, img_size), seedi) if i % 2 1: ct add_local_manipulation(base_ct, start_slice8, seedi) label 1 else: ct base_ct label 0 self.samples.append((ct, label)) def __len__(self): return len(self.samples) def __getitem__(self, idx): ct, label self.samples[idx] # 转成tensor但为了减少显存占用可以在主训练循环中逐包处理 return torch.from_numpy(ct).float(), label def collate_bag_to_list(batch): 由于每个包的切片数可能不同这里不直接堆叠。 返回一个合法的batch结构。 cts, labels zip(*batch) return list(cts), torch.tensor(labels, dtypetorch.long) def encode_ct_slices(ct_tensor, encoder, device): 将一个CT卷的切片逐张送入Encoder得到实例特征包。 ct_tensor ct_tensor.to(device) slices ct_tensor.unsqueeze(1) # (D, 1, H, W) slices slices.repeat(1, 3, 1, 1) # (D, 3, H, W) features [] batch_size 16 for i in range(0, slices.size(0), batch_size): batch slices[i:i batch_size] with torch.no_grad(): feat encoder(batch) features.append(feat) return torch.stack(features, dim0) # (D, 512) def train_one_epoch(model, encoder, dataloader, optimizer, criterion, device): model.train() encoder.eval() total_loss 0.0 for cts, labels in dataloader: optimizer.zero_grad() loss 0.0 for ct, label in zip(cts, labels): # ct: (D, H, W) bag_feat encode_ct_slices(ct, encoder, device) # (D, D_feat) bag_feat bag_feat.unsqueeze(0) # (1, D, D_feat) logits, attn_weights, head_weights model(bag_feat) label label.unsqueeze(0).to(device) loss loss criterion(logits, label) loss loss / len(cts) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(dataloader) def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) encoder SliceEncoder(out_dim256).to(device) model HierarchicalAttentionMIL(in_dim256, n_heads4, num_classes2).to(device) dataset CTBagDataset(num_samples60, seq_len24, img_size64) dataloader DataLoader(dataset, batch_size4, shuffleTrue, collate_fncollate_bag_to_list) optimizer optim.AdamW(list(model.parameters()) list(encoder.parameters()), lr1e-4) criterion nn.CrossEntropyLoss() for epoch in range(10): loss train_one_epoch(model, encoder, dataloader, optimizer, criterion, device) print(fEpoch {epoch 1}, Loss: {loss:.4f}) if __name__ __main__: main()这里有一个值得提醒的点我把Encoder放在torch.no_grad()下提取特征训练时只更新MIL模型参数。这样做的目的是防止特征提取器被少量模拟数据带偏。真实项目中通常建议先冻结Encoder做初步训练再逐步解冻微调。4.6 可解释结果可视化推理阶段我们需要把注意力权重映射到切片上。对单包输入模型返回的attn_weights[head_index]就是该注意力头对每个切片的权重。# 文件路径utils/visualize.py import matplotlib.pyplot as plt def visualize_attention(ct_slices, attention_weights, titleAttention): ct_slices: shape (D, H, W) attention_weights: shape (D,) 显示注意力权重最高的切片的紧凑汇总图。 top_indices attention_weights.argsort()[-3:][::-1] fig, axes plt.subplots(1, len(top_indices), figsize(12, 4)) for ax, idx in zip(axes, top_indices): ax.imshow(ct_slices[idx], cmapgray) ax.set_title(fSlice {idx}, weight{attention_weights[idx]:.3f}) ax.axis(off) plt.suptitle(title) plt.show()这样就能从包级预测倒推回切片级证据完成前摄可解释闭环。5. 常见问题与排查思路5.1 问题速查表问题现象常见原因解决思路GPU显存不足CT切片数过多或Batch Size过大降低单包切片数、减小Resize尺寸、使用梯度累积注意力权重接近均匀分布实例级注意力退化与平均池化效果相同引入多个注意力头增加注意力正则项调整初始化篡改区域切片权重仍然不高篡改区域过小模型不需要它们也能分类正确使用局部对比损失或对注意力权重增加稀疏约束训练loss下降但验证AUC低过拟合或数据划分不严格按患者级别划分数据集加入更强的数据增强篡改样本太少类别不平衡数据收集成本高使用Focal Loss、困难样本挖掘、合理的数据增强热图噪声大、不集中特征图分辨率低或注意力头关注了全局纹理提高输入分辨率结合Grad-CAM细化热图定位5.2 注意力退化的排查注意力退化是MIL模型最经典的问题之一。现象是训练后期模型认为所有切片权重都差不多包表示几乎等于“平均池化”。这会让可解释性失效因为无法定位可疑切片。排查顺序建议如下打印训练时的权重分布观察是否出现方差趋近于0检查是否有过强的Dropout或权重衰减尝试初始化偏置让注意力头初始偏向中间切片如果使用了多个注意力头检查是不是只有一个头参与训练在损失函数中加入注意力熵正则项惩罚过均匀的分布def attention_entropy_loss(weights, alpha0.1): weights: (B, N) 鼓励注意力权重不要过于均匀。 probs torch.softmax(weights, dim-1) entropy -torch.sum(probs * torch.log(probs 1e-8), dim-1).mean() return alpha * entropy需要注意的是熵正则的强度要控制好。太大会让模型只关注最显著的单个切片可能漏掉分布在多切片的篡改信号。5.3 数据泄漏问题医学影像实验中最常见也最影响结果可信度的问题是数据泄漏。如果一个患者做了多次扫描切片可能被切分到训练集和测试集模型实际上“见过”同一患者的影像分布测试指标就会虚高。避免方法是在数据划分阶段按患者ID分组而不是按文件路径分组。同时在预处理阶段记录原始数据来源保证任何增强操作都不跨越患者边界。5.4 弱监督标签噪声MIL框架下包级标签本身也可能存在噪声。比如某个CT卷实际未被篡改但标注系统标记错误。此时单个包的损失会误导模型。工程上可以引入标签平滑label smoothing来缓解criterion nn.CrossEntropyLoss(label_smoothing0.05)也可以采用软标签策略用多个标注者的一致性结果作为标签权重。6. 最佳实践与工程建议6.1 数据合规与安全边界医学影像数据属于高敏感数据。任何实验开始前必须确认数据来源合法、已脱敏、符合伦理审批要求。在生产环境中建议遵循最小权限原则数据仅存储在受限服务器训练代码不直接访问原始患者标识信息日志中不打印文件名或患者ID。绝对不能为了凑数据集而私自采集医院影像数据。篡改检测本身就是安全对抗场景如果不能保证数据链路可信模型的可信度也无从谈起。6.2 数据分层与训练验证强烈建议将整个训练流程拆成三个独立阶段特征提取器预训练在足够大的自然图像或医学图像数据集上预训练CNN。如果数据量不够直接用ImageNet预训练权重是合理起点但要意识到域差异。MIL模型训练先冻结Encoder训练注意力池化部分观察注意力分布是否合理。联合微调解冻Encoder浅层或全部层用小学习率微调。这种分阶段策略能显著提高训练稳定性尤其在样本量较小的医学影像场景中。6.3 可解释性的量化验证可解释性不只是一个“看起来漂亮”的热图。建议在测试集上量化评估注意力质量如果篡改区域有像素级标注可以计算注意力热图与真实篡改区域的Dice系数或IoU如果没有像素级标注可以设计“弱定位指标”检查最高注意力切片是否落在已知存在篡改的切片范围内记录各注意力头的注意力熵观察是否早停或受正则影响。只有可解释指标和分类指标同时提升模型的解释才有说服力。6.4 部署阶段的风险控制检测模型部署到实际系统时要关注几点固定预处理参数CT重采样、窗宽窗位、归一化数值都必须和训练时一致保留模型版本记录CT影像数据分布会随设备品牌、扫描参数变化模型需要定期用新数据做外部验证设置预测置信度阈值医学筛查场景通常需要低漏检建议单独设定敏感度阈值而不是直接用默认0.5人工复核机制任何AI告警都应进入人工复核流程模型产生的注意力热图只作为参考证据。7. 总结与学习路线现在回看HexMIL这套方案其实是在解决三个互相纠缠的问题A图像到底有没有被篡改、哪些切片区域最可疑、这个结论能否被人工复核。层次注意力MIL把这三个问题收敛到了一个模型里包级标签做监督信号实例注意力做局部定位类别级注意力在不同异常模式间做融合最终输出既分类又定位的解释结果。如果你对这个方向感兴趣下一步可以有计划地深入先从经典的MIL论文读起理解Attention-based Deep Multiple Instance Learning的基础公式然后阅读CLAM等更大规模的弱监督病理切片分类方法体会层次化注意力在更大包规模下的设计差异再结合Monai库把DICOM读取、预处理、切片组织标准化最后用一个小规模的篡改检测模拟数据集跑通完整流程。真实项目中最大的风险往往不在模型结构而在于数据分布漂移和标签质量。同一个篡改算法在不同CT设备上的表现差异可能非常大因此跨中心外部验证比单纯调模型结构更重要。如果你手头没有公开的CT篡改数据集可以先从合成实验开始。把正常CT和模拟篡改CT作为二分类数据跑通本文的代码闭环再逐步增加篡改方式的多样性。把基础流程吃透之后无论是更换更强的特征提取器还是引入对抗训练都会顺手很多。如果这篇笔记对你的项目有参考价值可以收藏备用。后续我再整理篡改数据增强和注意力可视化评估的进阶实践欢迎持续关注。
返回列表