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

基于全卷积网络的目标检测小数据集增强及复杂场景权重应用问询

Weighted Data Augmentation Solutions for Complex Object Detection with Small Datasets

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_range or width_shift_range to 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 09:01:33