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

如何在Keras ImageDataGenerator中高效实现图像中心裁剪?

Efficient Center Cropping with Keras/ImageDataGenerator

Great question! The center crop approach you’ve implemented works, but looping through every image in a batch with plain Python can definitely drag down performance—especially when you’re working with large datasets. Let’s go through a few much faster alternatives to fix this:

1. Use TensorFlow’s Built-in tf.image.central_crop (Top Performance)

Since Keras is tightly integrated with TensorFlow, using TF’s optimized, vectorized image operations will give you the biggest speed boost. tf.image.central_crop is purpose-built for this task and handles entire batches at once (no slow Python loops required).

Option A: Wrap with tf.data.Dataset

This is the most efficient approach because tf.data is optimized for parallel processing:

import tensorflow as tf
from tensorflow.keras.preprocessing.image import ImageDataGenerator

# Set up your base ImageDataGenerator (keep images at original 192x192 size)
datagen = ImageDataGenerator(
    # Add your regular preprocessing here (e.g., rescale=1./255)
)

original_generator = datagen.flow_from_directory(
    "your_data_directory",
    target_size=(192, 192),  # Don't scale—we'll crop instead
    batch_size=32,
    class_mode="categorical"  # Adjust based on your task
)

# Define the central crop logic (150/192 ≈ 0.78125 crop fraction)
def crop_batch(batch_x, batch_y):
    crop_fraction = 150 / 192
    batch_x = tf.image.central_crop(batch_x, crop_fraction)
    # Ensure exact 150x150 output (in case of minor rounding)
    batch_x = tf.image.resize(batch_x, (150, 150), method="bilinear")
    return batch_x, batch_y

# Convert to tf.data.Dataset and apply parallel processing
dataset = tf.data.Dataset.from_generator(
    lambda: original_generator,
    output_signature=(
        tf.TensorSpec(shape=(None, 192, 192, 3), dtype=tf.float32),
        tf.TensorSpec(shape=(None, num_classes), dtype=tf.float32)  # Replace num_classes with your actual count
    )
).map(crop_batch, num_parallel_calls=tf.data.AUTOTUNE)

# Train your model with the optimized dataset
model.fit(dataset, epochs=10)

Option B: Use preprocessing_function

If you prefer sticking closer to ImageDataGenerator, you can pass the TF crop logic directly to preprocessing_function:

def central_crop(img):
    crop_fraction = 150 / 192
    img = tf.image.central_crop(img, crop_fraction)
    img = tf.image.resize(img, (150, 150))
    return img

datagen = ImageDataGenerator(
    preprocessing_function=central_crop,
    # Add other preprocessing steps here
)

generator = datagen.flow_from_directory(
    "your_data_directory",
    target_size=(192, 192),
    batch_size=32
)

model.fit(generator, epochs=10)

2. Vectorize Your Numpy Operations (No TF Dependency)

If you want to stick with pure NumPy, you can eliminate the per-image loop by slicing the entire batch at once. NumPy batch operations run in optimized C code, which is way faster than Python loops:

import numpy as np
from tensorflow.keras import backend as K

def batch_central_crop(batch_x, crop_size):
    if K.image_data_format() == "channels_last":
        height, width = batch_x.shape[1], batch_x.shape[2]
        dy, dx = crop_size
        start_y = (height - dy) // 2
        start_x = (width - dx) // 2
        return batch_x[:, start_y:start_y+dy, start_x:start_x+dx, :]
    else:
        height, width = batch_x.shape[2], batch_x.shape[3]
        dy, dx = crop_size
        start_y = (height - dy) // 2
        start_x = (width - dx) // 2
        return batch_x[:, :, start_y:start_y+dy, start_x:start_x+dx]

def crop_generator(batches, crop_length):
    while True:
        batch_x, batch_y = next(batches)
        batch_crops = batch_central_crop(batch_x, (crop_length, crop_length))
        yield (batch_crops, batch_y)

3. Add CenterCrop Directly to Your Model

Another clean option is to integrate the cropping into your model architecture. This way, cropping is part of the model graph and optimized for GPU/TPU execution:

from tensorflow.keras.layers import Input, CenterCrop, Conv2D, Dense, Flatten
from tensorflow.keras.models import Model

# Define input shape as original 192x192
input_layer = Input(shape=(192, 192, 3))
# Add center crop layer first
cropped_input = CenterCrop(height=150, width=150)(input_layer)
# Build the rest of your model
x = Conv2D(32, (3,3), activation="relu")(cropped_input)
x = Flatten()(x)
output_layer = Dense(num_classes, activation="softmax")(x)

model = Model(inputs=input_layer, outputs=output_layer)

# Use ImageDataGenerator without custom cropping
datagen = ImageDataGenerator(rescale=1./255)
generator = datagen.flow_from_directory(
    "your_data_directory",
    target_size=(192, 192),
    batch_size=32
)

model.fit(generator, epochs=10)

Why Your Original Code Was Slow

The main bottleneck was the Python for loop processing each image individually. Python loops are inherently slow compared to vectorized operations (like TF/NumPy batch slicing), which can process the entire batch in parallel using optimized low-level code.


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:48:55