Keras训练自定义数据集VGG16内存占用过高,求高效加载方案
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:
Recommended Approach: ImageDataGenerator + flow_from_directory
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_directoryinfers 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

