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

如何用8通道图像训练zhixuhao/unet分类模型及适配数据生成器

针对8通道输入的Unet分类任务解决方案

一、输入数组的整理方式

1. numpy.concatenate合并(推荐新手)

  • 这是最直接的方式,提前将所有通道合并为8通道数组再喂给模型,避免网络层处理的复杂度。
  • 正确合并方法:假设3通道图像形状为(height, width, 3),每个单通道图像形状为(height, width, 1),需在通道维度(axis=-1,即最后一维)合并:
    import numpy as np
    # 示例:加载3通道和5个单通道图像
    img_rgb = np.load("rgb_img.npy")  # shape: (H,W,3)
    img_chan1 = np.load("chan1.npy")  # shape: (H,W,1)
    img_chan2 = np.load("chan2.npy")
    img_chan3 = np.load("chan3.npy")
    img_chan4 = np.load("chan4.npy")
    img_chan5 = np.load("chan5.npy")
    
    # 合并成8通道
    img_8chan = np.concatenate([img_rgb, img_chan1, img_chan2, img_chan3, img_chan4, img_chan5], axis=-1)
    # 最终shape: (H,W,8)
    
  • 你之前生成18通道是因为合并逻辑错误,比如误将3通道拆分为3个单通道重复合并,或使用了错误的合并维度(如axis=0/1),务必确保合并的是通道维度。

2. 网络层中合并(进阶)

  • 若不想提前合并,可在模型输入部分设置多个输入分支,再用Concatenate层合并:
    from keras.layers import Input, Concatenate
    # 定义输入分支
    input_rgb = Input(shape=(H,W,3))
    input_chan1 = Input(shape=(H,W,1))
    input_chan2 = Input(shape=(H,W,1))
    input_chan3 = Input(shape=(H,W,1))
    input_chan4 = Input(shape=(H,W,1))
    input_chan5 = Input(shape=(H,W,1))
    
    # 合并所有输入
    merged_input = Concatenate(axis=-1)([input_rgb, input_chan1, input_chan2, input_chan3, input_chan4, input_chan5])
    # 将merged_input传入原Unet的后续层
    
  • 这种方式需要修改原Unet的输入结构,同时调整数据生成器返回多个输入数组,对新手门槛较高,优先推荐numpy合并。

二、适配8通道的DataGenerator实现

ImageDataGenerator仅支持最多4通道(RGBA),必须自定义数据生成器。推荐使用keras.utils.Sequence类,这是Keras官方推荐的自定义数据加载方式,支持多线程,且能正确识别类别文件夹。

自定义Sequence生成器示例

import os
import numpy as np
from keras.utils import Sequence
from PIL import Image

class MultiChannelDataGenerator(Sequence):
    def __init__(self, base_dir, img_size=(256,256), batch_size=8, shuffle=True):
        self.base_dir = base_dir
        self.img_size = img_size
        self.batch_size = batch_size
        self.shuffle = shuffle
        
        # 读取类别文件夹及对应样本路径
        self.classes = sorted(os.listdir(base_dir))
        self.class_indices = {cls:i for i,cls in enumerate(self.classes)}
        self.samples = []
        for cls in self.classes:
            cls_dir = os.path.join(base_dir, cls)
            # 假设每个样本的3通道图像命名为xxx_rgb.jpg,单通道为xxx_chan1.png等
            sample_ids = [f.split("_rgb")[0] for f in os.listdir(cls_dir) if "_rgb.jpg" in f]
            for sid in sample_ids:
                self.samples.append((cls, sid))
        
        self.on_epoch_end()

    def __len__(self):
        return len(self.samples) // self.batch_size

    def __getitem__(self, idx):
        batch_samples = self.samples[idx*self.batch_size : (idx+1)*self.batch_size]
        x_batch = []
        y_batch = []
        
        for cls, sid in batch_samples:
            cls_dir = os.path.join(self.base_dir, cls)
            # 加载3通道图像并归一化
            rgb_path = os.path.join(cls_dir, f"{sid}_rgb.jpg")
            img_rgb = np.array(Image.open(rgb_path).resize(self.img_size)) / 255.0
            # 加载5个单通道图像,增加通道维度后归一化
            chan_paths = [
                os.path.join(cls_dir, f"{sid}_chan1.png"),
                os.path.join(cls_dir, f"{sid}_chan2.png"),
                os.path.join(cls_dir, f"{sid}_chan3.png"),
                os.path.join(cls_dir, f"{sid}_chan4.png"),
                os.path.join(cls_dir, f"{sid}_chan5.png")
            ]
            img_chans = []
            for p in chan_paths:
                img_chan = np.array(Image.open(p).resize(self.img_size))
                img_chan = np.expand_dims(img_chan, axis=-1)  # 转为(H,W,1)
                img_chans.append(img_chan / 255.0)
            
            # 合并成8通道
            img_8chan = np.concatenate([img_rgb] + img_chans, axis=-1)
            x_batch.append(img_8chan)
            # 生成类别标签(整数形式,若用交叉熵损失可转为one-hot)
            y_batch.append(self.class_indices[cls])
        
        return np.array(x_batch), np.array(y_batch)

    def on_epoch_end(self):
        if self.shuffle:
            np.random.shuffle(self.samples)

使用方法

# 假设数据集结构:
# train_dir/
#   class_A/
#     sample1_rgb.jpg
#     sample1_chan1.png
#     ...
#   class_B/
#     sample2_rgb.jpg
#     sample2_chan1.png
#     ...

train_generator = MultiChannelDataGenerator(base_dir="train_dir", batch_size=4)
val_generator = MultiChannelDataGenerator(base_dir="val_dir", batch_size=4, shuffle=False)

三、修改Unet模型的输入层

原zhixuhao/unet的输入层针对3通道设计,必须修改为8通道:找到原模型代码中的inputs = Input((img_rows, img_cols, 3)),改为inputs = Input((img_rows, img_cols, 8)),其余层无需调整(Unet卷积层可自适应通道数)。

内容的提问来源于stack exchange,提问作者Syuuuu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 12:07:05