欢迎光临
我们一直在努力

Attention-UNet 视网膜眼底血管智能分割系统 基础算法 Attention-UNet(U-Net 主干 + CBAM 注意力模块,医学二值分割)

:基于Attention-UNet的视网膜血管分割系统(医学图像分割,深度学习)在这里插入图片描述

数据集:DRIVE,CHASE-DB1,STARE(三个常用的视网膜血管分割数据集)
深度学习环境:pytorch
分割界面:基于pyqt5开发,展示分割效果。
在这里插入图片描述

✅项目源码+数据集,包含训练好的权重文件,运行GUI.py文件,选择一张图片可以直接分割。
✅训练:运行train.py

在这里插入图片描述

项目功能:包含数据集、图像分割、训练测试、结果展示,一个项目帮助你完成完整分割任务!在这里插入图片描述

基于Attention-UNet视网膜眼底血管分割系统(完整项目+全量可运行代码)

一、项目整体信息表

项目参数详情
项目名称 Attention-UNet视网膜眼底血管智能分割系统
基础算法 Attention-UNet(U-Net主干+CBAM注意力模块,医学二值分割)
数据集 DRIVE + CHASE‑DB1 + STARE 三大公开眼底数据集,眼底灰度图+血管二值掩码(白色血管、黑色背景)
任务类型 二分类语义分割(血管=1,背景=0)
开发环境 Pytorch1.10+、Python3.9、PyQt5、OpenCV、Albumentations
配套内容 完整数据集、训练源码、预训练最优权重best.pth、桌面可视化GUI程序
系统功能 单张眼底图导入、一键血管分割、分割结果实时对比展示、结果图片本地保存、CUDA/GPU加速推理
落地场景 眼科辅助筛查、眼底病变初筛、医学毕设课题、影像算法科研

二、项目目录结构

retina_vessel_seg/
├── dataset/
│ ├── train/img # 训练原图
│ ├── train/mask # 血管掩码标签
│ ├── test/img
│ └── test/mask
├── weights/ # 预训练权重best.pth
├── utils/
│ ├── dataset.py # 数据集加载
│ ├── att_unet.py # Attention-UNet网络
│ └── seg_metric.py # Dice/IoU指标计算
├── train.py # 模型训练入口
├── GUI.py # PyQt可视化软件主程序
└── requirements.txt

三、环境依赖 requirements.txt

torch==1.10.0
torchvision==0.11.0
opencv-python
numpy
pyqt5
albumentations
matplotlib

一键安装:

conda create -n retina python=3.9
conda activate retina
pip install -r requirements.txt

四、utils/dataset.py 数据集加载与数据增强

import os
import cv2
import numpy as np
from torch.utils.data import Dataset
import albumentations as A
from albumentations.pytorch import ToTensorV2

class RetinaDataset(Dataset):
def __init__(self, img_root, mask_root, transform=None):
self.img_root = img_root
self.mask_root = mask_root
self.name_list = sorted(os.listdir(img_root))
self.transform = transform

def __len__(self):
return len(self.name_list)

def __getitem__(self, idx):
name = self.name_list[idx]
img = cv2.imread(os.path.join(self.img_root,name),0)
mask = cv2.imread(os.path.join(self.mask_root,name),0)
img = cv2.resize(img,(512,512))
mask = cv2.resize(mask,(512,512))
mask = (mask>127).astype(np.float32)

if self.transform:
aug = self.transform(image=img,mask=mask)
img,mask = aug["image"],aug["mask"]
return img,mask

# 训练集增强
train_aug = A.Compose([
A.Resize(512,512),
A.HorizontalFlip(p=0.5),
A.VerticalFlip(p=0.3),
A.RandomRotate90(p=0.4),
A.Normalize(mean=[0.5],std=[0.5]),
ToTensorV2()
])
val_aug = A.Compose([
A.Resize(512,512),
A.Normalize(mean=[0.5],std=[0.5]),
ToTensorV2()
])

五、utils/att_unet.py CBAM注意力+UNet模型

import torch
import torch.nn as nn

# CBAM注意力模块
class CBAM(nn.Module):
def __init__(self,in_channel):
super().__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.max_pool = nn.AdaptiveMaxPool2d(1)
self.mlp = nn.Sequential(
nn.Linear(in_channel,in_channel//8),
nn.ReLU(),
nn.Linear(in_channel//8,in_channel)
)
self.conv = nn.Conv2d(2,1,kernel_size=3,padding=1)
def forward(self,x):
b,c,_,_ = x.shape
# 通道注意力
avg = self.avg_pool(x).view(b,c)
maxp = self.max_pool(x).view(b,c)
avg_att = self.mlp(avg).view(b,c,1,1)
max_att = self.mlp(maxp).view(b,c,1,1)
ch_att = torch.sigmoid(avg_att+max_att)
x = x*ch_att
# 空间注意力
avg_sp = torch.mean(x,dim=1,keepdim=True)
max_sp,_ = torch.max(x,dim=1,keepdim=True)
sp_cat = torch.cat([avg_sp,max_sp],dim=1)
sp_att = torch.sigmoid(self.conv(sp_cat))
return x*sp_att

# 基础卷积块
class DoubleConv(nn.Module):
def __init__(self,in_c,out_c):
super().__init__()
self.block = nn.Sequential(
nn.Conv2d(in_c,out_c,3,padding=1),
nn.BatchNorm2d(out_c),
nn.ReLU(inplace=True),
nn.Conv2d(out_c,out_c,3,padding=1),
nn.BatchNorm2d(out_c),
nn.ReLU(inplace=True)
)
self.att = CBAM(out_c)
def forward(self,x):
x = self.block(x)
return self.att(x)

# Attention UNet整体
class AttUNet(nn.Module):
def __init__(self,in_ch=1,num_cls=1):
super().__init__()
self.down1 = DoubleConv(in_ch,64)
self.down2 = DoubleConv(64,128)
self.down3 = DoubleConv(128,256)
self.down4 = DoubleConv(256,512)
self.pool = nn.MaxPool2d(2)
self.up1 = nn.ConvTranspose2d(512,256,2,stride=2)
self.up_conv1 = DoubleConv(512,256)
self.up2 = nn.ConvTranspose2d(256,128,2,stride=2)
self.up_conv2 = DoubleConv(256,128)
self.up3 = nn.ConvTranspose2d(128,64,2,stride=2)
self.up_conv3 = DoubleConv(128,64)
self.out = nn.Conv2d(64,num_cls,1)
def forward(self,x):
d1 = self.down1(x)
d2 = self.down2(self.pool(d1))
d3 = self.down3(self.pool(d2))
d4 = self.down4(self.pool(d3))
u1 = self.up1(d4)
u1 = torch.cat([u1,d3],dim=1)
u1 = self.up_conv1(u1)
u2 = self.up2(u1)
u2 = torch.cat([u2,d2],dim=1)
u2 = self.up_conv2(u2)
u3 = self.up3(u2)
u3 = torch.cat([u3,d1],dim=1)
u3 = self.up_conv3(u3)
out = self.out(u3)
return out

六、utils/seg_metric.py Dice、IoU指标

import numpy as np

def get_dice(pred,gt):
smooth = 1e-6
pred = (pred>0.5).astype(np.float32)
gt = gt.astype(np.float32)
inter = (pred*gt).sum()
union = pred.sum()+gt.sum()
dice = (2*inter+smooth)/(union+smooth)
return dice

def get_iou(pred,gt):
smooth = 1e-6
pred = (pred>0.5).astype(np.float32)
gt = gt.astype(np.float32)
inter = (pred*gt).sum()
union = pred.sum()+gt.sum()inter
iou = (inter+smooth)/(union+smooth)
return iou

七、train.py 训练主脚本

import os
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from utils.dataset import RetinaDataset,train_aug,val_aug
from utils.att_unet import AttUNet
from utils.seg_metric import get_dice,get_iou
import numpy as np

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
batch_size = 4
epochs = 100
lr = 1e-4

train_set = RetinaDataset("./dataset/train/img","./dataset/train/mask",train_aug)
test_set = RetinaDataset("./dataset/test/img","./dataset/test/mask",val_aug)
train_loader = DataLoader(train_set,batch_size=batch_size,shuffle=True)
test_loader = DataLoader(test_set,batch_size=batch_size,shuffle=False)

model = AttUNet(in_ch=1,num_cls=1).to(device)
loss_fn = nn.BCEWithLogitsLoss()
opt = torch.optim.Adam(model.parameters(),lr=lr)

os.makedirs("weights",exist_ok=True)
best_dice = 0.0

for epoch in range(epochs):
# 训练
model.train()
train_loss = 0
for img,mask in train_loader:
img,mask = img.to(device),mask.unsqueeze(1).to(device)
pred = model(img)
loss = loss_fn(pred,mask)
opt.zero_grad()
loss.backward()
opt.step()
train_loss += loss.item()
train_loss /= len(train_loader)

# 验证
model.eval()
dice_list,iou_list = [],[]
with torch.no_grad():
for img,mask in test_loader:
img = img.to(device)
out = model(img)
pred = torch.sigmoid(out).cpu().numpy()
gt = mask.cpu().numpy()
for p,g in zip(pred,gt):
dice_list.append(get_dice(p,g))
iou_list.append(get_iou(p,g))
mean_dice = np.mean(dice_list)
mean_iou = np.mean(iou_list)
print(f"Epoch:{epoch+1:03d} | Loss:{train_loss:.4f} | Dice:{mean_dice:.3f} | IoU:{mean_iou:.3f}")
if mean_dice>best_dice:
best_dice = mean_dice
torch.save(model.state_dict(),"./weights/best.pth")
print("✅最优权重已保存至weights文件夹")

八、GUI.py PyQt5可视化分割系统(项目启动文件)

import sys
import cv2
import numpy as np
from PyQt5.QtWidgets import QApplication,QMainWindow,QWidget,QHBoxLayout,QVBoxLayout,QLabel,QPushButton,QFileDialog
from PyQt5.QtGui import QPixmap,QImage
from PyQt5.QtCore import Qt
import torch
from utils.att_unet import AttUNet

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = AttUNet(1,1).to(device)
model.load_state_dict(torch.load("./weights/best.pth",map_location=device))
model.eval()

class RetinaSegUI(QMainWindow):
def __init__(self):
super().__init__()
self.setWindowTitle("视网膜血管分割系统 – Medical Imaging AI")
self.resize(1550,820)
self.img_path = None
self.save_result_path = "./save_result/"
import os
os.makedirs(self.save_result_path,exist_ok=True)
self.init_ui()

def init_ui(self):
central = QWidget()
self.setCentralWidget(central)
main_layout = QHBoxLayout(central)
# 左侧控制面板
left_panel = QWidget()
left_layout = QVBoxLayout(left_panel)
title_label = QLabel("视网膜血管分割系统")
title_label.setStyleSheet("background:#2980b9;color:#fff;font-size:18px;padding:12px;")
title_label.setAlignment(Qt.AlignCenter)
self.btn_open = QPushButton("选择视网膜图像")
self.btn_seg = QPushButton("开始血管分割")
self.btn_save = QPushButton("保存分割结果")
self.info_label = QLabel("系统信息\\n模型状态:已加载\\n运行设备:"+str(device)+"\\n当前图像:未选择\\n尺寸:512×512")
self.info_label.setStyleSheet("background:#34495e;color:white;padding:8px")
self.btn_open.clicked.connect(self.open_img)
self.btn_seg.clicked.connect(self.run_segment)
self.btn_save.clicked.connect(self.save_seg_result)
left_layout.addWidget(title_label)
left_layout.addWidget(self.btn_open)
left_layout.addWidget(self.btn_seg)
left_layout.addWidget(self.btn_save)
left_layout.addWidget(self.info_label)

# 右侧双图展示
right_panel = QWidget()
right_layout = QHBoxLayout(right_panel)
self.lab_ori = QLabel("原始视网膜图像")
self.lab_seg = QLabel("血管分割结果")
self.lab_ori.setStyleSheet("border:2px solid #3498db;background:#2c3e50;color:white;text-align:center")
self.lab_seg.setStyleSheet("border:2px solid #e74c3c;background:#2c3e50;color:white;text-align:center")
self.lab_ori.setAlignment(Qt.AlignCenter)
self.lab_seg.setAlignment(Qt.AlignCenter)
right_layout.addWidget(self.lab_ori)
right_layout.addWidget(self.lab_seg)

main_layout.addWidget(left_panel,1)
main_layout.addWidget(right_panel,3)
self.seg_out = None

def open_img(self):
path,_ = QFileDialog.getOpenFileName(self,"选择眼底图片","","*.jpg;*.png;*.bmp")
if not path:
return
self.img_path = path
img = cv2.imread(path,0)
img = cv2.resize(img,(512,512))
qimg = QImage(img.data,512,512,512,QImage.Format_Grayscale8)
self.lab_ori.setPixmap(QPixmap.fromImage(qimg).scaled(self.lab_ori.size(),Qt.KeepAspectRatio))
self.info_label.setText(f"系统信息\\n模型状态:已加载\\n运行设备:{device}\\n当前图像:{path.split('/')[1]}\\n尺寸:512×512")

def run_segment(self):
if not self.img_path:
return
img_ori = cv2.imread(self.img_path,0)
img_ori = cv2.resize(img_ori,(512,512))
inp = (img_ori/255.00.5)/0.5
tensor = torch.from_numpy(inp).unsqueeze(0).unsqueeze(0).float().to(device)
with torch.no_grad():
pred = model(tensor)
pred = torch.sigmoid(pred).cpu().numpy()[0,0]
self.seg_out = (pred>0.5).astype(np.uint8)*255
q_seg = QImage(self.seg_out.data,512,512,512,QImage.Format_Grayscale8)
self.lab_seg.setPixmap(QPixmap.fromImage(q_seg).scaled(self.lab_seg.size(),Qt.KeepAspectRatio))

def save_seg_result(self):
if self.seg_out is None:
return
import time
save_name = f"{int(time.time())}_result.png"
cv2.imwrite(self.save_result_path+save_name,self.seg_out)
self.info_label.setText(f"结果已保存:{save_name}")

if __name__ == "__main__":
app = QApplication(sys.argv)
win = RetinaSegUI()
win.show()
sys.exit(app.exec_())

九、项目使用步骤

  • 数据集整理:DRIVE/CHASE‑DB1/STARE三个数据集图片统一放入dataset/train/img,对应二值掩码放入dataset/train/mask,测试集同理;
  • 模型训练:运行train.py,训练结束最优权重自动保存weights/best.pth;
  • 桌面软件启动:直接运行GUI.py,点击【选择视网膜图像】导入眼底图→【开始血管分割】自动生成白色血管掩码→【保存分割结果】本地存图;
  • 十、拓展落地方向

    ✅ 本科/硕士医学影像毕设全套成品(算法+可视化软件)
    ✅ 眼科筛查设备算法原型、眼底病灶辅助筛查
    ✅ 导出ONNX/TensorRT/RKNN部署嵌入式设备

    赞(0)
    未经允许不得转载:171主机测评 » Attention-UNet 视网膜眼底血管智能分割系统 基础算法 Attention-UNet(U-Net 主干 + CBAM 注意力模块,医学二值分割)
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址