如何为猫狗二分类任务创建Python数据生成器解决内存不足问题
猫狗二分类内存优化:生成器实现方案
核心思路
原代码一次性将所有处理后的图像存入列表,导致内存溢出。生成器的本质是按需加载、处理并返回数据,不会一次性占用大量内存,完美适配大规模数据集。
步骤1:先收集图像路径与标签(低内存占用)
先把所有图像的路径和对应标签存成列表,这个列表仅存字符串和数字,内存占用极低:
import os from tqdm import tqdm import random train_dir = "../input/dog-cat/train" CATEGORIES = ["dog", "cat"] image_paths = [] labels = [] for category in CATEGORIES: path = os.path.join(train_dir, category) class_num = CATEGORIES.index(category) for img in tqdm(os.listdir(path)): try: img_full_path = os.path.join(path, img) image_paths.append(img_full_path) labels.append(class_num) except Exception as e: pass # 同步打乱路径和标签,避免数据与标签错位 combined = list(zip(image_paths, labels)) random.shuffle(combined) image_paths, labels = zip(*combined)
步骤2:实现数据生成器
用Python生成器函数,每次仅加载、处理一批图像,返回给模型使用:
import cv2 import numpy as np def data_generator(image_paths, labels, batch_size=32): total_samples = len(image_paths) while True: # 生成器需循环无限次,适配Keras/TensorFlow的训练逻辑 for idx in range(0, total_samples, batch_size): # 截取当前批次的路径和标签 batch_paths = image_paths[idx:idx+batch_size] batch_labels = labels[idx:idx+batch_size] batch_features = [] batch_targets = [] for path, label in zip(batch_paths, batch_labels): # 按需加载单张图像 img = cv2.imread(path) # 应用均值滤波(修复原代码中reduced_img_train未定义的问题) img_mean = cv2.blur(img, (9,9)) batch_features.append(img_mean) batch_targets.append(label) # 转换为模型可接收的numpy数组格式并返回 yield np.array(batch_features), np.array(batch_targets)
步骤3:用生成器训练模型
以Keras为例,直接将生成器传入训练函数即可:
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv2D, Flatten, Dense # 假设你已定义好模型结构(示例) model = Sequential([ Conv2D(32, (3,3), activation='relu', input_shape=(224, 224, 3)), # 替换为你的图像尺寸 Flatten(), Dense(1, activation='sigmoid') ]) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy']) # 初始化生成器 batch_size = 32 train_gen = data_generator(image_paths, labels, batch_size) # 计算每轮训练步数 steps_per_epoch = len(image_paths) // batch_size # 启动训练 model.fit(train_gen, steps_per_epoch=steps_per_epoch, epochs=10)
关键优化点
- 避免一次性加载所有图像到内存,仅在需要时加载处理
- 按批次返回数据,进一步控制单批次内存占用
- 修复原代码中
reduced_img_train未定义的错误
内容的提问来源于stack exchange,提问作者Meriem Chahinez Ben
相关产品推荐
相关产品推荐

