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

多类分割训练报错:无法广播数组形状(320,320,1,8)至(320,320,1)求解

多类分割任务中独热编码掩码的形状不匹配问题解决

问题背景

开展8类图像分割任务(0为背景,1-7为目标类):

  • 输入灰度图像取值范围0-255,形状为(320,320,1)
  • 掩码取值0-7,使用tf.one_hot做独热编码适配categorical_crossentropy损失
  • 训练时抛出错误:could not broadcast input array from shape (320,320,1,8) into shape (320,320,1)

错误原因

灰度掩码通过flow_from_directory加载后形状为(batch_size, 320, 320, 1),直接用tf.one_hot编码时,会在最后一个通道维度基础上再添加类别维度,导致掩码形状变为(batch_size, 320, 320, 1, 8),而模型输出形状是(batch_size, 320, 320, 8),二者维度不匹配引发错误。

解决方案

1. 调整独热编码后的掩码形状

在one_hot_encode_mask函数中,编码前先通过tf.squeeze去掉多余的单通道维度,将掩码形状从(320,320,1)转为(320,320),再做独热编码就能得到(320,320,8)的正确形状,与模型输出维度对齐:

def one_hot_encode_mask(mask):
    num_classes = 8
    # 压缩单通道维度后再编码
    one_hot_mask = tf.one_hot(tf.squeeze(mask, axis=-1), num_classes)
    return one_hot_mask

2. 修复UNet中的笔误

原UNet代码中跳过连接赋值错误:skips[level] = x-1应改为skips[level] = x,否则会导致特征图计算异常。

3. 无需调整输入图像形状

输入图像保持(320,320,1)即可,模型输入层已经适配该形状,无需修改。

修改后的完整代码

import os
import matplotlib.pyplot as plt
import tensorflow as tf
from tensorflow import keras
from keras_preprocessing.image import ImageDataGenerator
import numpy as np
import cv2


# 定义常量
IMAGE_HEIGHT = 320
IMAGE_WIDTH = 320
IMG_SIZE = (IMAGE_HEIGHT, IMAGE_WIDTH)
BATCH_SIZE_TRAIN = 16
BATCH_SIZE_TEST = 16
SEED = 909

data_dir = '/content/drive/MyDrive/data'
data_dir_train = os.path.join(data_dir, 'training')
data_dir_train_image = os.path.join(data_dir_train, 'img')
data_dir_train_mask = os.path.join(data_dir_train, 'mask')

data_dir_test = os.path.join(data_dir, 'test')
data_dir_test_image = os.path.join(data_dir_test, 'img')
data_dir_test_mask = os.path.join(data_dir_test, 'mask')

NUM_TRAIN = 3840
NUM_TEST = 960
NUM_OF_EPOCHS = 10


def create_segmentation_generator_train(img_path, mask_path, BATCH_SIZE):
    data_gen_args = dict(rescale=1./255)
    img_datagen = ImageDataGenerator(**data_gen_args)

    def one_hot_encode_mask(mask):
        num_classes = 8  # 7类目标+1类背景
        # 压缩单通道维度后做独热编码
        one_hot_mask = tf.one_hot(tf.squeeze(mask, axis=-1), num_classes)
        return one_hot_mask

    mask_datagen = ImageDataGenerator(dtype='float32', preprocessing_function=one_hot_encode_mask, **data_gen_args)
    img_generator = img_datagen.flow_from_directory(img_path, target_size=IMG_SIZE, class_mode=None, color_mode='grayscale', batch_size=BATCH_SIZE, seed=SEED)
    mask_generator = mask_datagen.flow_from_directory(mask_path, target_size=IMG_SIZE, class_mode=None, color_mode='grayscale', batch_size=BATCH_SIZE, seed=SEED)
    
    return zip(img_generator, mask_generator)


def create_segmentation_generator_test(img_path, mask_path, BATCH_SIZE):
    data_gen_args = dict(rescale=1./255)
    img_datagen = ImageDataGenerator(**data_gen_args)

    def one_hot_encode_mask(mask):
        num_classes = 8  # 7类目标+1类背景
        one_hot_mask = tf.one_hot(tf.squeeze(mask, axis=-1), num_classes)
        return one_hot_mask

    mask_datagen = ImageDataGenerator(dtype='float32', preprocessing_function=one_hot_encode_mask, **data_gen_args)
    img_generator = img_datagen.flow_from_directory(img_path, target_size=IMG_SIZE, class_mode=None, color_mode='grayscale', batch_size=BATCH_SIZE, seed=SEED)
    mask_generator = mask_datagen.flow_from_directory(mask_path, target_size=IMG_SIZE, class_mode=None, color_mode='grayscale', batch_size=BATCH_SIZE, seed=SEED)
    
    return zip(img_generator, mask_generator)


train_generator = create_segmentation_generator_train(data_dir_train_image, data_dir_train_mask, BATCH_SIZE_TRAIN)
test_generator = create_segmentation_generator_test(data_dir_test_image, data_dir_test_mask, BATCH_SIZE_TEST)


def display(display_list):
    plt.figure(figsize=(10, 10))
    title = ['输入图像', '真实掩码', '预测掩码']
    for i in range(len(display_list)):
        plt.subplot(1, len(display_list), i + 1)
        plt.title(title[i])
        plt.imshow(tf.keras.preprocessing.image.array_to_img(display_list[i]), cmap='gray')
    plt.show()


def show_dataset(datagen, num=1):
    for i in range(0, num):
        image, mask = next(datagen)
        mask = mask[0]  # 获取第一个掩码
        mask = np.argmax(mask, axis=-1)  # 将独热编码掩码转回整数标签

        plt.figure(figsize=(12, 6))
        plt.subplot(1, 2, 1)
        plt.title('输入图像')
        plt.imshow(image[0], cmap='gray')

        plt.subplot(1, 2, 2)
        plt.title('掩码(类别0-7)')
        plt.imshow(mask, cmap='gray')

        plt.show()


show_dataset(train_generator, 1)


def unet(n_levels, initial_features=32, n_blocks=2, kernel_size=3, pooling_size=2, in_channels=1, num_classes=8):
    # n_blocks:每个层级的卷积次数
    inputs = keras.layers.Input(shape=(IMAGE_HEIGHT, IMAGE_WIDTH, in_channels))
    x = inputs

    convpars = dict(kernel_size=kernel_size, activation='relu', padding='same')

    # 下采样路径
    skips = {}
    for level in range(n_levels):
        for _ in range(n_blocks):
            x = keras.layers.Conv2D(initial_features * 2 ** level, **convpars)(x)
        if level < n_levels - 1:
            # 修复笔误:将x-1改为x
            skips[level] = x
            x = keras.layers.MaxPool2D(pooling_size)(x)

    # 上采样路径
    for level in reversed(range(n_levels-1)):
        x = keras.layers.Conv2DTranspose(initial_features * 2 ** level, strides=pooling_size, **convpars)(x)
        x = keras.layers.Concatenate()([x, skips[level]])
        for _ in range(n_blocks):
            x = keras.layers.Conv2D(initial_features * 2 ** level, **convpars)(x)

    # 输出层
    x = keras.layers.Conv2D(num_classes, kernel_size=1, activation='softmax', padding='same')(x)

    return keras.Model(inputs=[inputs], outputs=[x], name=f'UNET-L{n_levels}-F{initial_features}')

EPOCH_STEP_TRAIN = NUM_TRAIN // BATCH_SIZE_TRAIN
EPOCH_STEP_TEST = NUM_TEST // BATCH_SIZE_TRAIN

model = unet(4)
model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.001), loss='categorical_crossentropy', metrics=['accuracy']) 

model.summary()

history = model.fit_generator(generator=train_generator, steps_per_epoch=EPOCH_STEP_TRAIN, validation_data=test_generator, validation_steps=EPOCH_STEP_TEST, epochs=NUM_OF_EPOCHS)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 16:15:53