You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改代码以支持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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 04:02:31