行业资讯

HexMIL:基于层次注意力多示例学习的CT体数据AI篡改检测

发布时间:2026/8/27 21:35:18
HexMIL:基于层次注意力多示例学习的CT体数据AI篡改检测 CT 体数据作为医院诊疗流程中最常见的高维影像之一正在越来越多地被生成模型当作篡改目标。传统医学影像防伪研究大多停留在 2D 平面图像而 CT 本质上是几十到数百张切片堆叠出的 3D 体数据单张切片的异常在整卷数据里往往只占很小比例。HexMIL 这个研究方向解决的正是这种场景下的检测问题用层次注意力多示例学习Hierarchical Attention MIL识别被 AI 篡改过的 CT 体数据并在给出风险预测的同时直接输出模型判定为可疑的切片位置和区域。这种模型自身携带解释能力的做法被称为 Ante-Hoc 可解释性。在实际工程里MIL 这个缩写有歧义汽车电子领域常指 Model-in-the-Loop 测试也就是模型在环仿真。本文的 MIL 严格指 Multiple Instance Learning即多示例学习。理解这个区别之后下面从问题背景、核心概念、方法骨架、数据构造、训练验证和工程落地几个维度把 HexMIL 这类方法的实现思路完整拆解一遍。1. 先理解 CT 体数据被 AI 篡改后检测难在哪里1.1 篡改不是整卷重写而是局部伪造CT 影像中的病变信息主要体现在组织密度差异上用 HU 值量化。AI 篡改 CT 的手段大致可以分为四类插入伪病变、删除真病变、重建伪影和结构替换。插入伪病变在原本正常的区域生成肿瘤、出血或结节用于影响诊断结果。删除真病变把真实存在的病灶平滑掉可能用于隐藏证据或骗取理赔。重建伪影模拟低剂量噪声、运动伪影或重建伪影干扰读片判断。结构替换改变器官边界或把一侧结构复制到另一侧。这些操作普遍只在体数据的三维局部发生。问题也随之而来医生面对的是数百张切片篡改区域可能只占整卷数据的 0.5% 以下。逐张切片检查效率极低普通 2D 网络又容易被大量正常切片淹没。篡改类型视觉表现检测难点插入伪病变局部出现异常密度团块可能与真实病变高度相似删除真病变原本结节被平滑填充没有直接的异常信号重建伪影类似运动伪影或低剂量噪声与真实噪声难以区分结构替换器官边界不连续需要上下文和先验知识1.2 为什么 2D 逐帧检测方法直接在体数据上失效最简单的处理方案是把每一帧切片当成独立图像训练二分类网络再对所有切片的预测取平均作为整个体数据的判断。这个方案在 CT 篡改检测上存在几个根本问题。第一标签比例失衡。整套 CT 中只有少数切片携带篡改痕迹逐帧训练时大量正常切片会被模型当成易分负样本。梯度被正常切片主导后模型很难对稀疏的异常区域形成敏感决策边界。第二缺少三维上下文。单张切片的异常信号很弱医生判断病灶是否真实存在时依赖相邻切片的连续性。2D 网络完全看不到这种连续性也无法感知“器官边缘在相邻层是否突然消失”这类三维特征。第三平均策略会稀释信号。假设一个体数据有 300 张切片其中只有 5 张被篡改逐帧预测的置信度本就不高取平均后大概率被归为正常。这也是为什么需要用多示例学习来处理这个任务。MIL 只需要把整个体数据标注为一个“包”包内至少有一张异常切片就认为包为正样本模型通过注意力机制自动找到关键切片而不是强制所有切片都参与最终决策。1.3 为什么医学场景需要 Ante-Hoc 解释医学影像检测不能只输出“这是篡改”的结论医生还需要知道模型为什么这样判断、具体哪个位置可疑、依据的是哪一段数据。Post-Hoc 解释是最常见的补救方案典型工具包括 Grad-CAM、SHAP、LIME。这类方法的共同问题是解释是在模型训练完成之后额外计算的模型本身不保证这些解释与决策过程一致。一旦解释与真实决策路径不符很难判断是模型学错了还是解释方法本身不合适。Ante-Hoc 可解释性要求在模型设计阶段就把解释能力嵌入架构。HexMIL 的层次注意力机制天然满足这个要求模型每一层注意力分配给哪些切片、哪些区域本身就是决策过程的一部分。输入一个新的体数据后模型输出的不只是风险概率还有一组切片级注意力权重。权重高的切片就是模型判断为可疑的依据。这种解释不是事后附加的而是训练目标的一部分。维度Post-Hoc 解释Ante-Hoc 解释解释产生时机模型训练后单独计算模型推理时同步输出与决策一致性不保证注意力参与决策一致性更强训练损失约束无可加入解释监督约束适用场景快速分析已有黑盒模型新模型设计阶段医学部署成本需要额外验证解释质量解释结构固定可审计2. 三个核心概念MIL、层次注意力、Ante-Hoc2.1 MIL 多示例学习包、示例和预测逻辑多示例学习的问题定义非常明确有一批“包”每个包包含任意数量的“示例”。如果包内至少一个示例为正包标签为正只有全部示例为负时包标签才为负。训练时只提供包标签不需要为每个示例单独标注。对应到 CT 篡改检测场景概念映射如下MIL 概念CT 场景映射包一次 CT 扫描得到的完整体数据示例采样得到的切片或局部 3D patch正包体数据中至少一处被 AI 篡改负包完整、未经篡改的正常体数据示例标签每张切片是否属于篡改区域训练时通常不提供包预测整个体数据是否被篡改MIL 的核心假设是“存在性”异常只需出现在局部就能通过注意力池化机制被模型捕获。训练过程中模型会学习为携带异常信号的切片分配较高注意力权重从而聚合成能够支撑正包判断的全局特征。2.2 层次注意力两级池化处理三维体数据层次注意力的出发点是CT 体数据既有空间结构又有序列结构。一种常见的处理方式是把体数据按切片序列切分为若干组每组包含连续若干张切片。第一级注意力在组内挑选有代表性的切片把组内信息聚合成局部特征第二级注意力在组之间挑选最重要的局部区域聚合成整个体数据的全局表示。这种设计有三个直接好处。第一降低计算复杂度。如果对几百张切片直接做全局注意力显存和计算量都可能失控。分组建模后注意力矩阵规模被限制在组内第二级只处理少量组特征。第二保留局部与全局两级语义。组内注意力可以定位“哪张切片可疑”组间注意力可以定位“哪个空间区域被篡改”。两者结合后输出结果更符合影像科医生的读片习惯。第三天然支持变长输入。不同 CT 的切片数量不固定模型可以通过组 padding 或采样策略统一长度而不需要把输入强行裁剪到固定切片数。需要强调一点这里说的“层次”核心思想是把注意力池化作为基础算子在不同尺度上重复使用而不是必须限制在两到三个注意力层。细粒度空间信息逐步聚合到粗粒度体数据表示才是层次结构的本质。2.3 Ante-Hoc 不等于加一个解释模块不少实现把模型设计成两段式结构后端是分类网络旁边再接一个 Grad-CAM 或热力图生成模块。这种设计并不是 Ante-Hoc因为热力图和分类结果之间没有结构化的约束。Ante-Hoc 模型的关键特征是解释输出和分类输出共用同一套前置特征且分类决策依赖解释结构。以 HexMIL 为例体数据的预测分数来自对注意力加权特征的进一步变换那么注意力权重本身就是解释内容。修改注意力权重后如果预测结果会直接改变说明解释和决策是绑定的而不是事后补上的。这一点在医学场景中很重要因为它保证了解释的可审计性。医生看到“模型认为第 87 到 92 层之间存在异常”时可以明确知道这个判断源自模型在这些层分配了高注意力而不是用一个额外的可视化工具反推。3. HexMIL 方法骨架与核心实现3.1 整体流程从一个输入 CT 体数据到最终输出HexMIL 风格的方法可以抽象成四个步骤。第一步解析并预处理 DICOM将体数据重采样、截断 HU 值、归一化。第二步对体数据按轴向切片采样 N 帧每帧缩放到固定尺寸通过一个共享权重的切片编码器提取特征。第三步把切片特征序列送入层次注意力模块第一级按组聚合得到局部特征第二级在组间聚合得到体数据全局特征。第四步全局特征经过分类头得到风险概率同时输出两级注意力权重作为异常定位线索。下面给出一个基于 PyTorch 的最小骨架用于说明层次注意力 MIL 的核心模块如何组合。这个代码不追求复现论文的每一项细节只体现方法结构。import torch import torch.nn as nn import torch.nn.functional as F from torchvision import models class SliceEncoder(nn.Module): 将一张 2D CT 切片编码为特征向量。 def __init__(self, out_dim512, backboneresnet18, pretrainedFalse): super().__init__() if backbone resnet18: if pretrained: base models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) else: base models.resnet18(weightsNone) self.features nn.Sequential(*list(base.children())[:-1]) in_features base.fc.in_features else: raise NotImplementedError(funsupported backbone: {backbone}) self.proj nn.Linear(in_features, out_dim) def forward(self, x): # x: [B, 1, H, W] 或 [B, 3, H, W] x self.features(x) x x.flatten(1) return self.proj(x) class AttentionPooling(nn.Module): 门控注意力池化对示例特征加权聚合。 def __init__(self, in_dim, hidden_dim128): super().__init__() self.att nn.Sequential( nn.Linear(in_dim, hidden_dim), nn.Tanh(), nn.Linear(hidden_dim, 1) ) def forward(self, features, maskNone): # features: [B, N, D] # mask: [B, N] 或 None1 表示有效位置 if mask is None: mask torch.ones(features.shape[:-1], devicefeatures.device) # 防止某些组全部是 padding 时softmax 出现 NaN valid_num mask.sum(dim-1, keepdimTrue) safe_mask mask.clone() safe_mask[valid_num.squeeze(-1) 0, 0] 1 scores self.att(features).squeeze(-1) # [B, N] scores scores.masked_fill(safe_mask 0, float(-inf)) weights torch.softmax(scores, dim-1) # [B, N] pooled torch.sum(weights.unsqueeze(-1) * features, dim1) # [B, D] return pooled, weights class HierarchicalAttentionMIL(nn.Module): 层次注意力 MIL 的最小实现。 输入一组切片特征 [B, T, D] 输出体数据风险分数、组级注意力、切片级注意力 def __init__(self, feat_dim512, group_size16, num_groups8, hidden_dim128): super().__init__() self.group_size group_size self.num_groups num_groups self.inner_att AttentionPooling(feat_dim, hidden_dim) self.outer_att AttentionPooling(feat_dim, hidden_dim) self.classifier nn.Linear(feat_dim, 1) def forward(self, slice_features): # slice_features: [B, T, D] B, T, D slice_features.shape assert T self.group_size * self.num_groups padded torch.zeros(B, self.group_size * self.num_groups, D) mask torch.zeros(B, self.group_size * self.num_groups) padded padded.to(slice_features.device) mask mask.to(slice_features.device) padded[:, :T] slice_features mask[:, :T] 1 padded padded.view(B, self.num_groups, self.group_size, D) mask mask.view(B, self.num_groups, self.group_size) # 第一级组内注意力 x padded.reshape(B * self.num_groups, self.group_size, D) m mask.reshape(B * self.num_groups, self.group_size) group_feats, slice_weights self.inner_att(x, m) group_feats group_feats.view(B, self.num_groups, D) # 判断每个组是否有真实切片 group_mask mask.view(B, self.num_groups, self.group_size).any(dim-1).long() # 第二级组间注意力 volume_feat, group_weights self.outer_att(group_feats, group_mask) logit self.classifier(volume_feat).squeeze(-1) return logit, group_weights, slice_weights