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

Keras训练自定义数据集VGG16内存占用过高,求高效加载方案

Solution to Memory Issues with Large Image Datasets in Keras

Hey there! Loading 3400+ images all at once is definitely going to choke your RAM—great catch on noticing small datasets work fine. The fix here is to use on-the-fly data loading with Keras' built-in tools, so you only load batches of images when you need them for training. Here's exactly how to adjust your code:

This is the simplest way to adapt your existing code, since it leverages Keras' utilities to read images directly from your folder structure (no need to manually load and stack arrays). Your dataset already has a standard class-based folder setup (data/Diseased and data/not diseased or YES/NO), which fits perfectly with this method.

Step 1: Modify Data Loading Code

Replace your entire data loading block (from Loading the training data to Split the dataset) with this:

import numpy as np
import os
import time
from vgg16 import VGG16
from keras.preprocessing import image
from keras.preprocessing.image import ImageDataGenerator  # Add this import
from imagenet_utils import preprocess_input, decode_predictions
from keras.layers import Dense, Activation, Flatten, Input
from keras.models import Model
from keras.utils import np_utils

# Define data path
PATH = os.getcwd()
data_path = PATH + '/data'

# Set up data generators with VGG16-compatible preprocessing
train_datagen = ImageDataGenerator(
    preprocessing_function=preprocess_input,
    validation_split=0.2  # Automatically split 20% of data for validation
)

# Generate training data (loads batches on-the-fly)
train_generator = train_datagen.flow_from_directory(
    data_path,
    target_size=(224, 224),  # Match VGG16's input size
    batch_size=32,
    class_mode='categorical',  # For 2-class classification
    subset='training'
)

# Generate validation data
val_generator = train_datagen.flow_from_directory(
    data_path,
    target_size=(224, 224),
    batch_size=32,
    class_mode='categorical',
    subset='validation'
)

# Check class mapping (matches your original 'YES'/'NO' labels)
print(f"Class mapping: {train_generator.class_indices}")

Step 2: Adjust Model Training

When training, use fit() directly with the generators (no need to pass X_train/y_train arrays anymore). Update both of your model training blocks like this:

For the first custom VGG model (freeze all except last layer):

# Custom_vgg_model_1 - Train only the final classifier
image_input = Input(shape=(224, 224, 3))
model = VGG16(input_tensor=image_input, include_top=True, weights='imagenet')
last_layer = model.get_layer('fc2').output
out = Dense(2, activation='softmax', name='output')(last_layer)
custom_vgg_model = Model(image_input, out)

# Freeze all layers except the final dense layer
for layer in custom_vgg_model.layers[:-1]:
    layer.trainable = False

custom_vgg_model.compile(loss='categorical_crossentropy', optimizer='rmsprop', metrics=['accuracy'])

t = time.time()
hist = custom_vgg_model.fit(
    train_generator,
    epochs=12,
    verbose=1,
    validation_data=val_generator
)
print(f'Training time: {time.time() - t:.2f} seconds')

# Evaluate on validation set
loss, accuracy = custom_vgg_model.evaluate(val_generator, verbose=1)
print(f"[INFO] loss={loss:.4f}, accuracy: {accuracy * 100:.4f}%")

For the second custom VGG model (unfreeze feature extraction layers):

# Custom_vgg_model2 - Train feature extraction + classifier
image_input = Input(shape=(224, 224, 3))
model = VGG16(input_tensor=image_input, include_top=True, weights='imagenet')
last_layer = model.get_layer('block5_pool').output
x = Flatten(name='flatten')(last_layer)
x = Dense(128, activation='relu', name='fc1')(x)
x = Dense(128, activation='relu', name='fc2')(x)
out = Dense(2, activation='softmax', name='output')(x)
custom_vgg_model2 = Model(image_input, out)

# Freeze all layers except the last 3 dense layers
for layer in custom_vgg_model2.layers[:-3]:
    layer.trainable = False

custom_vgg_model2.compile(loss='categorical_crossentropy', optimizer='adadelta', metrics=['accuracy'])

t = time.time()
hist = custom_vgg_model2.fit(
    train_generator,
    epochs=12,
    verbose=1,
    validation_data=val_generator
)
print(f'Training time: {time.time() - t:.2f} seconds')

# Evaluate on validation set
loss, accuracy = custom_vgg_model2.evaluate(val_generator, verbose=1)
print(f"[INFO] loss={loss:.4f}, accuracy: {accuracy * 100:.4f}%")

Step 3: Keep the Visualization Code

Your existing matplotlib visualization code works perfectly with the hist object from fit(), so you can leave that part unchanged.

Why This Works

  • Minimal RAM usage: The generator only loads batch_size (32) images at a time, so your RAM usage will drop drastically (no more 99% utilization!).
  • Automatic class labeling: flow_from_directory infers class labels directly from your folder names, eliminating the need to manually create label arrays.
  • Optional data augmentation: You can easily add augmentation (e.g., rotation, flipping) to improve model generalization by updating the ImageDataGenerator:
    train_datagen = ImageDataGenerator(
        preprocessing_function=preprocess_input,
        validation_split=0.2,
        rotation_range=10,
        horizontal_flip=True
    )
    

Alternative: tf.data.Dataset

If you want more flexibility (e.g., custom preprocessing pipelines), you can use TensorFlow's tf.data.Dataset API. However, ImageDataGenerator is simpler for your current use case and requires minimal code changes.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 07:29:41