cuibinge's picture
Sync YOLO training and evaluation utilities (part 2)
c2b1b26 verified
Raw
History Blame Contribute Delete
6.12 kB
"""
标签数据预处理脚本
将TIF格式的标签转换为PNG格式的二值mask
"""
import os
import numpy as np
from PIL import Image
import rasterio
from pathlib import Path
import shutil
def convert_tif_label_to_png_mask(tif_path, output_path, threshold=1):
"""将TIF标签转换为PNG二值mask"""
try:
with rasterio.open(tif_path) as src:
# 读取第一个波段
label_data = src.read(1)
# 创建二值mask (0: 背景, 255: 前景)
mask = np.zeros_like(label_data, dtype=np.uint8)
mask[label_data >= threshold] = 255
# 保存为PNG
mask_img = Image.fromarray(mask, mode='L')
mask_img.save(output_path)
print(f"转换: {tif_path} -> {output_path}")
print(f" 原始数据范围: {label_data.min()} - {label_data.max()}")
print(f" Mask唯一值: {np.unique(mask)}")
return True
except Exception as e:
print(f"转换失败 {tif_path}: {str(e)}")
return False
def prepare_dataset_labels():
"""准备数据集标签"""
# 创建mask目录
train_mask_dir = Path("data/train/masks")
val_mask_dir = Path("data/val/masks")
test_mask_dir = Path("data/test/masks")
train_mask_dir.mkdir(parents=True, exist_ok=True)
val_mask_dir.mkdir(parents=True, exist_ok=True)
test_mask_dir.mkdir(parents=True, exist_ok=True)
# 处理训练集标签
train_label_dir = Path("data/train/labels")
if train_label_dir.exists():
print("处理训练集标签...")
label_files = [f for f in train_label_dir.glob("*.TIF")]
for label_file in label_files:
# 获取对应的图像文件名
# 假设标签文件名是 labelXXXXX.TIF,图像是 patchXXXXX.TIF
label_name = label_file.stem # label10012
if label_name.startswith('label'):
# 提取数字部分
number = label_name[5:] # 10012
image_name = f"patch{number}.TIF"
# 检查对应的图像是否存在
image_path = Path("data/train/images") / image_name
if image_path.exists():
# 转换标签为mask
mask_name = f"patch{number}.png"
mask_path = train_mask_dir / mask_name
if convert_tif_label_to_png_mask(str(label_file), str(mask_path)):
print(f" 成功: {label_name} -> {mask_name}")
else:
print(f" 警告: 找不到对应的图像 {image_name}")
# 处理验证集标签
val_label_dir = Path("data/val/labels")
if val_label_dir.exists():
print("\n处理验证集标签...")
label_files = [f for f in val_label_dir.glob("*.TIF")]
for label_file in label_files:
label_name = label_file.stem
if label_name.startswith('label'):
number = label_name[5:]
image_name = f"patch{number}.TIF"
image_path = Path("data/val/images") / image_name
if image_path.exists():
mask_name = f"patch{number}.png"
mask_path = val_mask_dir / mask_name
if convert_tif_label_to_png_mask(str(label_file), str(mask_path)):
print(f" 成功: {label_name} -> {mask_name}")
# 处理测试集标签
test_label_dir = Path("data/test/labels")
if test_label_dir.exists():
print("\n处理测试集标签...")
label_files = [f for f in test_label_dir.glob("*.TIF")]
for label_file in label_files:
label_name = label_file.stem
if label_name.startswith('label'):
number = label_name[5:]
image_name = f"patch{number}.TIF"
image_path = Path("data/test/images") / image_name
if image_path.exists():
mask_name = f"patch{number}.png"
mask_path = test_mask_dir / mask_name
if convert_tif_label_to_png_mask(str(label_file), str(mask_path)):
print(f" 成功: {label_name} -> {mask_name}")
print("\n标签数据准备完成!")
def check_data_integrity():
"""检查数据完整性"""
print("\n检查数据完整性...")
datasets = [
("训练集", "data/train/images", "data/train/masks"),
("验证集", "data/val/images", "data/val/masks"),
("测试集", "data/test/images", "data/test/masks")
]
for name, image_dir, mask_dir in datasets:
image_path = Path(image_dir)
mask_path = Path(mask_dir)
if image_path.exists() and mask_path.exists():
image_files = list(image_path.glob("*.TIF"))
mask_files = list(mask_path.glob("*.png"))
print(f"{name}:")
print(f" 图像文件: {len(image_files)} 个")
print(f" Mask文件: {len(mask_files)} 个")
# 检查匹配情况
matched = 0
for image_file in image_files:
expected_mask = mask_path / (image_file.stem + ".png")
if expected_mask.exists():
matched += 1
print(f" 匹配文件: {matched} 个")
if matched < len(image_files):
print(f" ⚠️ 警告: 有 {len(image_files) - matched} 个图像缺少对应的mask")
else:
print(f"{name}: 目录不存在")
if __name__ == "__main__":
print("开始准备标签数据...")
prepare_dataset_labels()
check_data_integrity()
print("\n数据准备完成!")