cuibinge's picture
Sync YOLO training and evaluation utilities
dd0ae11 verified
Raw
History Blame Contribute Delete
5.13 kB
"""
检查训练数据,统计类别分布,帮助诊断问题
"""
import os
import numpy as np
from pathlib import Path
from tqdm import tqdm
from PIL import Image
import matplotlib.pyplot as plt
def check_dataset_statistics(image_dir, mask_dir):
"""检查数据集的统计信息"""
print("=" * 60)
print("检查训练数据统计信息")
print("=" * 60)
# 获取图像列表
images = [f for f in os.listdir(image_dir) if f.endswith('.TIF')]
print(f"\n找到 {len(images)} 张图像")
# 统计信息
total_pixels = 0
background_pixels = 0
foreground_pixels = 0
images_with_seaweed = 0
images_without_seaweed = 0
pixel_ratios = []
print("\n正在分析数据...")
for img_file in tqdm(images):
mask_name = img_file.replace('.TIF', '.png')
mask_path = os.path.join(mask_dir, mask_name)
if not os.path.exists(mask_path):
continue
# 读取mask
mask = Image.open(mask_path).convert('L')
mask = np.array(mask)
mask = (mask > 127).astype(np.uint8)
# 统计
total = mask.size
bg = np.sum(mask == 0)
fg = np.sum(mask == 1)
total_pixels += total
background_pixels += bg
foreground_pixels += fg
ratio = fg / total if total > 0 else 0
pixel_ratios.append(ratio)
if fg > 0:
images_with_seaweed += 1
else:
images_without_seaweed += 1
# 打印统计信息
print("\n" + "=" * 60)
print("数据集统计结果")
print("=" * 60)
print(f"总图像数: {len(images)}")
print(f" - 有浒苔的图像: {images_with_seaweed} ({images_with_seaweed/len(images)*100:.1f}%)")
print(f" - 无浒苔的图像: {images_without_seaweed} ({images_without_seaweed/len(images)*100:.1f}%)")
print(f"\n总像素数: {total_pixels:,}")
print(f" - 背景像素: {background_pixels:,} ({background_pixels/total_pixels*100:.2f}%)")
print(f" - 浒苔像素: {foreground_pixels:,} ({foreground_pixels/total_pixels*100:.2f}%)")
# 计算类别权重建议
bg_weight = total_pixels / (2 * background_pixels) if background_pixels > 0 else 1.0
fg_weight = total_pixels / (2 * foreground_pixels) if foreground_pixels > 0 else 1.0
print(f"\n建议的类别权重:")
print(f" - 背景权重: {bg_weight:.4f}")
print(f" - 前景权重: {fg_weight:.4f}")
print(f" - 权重比 (前景/背景): {fg_weight/bg_weight:.2f}")
# 浒苔比例分布
if pixel_ratios:
print(f"\n浒苔像素比例分布:")
print(f" - 最小值: {min(pixel_ratios):.4f}")
print(f" - 最大值: {max(pixel_ratios):.4f}")
print(f" - 平均值: {np.mean(pixel_ratios):.4f}")
print(f" - 中位数: {np.median(pixel_ratios):.4f}")
print(f" - 标准差: {np.std(pixel_ratios):.4f}")
# 绘制分布图
plt.figure(figsize=(10, 6))
plt.hist(pixel_ratios, bins=50, edgecolor='black')
plt.xlabel('浒苔像素比例')
plt.ylabel('图像数量')
plt.title('浒苔像素比例分布')
plt.axvline(np.mean(pixel_ratios), color='r', linestyle='--', label=f'平均值: {np.mean(pixel_ratios):.4f}')
plt.legend()
plt.grid(True, alpha=0.3)
plt.savefig('dataset_statistics.png', dpi=150, bbox_inches='tight')
print(f"\n分布图已保存到: dataset_statistics.png")
print("\n" + "=" * 60)
print("诊断建议:")
print("=" * 60)
if foreground_pixels / total_pixels < 0.01:
print("⚠️ 警告: 浒苔像素比例过低 (<1%),模型可能难以学习")
print(" 建议:")
print(" 1. 增加浒苔样本的数量")
print(" 2. 使用更高的前景权重 (foreground_weight >= 5.0)")
print(" 3. 考虑使用难例挖掘")
if images_without_seaweed / len(images) < 0.3:
print("⚠️ 警告: 负样本(无浒苔图像)比例过低 (<30%)")
print(" 建议: 增加负样本数量")
if fg_weight / bg_weight > 10:
print("⚠️ 警告: 类别不平衡非常严重,建议使用更高的前景权重")
print("\n推荐的训练配置:")
print(f' "background_weight": {bg_weight:.2f},')
print(f' "foreground_weight": {max(fg_weight, 3.0):.2f},')
print(f' "focal_alpha": [1.0, {max(fg_weight/bg_weight, 3.0):.2f}],')
if __name__ == "__main__":
# 检查训练数据
train_image_dir = "data/train/images"
train_mask_dir = "data/train/masks"
if os.path.exists(train_image_dir) and os.path.exists(train_mask_dir):
check_dataset_statistics(train_image_dir, train_mask_dir)
else:
print(f"错误: 找不到数据目录")
print(f" 图像目录: {train_image_dir}")
print(f" 标签目录: {train_mask_dir}")