如何修改代码以支持RGBX四通道图像的语义分割?
适配RGBX四通道语义分割的代码修改方案
核心问题说明
ImageDataGenerator仅支持单通道或三通道图像,无法直接处理四通道数据,因此需要替换为自定义数据生成器,并调整预处理逻辑以适配额外通道。
具体修改步骤
1. 替换数据生成器(使用Albumentations实现同步增强)
Albumentations支持多通道图像的增强操作,且能保证图像与掩码的增强同步,比手动实现更高效可靠。先安装依赖库:
pip install albumentations
2. 调整预处理函数
预训练骨干网络的预处理(如ResNet34的preprocess_input)仅针对RGB通道设计,额外通道需要单独做归一化处理,避免干扰RGB通道的预训练特征提取。
3. 修改生成器逻辑
自定义生成器负责加载四通道图像、掩码,应用同步增强,再执行预处理。
修改后的完整代码
import os import cv2 import numpy as np import glob from matplotlib import pyplot as plt import tensorflow as tf import splitfolders import segmentation_models as sm from sklearn.preprocessing import MinMaxScaler from keras.utils import to_categorical import albumentations as A # 分割数据集(原有逻辑保留) input_folder='path folder to my images and masks ' output_folder='path to output folder' splitfolders.ratio(input_folder, output=output_folder, seed=42, ratio=(.75,.25), group_prefix=None) seed=24 batch_size=16 n_classes=2 # 分别初始化RGB通道和额外通道的归一化器 rgb_scaler = MinMaxScaler() extra_scaler = MinMaxScaler() BACKBONE='resnet34' preprocess_input=sm.get_preprocessing(BACKBONE) def preprocess_data(img, mask, num_class): # 拆分RGB通道和额外通道 rgb_img = img[..., :3] extra_channel = img[..., 3:] # RGB通道:先归一化,再用预训练骨干的预处理 rgb_img = rgb_scaler.fit_transform(rgb_img.reshape(-1, 3)).reshape(rgb_img.shape) rgb_img = preprocess_input(rgb_img) # 额外通道:单独归一化 extra_channel = extra_scaler.fit_transform(extra_channel.reshape(-1, 1)).reshape(extra_channel.shape) # 合并通道 processed_img = np.concatenate([rgb_img, extra_channel], axis=-1) # 掩码处理 mask = to_categorical(mask, num_class) return (processed_img, mask) # 定义Albumentations增强管道 transform = A.Compose([ A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5), A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5, border_mode=cv2.BORDER_REFLECT) ], additional_targets={'mask': 'mask'}) def trainGenerator(train_img_dir, train_mask_dir, num_class): # 获取所有图像和掩码路径(假设图像和掩码文件名一一对应) img_paths = sorted(glob.glob(os.path.join(train_img_dir, '*'))) mask_paths = sorted(glob.glob(os.path.join(train_mask_dir, '*'))) while True: # 每个epoch打乱一次数据 indices = np.random.permutation(len(img_paths)) for i in range(0, len(indices), batch_size): batch_indices = indices[i:i+batch_size] batch_imgs = [] batch_masks = [] for idx in batch_indices: # 加载完整四通道图像 img = cv2.imread(img_paths[idx], cv2.IMREAD_UNCHANGED) if img.shape[-1] !=4: raise ValueError("图像不是四通道格式") # 转换为RGB+X格式(若原图像是BGRA) img = cv2.cvtColor(img, cv2.COLOR_BGRA2RGBA) # 加载单通道掩码 mask = cv2.imread(mask_paths[idx], cv2.IMREAD_GRAYSCALE) # 应用同步增强 augmented = transform(image=img, mask=mask) aug_img = augmented['image'] aug_mask = augmented['mask'] batch_imgs.append(aug_img) batch_masks.append(aug_mask) # 转换为numpy数组并预处理 batch_imgs = np.array(batch_imgs) batch_masks = np.array(batch_masks) batch_imgs, batch_masks = preprocess_data(batch_imgs, batch_masks, num_class) yield (batch_imgs, batch_masks) # 注意:路径需指向图像/掩码的直接存储目录 train_img_path='path for training images directory' train_mask_path='path for training masks directory' train_img_gen=trainGenerator(train_img_path, train_mask_path, num_class=2) val_img_path='path for validation images directory' val_mask_path='path for validation masks directory' val_img_gen=trainGenerator(val_img_path, val_mask_path, num_class=2) # 验证生成器输出格式 x, y=next(train_img_gen) print(f"输入图像形状:{x.shape}") # 应为(batch_size,256,256,4) print(f"掩码形状:{y.shape}") # 应为(batch_size,256,256,2) # 可视化样本 for i in range(0,3): # RGB通道需归一化到0-1才能显示(预训练预处理会改变值域) rgb_image = x[i][..., :3] rgb_image = (rgb_image - np.min(rgb_image)) / (np.max(rgb_image) - np.min(rgb_image)) extra_channel = x[i][..., 3] mask=np.argmax(y[i], axis=2) plt.figure(figsize=(12,4)) plt.subplot(1,3,1) plt.imshow(rgb_image) plt.title("RGB通道") plt.subplot(1,3,2) plt.imshow(extra_channel, cmap='viridis') plt.title("额外通道") plt.subplot(1,3,3) plt.imshow(mask, cmap='gray') plt.title("掩码") plt.show() # 计算训练步数 num_train_imgs=len(os.listdir(train_img_path)) num_val_images=len(os.listdir(val_img_path)) steps_per_epochs=num_train_imgs//batch_size val_steps_per_epoch=num_val_images//batch_size IMG_HEIGHT=x.shape[1] IMG_WIDTH=x.shape[2] IMG_CHANNELS=x.shape[3] # 模型定义(保留原有逻辑,确保输入通道为4) model=sm.Unet(BACKBONE, encoder_weights=None, input_shape=(IMG_HEIGHT,IMG_WIDTH,IMG_CHANNELS), classes=n_classes, activation='softmax') model.compile('Adam', loss=sm.losses.binary_crossentropy, metrics=[sm.metrics.iou_score, sm.metrics.FScore()]) # 启动训练 history=model.fit(train_img_gen, steps_per_epoch=steps_per_epochs, epochs=100, verbose=1, validation_data=val_img_gen, validation_steps=val_steps_per_epoch)
关键细节说明
- 通道拆分处理:将RGB通道与额外通道分开预处理,避免预训练的RGB预处理逻辑干扰额外通道的特征表达。
- 同步增强:通过Albumentations的
additional_targets参数,保证图像和掩码应用完全一致的增强操作,避免错位。 - 图像加载:使用
cv2.imread(..., cv2.IMREAD_UNCHANGED)读取完整四通道数据,防止丢失额外通道信息。 - 可视化调整:预训练骨干的预处理会改变RGB通道的值域,可视化前需重新归一化到0-1区间。
内容的提问来源于stack exchange,提问作者FF123456
相关产品推荐
相关产品推荐

