基于全卷积网络的目标检测小数据集增强及复杂场景权重应用问询
Great question! When working on complex object detection tasks with small datasets, the vanilla ImageDataGenerator setup you mentioned often isn't enough—especially if you're dealing with class imbalance, varying target importance (like small vs. large objects), or need to prioritize certain augmentations for harder samples. Below are practical, weighted approaches you can implement using Keras/TensorFlow:
1. Weighted Class Augmentation for Class Imbalance
If your dataset has skewed class distributions (e.g., rare objects that you need the model to learn better), you can assign higher augmentation weights to underrepresented classes. This ensures those samples get transformed more frequently during training.
Implementation Steps:
- First, calculate class weights based on your dataset's distribution.
- Create a custom generator that samples images with replacement, using class weights to prioritize minority classes, then applies augmentations.
from keras.preprocessing.image import ImageDataGenerator import numpy as np from sklearn.utils.class_weight import compute_class_weight # Assume you have your training data and labels loaded train_images = ... # Shape: (num_samples, height, width, channels) train_labels = ... # Shape: (num_samples,) - class indices # Calculate class weights class_weights = compute_class_weight('balanced', classes=np.unique(train_labels), y=train_labels) class_weight_dict = dict(enumerate(class_weights)) # Initialize base augmenter datagen = ImageDataGenerator( rotation_range=40, width_shift_range=0.2, height_shift_range=0.2, rescale=1./255, shear_range=0.2, zoom_range=0.2, horizontal_flip=True, fill_mode='nearest' ) # Custom weighted generator def weighted_aug_generator(images, labels, datagen, class_weight_dict, batch_size=32): while True: # Sample indices based on class weights sample_probs = [class_weight_dict[label] for label in labels] sample_probs = np.array(sample_probs) / np.sum(sample_probs) batch_indices = np.random.choice(len(images), size=batch_size, p=sample_probs) batch_images = images[batch_indices] batch_labels = labels[batch_indices] # Apply augmentations augmented_images = [] for img in batch_images: augmented_img = datagen.random_transform(img) augmented_images.append(augmented_img) yield np.array(augmented_images), np.array(batch_labels) # Use the generator in training train_generator = weighted_aug_generator(train_images, train_labels, datagen, class_weight_dict)
2. Object-Aware Weighted Augmentation (Target Detection-Specific)
For object detection, not all objects in an image are equal—small objects or rare targets often need more targeted augmentations (e.g., zooming in to make them more visible). You can assign weights to individual objects and adjust augmentation parameters based on those weights.
Implementation Idea:
- For each image, iterate over its bounding boxes/targets.
- For high-weight targets (e.g., small objects), increase the probability of augmentations like
zoom_rangeorwidth_shift_rangeto emphasize those targets.
def object_weighted_transform(img, bboxes, target_weights): # target_weights: list of weights for each object in the image (0-1) datagen = ImageDataGenerator() # Adjust augmentation parameters based on average target weight avg_weight = np.mean(target_weights) # Higher weight = more aggressive zoom/shift for small targets zoom_factor = 0.1 + avg_weight * 0.2 shift_factor = 0.1 + avg_weight * 0.1 # Apply random transform with adjusted params transform_params = datagen.get_random_transform(img.shape) transform_params['zoom_range'] = (1 - zoom_factor, 1 + zoom_factor) transform_params['width_shift_range'] = shift_factor transform_params['height_shift_range'] = shift_factor augmented_img = datagen.apply_transform(img, transform_params) # Also transform bounding boxes (critical for detection!) augmented_bboxes = [] for bbox in bboxes: x1, y1, x2, y2 = bbox # Apply same transform to bbox coordinates new_x1 = x1 * transform_params['zx'] + transform_params['tx'] new_y1 = y1 * transform_params['zy'] + transform_params['ty'] new_x2 = x2 * transform_params['zx'] + transform_params['tx'] new_y2 = y2 * transform_params['zy'] + transform_params['ty'] augmented_bboxes.append([new_x1, new_y1, new_x2, new_y2]) return augmented_img, np.array(augmented_bboxes) # Use this in your detection pipeline: # For each image and its bboxes/target weights, apply the weighted transform
3. Weighted Augmentation Operation Probabilities
Another approach is to assign weights to individual augmentation operations, so more useful operations (for your complex scenario) are applied more often. For example, if rotation harms your target's visibility but horizontal flip helps, you can adjust their probabilities.
Implementation:
class WeightedImageDataGenerator(ImageDataGenerator): def __init__(self, aug_weights=None, **kwargs): super().__init__(**kwargs) # aug_weights: dict like {'rotation': 0.3, 'horizontal_flip': 0.8, ...} self.aug_weights = aug_weights or {} def get_random_transform(self, img_shape): transform = super().get_random_transform(img_shape) # Override transform flags based on weights for aug, weight in self.aug_weights.items(): if np.random.rand() > weight: # Disable this augmentation for this sample if aug == 'rotation_range': transform['theta'] = 0.0 elif aug == 'horizontal_flip': transform['flip_horizontal'] = False elif aug == 'zoom_range': transform['zx'] = 1.0 transform['zy'] = 1.0 # Add handling for other augmentations as needed return transform # Usage: Prioritize horizontal flip and zoom over rotation weighted_datagen = WeightedImageDataGenerator( rotation_range=40, width_shift_range=0.2, height_shift_range=0.2, rescale=1./255, shear_range=0.2, zoom_range=0.2, horizontal_flip=True, fill_mode='nearest', aug_weights={'rotation': 0.2, 'horizontal_flip': 0.9, 'zoom_range': 0.7} ) # Use like a regular ImageDataGenerator generator = weighted_datagen.flow_from_directory(...)
All these approaches let you tailor augmentation to your complex detection scenario by adding weight-based logic—whether it's for class balance, target-specific needs, or operation priority.
内容的提问来源于stack exchange,提问作者Oblomov

