如何用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
相关产品推荐
相关产品推荐

