mmseg实现oversampling,undersampling,cutmix

MMSegmentation 的 Dataset 支持读取一个 .txt 文件(里面写满图片的名字,不带后缀)。你可以把这个 .txt 文件当作一个“过滤器(Filter)”。

oversampling

第一步:生成两个 .txt 列表文件

你需要写一个几十行的简单 Python 脚本,遍历你的掩码(Mask)文件夹。

import os
import cv2
import numpy as np
from tqdm import tqdm

# ================= 配置区域 =================
# 你的 Mask 文件夹路径 (请修改为你真实的路径)
mask_dir = 'data/your_data/masks' 

# 生成的 txt 文件保存路径
output_dir = 'data/your_data/splits'

# 定义什么是“背景像素” (通常背景是 0)
bg_value = 0 
# ============================================

def main():
    # 确保输出目录存在
    os.makedirs(output_dir, exist_ok=True)
    
    normal_txt_path = os.path.join(output_dir, 'normal_train.txt')
    defect_txt_path = os.path.join(output_dir, 'defect_train.txt')
    
    normal_count = 0
    defect_count = 0

    print(f"正在扫描掩码文件夹: {mask_dir}")
    
    # 获取所有的 mask 文件 (假设后缀是 .png, 也可以是 .jpg等)
    mask_files = [f for f in os.listdir(mask_dir) if f.endswith('.png') or f.endswith('.jpg')]
    
    with open(normal_txt_path, 'w') as f_normal, open(defect_txt_path, 'w') as f_defect:
        for mask_name in tqdm(mask_files, desc="处理进度"):
            mask_path = os.path.join(mask_dir, mask_name)
            
            # 使用 cv2 以灰度模式读取掩码图
            mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)
            
            if mask is None:
                print(f"警告: 无法读取图像 {mask_path},已跳过。")
                continue
            
            # 提取不带后缀的文件名 (例如: img_001.png -> img_001)
            file_id = os.path.splitext(mask_name)[0]
            
            # 判断逻辑:如果图像中的最大像素值等于背景值,说明全图都是背景
            # 或者使用 np.all(mask == bg_value)
            if np.all(mask == bg_value):
                # 纯背景图
                f_normal.write(f"{file_id}\n")
                normal_count += 1
            else:
                # 包含缺陷
                f_defect.write(f"{file_id}\n")
                defect_count += 1

    print("\n✅ 数据划分完成!")
    print(f"总计找到 {normal_count} 张纯背景图片,已写入 {normal_txt_path}")
    print(f"总计找到 {defect_count} 张缺陷图片,已写入 {defect_txt_path}")

if __name__ == '__main__':
    main()

如果一张 Mask 全是黑的(只有背景),就把它的名字写进 normal_train.txt。

如果一张 Mask 里有缺陷像素,就把它的名字写进 defect_train.txt。

你的 splits/ 文件夹下会有这两个文件,内容大概长这样:

# normal_train.txt
img_0001
img_0003
img_0004
...
# defect_train.txt
img_0002
img_0005
...

第二步:在 Config 文件中实现“混合路径下的分离过采样”

现在,你的正常图和缺陷图依然安安静静地躺在同一个文件夹里,但我们在配置代码中,通过 ann_file 让它们在逻辑上分道扬镳,并对缺陷数据套上 RepeatDataset!

请看这段极其优雅的 Config 代码:

dataset_type = 'YourDefectDataset'
data_root = 'data/your_data/' # 根目录
data_prefix = dict(img_path='images', seg_map_path='masks') # 核心:指向同一个混合文件夹!

# 1. 正常背景数据集(通过 normal_train.txt 过滤,只加载那80张)
dataset_normal = dict(
    type=dataset_type,
    data_root=data_root,
    ann_file='splits/normal_train.txt',  # 👈 核心参数:仅读取正常图片的列表
    data_prefix=data_prefix,             # 指向混合文件夹
    pipeline=train_pipeline)

# 2. 缺陷数据集(通过 defect_train.txt 过滤,加载那20张,并放大4倍)
dataset_defect = dict(
    type='RepeatDataset',
    times=4,  # 重复 4 次,实现过采样
    dataset=dict(
        type=dataset_type,
        data_root=data_root,
        ann_file='splits/defect_train.txt',  # 👈 核心参数:仅读取缺陷图片的列表
        data_prefix=data_prefix,             # 依然指向同一个混合文件夹!
        pipeline=train_pipeline)
)

# 3. 在 Dataloader 中把它们拼接起来
train_dataloader = dict(
    batch_size=8,
    num_workers=4,
    persistent_workers=True,
    sampler=dict(type='InfiniteSampler', shuffle=True),
    dataset=dict(
        type='ConcatDataset',  
        datasets=[dataset_normal, dataset_defect] # 逻辑合并:80张正常 + (20x4)张缺陷
    )
)

undersampling

第一步:修改 Python 脚本(加入随机丢弃逻辑)

基于我们刚才写的 split_dataset.py,我们只需要加入 random.sample 来对纯背景列表进行“抽样缩水”。你可以把这个新脚本另存为 split_dataset_undersample.py:

import os
import cv2
import numpy as np
import random
from tqdm import tqdm

# ================= 配置区域 =================
mask_dir = 'data/your_data/masks' 
output_dir = 'data/your_data/splits'
bg_value = 0 

# 欠采样核心参数:背景保留率 (例如 0.25 表示只保留 25% 的背景图)
bg_keep_ratio = 0.25 
# ============================================

def main():
    os.makedirs(output_dir, exist_ok=True)
    
    # 动态命名,把保留率写在文件名里,方便做消融实验!
    normal_txt_path = os.path.join(output_dir, f'normal_train_keep_{bg_keep_ratio}.txt')
    defect_txt_path = os.path.join(output_dir, 'defect_train.txt')
    
    all_normals = []
    all_defects = []

    print(f"正在扫描掩码文件夹: {mask_dir}")
    mask_files = [f for f in os.listdir(mask_dir) if f.endswith('.png') or f.endswith('.jpg')]
    
    for mask_name in tqdm(mask_files, desc="分类进度"):
        mask_path = os.path.join(mask_dir, mask_name)
        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)
        
        if mask is None:
            continue
            
        file_id = os.path.splitext(mask_name)[0]
        
        if np.all(mask == bg_value):
            all_normals.append(file_id) # 先存入列表,不急着写文件
        else:
            all_defects.append(file_id)

    # ============= 欠采样核心逻辑 =============
    # 计算需要保留的背景图数量
    keep_count = int(len(all_normals) * bg_keep_ratio)
    # 设定随机种子以保证绝对的可复现性 (Reproducibility)
    random.seed(42) 
    # 随机抽取指定数量的背景图
    undersampled_normals = random.sample(all_normals, keep_count)
    # ==========================================

    # 写入文件
    with open(normal_txt_path, 'w') as f_normal:
        for file_id in undersampled_normals:
            f_normal.write(f"{file_id}\n")

    with open(defect_txt_path, 'w') as f_defect:
        for file_id in all_defects:
            f_defect.write(f"{file_id}\n")

    print("\n✅ 欠采样数据划分完成!")
    print(f"原始背景图总数: {len(all_normals)},保留率 {bg_keep_ratio},实际写入 {keep_count} 张 -> {normal_txt_path}")
    print(f"缺陷图片总数: {len(all_defects)},全量写入 -> {defect_txt_path}")

if __name__ == '__main__':
    main()

第二步:极度清爽的 Config 配置

因为欠采样已经在上一步的 .txt 文件里做好了(列表变短了),所以在 MMSegmentation 的 Config 里,我们不再需要使用任何 Wrapper(不需要 RepeatDataset),直接把两个 .txt 文件拼起来就行!

dataset_type = 'YourDefectDataset'
data_root = 'data/your_data/' 
data_prefix = dict(img_path='images', seg_map_path='masks')

# 1. 欠采样后的正常背景数据集 (比如原本80张,现在txt里只有20张)
dataset_normal_undersampled = dict(
    type=dataset_type,
    data_root=data_root,
    ann_file='splits/normal_train_keep_0.25.txt',  # 👈 核心:读取被砍掉 75% 后的列表
    data_prefix=data_prefix,
    pipeline=train_pipeline)

# 2. 缺陷数据集 (20张原封不动)
dataset_defect = dict(
    type=dataset_type,
    data_root=data_root,
    ann_file='splits/defect_train.txt',
    data_prefix=data_prefix,
    pipeline=train_pipeline
)

# 3. 拼接
train_dataloader = dict(
    batch_size=8,
    dataset=dict(
        type='ConcatDataset',  
        # 物理合并:20张正常 + 20张缺陷 = 40张的极小数据集
        datasets=[dataset_normal_undersampled, dataset_defect] 
    )
)

cutmix

骤一:创建自定义 Data Preprocessor

在你项目代码的某个合适位置(例如 mmseg/models/cutmix_preprocessor.py),新建一个 Python 文件,并写入以下代码。

这个类继承了 MMSeg 原本的 SegDataPreProcessor,我们在它做完基础的归一化(Normalize)和填充(Pad)之后,再执行 CutMix 操作。

import torch
import numpy as np
from mmseg.registry import MODELS
from mmseg.models.data_preprocessor import SegDataPreProcessor

@MODELS.register_module()
class CutMixSegDataPreProcessor(SegDataPreProcessor):
    def __init__(self, cutmix_prob=0.5, alpha=1.0, **kwargs):
        """
        基于 MMSegmentation 1.x 的 CutMix 数据预处理器
        Args:
            cutmix_prob (float): 触发 CutMix 的概率,默认 0.5
            alpha (float): Beta 分布的参数,用于生成裁剪比例,默认 1.0
            **kwargs: 继承自 SegDataPreProcessor 的参数 (如 mean, std, bgr_to_rgb 等)
        """
        super().__init__(**kwargs)
        self.cutmix_prob = cutmix_prob
        self.alpha = alpha

    def forward(self, data: dict, training: bool = False) -> dict:
        # 1. 先执行父类的标准操作(将数据移动到 GPU、Normalize、Pad 等)
        data = super().forward(data, training)
        
        # 2. 如果是测试/验证阶段,或者没有触发概率,则直接返回原始数据
        if not training or torch.rand(1).item() > self.cutmix_prob:
            return data

        inputs, data_samples = data['inputs'], data['data_samples']
        batch_size = inputs.size(0)
        
        # 如果 batch_size 小于 2,无法进行混合操作
        if batch_size < 2:
            return data

        # ================= 核心 CutMix 逻辑 =================
        
        # 生成一个随机打乱的索引列表,用于寻找“另一张图 (Image B)”
        rand_index = torch.randperm(batch_size).to(inputs.device)

        # 3. 从 Beta 分布中采样一个 lambda 值 (决定保留的面积比例)
        lam = np.random.beta(self.alpha, self.alpha)
        
        # 获取图像的高和宽
        _, _, H, W = inputs.shape
        
        # 计算裁剪框的长宽比例
        cut_rat = np.sqrt(1. - lam)
        cut_w = int(W * cut_rat)
        cut_h = int(H * cut_rat)

        # 随机选择裁剪框的中心点
        cx = np.random.randint(W)
        cy = np.random.randint(H)

        # 计算裁剪框的四个边界 (防止越界)
        bbx1 = np.clip(cx - cut_w // 2, 0, W)
        bby1 = np.clip(cy - cut_h // 2, 0, H)
        bbx2 = np.clip(cx + cut_w // 2, 0, W)
        bby2 = np.clip(cy + cut_h // 2, 0, H)

        # 4. 混合图像 (Images):把 Image B 的方块贴到 Image A 上
        inputs[:, :, bby1:bby2, bbx1:bbx2] = inputs[rand_index, :, bby1:bby2, bbx1:bbx2]

        # 5. 混合掩码 (Masks / Ground Truth):把 Mask B 的方块贴到 Mask A 上
        for i in range(batch_size):
            # 获取原始 Mask 张量 (Mask A)
            mask_a = data_samples[i].gt_sem_seg.data
            
            # 获取打乱索引对应的 Mask 张量 (Mask B)
            mask_b = data_samples[rand_index[i]].gt_sem_seg.data
            
            # 替换对应区域的掩码
            mask_a[:, bby1:bby2, bbx1:bbx2] = mask_b[:, bby1:bby2, bbx1:bbx2]

        # 将修改后的 inputs 写回 data 字典
        data['inputs'] = inputs
        
        return data

步骤二:在 Config 文件中激活

写好 Python 代码后,你需要告诉 MMSegmentation 在构建模型时使用你写的这个 CutMixSegDataPreProcessor,而不是默认的。

打开你正在使用的配置文件(比如 configs/fcn/fcn...py),找到 data_preprocessor 部分并进行修改:

# 1. 确保在配置文件开头导入了你自定义的模块
# 这里的路径取决于你把 cutmix_preprocessor.py 放在了哪里
custom_imports = dict(imports=['mmseg.models.data_preprocessor.cutmix_preprocessor'], allow_failed_imports=False)

# 2. 定义包含 CutMix 逻辑的数据预处理器
data_preprocessor = dict(
    type='CutMixSegDataPreProcessor', # 👈 指向你刚刚写的类
    cutmix_prob=0.5,                  # 👈 设定 CutMix 触发的概率 (0.5 是常用默认值)
    alpha=1.0,                        # 👈 Beta 分布参数
    # 下面这些是基础的预处理参数,通常和默认的一致
    mean=[123.675, 116.28, 103.53],   
    std=[58.395, 57.12, 57.375],      
    bgr_to_rgb=True,
    pad_val=0,
    seg_pad_val=255)                  # 忽略标签的值,通常是 255

# 3. 将这个 data_preprocessor 传给模型
model = dict(
    type='EncoderDecoder', # 或者是你们特定的 segmentor 类型
    data_preprocessor=data_preprocessor, # 👈 应用到模型中
    # ... 其他模型配置 (backbone, decode_head 等) ...
)
最后编辑于
©著作权归作者所有,转载或内容合作请联系作者
【社区内容提示】社区部分内容疑似由AI辅助生成,浏览时请结合常识与多方信息审慎甄别。
平台声明:文章内容(如有图片或视频亦包括在内)由作者上传并发布,文章内容仅代表作者本人观点,简书系信息发布平台,仅提供信息存储服务。

友情链接更多精彩内容