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

如何优化Python循环避免内存耗尽?Colab Pro+代码崩溃求助

内存耗尽问题优化方案

原代码一次性将所有图片加载到内存并转换为大numpy数组,当数据集规模较大时会直接耗尽Colab内存。以下是几种高效替代方案:

方案1:使用TensorFlow流式数据集(推荐)

直接用TensorFlow的image_dataset_from_directory实现流式加载,仅在需要时加载批次数据,完全避免内存过载:

import tensorflow as tf

# 处理第一个目录,强制标签为1
dataset1 = tf.keras.utils.image_dataset_from_directory(
    data_dir1,
    image_size=(224, 224),
    batch_size=32,  # 可根据内存调整批次大小
    label_mode='int'
)
dataset1 = dataset1.map(lambda x, _: (x / 255.0, tf.constant(1, dtype=tf.int32)))

# 处理第二个目录,强制标签为2
dataset2 = tf.keras.utils.image_dataset_from_directory(
    data_dir2,
    image_size=(224, 224),
    batch_size=32,
    label_mode='int'
)
dataset2 = dataset2.map(lambda x, _: (x / 255.0, tf.constant(2, dtype=tf.int32)))

# 合并两个数据集
dataset = dataset1.concatenate(dataset2)

# 后续可直接用于模型训练:model.fit(dataset)

方案2:分批次处理并保存分片

将数据分成小批次处理后保存为numpy分片,后续按需加载或合并:

import cv2
import numpy as np
import os
from imutils import paths

def process_and_save_batch(image_paths, label, batch_size=32, save_dir='data_shards'):
    os.makedirs(save_dir, exist_ok=True)
    total = len(image_paths)
    for idx in range(0, total, batch_size):
        batch_paths = image_paths[idx:idx+batch_size]
        batch_imgs = []
        for path in batch_paths:
            img = cv2.imread(path)
            img = cv2.resize(img, (224, 224))
            batch_imgs.append(img)
        # 归一化并保存
        batch_data = np.array(batch_imgs) / 255.0
        batch_labels = np.full(len(batch_data), label, dtype=np.int32)
        np.save(os.path.join(save_dir, f'data_{idx//batch_size}.npy'), batch_data)
        np.save(os.path.join(save_dir, f'labels_{idx//batch_size}.npy'), batch_labels)

# 处理两个目录
process_and_save_batch(list(paths.list_images(data_dir1)), label=1)
process_and_save_batch(list(paths.list_images(data_dir2)), label=2)

# 后续加载所有分片(若内存允许)
data = []
labels = []
for file in sorted(os.listdir('data_shards')):
    if file.startswith('data_'):
        data.append(np.load(os.path.join('data_shards', file)))
    elif file.startswith('labels_'):
        labels.append(np.load(os.path.join('data_shards', file)))
data = np.concatenate(data, axis=0)
labels = np.concatenate(labels, axis=0)

方案3:使用Dask延迟加载超大数据集

对于超大规模数据集,用Dask数组实现延迟计算,仅在需要时加载对应块:

import dask.array as da
import cv2
from imutils import paths

def load_and_preprocess(path):
    img = cv2.imread(path)
    img = cv2.resize(img, (224, 224))
    return img / 255.0

# 构建Dask数组
paths1 = list(paths.list_images(data_dir1))
data1 = da.map_blocks(load_and_preprocess, paths1, dtype=np.float32, chunks=(1, 224, 224, 3))
labels1 = da.full(len(paths1), 1, dtype=np.int32, chunks=(1,))

paths2 = list(paths.list_images(data_dir2))
data2 = da.map_blocks(load_and_preprocess, paths2, dtype=np.float32, chunks=(1, 224, 224, 3))
labels2 = da.full(len(paths2), 2, dtype=np.int32, chunks=(1,))

# 合并数据集
data = da.concatenate([data1, data2], axis=0)
labels = da.concatenate([labels1, labels2], axis=0)

# 需要实际数据时调用compute(),比如取前100个样本
sample_data = data[:100].compute()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 21:27:14