U-net图像分割实战
·
一、U-net模型
代码实现了一个基于 U-net 架构的图像分割模型,主要功能是:
- 加载自定义的图像数据集(原图 + 掩码图)
- 构建 U-net 网络结构(编码 - 解码 + 跳跃连接)
- 动态调整批次大小进行模型训练
- 定期保存模型,最终加载训练好的模型
import numpy as np
import random
import os
from keras.models import save_model, load_model, Model
from keras.layers import Input, Dropout, BatchNormalization, LeakyReLU, concatenate
from keras.layers import Conv2D, MaxPooling2D, AveragePooling2D, Conv2DTranspose
import matplotlib.pyplot as plt
from skimage import io
from skimage.transform import resize
# 这里换成自己本地数据集的路径,同理下面的数据集路径自行替换
input_name = os.listdir('D:/PyCharm/dataset/unet/imgs/train/')
n = len(input_name)
batch_size = 8
input_size_1 = 256
input_size_2 = 256
### 批次数据生成函数
def batch_data(input_name, n, batch_size=8, input_size_1=256, input_size_2=256):
# 随机选1张图片作为批次初始值
rand_num = random.randint(0, n - 1)
# 调整图片尺寸为256×256×3(3通道RGB)
img1 = io.imread('D:/PyCharm/dataset/unet/imgs/train/' + input_name[rand_num]).astype("float")
img2 = io.imread('D:/PyCharm/dataset/unet/masks/train/' + input_name[rand_num]).astype("float")
img1 = resize(img1, [input_size_1, input_size_2, 3])
img2 = resize(img2, [input_size_1, input_size_2, 3])
# 增加维度(适配模型输入:[批次数, 高, 宽, 通道])
img1 = np.reshape(img1, (1, input_size_1, input_size_2, 3))
img2 = np.reshape(img2, (1, input_size_1, input_size_2, 3))
# 归一化(将像素值从0-255缩放到0-1,加速模型收敛)
img1 /= 255
img2 /= 255
# 初始化批次输入/输出
batch_input = img1
batch_output = img2
# 循环补充批次剩余图片
for batch_iter in range(1, batch_size):
rand_num = random.randint(0, n - 1)
img1 = io.imread('D:/PyCharm/dataset/unet/imgs/train/' + input_name[rand_num]).astype("float")
img2 = io.imread('D:/PyCharm/dataset/unet/masks/train/' + input_name[rand_num]).astype("float")
img1 = resize(img1, [input_size_1, input_size_2, 3])
img2 = resize(img2, [input_size_1, input_size_2, 3])
img1 = np.reshape(img1, (1, input_size_1, input_size_2, 3))
img2 = np.reshape(img2, (1, input_size_1, input_size_2, 3))
img1 /= 255
img2 /= 255
batch_input = np.concatenate((batch_input, img1), axis=0)
batch_output = np.concatenate((batch_output, img2), axis=0)
return batch_input, batch_output
# 封装卷积+批量归一化+LeakyReLU激活
def Conv2d_BN(x, nb_filter, kernel_size, strides=(1, 1), padding='same'):
x = Conv2D(nb_filter, kernel_size, strides=strides, padding=padding)(x)
x = BatchNormalization(axis=3)(x)# 批量归一化(axis=3对应通道维度)
x = LeakyReLU(alpha=0.1)(x)
return x
# 封装转置卷积+批量归一化+LeakyReLU激活(上采样用)
def Conv2dT_BN(x, filters, kernel_size, strides=(2, 2), padding='same'):
x = Conv2DTranspose(filters, kernel_size, strides=strides, padding=padding)(x)
x = BatchNormalization(axis=3)(x)
x = LeakyReLU(alpha=0.1)(x)
return x
# 输入层:定义输入尺寸为256×256×3
inpt = Input(shape=(input_size_1, input_size_2, 3))
# 编码器第一层(下采样):256×256 → 128×128
conv1 = Conv2d_BN(inpt, 8, (3, 3))# 8个3×3卷积核,输出256×256×8
conv1 = Conv2d_BN(conv1, 8, (3, 3)) # 再次卷积,强化特征提取
pool1 = MaxPooling2D(pool_size=(2, 2), strides=(2, 2), padding='same')(conv1)# 最大池化,尺寸减半
conv2 = Conv2d_BN(pool1, 16, (3, 3))
conv2 = Conv2d_BN(conv2, 16, (3, 3))
pool2 = MaxPooling2D(pool_size=(2, 2), strides=(2, 2), padding='same')(conv2)
conv3 = Conv2d_BN(pool2, 32, (3, 3))
conv3 = Conv2d_BN(conv3, 32, (3, 3))
pool3 = MaxPooling2D(pool_size=(2, 2), strides=(2, 2), padding='same')(conv3)
conv4 = Conv2d_BN(pool3, 64, (3, 3))
conv4 = Conv2d_BN(conv4, 64, (3, 3))
pool4 = MaxPooling2D(pool_size=(2, 2), strides=(2, 2), padding='same')(conv4)
conv5 = Conv2d_BN(pool4, 128, (3, 3))# 卷积核数128(最大值)
conv5 = Dropout(0.5)(conv5)# Dropout防止过拟合(随机丢弃50%神经元)
conv5 = Conv2d_BN(conv5, 128, (3, 3))
conv5 = Dropout(0.5)(conv5)
# 解码器第一层(上采样):16×16 → 32×32
convt1 = Conv2dT_BN(conv5, 64, (3, 3))# 转置卷积上采样,卷积核数64(对应编码器第四层)
concat1 = concatenate([conv4, convt1], axis=3)# 拼接编码器第四层特征(跳跃连接)
concat1 = Dropout(0.5)(concat1)
conv6 = Conv2d_BN(concat1, 64, (3, 3))
conv6 = Conv2d_BN(conv6, 64, (3, 3))
convt2 = Conv2dT_BN(conv6, 32, (3, 3))
concat2 = concatenate([conv3, convt2], axis=3)
concat2 = Dropout(0.5)(concat2)
conv7 = Conv2d_BN(concat2, 32, (3, 3))
conv7 = Conv2d_BN(conv7, 32, (3, 3))
convt3 = Conv2dT_BN(conv7, 16, (3, 3))
concat3 = concatenate([conv2, convt3], axis=3)
concat3 = Dropout(0.5)(concat3)
conv8 = Conv2d_BN(concat3, 16, (3, 3))
conv8 = Conv2d_BN(conv8, 16, (3, 3))
convt4 = Conv2dT_BN(conv8, 8, (3, 3))
concat4 = concatenate([conv1, convt4], axis=3)
concat4 = Dropout(0.5)(concat4)
conv9 = Conv2d_BN(concat4, 8, (3, 3))
conv9 = Conv2d_BN(conv9, 8, (3, 3))
conv9 = Dropout(0.5)(conv9)
# 输出层:256×256×3(与掩码维度一致)
outpt = Conv2D(filters=3, kernel_size=(1, 1), strides=(1, 1), padding='same', activation='sigmoid')(conv9)
model = Model(inpt, outpt)
model.compile(loss='mean_squared_error', optimizer='Nadam', metrics=['accuracy'])
model.summary()
itr = 3000
S = [] # 可用于存储损失/准确率
patience = 200 # 连续200次loss不下降就停止
min_loss = float('inf') # 初始最小loss设为无穷大
stop_count = 0 # 计数连续不下降的次数
for i in range(itr):
print("iteration = ", i + 1)
# 动态调整批次大小:前期小批次快速收敛,后期大批次稳定训练
if i < 500:
bs = 4
elif i < 2000:
bs = 8
elif i < 5000:
bs = 16
else:
bs = 32
# 生成当前批次的训练数据
train_X, train_Y = batch_data(input_name, n, batch_size=bs)
# 训练1个epoch(仅用当前批次数据)
# model.fit(train_X, train_Y, epochs=1, verbose=0)
#训练并获取loss
history = model.fit(train_X, train_Y, epochs=1, verbose=0)
current_loss = history.history['loss'][0] # 获取本次迭代的loss
S.append(current_loss) # 保存loss到列表
# 早停逻辑:判断是否需要停止训练
if current_loss < min_loss:
min_loss = current_loss # 更新最小loss
stop_count = 0 # 重置连续不下降计数
save_model(model, 'best_unet.h5') # 保存当前最优模型(loss最低)
print(f"更新最优模型,当前最小loss: {min_loss:.4f}")
else:
stop_count += 1 # 连续不下降次数+1
# 达到耐心值,提前停止训练
if stop_count >= patience:
print(f"\n连续{patience}次迭代loss未下降,提前停止训练!")
print(f"最优loss: {min_loss:.4f},最优模型已保存为 best_unet.h5")
break # 跳出循环,结束训练
# 每100次迭代保存一次模型
if i % 100 == 99:
save_model(model, 'unet.h5')
# 加载最终保存的模型(验证模型可正常加载)
model = load_model('unet.h5')
print("训练完成!")
print("\n训练完成!最终加载的是loss最低的最优模型 best_unet.h5")
注意调整图片是二值格式还是RGB格式。训练时间较长,训练结果:

二.U-net 模型测试与评估
代码核心目的:
- 加载训练好的 U-net 模型
- 加载测试数据集(原图 + 真实掩码)
- 用模型对测试数据进行预测
- 计算多种量化评估指标(MSE、SSIM、Dice 系数、IoU),衡量模型分割效果
- 可视化对比:输入图、真实掩码、预测掩码
import numpy as np
import random
import os
from keras.models import load_model
from skimage import io
from skimage.transform import resize
import matplotlib.pyplot as plt
from skimage.metrics import structural_similarity as ssim
from sklearn.metrics import f1_score, jaccard_score
def batch_data_test(input_name, n, batch_size=8, input_size_1=256, input_size_2=256):
rand_num = random.randint(0, n - 1)
img1 = io.imread('D:/PyCharm/dataset/unet/imgs/test/' + input_name[rand_num]).astype("float")
img2 = io.imread('D:/PyCharm/dataset/unet/masks/test/' + input_name[rand_num]).astype("float")
# 调整尺寸为256×256×3(与训练时输入尺寸一致)
img1 = resize(img1, [input_size_1, input_size_2, 3])
img2 = resize(img2, [input_size_1, input_size_2, 3])
# 增加批次维度(适配模型输入格式:[批次数, 高, 宽, 通道])
img1 = np.reshape(img1, (1, input_size_1, input_size_2, 3))
img2 = np.reshape(img2, (1, input_size_1, input_size_2, 3))
img1 /= 255
img2 /= 255
batch_input = img1
batch_output = img2
for batch_iter in range(1, batch_size):
rand_num = random.randint(0, n - 1)
img1 = io.imread('D:/PyCharm/dataset/unet/imgs/test/' + input_name[rand_num]).astype("float")
img2 = io.imread('D:/PyCharm/dataset/unet/masks/test/' + input_name[rand_num]).astype("float")
img1 = resize(img1, [input_size_1, input_size_2, 3])
img2 = resize(img2, [input_size_1, input_size_2, 3])
img1 = np.reshape(img1, (1, input_size_1, input_size_2, 3))
img2 = np.reshape(img2, (1, input_size_1, input_size_2, 3))
img1 /= 255
img2 /= 255
# 拼接图片到批次中
batch_input = np.concatenate((batch_input, img1), axis=0)
batch_output = np.concatenate((batch_output, img2), axis=0)
return batch_input, batch_output
# 加载已训练的模型
model = load_model('unet.h5')
test_name = os.listdir('D:/PyCharm/dataset/unet/imgs/test/')
n_test = len(test_name)
# 生成测试批次数据(batch_size=1,即只选1张测试图)
test_X, test_Y = batch_data_test(test_name, n_test, batch_size=1)
# 进行预测:输入测试图像test_X,输出预测掩码pred_Y
pred_Y = model.predict(test_X)
# 计算均方误差,衡量预测值与真实值的像素级误差,值越小越好
mse = np.mean((test_Y - pred_Y) ** 2)
# 计算结构相似性指数,衡量图像结构相似度,范围0-1,越接近1越好
ssim_score = np.mean([ssim(test_Y[i], pred_Y[i], multichannel=True) for i in range(test_Y.shape[0])])
# 计算Dice系数(F1分数):分割任务核心指标,先将像素值二值化
binary_test_Y = (test_Y > 0.5).astype(int)
binary_pred_Y = (pred_Y > 0.5).astype(int)
# flatten():将256×256×3的数组展平为一维,适配f1_score输入格式
dice_coefficient = np.mean([f1_score(binary_test_Y[i].flatten(), binary_pred_Y[i].flatten()) for i in range(binary_test_Y.shape[0])])
# 计算交并比,值越大越好
iou = np.mean([jaccard_score(binary_test_Y[i].flatten(), binary_pred_Y[i].flatten()) for i in range(binary_test_Y.shape[0])])
print(f"Mean Squared Error: {mse}")
print(f"Structural Similarity Index: {ssim_score}")
print(f"Dice Coefficient: {dice_coefficient}")
print(f"Intersection over Union: {iou}")
# 选择要显示的图像
image_index = 0
# 创建一个画布,并添加三个子图
plt.figure(figsize=(12, 4)) # 设置画布大小
# 显示输入图像
plt.subplot(1, 3, 1) # 1行3列,第1个子图
plt.imshow(test_X[image_index, :, :, :])
plt.title("Input Image")
plt.axis('off')
# 显示真实标签
plt.subplot(1, 3, 2)
plt.imshow(test_Y[image_index, :, :, :])
plt.title("True Mask")
plt.axis('off')
# 显示模型的预测结果
plt.subplot(1, 3, 3)
plt.imshow(pred_Y[image_index, :, :, :])
plt.title("Predicted Mask")
plt.axis('off')
plt.show()
运行结果:
预测效果

评价指标

更多推荐
所有评论(0)