拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

瞳孔虹膜分割实战:从UNet训练到推理避坑全指南

瞳孔虹膜分割实战:从UNet训练到推理避坑全指南 简介这是一份面向计算机视觉学习者的瞳孔虹膜分割数据集包含394张训练图片与112张测试图片分辨率统一为640×640mask标签以0、1、2灰度值分别对应背景、瞳孔与虹膜区域适合用于语义分割模型的训练与效果验证。包内训练集与测试集均按images和masks目录组织另附一个Python可视化脚本可随机抽取图片并同时展示原图、GT掩膜及蒙板叠加效果便于快速检查标签质量。资源共1014个文件包括506张jpg原图、507张png掩膜及1个py脚本压缩包大小20.72MB结构简洁、上手门槛低。目前已有266人学习使用适合刚接触分割任务的学生或研究者作为练手数据也方便在现有模型上进行迁移测试。1. 图像分割数据集选型瞳孔虹膜分割的训练集与测试集值不值得下做医学图像分割的人都有体会找数据集比调模型更耗时。这个瞳孔虹膜分割数据集把前后端都备齐了——训练集 394 张图带 394 张 mask测试集 112 张带 112 张 mask分辨率统一 640×640mask 是 0/1/2 三值灰度图。瞳孔和虹膜是典型的小目标加弱边界场景394 张训练量不算大配合数据增强跑 UNet 完全够用。它适合两类人刚入门图像分割、想找一份干净带标注数据练手的学生做眼动追踪、虹膜识别、瞳孔检测预研的工程师。这份资源帮你省掉最脏的数据清洗和标注环节拿到手直接面对模型训练。需要说明的是这里交付的是数据集加一个可视化脚本不是训练好的模型权重模型要自己跑。2. 数据集结构与标注规范394112 张图的目录组织和 0/1/2 灰度语义拿到压缩包先别急着解压就跑模型把目录结构和标注语义摸清楚后面能省掉大量排错时间。这一章从文件命名、mask 灰度值到可视化脚本逐个拆全是实际操作层面的细节新手照着做不会跑偏熟手也能从中确认这份数据的边界条件。2.1 目录组织与文件命名Roboflow 导出痕迹和同名映射解压后根目录下一般是 train 和 test 两个大目录每个目录里再分 images 和 masks。train/images 下 394 张 jpgtrain/masks 下 394 张对应 masktest 同理是 112 对。先别急着改目录结构PyTorch 的 Dataset 类里直接用相对路径拼接最省事后面如果要换框架也只要保证 images 和 masks 两个根路径不变迁移成本很低。文件命名一眼能看出来源比如5f88420adb41b5d5_jpg.rf.b2feee328e6a5ba8ee09bf83ecc6b975.jpg。中间的.rf.是 Roboflow 平台导出的标志后面的十六进制串是该样本在导出时的唯一标识_jpg后缀表示原始文件是 jpg 格式。这个信息不是没用的以后要合并其他数据集这个 hash 能帮你做去重如果哪天发现图片和 mask 对不上号先检查文件名里.rf.前面那段是否一致而不是依赖文件夹位置。mask 文件的扩展名以压缩包里的实际为准常见的是与图片同名的 .png。这里有个通用原则mask 一律用 png 或 tif 这类无损格式不要用 jpg。jpg 有损压缩会在类别边界产生伪影把 0/1/2 这种索引值压成 0.7、1.3训练时 one-hot 编码直接错位。这种错误特别难查——模型还在跑指标只是悄悄变差。我见过有人在这上面耗了两天最后发现是数据读取时把 jpg mask 当灰度图用了。注意训练和评估脚本里 mask 的扩展名替换规则必须完全一致解压后先ls看一眼实际扩展名再写进 Dataset 代码。2.2 mask 灰度语义0、1、2 分别代表什么mask 是单通道灰度图三个灰度值对应三个类别按 0/1/2 的顺序通常 0 是背景1 是瞳孔2 是虹膜具体以压缩包内说明为准。这里用的不是把 255 当前景的约定而是明确的索引语义值域只有 0/1/2天然适合直接作为分类任务的 target不需要额外二值化。读 mask 我一般用 OpenCV 的 imread 加IMREAD_UNCHANGED参数保证按原始位深读进来import cv2 img cv2.imread(train/images/5f88420adb41b5d5_jpg.rf.b2feee328e6a5ba8ee09bf83ecc6b975.jpg) mask cv2.imread(train/masks/5f88420adb41b5d5_jpg.rf.b2feee328e6a5ba8ee09bf83ecc6b975.png, cv2.IMREAD_UNCHANGED) print(mask.shape, mask.dtype) # (640, 640) uint8 print(set(mask.ravel().tolist())) # {0, 1, 2}这里有两个细节。第一mask 必须是单通道shape 是 (640, 640) 而不是 (640, 640, 3)如果读出来是三维说明保存时被转成了 RGB要先取单通道再核对灰度值集合。第二dtype 应该是 uint8值域就是 {0, 1, 2}。如果set()之后看到 {0, 255}说明 mask 被 jpg 压缩过或保存时做了二值化这份数据在这个环节就脏了得换文件。为什么用 0/1/2 三值图而不是三张二值 mask直观原因是省空间1 张 uint8 图只有 3 张二值图的四分之一更深层的原因是训练时交叉熵损失天然接受类别索引PyTorch 的CrossEntropyLoss直接吃[B, H, W]的整数 target不需要转 one-hot。只有用 Dice Loss 自己写 one-hot 时才要额外转换这一步后面讲。2.3 可视化脚本GT 与原始图像对齐的第一道关卡数据集附带的可视化脚本功能是随机抽一张图把原始图片、GT mask、GT 叠加在原图上的结果并排展示并保存到当前目录。这个脚本的价值被很多人低估——拿到数据集第一件事就该跑它而不是直接开训练。我一般会先看十几张叠加图确认三件事mask 与眼睛图像是否对齐、瞳孔和虹膜的边界是否干净、有没有样本标错类别。脚本核心逻辑和下面的写法等价核心是利用颜色映射把索引值变成可视化颜色import cv2 import numpy as np import glob, random, os img_path random.choice(glob.glob(train/images/*.jpg)) mask_path os.path.join(train/masks, os.path.basename(img_path).replace(.jpg, .png)) img cv2.imread(img_path) # BGR mask cv2.imread(mask_path, cv2.IMREAD_UNCHANGED) # 0/1/2 color_map np.array([[0, 0, 0], [255, 0, 0], [0, 0, 255]], dtypenp.uint8) overlay color_map[mask] # 索引查表 blended cv2.addWeighted(img, 0.6, overlay, 0.4, 0) # 原图与伪彩叠加 canvas np.hstack([img, np.stack([mask] * 3, axis-1), blended]) cv2.imwrite(check_visual.png, canvas)color_map[mask]这一步是 numpy 的索引查表mask 中等于 1 的位置被替换成color_map[1]这行颜色等于 0 的位置替换成黑色等于 2 的位置替换成红色一步完成类别到颜色的映射比循环遍历像素快几个数量级。addWeighted的两个权重 0.6 和 0.4 控制原图和伪彩的透明度太透明看不清边界太不透明又看不出类别我一般保持 0.6/0.4 不动。hstack把三张图拼成一行方便对比最后imwrite保存到当前目录。这个脚本唯一要改的就是 mask 扩展名的替换规则以目录里实际扩展名为准。虽然不是训练代码但它承担了数据质检的职责——每次拿到新数据集我都先跑一遍花五分钟确认标注质量能省掉后面调模型的半天时间。3. 用 UNet 训练瞳孔虹膜分割数据加载、Dice Loss 与训练参数数据和标注确认没问题之后进入训练环节。瞳孔虹膜分割这种二类目标任务最稳的基线结构就是 UNet——编码器下采样四次提取语义解码器上采样四次恢复细节跳跃连接保住边缘信息在 640×640 输入下参数量和显存占用都可控。这一章给出完整的数据加载、损失函数和训练配置全部围绕这份数据集的实际特点来写。3.1 数据加载为什么必须保持 mask 的索引语义写 Dataset 类时最关键的坑就是 mask 读取。很多人习惯用 PIL 读图但 PIL 的Image.open对灰度图有个隐蔽行为如果图片保存时用了调色板模式P 模式np.array之后得到的是调色板索引而不是灰度值值域可能完全不是 0/1/2。所以我的习惯是统一用cv2.imread(..., IMREAD_UNCHANGED)从源头避开这个问题。import os, glob import cv2 import numpy as np import torch from torch.utils.data import Dataset class IrisSegDataset(Dataset): def __init__(self, img_dir, mask_dir, trainTrue): self.img_paths sorted(glob.glob(os.path.join(img_dir, *.jpg))) self.mask_dir mask_dir self.train train def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img_path self.img_paths[idx] name os.path.basename(img_path).replace(.jpg, .png) mask_path os.path.join(self.mask_dir, name) img cv2.imread(img_path) # BGR uint8 img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # 转 RGB mask cv2.imread(mask_path, cv2.IMREAD_UNCHANGED) if self.train: # 随机水平翻转图片和 mask 必须共用同一个随机种子 if np.random.rand() 0.5: img cv2.flip(img, 1) mask cv2.flip(mask, 1) img img.astype(np.float32) / 127.5 - 1.0 # 归一化到 [-1,1] mask mask.astype(np.int64) # 交叉熵需要 long img torch.from_numpy(img).permute(2, 0, 1).float() # HWC - CHW mask torch.from_numpy(mask) # [H, W] return img, mask几个要点拆开说。第一mask 用astype(np.int64)而不是 float因为CrossEntropyLoss要求 target 是 LongTensorfloat 会直接报错。第二图像和 mask 做随机翻转时必须共享随机数各翻各的会让标注和内容错位这是新手最容易翻车的地方。第三归一化用/127.5 - 1把 RGB 压到 [-1, 1]配合默认初始化和 Adam 在医学图像上收敛比 [0, 1] 稍快一点虽然不是必须但我习惯这么写。图像增强这里我刻意没加颜色抖动。人眼图像的颜色分布比较统一对亮度敏感但对色相不敏感强加 HSV 扰动反而让网络学到错误的颜色不变性。要增强就做几何类翻转、旋转 ±15°、缩放 0.9~1.1这些对瞳孔这种近圆形目标非常友好。常见做法是直接用 Albumentations 库随机翻转、旋转、缩放一步到位它在内部帮你同步 img 和 mask但我上面故意用原生 cv2 写是为了让你看清这个同步的关键点。3.2 损失函数为什么单用交叉熵会偏向背景在 640×640 图像里瞳孔直径一般也就几十个像素面积占比可能只有 2%~5%虹膜稍大但也只占 10% 左右。这意味着背景像素占了 85% 以上如果直接用交叉熵网络只要学会输出全背景就能拿到 0.85 的准确率loss 看起来在降实际什么都没学到。解决的办法是 Dice Loss或者交叉熵和 Dice 加权组合。Dice 对类别不平衡不敏感因为它按区域重叠率算不受像素数量主导。我一般用 CE Dice 各一半权重收敛速度比纯 Dice 快边界也比纯 CE 干净。import torch import torch.nn.functional as F def dice_loss(pred, target, eps1e-6): # pred: [B, C, H, W]已经过 softmax # target: [B, H, W]值域 {0, 1, 2} num_classes pred.shape[1] target_onehot F.one_hot(target, num_classes).permute(0, 3, 1, 2).float() intersection (pred * target_onehot).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) dice (2 * intersection eps) / (union eps) # [B, C] return 1 - dice.mean()这里one_hot是核心F.one_hot把 [B, H, W] 的索引变成 [B, H, W, C]再用permute转成 [B, C, H, W] 才能和 pred 逐元素相乘。eps 加在分母上防止某个类别在 batch 里完全缺失时除零。最后dice.mean()是对所有类别取平均等价于 mDice比单独算每个类别再加权更稳。如果发现虹膜总是被吞可以把返回改成加权平均比如按像素占比的倒数给瞳孔更大权重效果直接反映在指标上。3.3 训练参数batch size、学习率与总轮数基于 640×640 输入和 UNet 基础结构我给的基线配置如下单张 24G 显存的卡能跑显存小的话输入缩到 512 也行但 640 是数据集的原始分辨率建议优先保持。参数值说明输入尺寸640×640保持原始分辨率不额外下采样batch size824G 显存刚好16 需要 32G优化器Adamlr1e-4weight_decay1e-5学习率1e-4cosine 衰减到 1e-6损失函数0.5×CE 0.5×Dice中和两者的偏好训练轮数120394 张图约 5900 次迭代增强策略翻转/旋转/缩放不加颜色抖动训练循环本身是标准的关键是每个 epoch 结束要在测试集上算 mIoU 和 mDice并且保存最优权重而不是最后一轮best_miou 0.0 for epoch in range(120): train_one_epoch(model, train_loader, optimizer, criterion) miou, mdice evaluate(model, test_loader, num_classes3) if miou best_miou: best_miou miou torch.save(model.state_dict(), best_iris.pth) print(fepoch {epoch} miou{miou:.4f} mdice{mdice:.4f})这里我特别强调一个容易被忽略的问题394 张训练图不足以支撑随机划分验证集因为人眼图像里同一个人的双眼可能非常相似随机划分会把同源样本分到两边验证指标虚高。这份数据既然给了独立的测试集就应该固定用测试集做最终评估不要把测试集混进训练更不要拿它做早停——早停应该在真正的验证集上做没有独立验证集就老老实实训练固定轮数测试集只做最终报告。这和 YOLO 训练自己的数据集时按比例分 val 的做法不一样分割小数据集的样本独立性更弱划分必须更保守。4. 瞳孔虹膜分割常见问题避坑五个真实翻车现场这一章写的都是实际跑分割数据集时摔过的跟头每条按现象、原因、解决三步给可以直接对照自己的报错信息和指标表现。分割任务不像检测任务能一眼看到框有没有画对mask 的错位、值域混淆、类别丢失都藏在指标和可视化里不排查到根源调参就是瞎调。4.1 现象训练 loss 下降正常测试集 mIoU 却卡在 0.3 不动原因不是模型问题是归一化不一致。有些眼睛图像有强烈的镜面反光角膜上的亮点像素值整体偏亮如果训练时每张图单独做 mean-std 归一化评估时用的却是全局统计量分布就错位了。更常见的是 mask 读取方式不同——训练代码用IMREAD_UNCHANGED读到 0/1/2评估代码顺手用默认imread读成 0/255 或三通道图类别错乱mIoU 直接崩掉。解决抽几对训练和测试图片分别打印像素值分布确认归一化路径完全一致mask 统一用IMREAD_UNCHANGED并在评估脚本里加一行断言assert set(np.unique(mask)) {0, 1, 2}。镜面反光的问题我一般加一步 CLAHE 对比度增强把局部高光压一压mIoU 常有 2~3 个点的提升。4.2 现象mask 读出来全是 255可视化一片白原因这份数据的 mask 是 0/1/2 索引图直接用看图软件打开是黑的因为 0/1/2 在 8bit 里几乎等于 0。有人想转成可见的拿 PIL 打开后转成 RGB 或 L 模式再保存保存时 1 和 2 被拉伸成 255索引语义彻底被破坏。另一种情况是用了imread的默认模式把 16 位深度的图压成 8 位值全变了。解决解码时永远用cv2.IMREAD_UNCHANGED不经过任何中间转换要保存副本就用 png 格式保存前先断言唯一值集合是 {0, 1, 2}。我用过一段自定义 load 函数做逐张检查后来发现开销太大改成数据集初始化时一次性校验全部 mask——初始化通过就代表这批数据是干净的后面训练直接信任它。4.3 现象瞳孔区域太小网络直接学丢预测结果里瞳孔消失原因瞳孔面积占比 2%~5%在 Dice Loss 里它只贡献很小一部分梯度被背景梯度淹没交叉熵更是如此全背景预测已经有 0.85 的准确率网络当然选择偷懒。这在小目标分割里非常典型不是模型结构的问题是损失函数和样本分布在打架。解决两个办法叠加。第一在 loss 里给类别加权重把返回改成1 - (dice[:, 1] * w1 dice[:, 2] * w2)权重按 1/面积占比归一化瞳孔可以给到 8~12。第二数据增强时做随机裁剪放大——把瞳孔区域附近裁出来缩放到 640×640相当于变相增加小目标样本。我用后者效果更明显因为瞳孔的相对大小在预测时更接近真实场景模型见过各种尺度的瞳孔泛化更稳。4.4 现象虹膜区域边缘碎、内部有黑洞mIoU 上不去原因虹膜和瞳孔的边界是弱边缘虹膜自身又有纹理和反光点网络把反光点学成了背景孔。另一个原因是训练时下采样太多UNet 的 encoder 在第 3、4 层把细节抹掉了恢复不到原分辨率。解决推理时用完整 640×640 分辨率不要在预测前 resize 到 256后处理加形态学闭运算第 6 章细说。另外可以尝试在 loss 里加边界项对 mask 求 Sobel 梯度作为边界权重加权到 CE 上。这个技巧对虹膜这种纹理目标很有效但对瞳孔这种光滑目标收益不大——所以先确认问题出在哪个类别再决定要不要上边界损失避免盲目加复杂度。4.5 现象自带可视化脚本运行报 RuntimeError或叠加图全黑原因脚本里如果用了 PIL 读 mask在 P 模式下np.array得到的是调色板索引而不是灰度值叠加时用 mask 直接当 RGB 第三维或乘了错误权重都会黑图。另一个隐蔽点不同 OpenCV 版本对 16 位 png 的返回 dtype 不同老版本返回 uint16新版本可能返回 float32直接当索引用没问题但如果脚本里有类型转换就会炸。解决按 2.3 节的写法重写核心是不经过 PIL、不用调色板、用color_map[mask]查表。Python 环境版本冲突时先升级 opencv-python 到 4.8 以上再跑一遍脚本大多数 RuntimeError 能直接消失。如果还有问题打印mask.dtype和mask.shape对照 2.2 节的检查方法多半是读取模式的问题。5. 推理与后处理从预测 mask 到瞳孔直径和评估指标训练完模型下一步是推理和后处理。这一章讲清楚三件事怎么从模型输出拿到类别 mask怎么从 mask 里提取瞳孔面积和直径这类业务指标以及怎么报测试集指标才能和论文公平对比。每一步都有对应的代码和参数说明。5.1 推理脚本argmax 与类别置信度模型输出是 [1, 3, 640, 640] 的 logits常规做法是 argmax 取类别索引model.eval() with torch.no_grad(): logits model(img_tensor) # [1, 3, H, W] pred logits.argmax(dim1).squeeze(0).cpu().numpy() # [H, W], {0, 1, 2}argmax 天然返回索引不需要先 softmax因为 argmax 在 softmax 前后结果不变省一次计算。如果要做置信度过滤才需要先 softmax 再取最大值。瞳孔分割场景里类别置信度通常用在两个地方一是评估时排除低置信度样本二是做半自动标注时用置信度提示人工复核但推理环节直接用 argmax 就够了。5.2 连通域分析与瞳孔直径估算预测 mask 里通常会有零星噪点直接用会污染面积和直径指标。我一般用 scipy.ndimage 做连通域标记按面积阈值剔除小区域from scipy import ndimage import numpy as np pred_pupil (pred 1).astype(np.uint8) labels, n ndimage.label(pred_pupil) if n 0: print(no pupil detected) else: sizes ndimage.sum(pred_pupil, labels, range(1, n 1)) keep_id int(np.argmax(sizes)) 1 # 取最大连通域 pupil_mask (labels keep_id) area int(sizes[keep_id - 1]) diameter 2 * np.sqrt(area / np.pi) # 等效圆直径ndimage.label默认按 8 连通标记对瞳孔这种实心目标够用。取最大连通域能同时去掉孤立的误检块。等效圆直径假设瞳孔近似圆形这个假设对正常眼成立但对严重变形的瞳孔会有偏差所以报告时建议同时给出 area 和 diameter 两个值让下游自己决定用哪个。对虹膜类别也可以做同样的连通域处理区别是虹膜是环形取最大连通域时可能会漏掉被瞳孔分隔的外环这种情况要先在 mask 上把瞳孔区域填掉再分析代码逻辑里要多一步。5.3 评估指标mIoU 和 Dice 的实际关系测试集报告指标最常用的是 mIoU 和 mDice。按类别分别算再平均就是 mIoU/mDice代码写起来就几行def compute_metrics(pred_all, gt_all, num_classes3): ious, dices [], [] for c in range(num_classes): p (pred_all c); g (gt_all c) inter (p g).sum(); union (p | g).sum() ious.append(inter / max(union, 1)) dices.append(2 * inter / max(p.sum() g.sum(), 1)) return np.mean(ious), np.mean(dices)注意 mIoU 和 mDice 的关系IoU Dice / (2 - Dice)所以 Dice 永远大于等于 IoU两者不要混着跟论文比。报指标时写清楚是 per-class 平均还是忽略背景——瞳孔数据集上忽略背景的 mIoU 会比包含背景的高 5~8 个点这不算造假但要在报告里注明口径否则别人没法复现你的对比。5.4 测试时增强水平翻转 TTA 的稳定收益小数据集上 TTA 是性价比最高的提分手段不用改模型、不用重训练只在推理时多做一次水平翻转并平均概率。瞳孔分割的测试图存在左右对称性翻转后预测结果理论上应该一致但网络对强反光点的响应略有差异平均后能压低噪声def predict_with_tta(model, img_tensor): logits model(img_tensor) # 原始方向 logits_flip model(torch.flip(img_tensor, dims[3])) # 水平翻转 logits_flip torch.flip(logits_flip, dims[3]) # 翻回原方向 prob (torch.softmax(logits, dim1) torch.softmax(logits_flip, dim1)) / 2 return prob.argmax(dim1)dims[3]表示翻转 W 维即水平翻转。flip 输入再 flip 输出保证两张图的空间位置一一对应才能做平均。TTA 在瞳孔分割上一般能带来 0.5~1 个点的 mIoU 增益代价是推理时间翻倍。如果业务对延迟敏感我会把 TTA 关闭只在离线评估和出报告时打开这个取舍要在文档里写清楚。6. 一个提升分割精度的技巧形态学闭运算补虹膜反光破洞虹膜预测结果最常见的缺陷是内部散落着小黑洞——这些洞大多对应角膜镜面反光点GT 里标成虹膜但网络在高亮区域容易犹豫。用形态学闭运算可以低成本补洞闭运算等于先膨胀再腐蚀恰好能填小洞而不明显改变外轮廓。import cv2 kernel cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15)) iris (pred 2).astype(np.uint8) iris_filled cv2.morphologyEx(iris, cv2.MORPH_CLOSE, kernel) pred_filled pred.copy() pred_filled[iris_filled 1] 2核大小 15×15 在 640×640 图上约等于瞳孔直径的三分之一能填掉大多数反光洞又不会把背景和虹膜糊成一片。核太大会把眼白区域也闭进来误提成虹膜核太小则填不掉大反光斑。用椭圆核而不是矩形核是因为虹膜本身是环形结构椭圆核更贴合纹理走向边界更自然。注意不要把闭运算用在瞳孔上——瞳孔要用来算直径闭运算会让边界外扩直径偏大 2~3 个像素。如果一定要对瞳孔做清洗用开运算先腐蚀再膨胀去毛刺它对尺寸的影响比闭运算小得多。更稳的做法是只在评估和报告指标前填洞训练数据保持原样让网络自己去学反光点其实长在虹膜上。从那以后我每次拿到新的分割数据集都先跑一遍可视化脚本确认 GT 边界再定后处理管线而不是一上来就调模型结构。这个顺序帮我少走了很多弯路希望帮到你。本文还有配套的精品资源点击获取
返回列表