背景

sahi是比较简单易行的小目标检测工具,yolo等工具都有了支持,官网上给出了非常明显的效果示意图(https://github.com/obss/sahi),自己实际测验效果也非常明显。

例如这是一张田野里的电线杆,不仔细看,你很难发现绝缘子上的一些破损。在这里插入图片描述
当采用sahi进行识别,可以精准的识别两个小缺口,这些缺口可能导致线路故障。(注意,这里对模型进行了特殊的训练,训练过程会在后面介绍)在这里插入图片描述

训练

sahi的论文里面提及了模型训练。个人认为,在对较大目标(例如上图中的绝缘子整体)进行sahi识别时,可以直接跑正常的yolo流程,但是对于小目标(例如绝缘子上的缺口),可能就需要额外的训练技巧了。

由于sahi是在640等尺寸上进行识别,因此我想到了在640的尺寸上进行单独的模型训练,这样能保证特征被充分学习(你也可以尝试放大图片,参考这个文章On the Importance of Large Objects in CNN Based Object Detection Algorithms)。

整体效果如下:

原始图像中有三个小汽车(假设小汽车是小目标)
在这里插入图片描述
处理结束后,整张图片被分成了两张子图片,可以利用yolo进行训练了。个人感觉这样训练效果更好。
在这里插入图片描述

代码

下面是我数据处理的代码,使用的时候请把640换成你需要的尺寸:

import os
import cv2
import xml.etree.ElementTree as ET

def read_xml(xml_path):
    tree = ET.parse(xml_path)
    root = tree.getroot()
    annotations = []
    for obj in root.findall('object'):
        name = obj.find('name').text
        bndbox = obj.find('bndbox')
        xmin = int(bndbox.find('xmin').text)
        ymin = int(bndbox.find('ymin').text)
        xmax = int(bndbox.find('xmax').text)
        ymax = int(bndbox.find('ymax').text)
        annotations.append({'name': name, 'bbox': (xmin, ymin, xmax, ymax)})
    return root, annotations

def process_images(image_folder, xml_folder, output_folder, target_label="qk", split_size=640):
    try:
        if not os.path.exists(output_folder):
            os.makedirs(output_folder)

        for filename in os.listdir(image_folder):
            print(filename)
            if filename.endswith(('.jpg', '.png')):
                image_path = os.path.join(image_folder, filename)
                xml_path = os.path.join(xml_folder, filename.replace('.jpg', '.xml').replace('.png', '.xml'))

                # 读取图像和XML标注
                img = cv2.imread(image_path)
                root, annotations = read_xml(xml_path)

                # 找出所有目标标签的位置
                target_labels = [ann for ann in annotations if target_label in ann['name']]
                if not target_labels:
                    continue

                # 每张图片维护一个processed_areas字典
                processed_areas=[]

                # 遍历所有目标标签进行切分处理
                for i, label in enumerate(target_labels):
                    xmin, ymin, xmax, ymax = label['bbox']

                    # 检查是否该标签已经被其他标签切分框覆盖
                    skip_label = False
                            
                    for processed_area in processed_areas:
                        # 获取已处理标签的位置
                        proc_xmin, proc_xmax, proc_ymin, proc_ymax = processed_area

                        # 检查当前标签与已处理标签是否重叠
                        if (xmin > proc_xmin and xmax < proc_xmax and ymin > proc_ymin and ymax < proc_ymax):
                            skip_label = True
                            break

                    if skip_label:
                        continue  # 跳过这个标签

                    # 计算切分区域,确保640x640区域包含标签
                    x_center, y_center = (xmin + xmax) // 2, (ymin + ymax) // 2
                    split_x_min = max(0, x_center - split_size // 2)
                    split_y_min = max(0, y_center - split_size // 2)
                    split_x_max = min(img.shape[1], split_x_min + split_size)
                    split_y_max = min(img.shape[0], split_y_min + split_size)

                    # 获取切割后的图像
                    cropped_img = img[split_y_min:split_y_max, split_x_min:split_x_max]
                    processed_areas.append([split_x_min,split_x_max, split_y_min,split_y_max])

                    # 重新计算标签位置
                    new_annotations = []
                    for ann in annotations:
                        if target_label in ann['name'] or 1:
                            old_xmin, old_ymin, old_xmax, old_ymax = ann['bbox']
                            if split_x_min <= old_xmin < split_x_max and split_y_min <= old_ymin < split_y_max:
                                # 标签位于切割区域内,计算新的坐标
                                new_xmin = old_xmin - split_x_min
                                new_ymin = old_ymin - split_y_min
                                new_xmax = old_xmax - split_x_min
                                new_ymax = old_ymax - split_y_min
                                new_annotations.append({'name': ann['name'], 'bbox': (new_xmin, new_ymin, new_xmax, new_ymax)})

                    # 为新文件生成符合规范的文件名
                    base_name = os.path.splitext(filename)[0]  # 获取原始文件名(不包含扩展名)
                    new_image_name = f"{base_name}_split_{i+1}.jpg"
                    new_xml_name = f"{base_name}_split_{i+1}.xml"

                    # 保存切割图像
                    output_image_path = os.path.join(output_folder, new_image_name)
                    cv2.imwrite(output_image_path, cropped_img)

                    # 复制原始XML并修改object节点
                    new_root = ET.Element("annotation")
                    for elem in root:
                        if elem.tag != "object":
                            new_root.append(elem)  # 保留原XML中非object的其他元素

                    # 只修改object节点
                    for new_ann in new_annotations:
                        object_elem = ET.SubElement(new_root, 'object')
                        name_elem = ET.SubElement(object_elem, 'name')
                        name_elem.text = new_ann['name']
                        bndbox_elem = ET.SubElement(object_elem, 'bndbox')
                        ET.SubElement(bndbox_elem, 'xmin').text = str(new_ann['bbox'][0])
                        ET.SubElement(bndbox_elem, 'ymin').text = str(new_ann['bbox'][1])
                        ET.SubElement(bndbox_elem, 'xmax').text = str(new_ann['bbox'][2])
                        ET.SubElement(bndbox_elem, 'ymax').text = str(new_ann['bbox'][3])

                    tree = ET.ElementTree(new_root)
                    output_xml_path = os.path.join(output_folder, new_xml_name)
                    tree.write(output_xml_path)
                    print("saved ", output_image_path, "and", output_xml_path)

    except Exception as e:
        print(e)

# 调用函数
image_folder = r'C:\Users\28715\Pictures\test'  # 请替换为你实际的图像文件夹路径
xml_folder = image_folder               # XML文件夹与图像文件夹相同
output_folder = r'C:\Users\28715\Pictures\testaaa'  # 请替换为你保存结果的文件夹路径
process_images(image_folder, xml_folder, output_folder, "car", 640)

更多推荐