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

使用tf.data.Dataset.from_generator报错:生成器输出形状不符

解决tf.data.Dataset与自定义生成器形状不匹配的TypeError

错误原因分析

报错TypeError: generator yielded an element of shape (32, 224, 224, 3) where an element of shape (224, 224, 3) was expected的核心问题有两点:

  1. 生成器输出与定义的形状不匹配:你的ImageSequence生成器每次yield的是一个批量(32张图片),但output_shapes定义的是单张图片的形状。
  2. 重复批量操作:生成器已经按batch_size=32返回批量数据,后续又调用train_data.batch(batch_size),导致再次对批量数据做打包操作,进一步加剧形状冲突。

解决方案

1. 修正ImageSequence类

主要调整__len__、__call__方法,并将返回的列表转为numpy数组:

import cv2
import numpy as np
import os
import pandas as pd
from sklearn.model_selection import train_test_split
import tensorflow as tf

class ImageSequence:
    def __init__(self, df, mode, img_size=(224, 224), num_channels=3, batch_size=32):
        self.df = df
        self.indices = np.arange(len(df))
        self.batch_size = batch_size
        self.img_dir = 'dataset'
        self.img_size = tuple(img_size)
        self.num_channels = num_channels
        self.mode = mode
        
    def __getitem__(self, idx):
        # 计算当前batch的样本索引,避免越界
        start_idx = idx * self.batch_size
        end_idx = min((idx + 1) * self.batch_size, len(self.df))
        sample_indices = self.indices[start_idx:end_idx]
        
        imgs = []
        genders = []
        for _, row in self.df.iloc[sample_indices].iterrows():
            img = cv2.imread(str(os.path.join(self.img_dir, row["img_paths"])))
            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
            img = cv2.resize(img, self.img_size)
            img = img.astype(np.float32) / 255.0
            
            imgs.append(img)
            genders.append(row["genders"])

        # 将列表转为numpy数组,便于TensorFlow处理
        return np.array(imgs), np.array(genders)
    
    def __len__(self):
        # 返回批次数,向上取整
        return (len(self.df) + self.batch_size - 1) // self.batch_size
    
    def __call__(self):
        for i in range(self.__len__()):
            yield self.__getitem__(i)
            
            if i == self.__len__() - 1:
                self.on_epoch_end()

    def on_epoch_end(self):
        np.random.shuffle(self.indices)

2. 修正调用代码

去掉重复的batch操作,并更新output_shapes为批量形状:

epochs = 20
batch_size = 32

csv_path = 'asian_dataset.csv'
df = pd.read_csv(str(csv_path))
train, val = train_test_split(df, random_state=42, test_size=0.1)

train_gen = ImageSequence(train, "train", batch_size=batch_size)
val_gen = ImageSequence(val, "val", batch_size=batch_size)

# 定义输出类型和形状,使用None适配最后一个batch的可变样本数
output_types = (tf.float32, tf.int32)
output_shapes = ((None, 224, 224, 3), (None,))

train_data = tf.data.Dataset.from_generator(
    train_gen,
    output_types=output_types,
    output_shapes=output_shapes
)
val_data = tf.data.Dataset.from_generator(
    val_gen,
    output_types=output_types,
    output_shapes=output_shapes
)

# 无需再调用batch(),生成器已返回批量数据
print(train_data)

关键修改说明

  • __len__方法:从返回样本总数改为返回批次数,避免生成器循环过多无效次数。
  • __getitem__:添加min处理最后一个batch的边界,防止索引越界;将返回的列表转为numpy数组,确保输出格式符合TensorFlow要求。
  • 调用代码:移除train_data.batch(batch_size),并将output_shapes改为批量维度(用None兼容最后一个batch的样本数不足情况)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 06:30:05