尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

可解释CT篡改检测:基于MIL与分层注意力的HexMIL实现

可解释CT篡改检测:基于MIL与分层注意力的HexMIL实现 CT 检查是肺癌、卒中、腹部外伤等场景最常用的影像学依据之一。过去几年生成式 AI 在医学影像上的能力提升很快已经可以生成足够逼真的二维 X 光片、三维 CT 断层序列。这让医学影像安全出现了一个新的攻防场景如果有攻击者用 AI 在真实 CT 体数据里悄悄插入一个病灶或者抹掉一个已经存在的结节辅助诊断系统会不会被骗现实是这类篡改很难靠肉眼识别而且攻击者可以让篡改痕迹在连续多层切片中保持大致协调传统的逐图检查基本失效。更麻烦的是标注。给一个 CT 体数据写“正常”或“异常”医生很快就能完成但要标注具体哪一层切片、哪个空间位置被动过往往要逐层逐区域看成本高、主观性大。HexMIL 这类方法的价值就在于用多实例学习把“体数据有没有被改”这个很粗的监督信号转化为“哪几层切片、哪些区域最可疑”的细粒度解释。从论文标题看它并不是又引入了一个全新的检测网络而是把三个关键思想组合在一起MIL、分层注意力、Ante-Hoc 可解释性。本文会按照“问题背景 → 核心概念 → 方法拆解 → 环境准备 → 完整代码 → 效果验证 → 常见问题 → 工程建议”的顺序展开。我会用 PyTorch 实现一个简化版 HexMIL跑通从数据构造到可解释输出的完整链路。这样即使你手头暂时没有真实 CT 数据也能把模型结构和训练流程先跑起来然后再替换成自己的数据集。1. 为什么需要可解释的 CT 体积篡改检测1.1 CT 体积数据比 2D 影像更容易被“悄然修改”二维影像篡改检测已经有大量研究比如检测 X 光片的局部拼接、重采样痕迹、生成痕迹等。但 CT 是三维体数据存储形式是连续断层切片叠加在一起。攻击者如果想要修改一个肺结节不会只改一张切片而是会在相邻十几层甚至几十层同步修改否则重建出来的三维形体不自然医生一滚动鼠标就能发现。这种跨层协同修改给检测带来了两个直接影响。第一单张切片的局部统计特征可能没有明显异常因为攻击者已经尽量让灰度过渡、纹理细节接近真实器官。第二篡改区域在体数据中占比极低把整卷 CT 作为输入做 3D 分类模型需要从约一亿个体素中找到几百个可疑体素样本极不均衡直接训练 3D CNN 很容易把注意力放在“全局灰度分布是否正常”上而忽略真正的篡改区域。1.2 弱监督是实际落地必须接受的现实如果要求像素级篡改标注训练数据很难获得。真实篡改 CT 本身就少就算有标注者也需要逐层勾勒异常区域而且不同医生对“改了什么、哪里是源头”判断不一定一致。更常见的场景是机构手里只有数据级标签这卷 CT 是正常的那卷 CT 是经过 AI 修改的但具体改了哪里没有标注或者标注非常粗糙。MIL 正好处理这种弱监督问题。我们把一个 CT 体数据看作一个“包”把切片或局部块看作“实例”。训练时只需要告诉模型这个包是正样本还是负样本不用告诉它哪个实例导致了正样本。模型被自动引导去发现最能解释正样本标签的证据。这听起来像是在模糊处理问题但注意力机制可以让“找证据”的过程显式化最终同时得到预测和证据位置。1.3 对比几种可能的技术路线方案类型监督级别是否输出定位主要问题全监督分割逐体素标注是标注成本极高泛化依赖标注质量2D CNN 分类数据级基本不能忽略层间连续性容易误判3D CNN 分类数据级可做后处理可视化黑盒解释依赖额外方法难以证伪MIL 分层注意力数据级是天然输出注意力不等于因果需要医学验证从材料看HexMIL 走的是最后一条路线。它的目标不是追求某个单一准确率指标而是让模型在给出“是否被篡改”结论的同时说明“我根据哪些切片、哪些局部区域做的判断”。这对于医疗安全场景尤其重要因为模型不能只是自信地给出一个数字还要能让医生去复核。2. 核心概念MIL、分层注意力与 Ante-Hoc 可解释性2.1 MIL把 CT 体数据看成“一袋切片”多实例学习Multiple Instance LearningMIL的经典设定是一个包bag包含多个实例instance。包的标签为正意味着包里至少有一个正实例包的标签为负意味着所有实例都是负的。训练时看起来是包级标签但模型要学习识别那些隐藏的负/正实例。在 CT 篡改检测场景里一个 CT 体数据就是包一张切片或一个局部块就是实例。如果这卷 CT 被篡改过那么至少有一张切片、至少有一个局部区域是异常的。医学上“病灶通常集中在局部”的先验和 MIL 的假设非常匹配。它不需要我们对每一层切片都做标注只需要体数据级别的正负标签就能训练出来。传统 MIL 的聚合方式很简单比如对实例特征取最大值或平均值。最大值池化过于极端对噪声敏感平均值池化又会把单点异常稀释掉。注意力池化则是给每个实例学习一个权重再对特征做加权求和。这样模型可以自动放大那些携带“篡改证据”的实例抑制无关的正常区域。2.2 注意力池化让模型自己找出“最可疑的切片”公式可以这样理解假设实例特征为 h_1, h_2 ... h_n注意力网络先对每个实例算一个得分 score_i再用 Softmax 归一化成权重 w_i最后聚合结果 z Σ w_i h_i。权重 w_i 越大说明第 i 个实例对最终判断的贡献越大。这个机制有两个好处。第一模型可以端到端学习不需要额外的区域标注。第二权重天然可以作为解释信息我们不需要事后用 Grad-CAM 再去“猜测”模型看哪里而是直接读取权重。当然注意力权重不等于因果比如模型可能因为某个区域灰度分布怪异而给高权重但这个区域不一定是攻击者真正修改的位置。所以在医疗场景中注意力只能作为“嫌疑区域提示”还需要医生复核。2.3 分层注意力先观察局部再观察全序列如果只做一层注意力对象通常是整张切片。但 CT 诊断需要跨尺度先看某一层里哪个局部区域可疑再看整个体数据中哪些层最可疑。HexMIL 的“Hierarchical Attention”对应的正是这种两级结构。第一级注意力把一张切片切成若干局部块对块特征做注意力池化得到“这一张切片的表示”同时给出块级权重。第二级注意力把所有切片的表示做注意力池化得到“整个体数据的表示”同时给出切片级权重。最终分类器使用体数据表示输出篡改概率。用医生的工作方式来类比医生看 CT 时不会一帧帧匀速滚动他会先快速扫一遍感觉某一层肺窗特别奇怪然后停下放大观察局部纹理和密度如果局部确实可疑他会再翻前后若干层确认这个异常是连续性存在还是孤立的偶然伪影。分层注意力也是这个逻辑底层负责局部高层负责全局。2.4 Ante-Hoc 可解释性解释是模型前向计算的一部分机器学习可解释性有两条路线。Post-Hoc 指模型训练好之后再用 Grad-CAM、SHAP 等方法解释模型“为什么这么判断”。这种做法实现简单但不一定忠实于模型内部计算而且解释结果可能不稳定。Ante-Hoc 则是在模型设计阶段就预留解释能力让解释信息在模型前向传播时自然产生通常是注意力权重或某个可解释特征向量。HexMIL 强调 Ante-Hoc意味着“预测 解释”是由同一个前向过程给出的。输入一个 CT 体数据模型同时输出篡改概率、切片级注意力权重和块级注意力权重。这种设计特别适合需要审计的场景临床医生看到一个输出的同时也能看到判断依据安全团队可以在上线前检查模型注意力是否集中在合理区域而不是学习到了某个数据集的无关伪影。3. 方法设计概览HexMIL 整体思路3.1 从 CT 体数据到可解释预测的完整链路一个典型的 HexMIL 流程可以分为五个阶段数据预处理读取 CT 体数据完成重采样、窗宽窗位调整、归一化。实例构建将体数据拆成“切片”再把每张切片拆成“局部块”。局部块是注意力机制的最小单元。特征提取用二维卷积网络把每个局部块编码成一个固定长度的特征向量。分层注意力聚合先对块特征做注意力池化得到切片表征再对切片表征做注意力池化得到体数据表征。分类与解释输出体数据表征过一层线性分类器输出篡改概率同时读出两级的注意力权重作为可疑切片和可疑区域的提示。这个结构里“体数据级标签”是唯一的监督信号而解释信息由注意力权重产生不需要额外的分割标注。这也是它在标注受限场景下更实用的原因。3.2 注意力权重为何能作为解释注意力池化的输出是 z Σ w_i h_i分类器只依赖 z。如果某个切片或某个局部块的权重很小那么即使它对应的特征向量携带了强信号经过加权求和后也会被稀释。反过来最终分类结果很大程度上由高权重实例决定。因此权重序列天然反映了模型在做决策时依赖哪些输入位置。当然这里有一个需要警惕的点注意力高权重不一定代表“真实篡改区域”它只代表“模型认为和标签相关的区域”。如果训练数据存在偏差比如所有正样本都来自同一台扫描设备、同一类型伪影模型可能把注意力集中在设备伪影上。因此工程上需要把注意力可视化交给医生或安全审计人员验证而不是盲目相信。3.3 损失函数与训练策略基础损失通常用二分类交叉熵。由于篡改体数据在实际中可能远少于正常体数据可以改用带权重的交叉熵或 Focal Loss让模型更关注少数类。为了让注意力更符合“局部篡改”的先验可以给注意力权重加上稀疏正则。例如希望体数据级注意力只集中在少数切片而不是平均分布可以计算注意力权重的熵惩罚过高的熵值。也可以增加一致性正则让相邻切片的注意力变化更平滑减少孤立尖刺。需要说明的是这两类正则手法属于通用 MIL 训练技巧。在正式复现 HexMIL 论文时应该以论文或开源代码中的实现为准。没有官方细节时不要用网上的二手转述代替原文。4. 环境准备与前置条件本文的演示代码基于 PyTorch建议用 conda 创建独立环境。你可以选择 GPU 版或 CPU 版如果只是跑通流程CPU 也可以要处理真实 CT 数据强烈建议 GPU。conda create -n hexmil python3.9 -y conda activate hexmil # 如果使用 GPU请根据本机 CUDA 版本选择 PyTorch 安装命令 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # CPU 版可直接执行 # pip install torch torchvision后续还需要医学影像读取、科学计算和可视化相关库pip install simpleitk pydicom numpy pandas matplotlib scikit-learn如果希望用更完整的医学图像处理工具集可以安装 MONAI。本文示例不强制依赖它所以可以按需安装pip install monai版本方面因为项目没有给出固定依赖建议使用截至本文写作时较新的稳定版本。实际复现时如果安装顺序冲突优先以 PyTorch 版本为基准再对齐其他库。5. 核心流程拆解5.1 读取 CT 体数据CT 数据常见格式是 DICOM 序列或 NIfTI。DICOM 是一个文件夹里有很多张断层切片每张切片包含患者信息、扫描参数和像素数据。读取时通常需要把整个序列加载成一个三维数组。NIfTI 更简单一个文件就是一整个体数据。使用 SimpleITK 可以同时处理这两种格式import SimpleITK as sitk def load_ct_volume(path_or_dir, is_dicomTrue): if is_dicom: reader sitk.ImageSeriesReader() dicom_names reader.GetGDCMSeriesFileNames(path_or_dir) if len(dicom_names) 0: raise RuntimeError(No DICOM series found in path: {}.format(path_or_dir)) reader.SetFileNames(dicom_names) image reader.Execute() else: image sitk.ReadImage(path_or_dir) volume sitk.GetArrayFromImage(image) # shape: (z, y, x) spacing image.GetSpacing() # (x_spacing, y_spacing, z_spacing) return volume, spacing需要特别注意SimpleITK 返回的数组顺序是(z, y, x)也就是第一维是切片序号。很多新手在这里把维度弄反导致预处理时出现“旋转后的 CT”。建议在写数据管线时先打印volume.shape和spacing并对其中一个公开样例做可视化确认方向。5.2 预处理重采样、归一化和窗宽窗位CT 的原始灰度单位是 Hounsfield UnitHU不同扫描设备之间数值范围差异不大。我们通常先按窗宽窗位截断比如肺部窗口常用窗宽 1500、窗位 -500软组织窗口常用窗宽 400、窗位 40。不同任务需要不同的窗口篡改检测更关心异常密度区域因此可以先用标准窗口截断。import numpy as np def preprocess_hu(volume, window_width400, window_level40): lower window_level - window_width / 2.0 upper window_level window_width / 2.0 volume np.clip(volume, lower, upper) volume (volume - lower) / (upper - lower) return volume.astype(np.float32)这一步的作用是压缩灰度范围让模型不用在 12-bit 甚至 16-bit 的宽动态范围里学习同时把空气、脂肪、软组织、骨骼等结构在数值上区分得更明显。真实项目中还应该根据体素间距做重采样。比如把层距和面内分辨率都重采样到 1mm避免不同中心的数据尺度不一致。重采样可以用 SimpleITK 的Resample也可以用 MONAI 的ResampleToMatch。5.3 构造 Bag 数据集体数据、切片、局部块MIL 的一层重要设计是“如何把体数据变成包”。如果直接使用所有切片不同体数据的层数可能不同会给 batch 训练带来麻烦。更实际的做法是训练时从每个体数据中随机采样固定数量的切片再对每张切片切出固定数量的局部块。这样每个 batch 的形状是固定的模型可以稳定训练。切片采样时要注意不要简单取前 N 层因为 CT 卷两端往往是空气或扫描床边缘没有诊断意义。应该先做体数据范围裁剪去掉灰度极低的背景层。如果目标是检测局部篡改采样应该尽量保留整个体数据范围避免只取中间区域导致模型学不到上下文。局部块切分可以使用规则网格。为了演示本文把每张 64×64 的切片切分成 4×4 个 16×16 的块。真实数据可以直接在原始分辨率上切也可以先缩放到固定尺寸。切块的坐标需要保留因为最后要把注意力权重映射回原始图像区域。5.4 分层注意力模型结构模型结构可以分成四段Patch Embedding输入是一张切片中的每个局部块输出一个固定长度向量。Slice Attention对同一张切片的所有块向量做注意力池化得到切片向量并输出块权重。Volume Attention对一个体数据的所有切片向量做注意力池化得到体数据向量并输出切片权重。Classifier体数据向量经过线性层输出一个 logit。关键点在于两类权重分别对应两种尺度的解释。块权重能告诉你“这张切片里最可疑的位置”切片权重能告诉你“该重点查看哪些切片”。生产中可以把这两级权重分别叠加到 CT 原始图像上生成可疑区域热图。5.5 训练、验证与解释输出训练时用包级标签计算损失反向传播会同时更新 Patch Embedding、注意力和分类器。验证时除了看体数据级别的 AUC、准确率还要检查注意力分布是否稳定模型是否总是只关注某一个固定的切片位置而不是根据体数据内容动态变化高权重的切片和真实攻击切片是否重叠如果一个正常体数据也被给出很高的正样本概率它关注的是哪些区域这些都是评估解释质量的重要信号。只看准确率不看解释等于没有利用 HexMIL 的核心能力。6. 完整示例代码实现下面我们用合成数据实现一个可运行的简化版 HexMIL。合成数据不包含真实医疗图像只用于验证模型结构和训练链路。实际使用时要替换成合规获取的真实 CT 数据并重新预处理。6.1 合成 CT 体数据生成与 Datasetimport random import numpy as np import torch from torch.utils.data import Dataset class SyntheticCTVolumes(Dataset): 合成 CT 体数据 - 每个样本是形状为 (S, P, C, H, W) 的张量 - S 表示切片数 - P 表示每张切片切出的局部块数 - C, H, W 表示单通道局部块的尺寸 def __init__(self, num_volumes64, num_slices16, patches_per_side4, patch_size16, manipulated_ratio0.5, seed0): self.num_volumes num_volumes self.num_slices num_slices self.patches_per_side patches_per_side self.patch_size patch_size self.manipulated_ratio manipulated_ratio self.seed seed random.seed(seed) np.random.seed(seed) def _make_patches(self, manipulated): # patch 基础值是随机噪声模拟正常组织纹理 num_patches self.patches_per_side * self.patches_per_side patches np.random.rand(self.num_slices, num_patches, 1, self.patch_size, self.patch_size).astype(np.float32) * 0.3 if manipulated: # 选择 1 到 3 张攻击切片 num_attack_slices np.random.randint(1, 4) attack_slices np.random.choice(self.num_slices, num_attack_slices, replaceFalse) for s in attack_slices: # 随机选择该切片中的某个局部块 patch_idx np.random.randint(num_patches) # 给这个局部块添加一个较强的高密度信号模拟异常区域 patches[s, patch_idx] 2.0 return patches def __len__(self): return self.num_volumes def __getitem__(self, idx): manipulated random.random() self.manipulated_ratio patches self._make_patches(manipulated) volume torch.from_numpy(patches) label torch.tensor(1.0 if manipulated else 0.0, dtypetorch.float32) return volume, label这个 Dataset 的核心是“正包里一定存在至少一个高值局部块”。这样注意力模型很容易学会把高权重分给攻击块从而让链路跑通。真实场景中伪造信号会更隐蔽模型也需要更强的特征提取器。6.2 Patch Embedding 与分层注意力模型import torch.nn as nn class PatchEmbedding(nn.Module): 把每个局部块编码成一个 embed_dim 维向量 def __init__(self, in_channels1, embed_dim128): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_channels, 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.AdaptiveAvgPool2d((1, 1)), ) self.proj nn.Linear(64, embed_dim) def forward(self, x): # x: (B, N, C, H, W) B, N, C, H, W x.shape x x.view(B * N, C, H, W) feat self.conv(x).flatten(1) return self.proj(feat).view(B, N, -1) class AttentionPooling(nn.Module): 注意力池化 输入 (B, N, D)输出 (B, D) 和权重 (B, N) def __init__(self, embed_dim128, hidden_dim64): super().__init__() self.attention nn.Sequential( nn.Linear(embed_dim, hidden_dim), nn.Tanh(), nn.Linear(hidden_dim, 1, biasFalse), ) def forward(self, x): scores self.attention(x).squeeze(-1) # (B, N) weights torch.softmax(scores, dim-1) # (B, N) pooled torch.sum(weights.unsqueeze(-1) * x, dim1) return pooled, weights class HexMILNet(nn.Module): 简化版 HexMIL 底层级特征 切片内注意力 切片间注意力 二分类头 def __init__(self, in_channels1, embed_dim128): super().__init__() self.patch_embedding PatchEmbedding(in_channels, embed_dim) self.slice_attention AttentionPooling(embed_dim) self.volume_attention AttentionPooling(embed_dim) self.classifier nn.Linear(embed_dim, 1) def forward(self, x): # x: (B, S, P, C, H, W) B, S, P, C, H, W x.shape x x.view(B, S * P, C, H, W) patch_feats self.patch_embedding(x) # (B, S*P, D) patch_feats patch_feats.view(B, S, P, -1) slice_vecs [] patch_weights [] for s in range(S): slice_vec, weights_s self.slice_attention(patch_feats[:, s]) slice_vecs.append(slice_vec) patch_weights.append(weights_s) slice_vecs torch.stack(slice_vecs, dim1) # (B, S, D) patch_weights torch.stack(patch_weights, dim1) # (B, S, P) vol_vec, slice_weights self.volume_attention(slice_vecs) # (B, D), (B, S) logits self.classifier(vol_vec).squeeze(-1) # (B,) return logits, slice_weights, patch_weights这里有两个值得注意的设计点。第一个是PatchEmbedding把(B, S * P, C, H, W)当成一个整体输入。好处是卷积计算可以批量完成不需要写嵌套循环。第二个是切片内注意力我们用一个 for 循环遍历切片代码更直观。如果切片数很多可以考虑把切片维度并入 Batch或使用更高效的分组实现但本文重点是讲清思路不是优化极致性能。6.3 训练流程与可解释性输出import torch.optim as optim from torch.utils.data import DataLoader def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss 0.0 correct 0 total 0 for volume, label in loader: volume volume.to(device) label label.to(device) optimizer.zero_grad() logits, _, _ model(volume) loss criterion(logits, label) loss.backward() optimizer.step() preds (torch.sigmoid(logits) 0.5).long() correct (preds label.long()).sum().item() total label.size(0) total_loss loss.item() * label.size(0) return total_loss / total, correct / total def evaluate(model, loader, criterion, device): model.eval() total_loss 0.0 correct 0 total 0 with torch.no_grad(): for volume, label in loader: volume volume.to(device) label label.to(device) logits, _, _ model(volume) loss criterion(logits, label) preds (torch.sigmoid(logits) 0.5).long() correct (preds label.long()).sum().item() total label.size(0) total_loss loss.item() * label.size(0) return total_loss / total, correct / total def show_explanation(model, volume, device): 输出预测标签、最可疑的切片和局部块索引 model.eval() with torch.no_grad(): logits, slice_weights, patch_weights model(volume.to(device)) probs torch.sigmoid(logits) # 取 batch 中第一个样本 slice_idx torch.argmax(slice_weights[0]).item() patch_idx torch.argmax(patch_weights[0][slice_idx]).item() patch_line patch_idx // 4 patch_col patch_idx % 4 print(Predicted probability: {:.4f}.format(probs[0].item())) print(Most suspicious slice index: {}.format(slice_idx)) print(Most suspicious patch at: row {}, col {}.format(patch_line, patch_col)) def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(Using device:, device) train_dataset SyntheticCTVolumes(num_volumes100, seed0) val_dataset SyntheticCTVolumes(num_volumes20, seed1) train_loader DataLoader(train_dataset, batch_size4, shuffleTrue) val_loader DataLoader(val_dataset, batch_size4, shuffleFalse) model HexMILNet(in_channels1, embed_dim128).to(device) criterion nn.BCEWithLogitsLoss() optimizer optim.Adam(model.parameters(), lr1e-3) for epoch in range(10): train_loss, train_acc train_one_epoch(model, train_loader, optimizer, criterion, device) val_loss, val_acc evaluate(model, val_loader, criterion, device) print(Epoch {:02d}/10 train_loss {:.4f} train_acc {:.4f} val_loss {:.4f} val_acc {:.4f}.format( epoch 1, train_loss, train_acc, val_loss, val_acc)) sample_volume, sample_label val_dataset[0] sample_volume sample_volume.unsqueeze(0).to(device) print(Ground truth label:, sample_label.item()) show_explanation(model, sample_volume, device) if __name__ __main__: main()这段代码可以直接复制到.py文件运行。合成数据的特征是“攻击块的值显著高于正常背景”因此模型跑几个 epoch 后准确率通常能明显上升。更重要的是show_explanation会输出最可疑的切片编号和局部块行列坐标让你直接看到解释结果。需要提醒的是这个示例没有使用真实 CT 数据它只负责验证代码链路。你不能把这里的准确率当作 HexMIL 在真实篡改 CT 上的性能。6.4 替换为真实 CT 数据的一般方法假如你已经有了一组真实 CT 体数据和数据级标签可以按下面思路改造 Dataset用load_ct_volume读取所有体数据保存为volume_id - volume array的映射。对每个体数据执行窗口截断、重采样和归一化保存为.npy或内存缓存。在__getitem__中对每个 volume 随机采样固定数量的切片。对每张切片用规则网格切出局部块并记录局部块在原始切片上的坐标。返回(volume_tensor, label, meta_info)其中meta_info用于把注意力权重映射回原始像素坐标。真实数据的难点通常不在模型而在数据清洗。不同 CT 设备的体素间距差异、扫描范围差异、噪声水平差异都会让注意力模型找到“捷径特征”。比如如果所有正样本都来自某一台设备模型会倾向于给这台设备的整体灰度分布高权重而不是真正去定位篡改区域。因此在评估阶段一定要做外部验证集至少保证设备来源是模型没见过的。7. 运行结果与效果验证7.1 预期输出形态运行上面训练脚本后你会看到类似下面的日志Using device: cpu Epoch 01/10 train_loss 0.6932 train_acc 0.4500 val_loss 0.7001 val_acc 0.4000 Epoch 02/10 train_loss 0.5123 train_acc 0.7200 val_loss 0.5834 val_acc 0.6500 Epoch 03/10 train_loss 0.3011 train_acc 0.8600 val_loss 0.4122 val_acc 0.7500 ... Epoch 10/10 train_loss 0.1021 train_acc 0.9800 val_loss 0.2311 val_acc 0.9000具体数值会因随机种子、batch size、学习率不同而变化。重点观察两个趋势训练 loss 和验证 loss 是否同步下降如果验证 loss 后期反弹说明过拟合。模型是否明显区分正负样本。合成数据里正样本的异常块信号很强如果准确率长期不增长大概率是代码链路有问题。7.2 想确认注意力是否真的指向攻击块训练完成后脚本最后一行会打印预测概率、最可疑切片编号和局部块坐标。你可以把注意力权重按如下方式检查对每个测试体数据获取真实攻击切片编号。如果show_explanation输出的最可疑切片就在攻击切片附近说明模型定位正确。对攻击切片看最可疑局部块是否等于我们生成时手动加了高值的块。查看注意力权重是否
返回列表