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 等) ...
)