创见博客
基于伪标签方法的语义分割模型的域适应
七崽爱吃小饼干2025/09/15阅读 5

方案预设

现在训练的语义分割模型是在ATR数据集上进行训练的,ATR数据集主要的特点是:中分辨率、个人室内或室外的时尚rgb照,而目前我进行分割的SYSU-MM01数据集的主要特点是:分别率低、个人室内或室外的rgb以及ir照,监控风格。存在比较大的分割差异,这也导致了模型的分割效果比较差。所以现在的目标就是让我的模型更具泛化能力,能够适应可见光红外行人重识别数据集的风格。

我提出的这种方法主要是参考了完全无监督行人重识别的伪标签方法。在完全无监督行人重识别中,主要通过聚类,生成伪标签,并且可以通过提高阈值来去掉不够可靠的标签。

具体来说,我预先在ATR数据集上训练好我的六分类语义分割模型,然后通过该语义分割模型对SYSU-MM01数据集进行分割处理。舍弃掉置信度较低的分割结果,将置信度比较高的结果作为伪标签,将数据混合到ATR数据集当中对模型进行训练,从而令模型适应可见光红外行人重识别数据集。

训练前后的模型,将在经过我从SYSU-MM01随机采样后标注的数据集上进行测试,对比前后的模型性能。

存在的一些问题:

  • 因为ATR数据集是RGB数据集,所以初始分割的IR图像得到的伪标签质量可能会很差。
  • 伪标签阈值的选择,如果阈值过高得到的伪标签数量就比较少,模型会难以学习到目标域的风格。如果阈值过低,会导致生成的伪标签的质量差,破坏模型的稳定性。
  • 测试集的样本是否够全面,能不能覆盖不同的场景以及情况。
  • 要注意测试集和训练集不要重合,这样会导致污染评估结果。

改进的方案:

  • 用CA技术,让ATR数据集的风格更倾向ir数据,从而让模型学习到与颜色无关的特征,从而能够得到更利好ir数据的伪标签。
  • 阈值的选择手动进行调整,目前没有特别好的阈值调整方案
  • 测试集是从不同cam下从不同的id下随机进行抽取的,共计ir和rgb各200张,覆盖了室内室外以及各种姿势。
  • 训练集排除了同一cam下的同一id,保证了两者不会重合,确保评估结果不被污染。然后从ir数据以及rgb数据下各按照生成测试集的方案(保证覆盖面比较全)抽取了3000张(考虑到根据阈值会弃用一部分样本),最后参与训练的伪标签大概会在4000多张,与测试集比例大概在10:1。

方案实施

实验一

先按照方案,从sysu_mm01数据集中抽取了4000张数据,避开了测试集。

训练集抽取脚本

python
import os
import random
import shutil
from collections import defaultdict


def load_test_cam_id_pairs(log_path):
    """提取测试集的 (相机, ID) 对,避免重复"""
    test_pairs = set()
    with open(log_path, 'r') as f:
        for line in f:
            parts = line.strip().split(' | ')
            if len(parts) < 3:
                continue
            cam = int(parts[1].split('相机')[1].split('(')[0])
            person_id = parts[2].split('ID')[1]
            test_pairs.add((cam, person_id))
    return test_pairs


def parse_sysu_mm01_structure(data_root, exclude_pairs):
    """解析数据集,仅排除 (相机, ID) 完全匹配的测试集样本"""
    cam_info = {
        1: {'type': 'rgb', 'scene': 'indoor'},
        2: {'type': 'rgb', 'scene': 'indoor'},
        3: {'type': 'ir', 'scene': 'indoor'},
        4: {'type': 'rgb', 'scene': 'outdoor'},
        5: {'type': 'rgb', 'scene': 'outdoor'},
        6: {'type': 'ir', 'scene': 'outdoor'}
    }

    cam_dict = defaultdict(lambda: defaultdict(list))  # cam -> id -> [img_paths]

    for cam in cam_info.keys():
        cam_dir = os.path.join(data_root, f'cam{cam}')
        if not os.path.exists(cam_dir):
            print(f"警告: 相机目录 {cam_dir} 不存在")
            continue

        for id_folder in os.listdir(cam_dir):
            if (cam, id_folder) in exclude_pairs:
                continue  # 排除测试集同一相机下的同一ID
            id_path = os.path.join(cam_dir, id_folder)
            if not os.path.isdir(id_path):
                continue

            # 收集该ID下的所有图像(用于多图抽取)
            for img_file in os.listdir(id_path):
                if img_file.endswith(('.jpg', '.png')):
                    img_path = os.path.join(id_path, img_file)
                    cam_dict[cam][id_folder].append(img_path)

    return cam_dict, cam_info


def select_pseudo_label_samples(cam_dict, cam_info, output_dir,
                                num_rgb=3000, num_ir=3000,
                                max_per_id=5):  # 每个ID最多抽5张(可调整)
    """从同一ID下抽取多张图像,提升伪标签总量"""
    os.makedirs(output_dir, exist_ok=True)
    os.makedirs(os.path.join(output_dir, 'rgb'), exist_ok=True)
    os.makedirs(os.path.join(output_dir, 'ir'), exist_ok=True)
    selection_log = []

    # 处理RGB相机 (1,2,4,5)
    rgb_cams = [1, 2, 4, 5]
    per_rgb_cam = num_rgb // len(rgb_cams)
    for cam in rgb_cams:
        cam_id_dict = cam_dict.get(cam, {})
        if not cam_id_dict:
            print(f"警告: 相机 {cam} 没有可用数据")
            continue

        collected = 0
        # 按ID遍历,每个ID抽取多张(直到满足该相机的目标数量)
        while collected < per_rgb_cam and cam_id_dict:
            # 随机选一个ID
            person_id = random.choice(list(cam_id_dict.keys()))
            img_paths = cam_id_dict[person_id]

            # 每个ID最多抽max_per_id张,或剩余需要的数量
            num_to_take = min(max_per_id, per_rgb_cam - collected, len(img_paths))
            if num_to_take <= 0:
                del cam_id_dict[person_id]  # 该ID已无可用图像,移除
                continue

            # 抽取num_to_take张图像
            selected_imgs = random.sample(img_paths, num_to_take)
            for img_path in selected_imgs:
                dest_path = os.path.join(output_dir, 'rgb',
                                         f'pseudo_cam{cam}_id{person_id}_{os.path.basename(img_path)}')
                shutil.copy2(img_path, dest_path)
                selection_log.append({
                    'type': 'rgb', 'cam': cam, 'scene': cam_info[cam]['scene'],
                    'person_id': person_id, 'src': img_path, 'dest': dest_path
                })

            collected += num_to_take
            print(f"相机{cam} RGB 已抽取 {collected}/{per_rgb_cam} 张")

        if collected < per_rgb_cam:
            print(f"警告: 相机{cam} 可用样本不足,仅抽取 {collected} 张")

    # 处理IR相机 (3,6)
    ir_cams = [3, 6]
    per_ir_cam = num_ir // len(ir_cams)
    for cam in ir_cams:
        cam_id_dict = cam_dict.get(cam, {})
        if not cam_id_dict:
            print(f"警告: 相机 {cam} 没有可用数据")
            continue

        collected = 0
        while collected < per_ir_cam and cam_id_dict:
            person_id = random.choice(list(cam_id_dict.keys()))
            img_paths = cam_id_dict[person_id]

            num_to_take = min(max_per_id, per_ir_cam - collected, len(img_paths))
            if num_to_take <= 0:
                del cam_id_dict[person_id]
                continue

            selected_imgs = random.sample(img_paths, num_to_take)
            for img_path in selected_imgs:
                dest_path = os.path.join(output_dir, 'ir',
                                         f'pseudo_cam{cam}_id{person_id}_{os.path.basename(img_path)}')
                shutil.copy2(img_path, dest_path)
                selection_log.append({
                    'type': 'ir', 'cam': cam, 'scene': cam_info[cam]['scene'],
                    'person_id': person_id, 'src': img_path, 'dest': dest_path
                })

            collected += num_to_take
            print(f"相机{cam} IR 已抽取 {collected}/{per_ir_cam} 张")

        if collected < per_ir_cam:
            print(f"警告: 相机{cam} 可用样本不足,仅抽取 {collected} 张")

    # 保存日志
    with open(os.path.join(output_dir, 'pseudo_label_log.txt'), 'w') as f:
        for entry in selection_log:
            f.write(
                f"{entry['type']} | 相机{entry['cam']}({entry['scene']}) | ID{entry['person_id']} | 来源: {entry['src']}\n")

    print(f"\n伪标签样本抽取完成! 共抽取 {len(selection_log)} 张图像")
    print(f"输出目录: {output_dir}")
    print(f"抽取日志: {os.path.join(output_dir, 'pseudo_label_log.txt')}")


if __name__ == "__main__":
    DATA_ROOT = "/Users/liujingmin/Desktop/project/Reid-SAM2/data/SYSU-MM01"
    TEST_LOG_PATH = "./selected_samples/selection_log.txt"
    OUTPUT_DIR = "./pseudo_label_samples"
    NUM_RGB = 3000  # 目标RGB总量(可调整)
    NUM_IR = 3000  # 目标IR总量(可调整)
    MAX_PER_ID = 5  # 每个ID最多抽5张(根据实际图像数量调整,建议3-5张)

    print("加载测试集 (相机, ID) 对...")
    test_pairs = load_test_cam_id_pairs(TEST_LOG_PATH)
    print(f"已排除 {len(test_pairs)} 个 (相机, ID) 对")

    print("解析数据集结构...")
    cam_dict, cam_info = parse_sysu_mm01_structure(DATA_ROOT, test_pairs)

    print("开始抽取伪标签样本(同一ID多图)...")
    select_pseudo_label_samples(cam_dict, cam_info, OUTPUT_DIR,
                                NUM_RGB, NUM_IR, MAX_PER_ID)

实验二

用预训练好的六分类模型,对训练集进行推理,根据阈值保留伪标签。

实验结果:

(ADCA) sh-5.1$ python generate_pseudo_labels.py 
开始处理伪标签样本,总样本数: 5458
使用设备: cuda,置信度阈值: 0.7
处理进度: 100%|█████████████████████████████████████████████████████| 683/683 [01:49<00:00,  6.24it/s]

处理完成!
总处理样本数: 5458
保留样本数: 5458 (占比 100.00%)
结果保存至: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels
统计信息: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels/filter_stats.txt

在置信度设置为0.7的时候,样本居然全部保留了,最后推测是因为背景占比比较大,背景的分割又比较简单,所以导致整体置信度就很高。

优化方案:

  • 计算置信度时,把背景类去除
  • 并且会计算前景像素的占比,如果占比太少,视为低质量样本(去除)
  • 把置信度提升到0.8

实验三

实验结果:

codeType
开始处理伪标签,总样本数: 5458
设备: cuda | 置信度阈值: 0.8
前景要求: 占比≥5% 且 像素数≥100
处理进度: 100%|██████████████████████████████████████████████████████| 683/683 [02:15<00:00,  5.03it/s]
处理完成!
保留样本: 2855/5458 (52.31%)
结果目录: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels
统计文件: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels/filter_stats.txt

检查了一下生成的mask的质量,比较一般,而且加了ca以后对可见光图像的分割效果感觉变差不少。

优化方案:

  • 进一步提高阈值,并且ir和rgb图片分别设置阈值
  • 不要对所有图像都做ca增强,随机对一些图像进行ca增强

实验四

经过上述调整后,重新进行了伪标签的生成,结果如下,样本保留数量和之前差不多,分割质量明显增强。但是存在一部分人体分割缺失特别多的,

codeType
(ADCA) sh-5.1$ python script/generate_pseudo_labels.py 
开始处理伪标签,总样本数: 5458
设备: cuda
RGB阈值 - 置信度: 0.85, 前景占比: 5%, 像素数: 100
IR阈值  - 置信度: 0.8, 前景占比: 3%, 像素数: 80
处理进度: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 683/683 [01:27<00:00,  7.76it/s]

处理完成!
RGB保留样本: 1357
IR保留样本: 1443
总保留比例: 51.30%
结果目录: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels
统计文件: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels/filter_stats.txt

实验五

由于一部分样本存在前景缺失严重的情况,所以要提高前景占比的阈值。

虽然抽取检查的样本里没有发现前景占比过高的情况,但是还是添加前景占比最高的阈值,提高留下的伪标签的质量。

实验结果

样本保留的比例还是比较高的,筛选掉了一些极端结果。

codeType
开始处理伪标签,总样本数: 5458
设备: cuda
RGB阈值 - 置信度: 0.85, 前景占比范围: 30%~95%
IR阈值  - 置信度: 0.8, 前景占比范围: 30%~95%
处理进度: 100%|██████████████████████████████████████████| 683/683 [01:24<00:00,  8.06it/s]

处理完成!
RGB保留样本: 1199
IR保留样本: 1008
总保留比例: 40.44%
结果目录: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels
统计文件: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels/filter_stats.txt
(ADCA) sh-5.1$ python script/mask_check.py 

实验六

在前面实验的基础上,这一次决定提高抽取的初始样本数量,然后进一步提高伪标签生成的阈值,获取更多质量更高的伪标签样本。

按照前面说的要和测试集比例为10:1,测试集规模为400个样本,所以最后需要4000的伪标签样本,目前的保留比例大概是40%,所以需要10000左右的初始样本。

实验结果: 抽取了10000张初始样本

codeType
伪标签样本抽取完成! 共抽取 8470 张图像
输出目录: ./pseudo_label_samples
抽取日志: ./pseudo_label_samples/pseudo_label_log.txt

最后伪标签样本数量大概3000,应该足够使用了。

codeType
(ADCA) sh-5.1$ python script/generate_pseudo_labels.py 
开始处理伪标签,总样本数: 8470
设备: cuda
RGB阈值 - 置信度: 0.85, 前景占比范围: 35%~90%
IR阈值  - 置信度: 0.8, 前景占比范围: 35%~90%
处理进度: 100%|█████████████████████████████████████████████████| 1059/1059 [02:11<00:00,  8.05it/s]

处理完成!
RGB保留样本: 1846
IR保留样本: 1292
总保留比例: 37.05%
结果目录: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels
统计文件: /home/Minda2024-9/liujingmin/datasets/sysu/filtered_pseudo_labels/filter_stats.txt

实验七

让模型在atr数据集以及伪标签数据集上进行训练。并且在ir图像上不使用通道变换增强。

进行风格训练后,mIOU比之前还要上涨了,对比之前最好的不加CA增强的0.836,以及随机CA增强的0.829,上升到了0.8823,剩下的就是等测试集标注出来以后,进行实验。

codeType
Epoch 14/20
Train Loss: 0.052842 | Val Loss: 0.055395
当前学习率: 0.000018
  Background IOU: 0.9743
  Head IOU: 0.8803
  UpperBody IOU: 0.8870
  ArmsHands IOU: 0.8307
  LowerBody IOU: 0.8839
  Shoes IOU: 0.8375
验证集mIOU: 0.8823
保存最佳模型...
学习率从 0.000018 调整为 0.000026
用肉眼看的话,在之前一直使用的两个例子上的效果明显都有上涨
评论
0/100