YOLOv9 pandas数据处理:标签统计与分析实战
YOLOv9 pandas数据处理:标签统计与分析实战
你是不是也遇到过这样的问题?好不容易收集了一大堆标注数据,准备用YOLOv9训练模型,结果训练效果总是不理想?模型在某些类别上表现很好,在另一些类别上却一塌糊涂?
很多时候,问题并不在模型本身,而是在数据上。数据分布不均衡、标签错误、样本数量差异过大,这些都会直接影响模型的训练效果。今天我就来分享一个实战技巧:用pandas对YOLOv9数据集进行标签统计与分析。
这个方法简单实用,不需要复杂的工具,只需要几行Python代码,就能让你对数据集了如指掌。通过分析标签分布,你能发现数据中的潜在问题,为后续的数据清洗、增强和模型训练提供有力指导。
1. 为什么需要标签统计分析?
在开始具体操作之前,我们先聊聊为什么标签分析这么重要。
1.1 数据质量直接影响模型效果
YOLOv9虽然强大,但它也是"巧妇难为无米之炊"。如果数据本身有问题,再好的模型也发挥不出应有的效果。常见的数据问题包括:
- 类别不平衡:某些类别的样本数量远远多于其他类别
- 标签错误:标注时的人为错误,比如把猫标成了狗
- 标注质量差:边界框不准确、漏标、多标等问题
- 样本多样性不足:所有图片都在相似的光线、角度、背景下拍摄
1.2 pandas:数据分析的瑞士军刀
pandas是Python中最流行的数据分析库,它提供了丰富的数据处理功能。对于YOLO标签分析来说,pandas有几个特别有用的特性:
- DataFrame结构:可以方便地存储和操作表格数据
- 强大的分组统计:按类别、按图片进行各种统计
- 可视化集成:可以配合matplotlib、seaborn进行数据可视化
- 高效处理:即使处理数万条标注数据,速度也很快
2. 环境准备与数据理解
2.1 检查镜像环境
如果你使用的是YOLOv9官方版训练与推理镜像,环境已经配置好了。我们快速确认一下:
# 激活YOLOv9环境
conda activate yolov9
# 检查pandas版本
python -c "import pandas as pd; print(f'pandas版本: {pd.__version__}')"
# 检查其他必要库
python -c "import numpy as np; import matplotlib.pyplot as plt; print('环境就绪')"
2.2 理解YOLO标签格式
YOLO的标签文件是纯文本格式,每个文件对应一张图片。格式如下:
<类别ID> <中心点x坐标> <中心点y坐标> <宽度> <高度>
例如:
0 0.5 0.5 0.3 0.4
1 0.2 0.3 0.1 0.2
这些坐标都是归一化后的值(0-1之间),相对于图片的宽度和高度。
2.3 准备示例数据集
为了演示,我准备了一个简单的车辆检测数据集,包含3个类别:
- 0: car(汽车)
- 1: bus(公交车)
- 2: truck(卡车)
数据集结构如下:
dataset/
├── images/
│ ├── train/
│ │ ├── 001.jpg
│ │ ├── 002.jpg
│ │ └── ...
│ └── val/
│ ├── 101.jpg
│ └── ...
└── labels/
├── train/
│ ├── 001.txt
│ ├── 002.txt
│ └── ...
└── val/
├── 101.txt
└── ...
3. 基础标签统计分析
3.1 读取所有标签文件
首先,我们写一个函数来读取指定目录下的所有标签文件:
import os
import pandas as pd
from pathlib import Path
def read_yolo_labels(label_dir):
"""
读取YOLO格式的标签文件
参数:
label_dir: 标签文件目录路径
返回:
DataFrame,包含所有标注信息
"""
all_labels = []
# 获取所有.txt文件
label_files = list(Path(label_dir).glob("*.txt"))
print(f"找到 {len(label_files)} 个标签文件")
for label_file in label_files:
# 获取对应的图片文件名(不含扩展名)
image_name = label_file.stem
try:
with open(label_file, 'r') as f:
lines = f.readlines()
for line in lines:
line = line.strip()
if line: # 跳过空行
parts = line.split()
if len(parts) == 5: # 标准YOLO格式
class_id = int(parts[0])
x_center = float(parts[1])
y_center = float(parts[2])
width = float(parts[3])
height = float(parts[4])
all_labels.append({
'image_name': image_name,
'class_id': class_id,
'x_center': x_center,
'y_center': y_center,
'width': width,
'height': height,
'label_file': str(label_file)
})
except Exception as e:
print(f"读取文件 {label_file} 时出错: {e}")
# 转换为DataFrame
df = pd.DataFrame(all_labels)
if not df.empty:
# 计算边界框面积(相对面积)
df['area'] = df['width'] * df['height']
# 计算边界框中心点位置
df['center_x_pixel'] = df['x_center'] * 100 # 假设图片宽度为100像素,用于可视化
df['center_y_pixel'] = df['y_center'] * 100 # 假设图片高度为100像素,用于可视化
return df
# 使用示例
train_labels_dir = "/path/to/your/dataset/labels/train"
df_train = read_yolo_labels(train_labels_dir)
print(f"训练集共有 {len(df_train)} 个标注框")
print(df_train.head()) # 查看前几行数据
3.2 基础统计信息
有了DataFrame,我们可以轻松计算各种统计信息:
def basic_statistics(df, class_names=None):
"""
计算基础统计信息
参数:
df: 包含标签数据的DataFrame
class_names: 类别名称列表,如 ['car', 'bus', 'truck']
"""
print("=" * 50)
print("基础统计信息")
print("=" * 50)
# 总体统计
print(f"总标注框数量: {len(df)}")
print(f"涉及图片数量: {df['image_name'].nunique()}")
print(f"平均每张图片标注框数: {len(df) / df['image_name'].nunique():.2f}")
# 按类别统计
print("\n按类别统计:")
class_counts = df['class_id'].value_counts().sort_index()
for class_id, count in class_counts.items():
class_name = class_names[class_id] if class_names and class_id < len(class_names) else f"class_{class_id}"
percentage = count / len(df) * 100
print(f" 类别 {class_name} (ID: {class_id}): {count} 个框 ({percentage:.1f}%)")
# 边界框尺寸统计
print("\n边界框尺寸统计:")
print(f" 平均宽度: {df['width'].mean():.4f}")
print(f" 平均高度: {df['height'].mean():.4f}")
print(f" 平均面积: {df['area'].mean():.6f}")
print(f" 最小面积: {df['area'].min():.6f}")
print(f" 最大面积: {df['area'].max():.6f}")
# 每张图片的标注框数量分布
boxes_per_image = df.groupby('image_name').size()
print(f"\n每张图片标注框数量:")
print(f" 最少: {boxes_per_image.min()} 个")
print(f" 最多: {boxes_per_image.max()} 个")
print(f" 平均: {boxes_per_image.mean():.2f} 个")
return class_counts
# 使用示例
class_names = ['car', 'bus', 'truck']
class_counts = basic_statistics(df_train, class_names)
4. 深入分析与可视化
4.1 类别分布可视化
文字统计虽然清晰,但可视化能让我们更直观地看到数据分布:
import matplotlib.pyplot as plt
import seaborn as sns
def visualize_class_distribution(df, class_names=None, save_path=None):
"""
可视化类别分布
参数:
df: 包含标签数据的DataFrame
class_names: 类别名称列表
save_path: 图片保存路径(可选)
"""
plt.figure(figsize=(12, 5))
# 子图1:柱状图
plt.subplot(1, 2, 1)
class_counts = df['class_id'].value_counts().sort_index()
# 准备x轴标签
if class_names:
x_labels = [class_names[i] if i < len(class_names) else f'class_{i}'
for i in class_counts.index]
else:
x_labels = [f'class_{i}' for i in class_counts.index]
bars = plt.bar(x_labels, class_counts.values)
plt.title('各类别标注框数量', fontsize=14, fontweight='bold')
plt.xlabel('类别', fontsize=12)
plt.ylabel('数量', fontsize=12)
plt.xticks(rotation=45)
# 在柱子上显示数量
for bar, count in zip(bars, class_counts.values):
height = bar.get_height()
plt.text(bar.get_x() + bar.get_width()/2., height + 0.1,
f'{count}', ha='center', va='bottom')
# 子图2:饼图
plt.subplot(1, 2, 2)
plt.pie(class_counts.values, labels=x_labels, autopct='%1.1f%%', startangle=90)
plt.title('类别分布比例', fontsize=14, fontweight='bold')
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=300, bbox_inches='tight')
print(f"图表已保存至: {save_path}")
plt.show()
# 使用示例
visualize_class_distribution(df_train, class_names, save_path='class_distribution.png')
4.2 边界框尺寸分析
边界框的尺寸分布对目标检测很重要,特别是对于YOLO这类需要预设anchor的模型:
def analyze_bbox_size(df, save_path=None):
"""
分析边界框尺寸分布
参数:
df: 包含标签数据的DataFrame
save_path: 图片保存路径(可选)
"""
plt.figure(figsize=(15, 5))
# 子图1:宽度和高度分布
plt.subplot(1, 3, 1)
plt.scatter(df['width'], df['height'], alpha=0.5, s=10)
plt.xlabel('宽度 (归一化)', fontsize=12)
plt.ylabel('高度 (归一化)', fontsize=12)
plt.title('边界框尺寸分布', fontsize=14, fontweight='bold')
plt.grid(True, alpha=0.3)
# 添加平均线
mean_width = df['width'].mean()
mean_height = df['height'].mean()
plt.axvline(mean_width, color='r', linestyle='--', alpha=0.7, label=f'平均宽度: {mean_width:.3f}')
plt.axhline(mean_height, color='g', linestyle='--', alpha=0.7, label=f'平均高度: {mean_height:.3f}')
plt.legend()
# 子图2:面积分布直方图
plt.subplot(1, 3, 2)
plt.hist(df['area'], bins=50, edgecolor='black', alpha=0.7)
plt.xlabel('边界框面积 (归一化)', fontsize=12)
plt.ylabel('频数', fontsize=12)
plt.title('边界框面积分布', fontsize=14, fontweight='bold')
plt.grid(True, alpha=0.3)
# 标注关键统计值
median_area = df['area'].median()
mean_area = df['area'].mean()
plt.axvline(median_area, color='r', linestyle='--', label=f'中位数: {median_area:.4f}')
plt.axvline(mean_area, color='g', linestyle='--', label=f'平均值: {mean_area:.4f}')
plt.legend()
# 子图3:宽高比分布
plt.subplot(1, 3, 3)
df['aspect_ratio'] = df['width'] / df['height']
plt.hist(df['aspect_ratio'], bins=50, edgecolor='black', alpha=0.7)
plt.xlabel('宽高比 (宽度/高度)', fontsize=12)
plt.ylabel('频数', fontsize=12)
plt.title('边界框宽高比分布', fontsize=14, fontweight='bold')
plt.grid(True, alpha=0.3)
# 标注关键统计值
median_ratio = df['aspect_ratio'].median()
mean_ratio = df['aspect_ratio'].mean()
plt.axvline(median_ratio, color='r', linestyle='--', label=f'中位数: {median_ratio:.2f}')
plt.axvline(mean_ratio, color='g', linestyle='--', label=f'平均值: {mean_ratio:.2f}')
plt.legend()
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=300, bbox_inches='tight')
print(f"图表已保存至: {save_path}")
plt.show()
# 打印统计信息
print("\n边界框尺寸详细统计:")
print(f" 宽度范围: [{df['width'].min():.4f}, {df['width'].max():.4f}]")
print(f" 高度范围: [{df['height'].min():.4f}, {df['height'].max():.4f}]")
print(f" 面积范围: [{df['area'].min():.6f}, {df['area'].max():.6f}]")
print(f" 宽高比范围: [{df['aspect_ratio'].min():.2f}, {df['aspect_ratio'].max():.2f}]")
return df
# 使用示例
df_train = analyze_bbox_size(df_train, save_path='bbox_analysis.png')
4.3 目标位置分布分析
目标在图片中的位置分布也很重要,这能反映数据集的拍摄角度和构图特点:
def analyze_object_position(df, save_path=None):
"""
分析目标在图片中的位置分布
参数:
df: 包含标签数据的DataFrame
save_path: 图片保存路径(可选)
"""
plt.figure(figsize=(12, 10))
# 创建模拟的100x100画布(对应归一化坐标)
plt.figure(figsize=(10, 10))
# 绘制散点图,使用透明度显示密度
plt.scatter(df['x_center'] * 100, df['y_center'] * 100,
alpha=0.3, s=10, c='blue')
# 添加网格和标签
plt.xlim(0, 100)
plt.ylim(0, 100)
plt.xlabel('水平位置 (0-100)', fontsize=12)
plt.ylabel('垂直位置 (0-100)', fontsize=12)
plt.title('目标中心点位置分布', fontsize=14, fontweight='bold')
plt.grid(True, alpha=0.3)
# 添加九宫格参考线
for i in range(1, 3):
plt.axvline(i * 33.33, color='gray', linestyle='--', alpha=0.5)
plt.axhline(i * 33.33, color='gray', linestyle='--', alpha=0.5)
# 添加区域标签
plt.text(16.5, 16.5, '左上', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(50, 16.5, '中上', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(83.5, 16.5, '右上', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(16.5, 50, '左中', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(50, 50, '中心', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(83.5, 50, '右中', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(16.5, 83.5, '左下', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(50, 83.5, '中下', ha='center', va='center', fontsize=10, alpha=0.7)
plt.text(83.5, 83.5, '右下', ha='center', va='center', fontsize=10, alpha=0.7)
# 计算每个区域的物体数量
df['x_region'] = pd.cut(df['x_center'], bins=[0, 0.333, 0.667, 1], labels=['左', '中', '右'])
df['y_region'] = pd.cut(df['y_center'], bins=[0, 0.333, 0.667, 1], labels=['上', '中', '下'])
df['region'] = df['y_region'].astype(str) + df['x_region'].astype(str)
region_counts = df['region'].value_counts()
print("\n目标位置区域分布:")
for region, count in region_counts.items():
percentage = count / len(df) * 100
print(f" {region}: {count} 个目标 ({percentage:.1f}%)")
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=300, bbox_inches='tight')
print(f"图表已保存至: {save_path}")
plt.show()
return df
# 使用示例
df_train = analyze_object_position(df_train, save_path='position_analysis.png')
5. 高级分析与问题检测
5.1 检测潜在的数据问题
通过统计分析,我们可以自动检测一些常见的数据问题:
def detect_data_issues(df, class_names=None):
"""
检测数据集中的潜在问题
参数:
df: 包含标签数据的DataFrame
class_names: 类别名称列表
"""
print("=" * 50)
print("数据问题检测报告")
print("=" * 50)
issues = []
# 1. 检查类别不平衡
class_counts = df['class_id'].value_counts()
total_boxes = len(df)
if len(class_counts) > 1: # 多类别情况
max_count = class_counts.max()
min_count = class_counts.min()
imbalance_ratio = max_count / min_count
if imbalance_ratio > 10:
issues.append(f"⚠️ 严重类别不平衡: 最多/最少 = {imbalance_ratio:.1f}倍")
elif imbalance_ratio > 5:
issues.append(f"⚠️ 明显类别不平衡: 最多/最少 = {imbalance_ratio:.1f}倍")
# 找出具体的不平衡类别
for class_id, count in class_counts.items():
percentage = count / total_boxes * 100
class_name = class_names[class_id] if class_names and class_id < len(class_names) else f"class_{class_id}"
if percentage < 5:
issues.append(f"⚠️ 类别 {class_name} 样本过少: 仅占 {percentage:.1f}%")
elif percentage > 50:
issues.append(f"⚠️ 类别 {class_name} 样本过多: 占比 {percentage:.1f}%")
# 2. 检查边界框尺寸异常
small_boxes = df[df['area'] < 0.001] # 面积小于0.1%
if len(small_boxes) > 0:
issues.append(f"⚠️ 发现 {len(small_boxes)} 个过小边界框 (面积 < 0.001)")
large_boxes = df[df['area'] > 0.5] # 面积大于50%
if len(large_boxes) > 0:
issues.append(f"⚠️ 发现 {len(large_boxes)} 个过大边界框 (面积 > 0.5)")
# 3. 检查宽高比异常
extreme_ratios = df[(df['aspect_ratio'] > 10) | (df['aspect_ratio'] < 0.1)]
if len(extreme_ratios) > 0:
issues.append(f"⚠️ 发现 {len(extreme_ratios)} 个极端宽高比 (>10:1 或 <1:10)")
# 4. 检查位置异常(中心点在边界附近)
edge_boxes = df[
(df['x_center'] < 0.05) | (df['x_center'] > 0.95) |
(df['y_center'] < 0.05) | (df['y_center'] > 0.95)
]
if len(edge_boxes) > 0:
issues.append(f"⚠️ 发现 {len(edge_boxes)} 个边界框中心点过于靠近图片边缘")
# 5. 检查每张图片的标注框数量
boxes_per_image = df.groupby('image_name').size()
images_with_many_boxes = boxes_per_image[boxes_per_image > 20]
if len(images_with_many_boxes) > 0:
issues.append(f"⚠️ 发现 {len(images_with_many_boxes)} 张图片标注框过多 (>20个)")
images_with_few_boxes = boxes_per_image[boxes_per_image == 1]
if len(images_with_few_boxes) > 0:
issues.append(f"⚠️ 发现 {len(images_with_few_boxes)} 张图片只有1个标注框")
# 输出检测结果
if issues:
print("发现以下潜在问题:")
for i, issue in enumerate(issues, 1):
print(f"{i}. {issue}")
# 生成建议
print("\n建议:")
if any("类别不平衡" in issue for issue in issues):
print(" - 考虑对少数类别进行数据增强")
print(" - 使用类别权重或Focal Loss")
if any("边界框" in issue for issue in issues):
print(" - 检查异常边界框的标注质量")
print(" - 考虑过滤或修正异常标注")
if any("标注框过多" in issue for issue in issues):
print(" - 检查密集场景的标注完整性")
if any("只有1个标注框" in issue for issue in issues):
print(" - 检查单目标图片是否需要更多样本")
else:
print("✅ 未发现明显数据问题")
return issues
# 使用示例
issues = detect_data_issues(df_train, class_names)
5.2 生成完整的数据分析报告
最后,我们可以把所有分析整合成一个完整的报告:
def generate_analysis_report(df, dataset_name="训练集", class_names=None, save_path=None):
"""
生成完整的数据分析报告
参数:
df: 包含标签数据的DataFrame
dataset_name: 数据集名称
class_names: 类别名称列表
save_path: 报告保存路径(可选)
"""
print("=" * 60)
print(f"YOLO数据集分析报告 - {dataset_name}")
print("=" * 60)
# 基本信息
print(f"\n📊 基本信息")
print(f" 数据集: {dataset_name}")
print(f" 总标注框数: {len(df):,}")
print(f" 涉及图片数: {df['image_name'].nunique():,}")
print(f" 类别数量: {df['class_id'].nunique()}")
# 类别分布
print(f"\n📈 类别分布")
class_counts = df['class_id'].value_counts().sort_index()
for class_id, count in class_counts.items():
class_name = class_names[class_id] if class_names and class_id < len(class_names) else f"class_{class_id}"
percentage = count / len(df) * 100
print(f" {class_name}: {count:,} 个框 ({percentage:.1f}%)")
# 边界框统计
print(f"\n📏 边界框尺寸统计")
print(f" 平均宽度: {df['width'].mean():.4f}")
print(f" 平均高度: {df['height'].mean():.4f}")
print(f" 平均面积: {df['area'].mean():.6f}")
print(f" 平均宽高比: {(df['width'] / df['height']).mean():.2f}")
# 每张图片统计
boxes_per_image = df.groupby('image_name').size()
print(f"\n🖼️ 每张图片统计")
print(f" 平均标注框数: {boxes_per_image.mean():.2f}")
print(f" 最少标注框数: {boxes_per_image.min()}")
print(f" 最多标注框数: {boxes_per_image.max()}")
# 位置分布
print(f"\n📍 目标位置分布")
center_x_mean = df['x_center'].mean()
center_y_mean = df['y_center'].mean()
if center_x_mean < 0.4:
x_pos = "偏左"
elif center_x_mean > 0.6:
x_pos = "偏右"
else:
x_pos = "居中"
if center_y_mean < 0.4:
y_pos = "偏上"
elif center_y_mean > 0.6:
y_pos = "偏下"
else:
y_pos = "居中"
print(f" 中心点平均位置: ({center_x_mean:.3f}, {center_y_mean:.3f})")
print(f" 整体分布: {x_pos}{y_pos}")
# 数据质量评估
print(f"\n🔍 数据质量评估")
# 计算质量分数(简单示例)
quality_score = 100
# 扣分项
if len(df) < 1000:
quality_score -= 10
print(f" ⚠️ 数据量较少 (<1000)")
if df['class_id'].nunique() > 1:
imbalance_ratio = class_counts.max() / class_counts.min()
if imbalance_ratio > 5:
quality_score -= 15
print(f" ⚠️ 类别不平衡严重 (最多/最少 = {imbalance_ratio:.1f}倍)")
small_boxes_ratio = len(df[df['area'] < 0.001]) / len(df)
if small_boxes_ratio > 0.1:
quality_score -= 10
print(f" ⚠️ 过小边界框较多 ({small_boxes_ratio*100:.1f}%)")
# 质量评级
print(f"\n📊 数据质量评分: {quality_score}/100")
if quality_score >= 90:
print(" ✅ 优秀 - 数据质量很好,适合直接用于训练")
elif quality_score >= 70:
print(" ⚠️ 良好 - 数据质量不错,但有一些小问题需要注意")
elif quality_score >= 50:
print(" ⚠️ 一般 - 建议进行数据清洗和增强")
else:
print(" ❌ 较差 - 需要大量数据清洗和增强工作")
# 训练建议
print(f"\n💡 训练建议")
if len(df) < 5000:
print(" - 建议增加数据量,或使用数据增强技术")
if df['class_id'].nunique() > 1 and class_counts.max() / class_counts.min() > 3:
print(" - 建议使用类别权重或Focal Loss处理类别不平衡")
if df['area'].mean() < 0.05:
print(" - 目标较小,建议使用更高的输入分辨率")
print(" - 建议进行数据增强:随机裁剪、旋转、色彩调整等")
# 保存报告到文件
if save_path:
import sys
original_stdout = sys.stdout
with open(save_path, 'w', encoding='utf-8') as f:
sys.stdout = f
generate_analysis_report(df, dataset_name, class_names, save_path=None)
sys.stdout = original_stdout
print(f"\n📄 完整报告已保存至: {save_path}")
return quality_score
# 使用示例
quality_score = generate_analysis_report(
df_train,
dataset_name="车辆检测训练集",
class_names=class_names,
save_path="dataset_analysis_report.txt"
)
6. 实战应用与建议
6.1 在YOLOv9训练前的实际应用
现在我们已经有了完整的分析工具,来看看如何在YOLOv9训练前实际应用:
def prepare_yolov9_training(df, class_names, output_dir="./data_analysis"):
"""
为YOLOv9训练准备数据分析和建议
参数:
df: 包含标签数据的DataFrame
class_names: 类别名称列表
output_dir: 输出目录
"""
import os
os.makedirs(output_dir, exist_ok=True)
print("开始YOLOv9训练前数据分析...")
# 1. 生成完整的分析报告
report_path = os.path.join(output_dir, "training_report.txt")
quality_score = generate_analysis_report(df, "训练集", class_names, report_path)
# 2. 保存可视化图表
visualize_class_distribution(df, class_names,
os.path.join(output_dir, "class_distribution.png"))
analyze_bbox_size(df, os.path.join(output_dir, "bbox_analysis.png"))
analyze_object_position(df, os.path.join(output_dir, "position_analysis.png"))
# 3. 检测数据问题
issues = detect_data_issues(df, class_names)
# 4. 生成YOLOv9训练配置建议
print("\n" + "=" * 60)
print("YOLOv9训练配置建议")
print("=" * 60)
# 基于分析结果给出建议
avg_boxes_per_image = len(df) / df['image_name'].nunique()
avg_area = df['area'].mean()
print(f"\n📋 基础配置建议:")
print(f" 1. 输入尺寸: 建议从640x640开始")
if avg_area < 0.02:
print(f" 2. 小目标检测: 检测到平均面积较小 ({avg_area:.4f}),建议:")
print(f" - 使用更高的输入分辨率(如1280x1280)")
print(f" - 调整anchor尺寸")
print(f" - 增加小目标数据增强")
if avg_boxes_per_image > 10:
print(f" 3. 密集目标: 平均每图 {avg_boxes_per_image:.1f} 个目标,建议:")
print(f" - 使用更大的batch size")
print(f" - 调整NMS参数")
print(f" - 考虑使用YOLOv9的密集场景优化")
# 类别不平衡处理建议
if len(class_names) > 1:
class_counts = df['class_id'].value_counts()
imbalance_ratio = class_counts.max() / class_counts.min()
if imbalance_ratio > 3:
print(f"\n⚖️ 类别不平衡处理建议 (不平衡比: {imbalance_ratio:.1f}):")
print(f" 1. 数据层面:")
print(f" - 对少数类别进行过采样")
print(f" - 使用数据增强增加少数类别样本")
print(f" 2. 损失函数层面:")
print(f" - 使用Focal Loss")
print(f" - 设置类别权重")
print(f" 3. 训练策略:")
print(f" - 调整class_weights参数")
print(f" - 使用渐进式训练策略")
# 数据增强建议
print(f"\n🔄 数据增强建议:")
print(f" 1. 基础增强: Mosaic, MixUp, RandomAffine")
# 基于位置分布建议
center_x_mean = df['x_center'].mean()
center_y_mean = df['y_center'].mean()
if abs(center_x_mean - 0.5) > 0.1 or abs(center_y_mean - 0.5) > 0.1:
print(f" 2. 位置增强: 检测到目标位置偏置,建议增加:")
print(f" - 随机平移增强")
print(f" - 随机裁剪增强")
# 基于尺寸分布建议
if df['aspect_ratio'].std() > 0.5:
print(f" 3. 形状增强: 宽高比变化较大,建议增加:")
print(f" - 随机缩放增强")
print(f" - 随机旋转增强")
print(f"\n✅ 分析完成!所有结果已保存到: {output_dir}")
print(f" - 分析报告: {report_path}")
print(f" - 可视化图表: {output_dir}/")
return {
'quality_score': quality_score,
'avg_boxes_per_image': avg_boxes_per_image,
'avg_area': avg_area,
'issues': issues
}
# 使用示例
analysis_results = prepare_yolov9_training(df_train, class_names)
6.2 验证集分析对比
不要忘记分析验证集,确保训练集和验证集分布一致:
def compare_train_val(train_df, val_df, class_names, output_dir="./data_comparison"):
"""
比较训练集和验证集的分布
参数:
train_df: 训练集DataFrame
val_df: 验证集DataFrame
class_names: 类别名称列表
output_dir: 输出目录
"""
import os
os.makedirs(output_dir, exist_ok=True)
print("比较训练集和验证集分布...")
# 类别分布对比
train_class_counts = train_df['class_id'].value_counts().sort_index()
val_class_counts = val_df['class_id'].value_counts().sort_index()
plt.figure(figsize=(10, 6))
x = range(len(class_names))
width = 0.35
plt.bar([i - width/2 for i in x], train_class_counts.values, width, label='训练集', alpha=0.8)
plt.bar([i + width/2 for i in x], val_class_counts.values, width, label='验证集', alpha=0.8)
plt.xlabel('类别', fontsize=12)
plt.ylabel('标注框数量', fontsize=12)
plt.title('训练集 vs 验证集 - 类别分布对比', fontsize=14, fontweight='bold')
plt.xticks(x, class_names)
plt.legend()
plt.grid(True, alpha=0.3)
# 添加数量标签
for i, (train_count, val_count) in enumerate(zip(train_class_counts.values, val_class_counts.values)):
plt.text(i - width/2, train_count + 0.1, str(train_count), ha='center', va='bottom')
plt.text(i + width/2, val_count + 0.1, str(val_count), ha='center', va='bottom')
plt.tight_layout()
plt.savefig(os.path.join(output_dir, "train_val_class_comparison.png"), dpi=300)
plt.show()
# 计算分布差异
print("\n📊 分布差异分析:")
for i, class_name in enumerate(class_names):
if i in train_class_counts.index and i in val_class_counts.index:
train_percent = train_class_counts[i] / len(train_df) * 100
val_percent = val_class_counts[i] / len(val_df) * 100
diff = abs(train_percent - val_percent)
if diff > 5:
print(f" ⚠️ {class_name}: 训练集 {train_percent:.1f}% vs 验证集 {val_percent:.1f}% (差异: {diff:.1f}%)")
else:
print(f" ✅ {class_name}: 训练集 {train_percent:.1f}% vs 验证集 {val_percent:.1f}% (差异: {diff:.1f}%)")
# 边界框尺寸对比
print(f"\n📏 边界框尺寸对比:")
print(f" 训练集平均面积: {train_df['area'].mean():.6f}")
print(f" 验证集平均面积: {val_df['area'].mean():.6f}")
print(f" 差异: {abs(train_df['area'].mean() - val_df['area'].mean()):.6f}")
# 位置分布对比
print(f"\n📍 中心点位置对比:")
print(f" 训练集平均位置: ({train_df['x_center'].mean():.3f}, {train_df['y_center'].mean():.3f})")
print(f" 验证集平均位置: ({val_df['x_center'].mean():.3f}, {val_df['y_center'].mean():.3f})")
print(f"\n✅ 对比完成!图表已保存到: {output_dir}")
7. 总结
通过今天分享的pandas标签分析方法,你现在应该能够:
7.1 掌握的核心技能
- 快速了解数据集全貌:几行代码就能知道数据集的规模、类别分布、标注质量
- 发现潜在问题:自动检测类别不平衡、异常标注、数据偏置等问题
- 可视化分析结果:生成直观的图表,帮助理解数据特征
- 生成训练建议:基于分析结果给出针对性的YOLOv9训练建议
7.2 实际应用价值
这个方法在实际项目中有几个重要的应用场景:
数据清洗阶段:在开始训练前,先用这个工具分析数据,发现并修复问题,能节省大量调试时间。
模型调试阶段:当模型在某些类别上表现不佳时,用这个工具分析训练数据,看看是不是数据本身有问题。
数据收集指导:分析现有数据的分布,指导后续数据收集应该关注哪些方面。
团队协作:生成的分析报告和可视化图表,能帮助团队成员快速理解数据集特点。
7.3 后续优化方向
如果你觉得这个工具好用,还可以进一步扩展:
- 集成到训练流程:把分析脚本集成到YOLOv9的训练脚本中,自动在训练前分析数据
- 添加更多分析维度:比如分析标注框的纵横比分布、目标重叠情况等
- 生成自动化报告:自动生成HTML格式的详细分析报告
- 批量处理:支持批量分析多个数据集,方便对比
记住,好的数据是成功训练的一半。花点时间分析数据,往往比盲目调整模型参数更有效。希望这个工具能帮助你在YOLOv9的训练中少走弯路,更快获得好结果!
获取更多AI镜像
想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。
更多推荐
所有评论(0)