| """
|
| 标签数据预处理脚本
|
| 将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 = np.zeros_like(label_data, dtype=np.uint8)
|
| mask[label_data >= threshold] = 255
|
|
|
|
|
| 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():
|
| """准备数据集标签"""
|
|
|
|
|
| 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:
|
|
|
|
|
| label_name = label_file.stem
|
| if label_name.startswith('label'):
|
|
|
| number = label_name[5:]
|
| image_name = f"patch{number}.TIF"
|
|
|
|
|
| image_path = Path("data/train/images") / image_name
|
| if image_path.exists():
|
|
|
| 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数据准备完成!")
|
|
|