如何修改CNN图像分类数据集批次生成函数以返回指定大小的唯一批次?
批次生成函数修改方案
你当前的代码会一次性加载所有数据集到内存并全部返回,无法实现按batch_size逐批次返回的需求,我们可以将其改造为Python生成器,并增加每个epoch的数据集打乱逻辑,具体修改如下:
修改要点
- 每个epoch启动时先打乱所有图像路径,保证不同epoch的批次组合唯一,降低过拟合风险
- 用
yield关键字替代原有的return,将函数改造为生成器,每次迭代仅返回单批次数据,同时大幅降低内存占用 - 自动适配最后一个批次样本量不足
batch_size的场景,若需要丢弃不足批次可自行调整判断逻辑
修改后完整代码
import cv2 import random def batch_generator(dataset, input_shape=(256, 256), batch_size=32, shuffle=True): """ 批次生成器,每个epoch迭代返回指定大小的图像+标签批次 :param dataset: 所有图像路径列表 :param input_shape: 输出图像尺寸,默认256×256 :param batch_size: 单批次样本量,默认32 :param shuffle: 每个epoch是否打乱数据集,默认True """ # 每个epoch启动时先打乱数据集 if shuffle: random.shuffle(dataset) # 按步长batch_size遍历所有样本 for i in range(0, len(dataset), batch_size): batch_paths = dataset[i:i+batch_size] batch_images = [] batch_labels = [] for img_path in batch_paths: # 读取并resize图像 img = cv2.resize(cv2.imread(img_path, cv2.IMREAD_COLOR), input_shape, interpolation=cv2.INTER_AREA) batch_images.append(img) # 读取对应标签 label = labels[img_path.split('/')[-2]] batch_labels.append(label) # 返回单批次数据,暂停等待下一次迭代 yield batch_images, batch_labels
调用方式
训练时每个epoch直接迭代生成器即可自动获取所有批次:
# 示例训练循环 for epoch in range(total_epochs): # 每个epoch自动生成所有唯一批次 for batch_imgs, batch_labels in batch_generator(dataset, input_shape=(256,256), batch_size=32): # 单批次训练逻辑 model.train_on_batch(batch_imgs, batch_labels)
可选优化点
- 可以将全局变量
labels改为入参传入,提升函数复用性 - 若数据集过大,可提前将resize后的图像保存为本地缓存文件,避免每个epoch重复读取resize的IO开销
- 可在生成器中增加图像归一化、数据增强逻辑,减少后续预处理步骤
内容的提问来源于stack exchange,提问作者Ashar
相关产品推荐
相关产品推荐

