在计算机视觉任务中,关键点检测(如人体姿态估计、面部关键点检测等)需要将标注数据转换为模型训练所需的格式。本文将详细介绍如何将JSON格式的关键点数据集转换为TXT格式,并合理划分训练集、验证集和测试集。

一、JSON转TXT格式详解

1.JSON格式解析

以常见的COCO关键点标注格式为例,JSON文件通常包含以下字段:

images: 图像信息(文件名、尺寸等)
annotations: 标注信息(边界框、关键点坐标及可见性)
categories: 类别定义(如关键点名称、骨架连接关系)

import os
import cv2
import numpy as np
import glob
import json
import tqdm

# 物体类别
class_list = ['person']
# 关键点的顺序
keypoint_list = ['1', '2']

# 新的txt文件保存路径(假设为当前目录下的"output_txts/")
output_folder = "./Yolo_Data/"


def create_output_folder_if_not_exists(folder_path):
    if not os.path.exists(folder_path):
        os.makedirs(folder_path)


def json_to_yolo(img_data, json_data):
    h, w = img_data.shape[:2]

    rectangles = {}
    # 遍历初始化
    for shape in json_data["shapes"]:
        label = shape["label"]  # pen, head, tail
        points = shape["points"]  # x,y coordinates
        shape_type = shape["shape_type"]

        # 只处理矩形
        if shape_type == "rectangle":
            if label in class_list:
                rect_id = len(rectangles)
                rectangles[rect_id] = {
                    "label": label,
                    "rect": points[0] + points[1],  # Rectangle [x1, y1, x2, y2]
                    "keypoints_list": []
                }

    # 遍历更新,将点加入对应的目标框中
    for shape in json_data["shapes"]:
        label = shape["label"]
        points = shape["points"]
        # 如果匹配到了对应的keypoint
        if label in keypoint_list:
            for rect_id, rectangle in rectangles.items():
                x1, y1, x2, y2 = rectangle["rect"]
                if x1 <= points[0][0] <= x2 and y1 <= points[0][1] <= y2:
                    rectangles[rect_id]["keypoints_list"].append(points[0])
                    break

    # 转为yolo格式
    yolo_list = []
    for id, rectangle in rectangles.items():
        result_list = []
        label_id = class_list.index(rectangle["label"])
        # x1,y1,x2,y2
        x1, y1, x2, y2 = rectangle["rect"]
        # center_x, center_y, width, height
        center_x = (x1 + x2) / 2
        center_y = (y1 + y2) / 2
        width = abs(x1 - x2)
        height = abs(y1 - y2)
        # normalize
        center_x /= w
        center_y /= h
        width /= w
        height /= h

        # 保留6位小数
        center_x = round(center_x, 6)
        center_y = round(center_y, 6)
        width = round(width, 6)
        height = round(height, 6)

        # 添加 label_id, center_x, center_y, width, height
        result_list = [label_id, center_x, center_y, width, height]

        # 添加 p1_x, p1_y, p1_v, p2_x, p2_y, p2_v
        for point in rectangle["keypoints_list"]:
            x, y = point
            x, y = int(x), int(y)
            # normalize
            x /= w
            y /= h
            # 保留6位小数
            x = round(x, 6)
            y = round(y, 6)

            result_list.extend([x, y, 2])

        yolo_list.append(result_list)

    return yolo_list


# 获取所有的图片
img_list = glob.glob("./allimage/*.jpg")

# 创建新的txt文件保存文件夹(如果不存在)
create_output_folder_if_not_exists(output_folder)

for img_path in tqdm.tqdm(img_list):

    img = cv2.imread(img_path)
    print(img_path)
    json_file = img_path.replace('jpg', 'json')
    with open(json_file) as json_file:
        json_data = json.load(json_file)

    yolo_list = json_to_yolo(img, json_data)

    # 更新txt文件保存路径,将其保存到新的文件夹下
    base_filename = os.path.basename(img_path).replace('jpg', 'txt')
    new_yolo_txt_path = os.path.join(output_folder, base_filename)

    with open(new_yolo_txt_path, "w") as f:
        for yolo in yolo_list:
            for i in range(len(yolo)):
                if i == 0:
                    f.write(str(yolo[i]))
                else:
                    f.write(" " + str(yolo[i]))
            f.write("\n")

二、数据集划分方法

1. 划分比例建议

训练集:80%
验证集:20%

import os
import random
import shutil


def create_folders(root_folder):
    images_folder = os.path.join(root_folder, 'images')
    labels_folder = os.path.join(root_folder, 'labels')

    os.makedirs(images_folder, exist_ok=True)
    os.makedirs(labels_folder, exist_ok=True)

    train_images_folder = os.path.join(images_folder, 'train')
    val_images_folder = os.path.join(images_folder, 'val')
    train_labels_folder = os.path.join(labels_folder, 'train')
    val_labels_folder = os.path.join(labels_folder, 'val')

    os.makedirs(train_images_folder, exist_ok=True)
    os.makedirs(val_images_folder, exist_ok=True)
    os.makedirs(train_labels_folder, exist_ok=True)
    os.makedirs(val_labels_folder, exist_ok=True)

    return train_images_folder, val_images_folder, train_labels_folder, val_labels_folder


def split_dataset(image_folder, label_folder, root_folder, train_ratio=0.8):
    images = os.listdir(image_folder)
    random.shuffle(images)
    split_index = int(train_ratio * len(images))

    train_images = images[:split_index]
    val_images = images[split_index:]

    for img in train_images:
        img_path = os.path.join(image_folder, img)
        label_filename = os.path.basename(img).replace('.jpg', '.txt')
        label_path = os.path.join(label_folder, label_filename)

        shutil.copy(img_path, os.path.join(root_folder, 'images', 'train'))
        shutil.copy(label_path, os.path.join(root_folder, 'labels', 'train'))

    for img in val_images:
        img_path = os.path.join(image_folder, img)
        label_filename = os.path.basename(img).replace('.jpg', '.txt')
        label_path = os.path.join(label_folder, label_filename)

        shutil.copy(img_path, os.path.join(root_folder, 'images', 'val'))
        shutil.copy(label_path, os.path.join(root_folder, 'labels', 'val'))


root_folder = 'my_yolo'
image_folder = 'mydata/image'
label_folder = 'mydata/labels'

train_images_folder, val_images_folder, train_labels_folder, val_labels_folder = create_folders(root_folder)
split_dataset(image_folder, label_folder, root_folder)

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐