1. 环境准备与项目理解
嘿,朋友们,今天咱们来聊聊一个非常实用的任务:用PyTorch版的DeepLabV3+来训练你自己的数据集。我知道,很多朋友在网上找教程,照着做,但总会遇到各种奇奇怪怪的问题,比如数据格式不对、代码跑不通、训练结果一片黑。我自己在图像分割这个领域摸爬滚打了好些年,从最早的FCN到现在的各种变体,DeepLabV3+绝对是一个兼顾精度和效率的“老朋友”。它那个ASPP模块和编码器-解码器结构,对付复杂场景的物体边缘分割,效果一直很稳。
那么,我们为什么要自己训练呢?官方的预训练模型虽然强大,但它是基于PASCAL VOC、Cityscapes这些通用数据集训练的。如果你的任务是识别工厂里的特定零件、医疗影像中的特殊病灶,或者像我之前做过的农业场景里的成熟果实,那通用模型的效果可能就不尽如人意了。这时候,用自己的数据“喂”出一个专属模型,就成了必经之路。
这个过程听起来有点技术门槛,但别怕,我今天就带你走一遍完整的流程。我们会从一个最常见的起点开始:你手里有一批用YOLO格式标注好的图片。很多标注工具,包括一些高效的在线平台,默认输出就是YOLO的txt格式。我们的目标,就是把这批“原材料”,一步步加工成DeepLabV3+官方代码能“消化”的格式,然后训练、测试,最终得到一个能用的模型。我会把每一步的代码、可能踩的坑、以及背后的“为什么”都讲清楚,保证你跟着做就能出结果。
我的实验环境是Windows 11,PyTorch 2.2,Python 3.8,用VSCode编辑。你用Linux或者Mac也没问题,命令基本是通用的。好了,废话不多说,咱们开始动手。
2. 数据集准备:从YOLO到标准分割格式
这是最核心、也最容易出错的一步。DeepLabV3+的PyTorch官方实现(比如广泛使用的 pytorch-deeplab-xception 这个仓库)通常期望数据集结构模仿PASCAL VOC。所以,我们的任务就是把YOLO格式的数据,“翻译”成VOC格式。
2.1 理解数据集的“标准长相”
首先,我们得知道目标长什么样。一个标准的VOC格式分割数据集,目录结构通常如下:
你的数据集根目录(例如:MyDataset)
├── JPEGImages/ # 存放所有的原始图片,如 .jpg 文件
├── SegmentationClass/ # 存放所有对应的分割标签图,.png格式,像素值代表类别
└── ImageSets/
└── Segmentation/ # 存放划分好的文件列表
├── train.txt
└── val.txt
JPEGImages和SegmentationClass里的文件必须一一对应,且文件名(不含后缀)要完全相同。train.txt和val.txt里面就只写文件名,一行一个,不要后缀。比如:
image_001
image_002
而YOLO格式的标注呢?每张图片对应一个同名的.txt文件,里面内容可能是:
0 0.5 0.5 0.2 0.3
1 0.7 0.2 0.1 0.1
这代表:类别0,中心点(x,y)和宽高(w,h),都是相对于图片宽高归一化的值。这和分割任务需要的像素级标签图(Mask)完全是两码事。所以,我们需要进行格式转换。
2.2 第一步:YOLO格式转JSON格式
直接从YOLO的bbox转到Mask比较困难,我们可以借助一个中间格式——JSON。这里我采用LabelMe工具使用的JSON结构,因为它结构清晰,后面有现成的工具可以转Mask。
假设你的YOLO标注文件在一个叫 txt/ 的文件夹里,对应的原图在 img/ 文件夹里。下面这个脚本,会把每个.txt文件转换成一个包含多边形轮廓信息的JSON文件。这里有个关键点:YOLO格式是边界框,而分割需要的是物体的轮廓。一种常见的做法是,把边界框的四个顶点当作一个粗略的多边形。对于形状规则的物体,这勉强可用;但对于不规则物体,这会导致标签不精确,影响最终效果。如果你的项目对精度要求高,强烈建议用分割工具重新标注。 这里为了流程演示,我们先用边界框模拟。
import json
import os
# 配置你的路径
txt_folder_path = "你的路径/Seg552/txt/"
json_folder_path = "你的路径/Seg552/json/"
img_folder_path = "你的路径/Seg552/img/"
os.makedirs(json_folder_path, exist_ok=True)
# 你的类别映射关系,根据你YOLO标注的类别ID来修改
label_mapping = {0: "background", 1: "cat", 2: "dog"} # 示例
for txt_file in os.listdir(txt_folder_path):
if txt_file.endswith('.txt'):
base_name = os.path.splitext(txt_file)[0]
img_file = base_name + '.jpg' # 假设图片是jpg格式
img_path = os.path.join(img_folder_path, img_file)
# 获取图片尺寸,JSON里需要
from PIL import Image
try:
with Image.open(img_path) as img:
width, height = img.size
except FileNotFoundError:
print(f"警告:找不到图片 {img_path},跳过该标注文件。")
continue
json_data = {
"version": "5.2.1",
"flags": {},
"shapes": [],
"imagePath": img_file,
"imageHeight": height,
"imageWidth": width
}
txt_file_path = os.path.join(txt_folder_path, txt_file)
with open(txt_file_path, 'r') as f:
for line in f:
parts = line.strip().split()
if len(parts) != 5:
continue
class_id = int(parts[0])
x_center, y_center, w, h = map(float, parts[1:])
# 将归一化坐标转换为绝对像素坐标
x_center_abs = x_center * width
y_center_abs = y_center * height
w_abs = w * width
h_abs = h * height
# 计算边界框的四个顶点(左上、右上、右下、左下)
x1 = x_center_abs - w_abs / 2
y1 = y_center_abs - h_abs / 2
x2 = x_center_abs + w_abs / 2
y2 = y_center_abs + h_abs / 2
points = [[x1, y1], [x2, y1], [x2, y2], [x1, y2]]
shape_data = {
"label": label_mapping.get(class_id, "unknown"),
"points": points,
"group_id": None,
"shape_type": "polygon",
"flags": {}
}
json_data["shapes"].append(shape_data)
json_file_path = os.path.join(json_folder_path, base_name + '.json')
with open(json_file_path, 'w') as jf:
json.dump(json_data, jf, indent=2)
print(f"转换完成: {txt_file} -> {base_name}.json")
运


4841

被折叠的 条评论
为什么被折叠?



