尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

19种器官细胞图像识别:PyTorch医学图像分类实战指南

发布时间:2026/9/24 19:28:50

资讯中心
01
ARTICLE

19种器官细胞图像识别:PyTorch医学图像分类实战指南

19种器官细胞图像识别:PyTorch医学图像分类实战指南
简介面向医学图像分类任务的中型数据集整合19类器官细胞图像覆盖肾上腺、子宫、甲状腺、食道等类别训练集2100张、测试集500张已按文件夹划分可直接用于CNN分类网络或基于yolov5的分类项目。包体共2000个文件其中1998张png图片构成数据主体附1个py可视化脚本和1个json类别字典文件压缩包大小约260MB。已有176人学习适合医学影像入门实践、算法验证及课程设计。借助附带的show脚本可快速查看各类别图像json文件可帮助对齐类别索引缩短数据预处理周期便于聚焦模型训练与调优。1. 19种器官细胞图像为什么我建议直接拿它当训练基座很多人以为拿到一份划分好的医学图像分类数据集开箱就能直接获得理想准确率但实际体验往往是“数据集越规整翻车的点越隐蔽”。这个 19 种器官细胞图像识别数据集好就好在把三件容易让人头大的事提前解决了按类别分好的文件结构训练、验证、测试一眼能看清类别字典文件把 19 类器官/细胞名字映射好不用自己猜标签图片按文件夹保存训练框架能直接引用。对刚入门医学图像分类的人来说这份数据能把从下载到跑通第一个模型的窗口压到半小时以内对做过几轮医疗影像项目的人来说它更像一块校准数据管线、测试图像识别模型的试验田——换预训练模型、对比增强策略都不必再花时间去清理原始切片。下文不会替你把模型训到 SOTA但会用第一人称方式带你走完“检查数据、接入管线、调参、排错、验证”的完整闭环照做就行。2. 先看清你手里有什么文件夹结构、类别字典文件与数据划分很多同学下完数据集第一步就直接解压开训结果跑到一半发现验证集标签和训练集对不上。我的建议是任何数据集先花 10 分钟把它拆开看清楚再决定怎么用。2.1 目录组织方式为什么说它是 ImageFolder 兼容结构这个数据集的落地方式通常是“主目录 三个子目录 类别子文件夹”。下载解压后进入根目录用一行命令就能看清结构# 在数据集根目录下执行 tree -d -L 2 .正常情况下你会看到类似这样的输出. ├── train │ ├── class_001_brain │ ├── class_002_breast │ ├── class_003_lung │ └── ... ├── val │ ├── class_001_brain │ └── ... └── test ├── class_001_brain └── ...每个类别名我在这里是用占位符写的你打开实际文件夹时看到的会是真正的器官/细胞英文名或代号但这不影响逻辑。这种“每个类别一个文件夹”的结构就是 PyTorch 官方torchvision.datasets.ImageFolder直接能吃进肚子的格式你甚至不需要自己写载入逻辑ImageFolder会按目录名字典序自动生成 class_to_idx。我一般会把这个 tree 的输出重定向到文件里留作档案方便和后面训练脚本里的 class_to_idx 做对照。要是机器上没装 tree用find . -maxdepth 2 -type d | sort也能看个大概。值得注意“划分好的数据【文件夹保存】”这句话的真实含义train/val/test 不是同一个文件夹里的一条记录而是实实在在的三套物理副本。好处是训练时不需要再写数据划分脚本、也不用担心样本泄漏坏处是磁盘占用会变大一个 19 类、每类几千张的显微镜图像数据集解压后可能有几十 GB。如果磁盘吃紧我的做法是先把 test 集打包留档训练过程中只用 train 和 val。注意如果你的目的是把模型和别人的结果对比那 test 集碰都不要碰最好连中间特征都别提出来。所有调参、选模型的工作都只放在 val 上进行最后才允许评测一次 test。2.2 类别字典文件里到底存了啥类别字典文件通常是一个 JSON 或 txt名字可能是 class_dict.json、classes.txt 或 label_map.json。它解决的问题很实际模型最终输出的是一串 0 到 18 的数字你必须知道数字 7 对应哪个器官才能把准确率转换成“医生看得懂的结论”。先看 JSON 版本# 不要直接双击打开命令行看一眼结构 head -n 20 class_dict.json典型内容长这样{ 0: brain, 1: breast, 2: colon, 3: kidney, 4: liver, 5: lung, ... }有些数据集会反过来把类别名当 key、索引当 value这都不影响使用。关键是你要在训练开始前把这份映射加载进内存并打印出来核对一次。import json with open(class_dict.json, r, encodingutf-8) as f: class_dict json.load(f) # 把 {索引: 类名} 转成 {类名: 索引}后面会给 Dataset 用 idx_to_name {int(k): v for k, v in class_dict.items()} name_to_idx {v: k for k, v in idx_to_name.items()} print(f共 {len(idx_to_name)} 个类别) for i in range(min(5, len(idx_to_name))): print(i, -, idx_to_name[i])这段代码的意义不只是转格式。很多数据集在传播过程中会被人为替换过文件夹名或者有人重新排过删除过某些类导致字典文件和实际目录顺序对不上。打印前五个只是为了快速嗅探异常如果 0 对应的名字很怪或者 name_to_idx 的长度不等于 19说明这份字典文件和你下载的文件夹版本不一致后面所有指标都不可信。2.3 拿到数据集后先做的三件事校验、统计、可视化不要急着开训先花十几分钟做三件事能省掉后面反复重启实验的麻烦。第一步校验图片能不能被正常读取。医学图像数据在打包传输过程中偶尔会混入损坏文件训练时遇到一张损坏图会直接让 DataLoader 报错中断半夜跑实验尤其难受。import os from PIL import Image root datasets for split in [train, val, test]: bad 0 for cls in sorted(os.listdir(os.path.join(root, split))): cls_dir os.path.join(root, split, cls) for img_name in sorted(os.listdir(cls_dir)): img_path os.path.join(cls_dir, img_name) try: with Image.open(img_path) as img: img.verify() except Exception: bad 1 print(损坏:, img_path) print(f{split} 损坏文件数: {bad})verify()不会把整张图解码进内存只检查文件头与整体完整性速度很快。如果有坏文件直接删掉或者单独移到一个_corrupted文件夹里不要在训练脚本里写“跳过损坏图片”的临时逻辑那样会让后续调试更混乱。第二步统计每个类别的样本数量。这一步的产出能直接指导你决定损失函数和评估指标19 类医学图像里“某类样本特别少”几乎是常态。import os import matplotlib.pyplot as plt counts {} for split in [train, val, test]: row {} for cls in sorted(os.listdir(os.path.join(root, split))): cls_dir os.path.join(root, split, cls) n len([f for f in os.listdir(cls_dir) if not f.startswith(.)]) row[cls] n counts[split] row # 只看训练集 fig, ax plt.subplots(figsize(12, 5)) names list(counts[train].keys()) nums [counts[train][n] for n in names] ax.bar(range(len(names)), nums) ax.set_xticks(range(len(names))) ax.set_xticklabels(names, rotation90) plt.tight_layout() plt.savefig(train_dist.png, dpi150)如果你看到某类样本数是其他类的 10 倍以上那“准确率”这个指标基本不可信后面我会讲怎么处理。第三步是随机抽样可视化人为确认标签没有错乱。这一步对医学图像特别值得细胞图像有时长得极其相似一张错标的训练图就能把模型拉偏。from PIL import Image classes sorted(os.listdir(datasets/train)) grid_cols len(classes) grid_rows 4 thumb 64 canvas Image.new(RGB, (grid_cols * thumb, grid_rows * thumb), white) for col, cls in enumerate(classes): cls_dir datasets/train/ cls images sorted(os.listdir(cls_dir))[:grid_rows] for row, img_name in enumerate(images): img Image.open(os.path.join(cls_dir, img_name)).resize((thumb, thumb)) canvas.paste(img, (col * thumb, row * thumb)) canvas.save(grid_preview.png)这张网格图的预览方式很直接纵向是每类的前 4 张图横向跨 19 类。肉眼扫一遍重点看有没有跟同类不搭边的图混进来。如果某类出现的是器官组织切片图照片中间突然混进一张带标尺、带文字注记的截图那这类数据就需要特殊处理。四张图全异常时我会直接把这个类从数据集中剔除比让模型硬学正确得多。3. 用 PyTorch 把数据集接入训练管线从自定义 Dataset 到一键训练做完第 2 章的检查数据本身已经可信了。这一章我直接用 PyTorch 写一套最小但“可迁移”的训练链路这里的代码你改改路径就能用。3.1 自定义 Dataset 还是直接用 ImageFolder既然文件夹结构已经是 ImageFolder 兼容格式最快的方式是直接from torchvision import datasets train_dataset datasets.ImageFolder( rootdatasets/train, )要不要为这个数据集写自定义的 Dataset 类我的判断标准是只做分类、不需要额外的 CSV 信息就没有必要自己写ImageFolder 足够。但如果你是以下情况之一自定义 Dataset 才更合适你需要在__getitem__里同时返回图像、标签、文件名用来做错题分析你要对每张图做特殊预处理比如病理切片的背景裁切、染色归一化你想把类别字典文件的权威映射直接传给 Dataset而不是依赖 ImageFolder 自己按目录排序生成。考虑到这个数据集自带类别字典文件我倾向于用自定义 Dataset把字典文件的读取和训练数据强绑定。这样以后换人维护脚本也不会因为字典重新对齐而出错。import os import json from PIL import Image from torch.utils.data import Dataset class OrganCellDataset(Dataset): def __init__(self, root, class_json, transformNone): self.root root self.transform transform # 读取类别字典这里按 {索引: 类名} 解析 with open(class_json, r, encodingutf-8) as f: raw json.load(f) self.name_to_idx {v: int(k) for k, v in raw.items()} self.samples [] for cls_name in sorted(os.listdir(root)): cls_dir os.path.join(root, cls_name) if not os.path.isdir(cls_dir): continue label self.name_to_idx[cls_name] for img_name in sorted(os.listdir(cls_dir)): if img_name.lower().endswith((.jpg, .jpeg, .png, .tif)): self.samples.append((os.path.join(cls_dir, img_name), label)) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] image Image.open(path).convert(RGB) if self.transform: image self.transform(image) return image, label这段代码有几个值得细说的设计点。第一索引转换用的是int(k)因为 JSON 的 key 是字符串不能直接拿去做训练标签。第二构建 self.samples 时用了 sorted 来保证任何机器上的读取顺序一致这样断点续训、多卡采样时结果可复现。第三把所有图片路径一次性列进内存而不是在__getitem__里临时 listdir一旦开始训练每次取样本只是一次下标访问不会有 IO 抖动。3.2 数据增强哪些能用哪些要小心医学图像分类最反直觉的一点是通用数据增强的某些手段在自然图像上很好用在细胞图像上反而会伤指标。我的默认配置是from torchvision import transforms train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.1, contrast0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) val_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])RandomResizedCrop 的 scale 下限我放在 0.7 而不是默认的 0.08。因为细胞图像中判别信息往往占据画面很大比例裁太狠会让模型学到“局部纹理骗局”而不是器官结构。水平和垂直翻转对细胞图像通常安全因为细胞没有明确的方向性——除非你的数据集是带组织的切片有解剖方向约束才需要把翻转关掉。ColorJitter 的幅度我收得很小亮度扰动 0.1、对比度 0.1。原因在后面的避坑章节细说简单说就是染色方案差异带来的颜色偏移和自然界的光照变化不是一回事你用自然图像的亮度增强尺度去模拟容易把模型教坏。3.3 最小可运行的训练脚本完整代码与参数说明这里给一个能在单卡上直接跑的分类训练脚本核心片段。省去断点续训和精排日志保留最关键的控制流import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import models # 1. 构建数据 train_ds OrganCellDataset(datasets/train, class_dict.json, train_transform) val_ds OrganCellDataset(datasets/val, class_dict.json, val_transform) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue) # 2. 载入预训练模型替换分类头 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) num_classes len(train_ds.name_to_idx) model.fc nn.Linear(model.fc.in_features, num_classes) # 3. 损失与优化器 criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30) # 4. 验证函数 def evaluate(model, loader, device): model.eval() correct total 0 with torch.no_grad(): for images, labels in loader: images, labels images.to(device), labels.to(device) preds model(images).argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return correct / total # 5. 训练循环 device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) best_acc 0.0 for epoch in range(30): model.train() total_loss, correct, total 0.0, 0, 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) preds logits.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) train_acc correct / total val_acc evaluate(model, val_loader, device) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pt) print(fepoch {epoch:02d} | loss {total_loss / total:.4f} f| train_acc {train_acc:.4f} | val_acc {val_acc:.4f})逻辑说明model.fc 是 ResNet18 最后一层全连接原输出是 1000 类这里替换成我们的 19 类weights... 表示用 ImageNet 预训练权重而不是随机初始化。AdamW 的 weight_decay 打在权重上而不是像 Adam 那样打在梯度上对医学图像这种小数据集能明显压过拟合。30 个 epoch 配 CosineAnnealingLR 是多数病理图像任务的合理起步配置——前几个 epoch 用相对大的步长找方向后面逐步降到接近 0不容易训出震荡。参数说明batch size 32 是对 224 分辨率、显存 8GB 级别的显卡设的安全值如果你的卡显存只有 6GB优先降到 16而不是减小图片分辨率——分辨率对医学图像的细节判别至关重要。num_workers 设成 4 通常够用Windows 下如果报和线程相关的错把它设成 0。pin_memoryTrue 在你有 GPU 时能少一点 CPU 到 GPU 的拷贝时间纯 CPU 训练时不需要。每次 epoch 结束时保存的是在 val 上表现最好的权重而不是最后一个 epoch 的权重这个习惯能避免误把过拟合最后的模型留下来。4. 19 类细胞分类的训练参数与调优从迁移学习到类别不平衡代码能跑只是及格训练不出好的验证指标才是真正的问题。这一章讲我在这类医学图像分类数据集上实际调参的顺序和理由照着这个顺序做通常会少走几轮实验。4.1 预训练模型怎么选ResNet 起步EfficientNet 提速在这个 19 类器官细胞数据集上我一般不直接上 ViT 或最新的大模型。原因不是模型不够强而是医学图像数据集大多是“类别多、单类样本有限”Transformer 一类模型没有归纳偏置小数据上要配很重的增强和更长训练时间才追得上 CNN。起步阶段我会用 ResNet50然后看指标表现的瓶颈在哪。一张简单的对照表是我自己做类似任务时的经验感受不是精密 benchmark模型大致参数量显存占用小数据集表现备注ResNet1811M~1.5GB快容易欠拟合拿来验证管线ResNet5025M~3.5GB稳首选默认选择EfficientNet-B312M~2GB好但训练更慢需要单独调学习率ViT-B/1686M~6GB不稳容易欠拟合数据量不够不建议我的选择逻辑先用 ResNet50 做一轮 30 epoch 的基线实验把 val 的 F1 记录下来。如果训练准确率远高于验证说明过拟合向 EfficientNet 或加正则方向调如果训练准确率都不高说明欠拟合再考虑更大模型或更长训练。建议不要一上来就并行跑四五个模型医学图像的训练成本和数据清洗成本都不低一次跑一个、改一个变量反而更快。输入尺寸方面224x224 是通用配置但如果显卡允许我强烈建议试一次 384 分辨率。细胞形态细节如核浆比、染色纹理在 224 下会被池化层抹掉相当一部分信息。我踩过的最明显例子是同一套 ResNet50从 224 换到 384F1 提升接近 4 个点训练时间只多了不到 50%。对医学图像分辨率红利往往比模型换大一档更值。4.2 学习率、batch size 与类别不平衡的三个联动参数这一节不单独讲参数因为这三个参数是互相牵扯的改一个其他两个必须跟着动。学习率迁移学习场景我不直接用 ImageNet 常用的 1e-3 起步而是用 1e-4 起步预热 3~5 个 epoch 后再切到正式学习率。原因是预训练特征已经很强太大的初始学习率会破坏早期特征尤其当数据集和 ImageNet 的图像分布差异较大时损失会在第 1 个 epoch 直接爆炸。如果你的 GPU 显存只允许 batch size 16学习率也要相应减半经验上 batch size 翻倍学习率大概乘 1.8~2 是安全的。类别不平衡如果统计时发现某类别的样本数特别少第一个要改的并不是损失函数而是采样器。先让模型在每个 batch 都能稳定看到少数类样本from torch.utils.data import WeightedRandomSampler labels [s[1] for s in train_ds.samples] class_counts [labels.count(c) for c in range(num_classes)] weights [1.0 / class_counts[l] for l in labels] sampler WeightedRandomSampler(weights, num_sampleslen(labels), replacementTrue) train_loader DataLoader(train_ds, batch_size32, samplersampler)WeightedRandomSampler 会让少数类被抽到的概率按权重放大但副作用是训练时同一张少数类图像可能多次进入不同 batch模型容易在这几类上过拟合。所以我的习惯是采样器加轻微的损失函数加权一起用而不是寄托于其中任何一个。比如nn.CrossEntropyLoss(weighttorch.tensor([...]))权重取1/sqrt(count)而不是1/count这个折中在医学图像上通常更稳。1/count 会把权重拉到极端值让模型对少数类输出过于激进验证集上反而翻车。除了这两个手段还可以对少数类做在线增强补偿训练时对同一张少数类图像做多个随机裁剪放进同一批。这一招在细胞图像上很有效因为细胞的旋转、平移不变性本质上允许你无中生有地造出合理样本等于把这个类别的采样频率又提起来一截。4.3 评估指标19 类任务里 accuracy 是最没用的那个医学图像分类的评价我从来不只看一句准确率。19 类任务里如果某几类样本特别多模型全猜那些类也能拿到一个体面的准确率但这个模型到了真实场景就是废的。至少要打印三样东西from sklearn.metrics import classification_report, balanced_accuracy_score # 假设已经在验证集上收集好了 y_true 和 y_pred print(classification_report(y_true, y_pred, target_nameslist(idx_to_name.values()))) print(balanced acc:, balanced_accuracy_score(y_true, y_pred))classification_report 会按类给 precision、recall、f1-score一眼能看到哪些器官类别在混淆。balanced_accuracy_score 是简单准确率的类别平均版本样本不平衡时它比 accuracy 诚实得多。此外我还会把混淆矩阵保存下来做错题分析这一步对医学图像特别重要两个器官大概率互相误判可能是图像本身相似比如乳腺和甲状腺在某些染色下有相似形态也可能是标签错误这两者的处理方向完全不同。5. 医学图像分类常见问题与避坑染色差异、过拟合与脏数据这一章集中写我在这类数据集上遇到过的四个高频坑按“现象 → 原因 → 解决”的口径总结。强烈建议把这些检查点写进核对清单每次换数据集、换模型都过一遍。5.1 现象训练损失一路下降验证指标却剧烈抖动在 19 类细胞数据集上训练时常见情况是 train loss 从 1.8 稳定降到 0.3但 val 的 F1 在 0.6 到 0.75 之间来回跳有时甚至比 epoch 15 还低。原因主要是两个一是医学图像数据集的 val 集太小部分类别只有几十张随机抽样的一个 batch 内如果混入一两张分布外的图像单 epoch 指标立马波动二是模型已经过拟合学到了训练集里和类别无关的噪声纹理验证时碰到稍微偏离分布的图就开始瞎猜。解决步骤我通常按顺序试先把 val 的评估改成“全量验证 重复 3 次取平均”的机制排除抽样噪声对结论的干扰然后在训练循环里加早停与最优权重保存——只认验证指标最好的 checkpoint不认最后一个 epoch最后回头看增强配置里的 ColorJitter 是否过强细胞图像的染色颜色是重要诊断特征你把对比度扰动开到 0.5模型当然会学到不稳定的颜色映射。我最终会把亮度和对比度各降到 0.1 以下并在推理阶段完全关闭增强。5.2 现象换了一批染色切片后模型准确率掉了十几个点这是医学图像经典“玄学”训练时模型学到的颜色分布来自数据集的染色方案测试时换一个不同染色强度、不同批次的切片你会发现所有类别置信度崩塌尤其是深染的细胞。原因本质是模型把染色颜色当成了辨别依据而不是细胞形态。解决可以分几层做。最省事的是推理和训练用同一种染色方案或统一转成灰度再喂给模型但这会损失诊断细节。常见做法是做染色归一化把每张图的颜色分布对齐到模板切片上。我一般用简单的 Reinhard 颜色归一化起步不需要引入其他大库已经能缓解大部分问题import numpy as np from PIL import Image def reinhard_normalize(image, target_mean, target_std): 将当前图像的 LAB 颜色统计对齐到目标统计。 target_mean / target_std 取自一张你认为染色标准的切片图。 im image.convert(LAB) arr np.array(im, dtypenp.float32) / 255.0 for c in range(3): mean arr[..., c].mean() std arr[..., c].std() 1e-8 arr[..., c] (arr[..., c] - mean) / std * target_std[c] target_mean[c] return Image.fromarray((np.clip(arr, 0, 1) * 255).astype(np.uint8), LAB).convert(RGB)target_mean 和 target_std 怎么定我的做法是从训练集里挑一张视觉上染色最均衡、细胞结构最清晰的图用同样代码算出它的 LAB 均值和标准差当全局模板。之后训练、验证、测试三套流程全部先过这个函数再进网络。注意如果只在训练集上做了归一化而测试集没做模型面对的反而是一个被扭曲过的分布掉点会更快。5.3 现象类别字典文件与文件夹名对不上训练时报“key not found”这个坑我吃过两次亏。一次是数据集的文件夹名是英文全名字典文件里用的却是缩写另一次是字典里有 19 个条目但实际文件夹只有 18 个某类样本被单独放在_others里。训练都能启动但标签错位从头到尾没人发现。原因并不一定是数据集的制作者有问题更常见的是传播过程中有人重新整理过目录、改过文件夹名但忘了同步更新字典文件。解决在真正训练前做一次全量交叉校验核心是“文件夹实际类名集合”和“字典文件类名集合”完全一致import json, os root datasets/train with open(class_dict.json, r, encodingutf-8) as f: raw json.load(f) folder_names {d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))} dict_names set(raw.values()) print(文件夹有而字典没有:, folder_names - dict_names) print(字典有而文件夹没有:, dict_names - folder_names) assert folder_names dict_names, 类别字典和文件夹不一致禁止训练这样的校验逻辑应该放在训练脚本的最开头而不是等到某张图的标签跑飞了才回头查。这种“错位”类 bug 一旦发生模型会像黑匣子一样输出看似合理但全部偏移的结果人工几乎不可能通过训练日志发现。5.4 现象某类在 val 里只有 20 张指标全被它拉垮19 类数据集的 val 划分如果用的是简单随机抽样少数类很可能在 val 里只有极少量代表甚至某个 epoch 的验证集合里缺了其中一两个类。这会带来两个问题指标波动大而且你没办法判断模型在这类上是有效还是盲猜。解决手段有两步。第一步是分层抽样确保每个类别在 val 中都有固定配额例如每类至少留 10% 或 30 张进 val如果数据集没有这样做我会自己重新划分一份 val把原始 val 当作 extra test。第二步是对少数类单独看置信度分布如果模型对某类的平均置信度低于 0.4说明这一类的判别特征根本没学会此时与其调参数不如回去检查训练集该类图像是否存在质量问题。不要忽略小样本类19 类里只要有一类学废了放到真实病理场景就是漏诊级别的问题。6. 用 Grad-CAM 检查模型到底看了哪里细胞分类上线前最后一道验证训练结束模型的准确率、F1 都做得不错这时我不急着上线而是先做一轮模型行为审查核心工具是 Grad-CAM。细胞分类在医生眼里必须“可解释”你的模型凭什么说这是某器官的细胞它看的是细胞核、细胞质还是切片边缘的脏痕后者本质是荒谬的捷径必须尽早发现。Grad-CAM 的原理一句话能讲清把最后一层卷积特征图按其对目标类别的梯度加权求和得到热力图。热力图越亮的区域就是模型做决策时依赖的图像区域。前面训练时保存的 best_model.pt 此时真正派上用场import torch from torchvision import models, transforms from PIL import Image model models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1) model.fc torch.nn.Linear(model.fc.in_features, 19) model.load_state_dict(torch.load(best_model.pt, map_locationcpu)) model.eval() img Image.open(sample.jpg).convert(RGB) t transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) x t(img).unsqueeze(0) # 钩子取最后一层卷积的输出和梯度 activation {} gradient {} def forward_hook(module, input, output): activation[value] output.detach() def backward_hook(module, grad_input, grad_output): gradient[value] grad_output[0].detach() h1 model.layer4[-1].register_forward_hook(forward_hook) h2 model.layer4[-1].register_full_backward_hook(backward_hook) out model(x) one_hot torch.zeros_like(out) one_hot[0, out.argmax(dim1).item()] 1 model.zero_grad() out.backward(one_hot) h1.remove() h2.remove() # 梯度全局平均池化后与特征图加权得到热力图 weights gradient[value].mean(dim(2, 3), keepdimTrue) cam (weights * activation[value]).sum(dim1, keepdimTrue) cam torch.relu(cam) import torch.nn.functional as F cam F.interpolate(cam, size(224, 224), modebilinear, align_cornersFalse) cam cam.squeeze().numpy() cam (cam - cam.min()) / (cam.max() - cam.min() 1e-8) # 到这里 cam 就是一个 0~1 的二维矩阵把 cam 叠加到原图上即可这里选的是模型预测概率最高的类作为目标做出来热力图就是“模型为什么把它判成这一类”的依据。我现在的习惯是每次训练完一批模型先跑 50 张训练集干净样本、再跑 50 张外部样本的 Grad-CAM 对比然后才决定是否进入下一轮调参。如果高亮区域集中在切片边缘、折叠处、气泡和其他伪影上那这个模型只是“在数据集上表现良好”真实场景不可用直接回到数据处理重新来。用 Grad-CAM 检查模型是一项不花多少时间却回报极高的习惯它相当于给模型做一次体检把黑匣子打开看一眼。数据检查、参数调整、行为验证三步做完你基本就能安心把模型用起来或继续往下迭代希望帮到你。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。