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

如何修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 13:36:03