计算YOLO数据集中每个类的目标数

这里写自定义目录标题

运行命令:

bash 复制代码
python count_labels.py --data data.yaml --labels labels

计算YOLO数据集中每个类的目标数。类名从data.yaml读取。

bash 复制代码
import argparse
from collections import Counter
from pathlib import Path

import yaml


def count_dir(label_dir: Path):
    """Count (class_id -> object count) and (class_id -> image count) under a dir."""
    obj_count = Counter()
    img_count = Counter()
    files = sorted(label_dir.rglob('*.txt'))
    for f in files:
        cls_in_img = set()
        if f.stat().st_size == 0:
            continue
        for line in f.read_text().strip().splitlines():
            parts = line.split()
            if len(parts) != 5:
                continue
            cls = int(float(parts[0]))
            obj_count[cls] += 1
            cls_in_img.add(cls)
        for c in cls_in_img:
            img_count[c] += 1
    return obj_count, img_count, len(files)


def main():
    p = argparse.ArgumentParser(description='Count YOLO objects per class from data.yaml')
    p.add_argument('--data', required=True, help='/mnt/sda/rdd/datasets_big2small/data640_640/data.yaml')
    p.add_argument('--labels', required=True, help='/mnt/sda/rdd/datasets_big2small/data640_640/labels/train')
    args = p.parse_args()

    with open(args.data) as f:
        cfg = yaml.safe_load(f)
    names = cfg.get('names') or [str(i) for i in range(cfg.get('nc', 1))]
    print(f'classes (nc={len(names)}): {names}\n')

    root = Path(args.labels)
    splits = [d for d in root.iterdir() if d.is_dir()] if root.is_dir() else []
    if splits:
        targets = sorted(splits)
        print(f'found splits: {[d.name for d in targets]}')
        totals = Counter()
        img_totals = Counter()
        for d in targets:
            oc, ic, nf = count_dir(d)
            print(f'\n--- {d.name} ({nf} label files) ---')
            for i, name in enumerate(names):
                print(f'  {i:>2} {name:<12} objects={oc[i]:>6}  images={ic[i]:>6}')
            totals += oc
            img_totals += ic
        print('\n--- TOTAL ---')
        for i, name in enumerate(names):
            print(f'  {i:>2} {name:<12} objects={totals[i]:>6}  images={img_totals[i]:>6}')
    else:
        oc, ic, nf = count_dir(root)
        print(f'--- {root} ({nf} label files) ---')
        for i, name in enumerate(names):
            print(f'  {i:>2} {name:<12} objects={oc[i]:>6}  images={ic[i]:>6}')


if __name__ == '__main__':
    main()
相关推荐
倒头就睡的小比特1 天前
算法竞赛C++常用的STL
c++·算法
小羊没烦恼!1 天前
初探性能优化——2个月到4小时的性能提升
java·开发语言·windows·算法·c#
猎头南楼1 天前
知识社区推荐系统实践:新用户冷启动与长短期兴趣建模的挑战 资深推荐算法工程师
人工智能·深度学习·算法·机器学习
高洁011 天前
未来三年,AI 落地的五个确定性判断一
人工智能·深度学习·机器学习·transformer·知识图谱
旖旎夜光1 天前
力控面试题 01.01: 判定字符是否唯一(位运算) —— 题解
c++·学习·算法·leetcode·力控
wzdark1 天前
大规模并行计算中的负载均衡算法研究4
算法
Because_of_Her11 天前
并查集-听课笔记
笔记·算法·并查集
码流子1 天前
高速公路安全监测实践:碰撞监测预警+物联网底座,从感知到处置的闭环
大数据·人工智能·物联网·算法·架构
another heaven1 天前
【算法/C++ MD5算法能否逆解码?原理、C++实现与同类哈希算法对比】
c++·算法·哈希算法
wzdark1 天前
从算法设计模式看编程思维的抽象能力4
算法