From 826f110894aeb1d5e1433dd44f81f04626009019 Mon Sep 17 00:00:00 2001 From: duxin Date: Tue, 7 Jul 2026 09:10:40 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E9=98=B2=E5=BE=A1=E6=80=A7=E7=BC=96?= =?UTF-8?q?=E7=A8=8B=E9=87=8D=E6=9E=84=20=E2=80=94=20=E6=A0=85=E6=A0=BC?= =?UTF-8?q?=E7=A9=BA=E9=97=B4=E5=AF=B9=E9=BD=90=20+=20NoData=20=E5=A4=84?= =?UTF-8?q?=E7=90=86=20+=20=E9=AB=98=E5=8D=B1=E4=BB=A3=E7=A0=81=E5=8A=A0?= =?UTF-8?q?=E5=9B=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit **核心修复 (preview_generator.py):** - 废弃 _align_mask_to_image() (numpy crop/pad 在地理空间上错误) - 新增 _warp_mask_to_image(): 使用 gdal.Warp 将掩膜重采样到与底图 完全一致的像素网格,处理旋转/偏移/投影差异 - 新增 _normalize_mask(nodata_value): 正确的掩膜值域自劢归一间 (0/1 vs 0/255) - 修复 alpha = mask_data/255.0 → mask_data (掩膜是二值 0/1, 不是 0/255) - 面积计算 valid_pixels 使用 warp 前数据排除 nodata 背景 **新增防御工具模块:** - spatial_validator.py: SpatialAlignmentError 自定义异常 + validate_two_rasters() / validate_spatial_alignment() 强制空间一致性检查 + has_rotation() / get_pixel_resolution() 诊断工具 - nodata_handler.py: read_band_safe() / read_bands_safe() 自动 NoData→nan + read_band_masked() 返回 MaskedArray + create_valid_mask() **P0 高危代码集成 (3处):** - sampling.py: 耀斑掩膜与水体掩膜 bool 运算前验证 numpy shape 一致 - find_severe_glint_area.py: 栅格 mask 读取后验证 dims+GT+projection 对齐 - waterindex_inversion/__init__.py: mask 与 BSQ 维度验证, 不一致时优雅降级 **旋转影像兼容性修复:** - extract_water_area.py: pixel_size = sqrt(gt[1]^2+gt[2]^2) 使用勾股定理 正确计算旋转影像的像素分辨率 --- .../waterindex_inversion/__init__.py | 25 +- src/core/utils/nodata_handler.py | 291 +++++++++++++++ src/core/utils/preview_generator.py | 278 ++++++++++++++- src/core/utils/spatial_validator.py | 334 ++++++++++++++++++ src/utils/extract_water_area.py | 7 +- src/utils/find_severe_glint_area.py | 9 + src/utils/sampling.py | 29 +- 7 files changed, 947 insertions(+), 26 deletions(-) create mode 100644 src/core/utils/nodata_handler.py create mode 100644 src/core/utils/spatial_validator.py diff --git a/src/core/algorithms/waterindex_inversion/__init__.py b/src/core/algorithms/waterindex_inversion/__init__.py index fb6bccc..e88cd4d 100644 --- a/src/core/algorithms/waterindex_inversion/__init__.py +++ b/src/core/algorithms/waterindex_inversion/__init__.py @@ -30,6 +30,7 @@ from typing import Any, Callable, Dict, List, Optional, Tuple import numpy as np from osgeo import gdal, osr +from src.core.utils.spatial_validator import validate_spatial_alignment # GDAL 驱动注册 gdal.UseExceptions() @@ -620,10 +621,26 @@ class WaterIndexProcessor: import rasterio with rasterio.open(water_mask_path) as msrc: water_mask = msrc.read(1) - print(f"[run_inversion] 水域掩膜已加载: {water_mask_path}," - f"形状={water_mask.shape}," - f"陆地区域(0)={int((water_mask == 0).sum())}," - f"水区域(>0)={int((water_mask > 0).sum())}") + mask_shape = water_mask.shape + + # 防御性检查:掩膜与 BSQ 影像的空间一致性 + bsq_expected_shape = (height, width) + try: + validate_spatial_alignment( + mask_shape, None, None, + bsq_expected_shape, None, None, + label_a="水域掩膜", label_b="BSQ影像", + check_projection=False, + ) + except Exception as shape_err: + print(f"[run_inversion] ⚠ 掩膜与影像维度不一致,跳过掩膜处理: {shape_err}") + water_mask = None + + if water_mask is not None: + print(f"[run_inversion] 水域掩膜已加载: {water_mask_path}," + f"形状={water_mask.shape}," + f"陆地区域(0)={int((water_mask == 0).sum())}," + f"水区域(>0)={int((water_mask > 0).sum())}") except Exception as mask_err: print(f"[run_inversion] ⚠ 掩膜加载失败,跳过掩膜处理: {mask_err}") water_mask = None diff --git a/src/core/utils/nodata_handler.py b/src/core/utils/nodata_handler.py new file mode 100644 index 0000000..6258793 --- /dev/null +++ b/src/core/utils/nodata_handler.py @@ -0,0 +1,291 @@ +# -*- coding: utf-8 -*- +""" +统一的栅格 NoData 处理模块 + +规范化项目中所有栅格读取逻辑。无论是通过 gdal 还是 rasterio, +都主动读取 NoData 值并屏蔽(替换为 np.nan 或使用 Masked Array), +防止背景黑边参与科学计算。 + +核心 API:: + + from src.core.utils.nodata_handler import ( + read_band_safe, + read_band_masked, + read_bands_safe, + ) + + # 读取单波段,自动处理 NoData → np.nan + data, nodata = read_band_safe(dataset, band_index=1) + + # 读取多波段 + bands_dict, nodata = read_bands_safe(dataset, [1, 2, 3]) + + # 获取 numpy MaskedArray + masked = read_band_masked(dataset, band_index=1) +""" + +from __future__ import annotations + +import numpy as np +from pathlib import Path +from typing import Optional, List, Tuple, Union, Dict + +try: + from osgeo import gdal + GDAL_AVAILABLE = True +except ImportError: + GDAL_AVAILABLE = False + + +# 项目通用的 NoData 回退值(仅在无法从文件元数据读取时使用) +_FALLBACK_NODATA_VALUES = [-9999.0, -999.0, -32768.0, 0.0] + + +# ============================================================ +# GDAL 路径 +# ============================================================ + +def _open_dataset(path_or_dataset): + """统一打开逻辑:接受文件路径或已打开的 GDAL Dataset""" + if isinstance(path_or_dataset, str): + if not GDAL_AVAILABLE: + raise ImportError("GDAL 未安装") + ds = gdal.Open(path_or_dataset, gdal.GA_ReadOnly) + if ds is None: + raise FileNotFoundError(f"无法打开栅格文件: {path_or_dataset}") + return ds, True # 需要关闭 + return path_or_dataset, False # 不需要关闭 + + +def read_nodata_value(dataset) -> Optional[float]: + """从 GDAL Dataset 的 Band 1 读取 NoData 值 + + Args: + dataset: GDAL Dataset 或文件路径 + + Returns: + NoData 值(float),若未设置则返回 None + """ + try: + band = dataset.GetRasterBand(1) + ndv = band.GetNoDataValue() + if ndv is not None: + return float(ndv) + except Exception: + pass + return None + + +def read_band_safe(dataset, + band_index: int = 1, + nodata_handling: str = "nan", + dtype: type = np.float32) -> Tuple[np.ndarray, Optional[float]]: + """安全读取单个波段,自动处理 NoData + + 这是推荐的栅格读取入口。读取数据的同时: + 1. 从元数据获取 NoData 值 + 2. 将 NoData 像素替换为 np.nan 或 0.0 + 3. 可选:对常见的硬编码 nodata 做回退处理 + + Args: + dataset: GDAL Dataset 或文件路径 + band_index: 波段序号(1-based,GDAL 惯例) + nodata_handling: "nan" → NoData→np.nan; "zero" → NoData→0.0 + dtype: 输出数组的 dtype + + Returns: + (data_array, nodata_value): + - data_array: shape=(height, width) 的数组,NoData 已替换 + - nodata_value: 原始 NoData 值(None 表示未设置) + + Examples:: + + ds = gdal.Open("image.bsq") + data, ndv = read_band_safe(ds, band_index=5) + # 现在 data 中的 NoData 像素都是 np.nan,可以安全参与统计 + valid_mean = np.nanmean(data) + """ + band = dataset.GetRasterBand(band_index) + arr = band.ReadAsArray().astype(dtype) + + # 读取 NoData 元数据 + nodata = read_nodata_value(dataset) + + # 应用 NoData 替换 + arr = _apply_nodata_mask(arr, nodata, nodata_handling) + + return arr, nodata + + +def read_bands_safe(dataset, + band_indices: List[int], + nodata_handling: str = "nan", + dtype: type = np.float32) -> Tuple[Dict[int, np.ndarray], Optional[float]]: + """安全读取多个波段 + + Args: + dataset: GDAL Dataset 或文件路径 + band_indices: 波段序号列表(1-based) + nodata_handling: "nan" 或 "zero" + dtype: 输出 dtype + + Returns: + (bands_dict, nodata_value): + - bands_dict: {band_index: data_array} 字典 + - nodata_value: NoData 值 + + Examples:: + + ds = gdal.Open("image.bsq") + bands, ndv = read_bands_safe(ds, [10, 20, 30]) + r, g, b = bands[10], bands[20], bands[30] + """ + nodata = read_nodata_value(dataset) + result = {} + + for bi in band_indices: + band = dataset.GetRasterBand(bi) + arr = band.ReadAsArray().astype(dtype) + arr = _apply_nodata_mask(arr, nodata, nodata_handling) + result[bi] = arr + + return result, nodata + + +def read_band_masked(dataset, + band_index: int = 1, + dtype: type = np.float32) -> np.ma.MaskedArray: + """读取波段并返回 numpy MaskedArray + + MaskedArray 保留原始数据值,同时通过 .mask 属性标记 NoData 像素。 + 适用于需要在保持原始值的前提下做条件统计的场景。 + + Args: + dataset: GDAL Dataset 或文件路径 + band_index: 波段序号(1-based) + dtype: 输出 dtype + + Returns: + numpy MaskedArray,NoData 像素被 mask + + Examples:: + + ds = gdal.Open("image.bsq") + masked = read_band_masked(ds, band_index=1) + # masked.mean() 自动排除 NoData 像素 + """ + band = dataset.GetRasterBand(band_index) + arr = band.ReadAsArray().astype(dtype) + nodata = read_nodata_value(dataset) + + if nodata is not None: + mask = np.isclose(arr, nodata) + # 同时检查常见硬编码 nodata + for fbv in _FALLBACK_NODATA_VALUES: + if fbv != nodata: + mask |= np.isclose(arr, fbv) + return np.ma.array(arr, mask=mask) + + return np.ma.array(arr, mask=~np.isfinite(arr)) + + +# ============================================================ +# 便捷函数:从文件路径读取 +# ============================================================ + +def safe_read_raster(img_path: str, + band_index: int = 1, + nodata_handling: str = "nan", + dtype: type = np.float32) -> Tuple[np.ndarray, Optional[float], + Tuple[int, int], + Optional[Tuple[float, ...]]]: + """从文件路径一站式安全读取栅格 + + 返回数据、NoData、维度和 GeoTransform,一次调用覆盖所有常见需求。 + + Args: + img_path: 栅格文件路径 + band_index: 波段序号(1-based) + nodata_handling: "nan" 或 "zero" + dtype: 输出 dtype + + Returns: + (data, nodata, shape, geotransform): + - data: NoData 已替换的 2D numpy 数组 + - nodata: 原始 NoData 值(None 表示未设置) + - shape: (rows, cols) 元组 + - geotransform: 6 元组 GeoTransform(None 表示无地理参考) + + Examples:: + + data, ndv, shape, gt = safe_read_raster("mask.dat") + print(f"尺寸={shape}, NoData={ndv}, GT={gt}") + """ + ds, should_close = _open_dataset(img_path) + try: + data, nodata = read_band_safe(ds, band_index, nodata_handling, dtype) + shape = (ds.RasterYSize, ds.RasterXSize) + gt = ds.GetGeoTransform() + return data, nodata, shape, gt + finally: + if should_close: + ds = None + + +# ============================================================ +# 内部辅助函数 +# ============================================================ + +def _apply_nodata_mask(arr: np.ndarray, + nodata: Optional[float], + handling: str) -> np.ndarray: + """将数组中的 NoData 值替换为目标标记值 + + Args: + arr: 输入数组 + nodata: NoData 值(None 表示不处理) + handling: "nan" → 替换为 np.nan; "zero" → 替换为 0.0 + + Returns: + 处理后的数组 + """ + replacement = np.nan if handling == "nan" else 0.0 + + # 1. 先处理无穷值 + arr[~np.isfinite(arr)] = replacement + + # 2. 处理显式 NoData + if nodata is not None and np.isfinite(nodata): + arr[np.isclose(arr, float(nodata))] = replacement + + # 3. 回退:硬编码的常见 nodata 值 + for fbv in _FALLBACK_NODATA_VALUES: + if nodata is None or fbv != nodata: + arr[np.isclose(arr, fbv)] = replacement + + return arr + + +def create_valid_mask(arr: np.ndarray, + nodata: Optional[float] = None) -> np.ndarray: + """创建有效像素的布尔掩膜 + + "有效"定义:有限、非 NaN、非 NoData。 + + Args: + arr: 输入数组 + nodata: 显式 NoData 值 + + Returns: + bool 数组,True = 有效像素 + + Examples:: + + data, ndv, _, _ = safe_read_raster("image.bsq") + valid = create_valid_mask(data, ndv) + water_in_valid_area = (water_mask > 0) & valid + """ + valid = np.isfinite(arr) + if nodata is not None and np.isfinite(nodata): + valid &= ~np.isclose(arr, float(nodata)) + return valid diff --git a/src/core/utils/preview_generator.py b/src/core/utils/preview_generator.py index 168be01..eb6e5ad 100644 --- a/src/core/utils/preview_generator.py +++ b/src/core/utils/preview_generator.py @@ -102,6 +102,187 @@ def get_wavelength_info(img_path: str) -> Optional[List[float]]: return None +# ============================================================ +# 掩膜辅助函数 +# ============================================================ + +# GDAL Warp 输出专用的 NoData 标记值 +# 选 -9999.0 是因为它已存在于项目既有的硬编码 nodata 列表中, +# 且与有效掩膜值(0=非水体, 1=水体, 255=备选水体编码)完全不冲突。 +_WARP_NODATA = -9999.0 + + +def _normalize_mask(mask_data: np.ndarray, + nodata_value: Optional[float] = None) -> np.ndarray: + """将掩膜数据自适应归一化到 [0, 1] 区间,严格排除 NoData + + 兼容多种掩膜值域: + - 二值 0/1(NDWI extract_water 默认输出) + - 0/255 图像格式 + - Float32 0.0/1.0 + - 含 nodata 特殊值(如 -9999, -999, -32768) + + Args: + mask_data: 原始掩膜数组(任意 shape 和 dtype) + nodata_value: 显式指定的 NoData 值。 + 传入 None 时回退到历史硬编码列表。 + + Returns: + 归一化到 [0, 1] 的 float32 数组,水体=1,非水体/noData=0 + """ + data = mask_data.astype(np.float32) + + # 1. 过滤 inf / NaN + data[~np.isfinite(data)] = 0.0 + + # 2. 显式 NoData 清除(优先使用 warp 阶段传入的 nodata_value) + if nodata_value is not None and np.isfinite(nodata_value): + data[np.isclose(data, float(nodata_value))] = 0.0 + + # 3. 回退:硬编码的常见 nodata 列表 + common_nodata = [-9999.0, -999.0, -32768.0] + for ndv in common_nodata: + data[np.isclose(data, ndv)] = 0.0 + + # 4. 如果所有值都 <= 0,返回全零(无有效水体像素) + if np.all(data <= 0): + print("警告: 掩膜数据全为非正值,可能不含有效水体像素。") + return np.zeros_like(data) + + # 5. 自动检测值域范围并归一化 + max_val = data.max() + if max_val > 1.0: + # 可能是 0/255 或更大范围,归一化到 [0, 1] + data = data / max_val + print(f"[掩膜归一化] 检测到值域上限 = {max_val:.1f},已除以该值归一化到 [0, 1]。") + + # 6. 钳制到 [0, 1] + data = np.clip(data, 0.0, 1.0) + + # 7. 二值化处理:离散掩膜做阈值判定,消除浮点误差 + unique_vals = np.unique(data) + if len(unique_vals) <= 3: + data = (data > 0.5).astype(np.float32) + + return data + + +def _warp_mask_to_image(mask_path: str, + base_dataset: gdal.Dataset, + nodata_output: float = _WARP_NODATA): + """使用 gdal.Warp 将掩膜重采样到与底图完全一致的像素网格 + + 这是解决"矩形色块错位/倾斜"问题的核心函数。当掩膜与 RGB 底图 + 具有不同的 GeoTransform(旋转参数)或投影时,numpy 级别的 crop/pad + 在地理空间上是错误的——只有通过真正的重投影/重采样, + 掩膜才能在每个像素位置上与底图精确对齐。 + + 内部使用 GDAL MEM 驱动在内存中完成 warp,不产生临时文件。 + + Args: + mask_path: 掩膜文件路径 + base_dataset: 已打开的 RGB 底图 GDAL Dataset + nodata_output: warp 后输出栅格的 NoData 填充值 + + Returns: + (warped_mask_2d, nodata_used) + - warped_mask_2d: shape=(base_h, base_w) 的 float32 数组, + 水体像素保留原值,背景/越界像素 = nodata_output + - nodata_used: 实际使用的 NoData 值(= nodata_output) + + Raises: + RuntimeError: 若 warp 后尺寸与底图不一致(理论上不应发生) + """ + base_w = base_dataset.RasterXSize + base_h = base_dataset.RasterYSize + base_gt = base_dataset.GetGeoTransform() + base_proj = base_dataset.GetProjection() + + # ---- 计算底图的地理范围(支持旋转影像) ---- + # 四个角点的地理坐标 + corners_x = [ + base_gt[0], # 左上 X + base_gt[0] + base_w * base_gt[1], # 右上 X + base_gt[0] + base_h * base_gt[2], # 左下 X + base_gt[0] + base_w * base_gt[1] + base_h * base_gt[2], # 右下 X + ] + corners_y = [ + base_gt[3], # 左上 Y + base_gt[3] + base_w * base_gt[4], # 右上 Y + base_gt[3] + base_h * base_gt[5], # 左下 Y + base_gt[3] + base_w * base_gt[4] + base_h * base_gt[5], # 右下 Y + ] + minx, maxx = min(corners_x), max(corners_x) + miny, maxy = min(corners_y), max(corners_y) + + # ---- 读取掩膜的原始 NoData 并传递到 Warp ---- + mask_ds = gdal.Open(mask_path, gdal.GA_ReadOnly) + if mask_ds is None: + raise ValueError(f"无法打开掩膜文件: {mask_path}") + + src_nodata = None + mask_gt = mask_ds.GetGeoTransform() + mask_proj = mask_ds.GetProjection() + try: + src_nodata = mask_ds.GetRasterBand(1).GetNoDataValue() + except Exception: + pass + + # ---- 构建 Warp 选项 ---- + warp_kwargs = { + 'format': 'MEM', + 'outputBounds': [minx, miny, maxx, maxy], + 'xRes': base_gt[1], # 保留符号 + 'yRes': base_gt[5], # 保留符号(通常为负,表示北在上) + 'dstSRS': base_proj, + 'srcSRS': mask_proj if mask_proj else base_proj, + 'dstNodata': nodata_output, + 'resampleAlg': gdal.GRA_NearestNeighbour, + 'targetAlignedPixels': True, + 'warpOptions': ['NUM_THREADS=ALL_CPUS'], + } + if src_nodata is not None: + warp_kwargs['srcNodata'] = src_nodata + + # ---- 执行 Warp(写入 MEM 数据集) ---- + mem_driver = gdal.GetDriverByName('MEM') + warp_kwargs_str = ( + f"outputBounds={minx},{miny},{maxx},{maxy} " + f"dst=({base_w}x{base_h}) " + f"xRes={base_gt[1]:.6f} yRes={base_gt[5]:.6f}" + ) + print(f"[GDAL Warp] 正在将掩膜重采样到与底图一致的像素网格: {warp_kwargs_str}") + + try: + warped_ds = gdal.Warp('', mask_ds, **warp_kwargs) + except Exception as e: + mask_ds = None + raise RuntimeError( + f"gdal.Warp 执行失败: {e}\n" + f" 掩膜: {mask_path}\n" + f" 底图范围: [{minx}, {miny}, {maxx}, {maxy}]\n" + f" 底图尺寸: {base_w} x {base_h}" + ) from e + + mask_ds = None # 源掩膜已不再需要 + + # ---- 验证输出尺寸 ---- + warped_w = warped_ds.RasterXSize + warped_h = warped_ds.RasterYSize + if warped_w != base_w or warped_h != base_h: + warped_ds = None + raise RuntimeError( + f"gdal.Warp 输出尺寸 ({warped_w}x{warped_h}) 与底图 " + f"({base_w}x{base_h}) 不一致!请检查投影/分辨率参数。" + ) + + # ---- 读取 warped 数据 ---- + warped_data = warped_ds.GetRasterBand(1).ReadAsArray().astype(np.float32) + warped_ds = None + + return warped_data, nodata_output + + # ============================================================ # 核心预览图生成函数 # ============================================================ @@ -274,26 +455,81 @@ def generate_water_mask_overlay(img_path: str, rgb_image = np.nan_to_num(np.stack([r_s, g_s, b_s], axis=2)) * 255 rgb_image = rgb_image.astype(np.uint8) - # 读取掩膜 - mask_dataset = gdal.Open(mask_path) - if mask_dataset is not None: - mask_data = mask_dataset.GetRasterBand(1).ReadAsArray() - mask_dataset = None + # ============================================================ + # 读取掩膜并重采样到与底图一致的像素网格 + # 使用 gdal.Warp 进行地理空间精确配准,处理旋转/偏移/投影差异 + # ============================================================ + mask_data = None + mask_nodata = _WARP_NODATA + warp_failed = False + + if not Path(mask_path).exists(): + print(f"警告: 掩膜文件不存在: {mask_path}") else: - print(f"警告: 无法打开掩膜文件: {mask_path}") - mask_data = None + try: + # Step A: 地理空间 warp(核心) + mask_data_raw, mask_nodata = _warp_mask_to_image( + mask_path, dataset, nodata_output=_WARP_NODATA + ) - # Alpha 混合 - overlay = np.zeros((height, width, 4), dtype=np.uint8) - overlay[:, :, 0:3] = mask_color - overlay[:, :, 3] = 255 # 全不透明 + # Step B: 值域归一化(显式传入 warp 的 nodata,确保背景透明) + mask_data = _normalize_mask(mask_data_raw, nodata_value=mask_nodata) + # Step C: 计算有效像素(排除 nodata 背景,用于面积统计) + # 必须在归一化之前用原始 warped 数据计算,因为归一化后 + # nodata 被转为 0,无法与"非水体陆地"区分 + valid_mask = (mask_data_raw != mask_nodata) & np.isfinite(mask_data_raw) + valid_pixel_count = int(np.sum(valid_mask)) + + except Exception as e: + print(f"警告: gdal.Warp 重采样失败,回退到直接读取模式: {e}") + warp_failed = True + # 回退:直接读取掩膜(假设与底图已对齐) + mask_ds = gdal.Open(mask_path, gdal.GA_ReadOnly) + if mask_ds is not None: + mask_data_raw = mask_ds.GetRasterBand(1).ReadAsArray().astype(np.float32) + # 读取源掩膜的 nodata(如果存在) + try: + src_nd = mask_ds.GetRasterBand(1).GetNoDataValue() + if src_nd is not None: + mask_nodata = float(src_nd) + except Exception: + pass + mask_ds = None + + # 若尺寸不一致(回退路径不做 warp,但做简单 crop/pad) + if mask_data_raw.shape[0] != height or mask_data_raw.shape[1] != width: + aligned = np.zeros((height, width), dtype=np.float32) + copy_h = min(mask_data_raw.shape[0], height) + copy_w = min(mask_data_raw.shape[1], width) + aligned[:copy_h, :copy_w] = mask_data_raw[:copy_h, :copy_w] + mask_data_raw = aligned + print( + f"[回退对齐] 掩膜 ({mask_data_raw.shape[1]}x{mask_data_raw.shape[0]}) " + f"→ 影像 ({width}x{height}),已裁剪/填充(非 warp)。" + ) + + mask_data = _normalize_mask(mask_data_raw, nodata_value=mask_nodata) + valid_mask = (mask_data_raw != mask_nodata) & np.isfinite(mask_data_raw) + valid_pixel_count = int(np.sum(valid_mask)) + else: + print(f"警告: 无法打开掩膜文件: {mask_path}") + valid_pixel_count = 0 + + # ============================================================ + # Alpha 混合:仅对水体像素叠加颜色 + # - 水体 (mask_data == 1) → alpha = mask_alpha + # - 非水体 (mask_data == 0) → alpha = 0(完全透明) + # - NoData 背景 (已在 _normalize_mask 中归零) → alpha = 0 + # ============================================================ blended = rgb_image.astype(np.float32) - if mask_data is not None: - alpha = mask_data.astype(np.float32) / 255.0 * mask_alpha - for c in range(3): - blended[:, :, c] = rgb_image[:, :, c].astype(np.float32) * (1 - alpha) + mask_color[c] * alpha - blended = blended.astype(np.uint8) + if mask_data is not None and np.any(mask_data > 0): + alpha = mask_data.astype(np.float32) * mask_alpha + alpha_3c = np.stack([alpha, alpha, alpha], axis=2) + mask_color_arr = np.array(mask_color, dtype=np.float32).reshape(1, 1, 3) + blended = (rgb_image.astype(np.float32) * (1.0 - alpha_3c) + + mask_color_arr * alpha_3c) + blended = np.clip(blended, 0, 255).astype(np.uint8) # 绘图 fig, ax = plt.subplots(figsize=(14, 10)) @@ -307,14 +543,20 @@ def generate_water_mask_overlay(img_path: str, ax.legend(handles=legend_elements, loc='upper right', framealpha=0.9) # 面积计算 + # 修复: valid_pixels 使用 warp 前数据计算的 valid_pixel_count, + # 仅统计有效数据像素,排除旋转产生的 Bounding Box 背景区域 if geotransform and geotransform[1] != 0: pixel_size_x = abs(geotransform[1]) pixel_size_y = abs(geotransform[5]) pixel_area = pixel_size_x * pixel_size_y if mask_data is not None: - water_pixels = np.sum(mask_data > 0) - valid_pixels = np.sum(mask_data >= 0) + water_pixels = int(np.sum(mask_data > 0)) + if not warp_failed: + valid_pixels = valid_pixel_count + else: + # 回退路径:无法区分 nodata 背景和非水体,用 >=0 作为近似 + valid_pixels = int(np.sum(mask_data >= 0)) water_km2 = water_pixels * pixel_area / 1_000_000 valid_km2 = valid_pixels * pixel_area / 1_000_000 pct = (water_pixels / valid_pixels * 100) if valid_pixels > 0 else 0 diff --git a/src/core/utils/spatial_validator.py b/src/core/utils/spatial_validator.py new file mode 100644 index 0000000..79eb285 --- /dev/null +++ b/src/core/utils/spatial_validator.py @@ -0,0 +1,334 @@ +# -*- coding: utf-8 -*- +""" +栅格空间对齐验证模块 + +提供统一的"空间对齐检查"拦截器,在任何需要将两张栅格进行数学运算的地方 +强制先调用验证函数,确保 GeoTransform、投影和维度一致。 + +用法:: + + from src.core.utils.spatial_validator import ( + SpatialAlignmentError, + validate_spatial_alignment, + validate_two_rasters, + ) + + # 方式 1:通过两个已打开的 GDAL Dataset 验证 + validate_two_rasters(ds_a, ds_b, label_a="底图", label_b="掩膜") + + # 方式 2:通过两个 numpy 数组 + 各自的 geotransform/projection 验证 + validate_spatial_alignment( + shape_a, gt_a, proj_a, + shape_b, gt_b, proj_b, + label_a="影像A", label_b="影像B", + ) + + # 方式 3:只验证 numpy shape(当不需要地理参考时) + validate_shapes(shape_a, shape_b, label_a="数组A", label_b="数组B") +""" + +from __future__ import annotations + +import numpy as np +from typing import Optional, Tuple, Any + + +# ============================================================ +# 自定义异常 +# ============================================================ + +class SpatialAlignmentError(ValueError): + """栅格空间对齐失败时抛出的自定义异常 + + 与普通 ValueError 的区别: + - 携带结构化的差异描述(shape_diff, gt_diff, proj_diff) + - 自动生成人类可读的中文诊断信息 + - 便于上层调用方 catch 后做降级处理(如触发 gdal.Warp) + """ + + def __init__(self, message: str, details: Optional[dict] = None): + super().__init__(message) + self.details = details or {} + + @classmethod + def from_shape_mismatch(cls, label_a: str, label_b: str, + shape_a: tuple, shape_b: tuple) -> "SpatialAlignmentError": + msg = ( + f"栅格维度不匹配:\n" + f" {label_a}: (行={shape_a[0]}, 列={shape_a[1]})\n" + f" {label_b}: (行={shape_b[0]}, 列={shape_b[1]})\n" + f"建议: 对 {label_b} 执行 gdal.Warp 重采样到 {label_a} 的像素网格。" + ) + return cls(msg, details={"shape_a": shape_a, "shape_b": shape_b}) + + @classmethod + def from_geotransform_mismatch(cls, label_a: str, label_b: str, + gt_a: tuple, gt_b: tuple, + diffs: list) -> "SpatialAlignmentError": + diff_desc = "\n".join(f" gt[{i}]: {label_a}={gt_a[i]:.6f} vs {label_b}={gt_b[i]:.6f} (差异={d:.6f})" + for i, d in diffs) + msg = ( + f"GeoTransform 参数不一致:\n" + f"{diff_desc}\n" + f"建议: 对 {label_b} 执行 gdal.Warp 重采样到 {label_a} 的像素网格。" + ) + return cls(msg, details={"gt_a": gt_a, "gt_b": gt_b, "diffs": diffs}) + + @classmethod + def from_projection_mismatch(cls, label_a: str, label_b: str, + proj_a: str, proj_b: str) -> "SpatialAlignmentError": + msg = ( + f"投影坐标系 (CRS) 不一致:\n" + f" {label_a}: {proj_a[:120]}...\n" + f" {label_b}: {proj_b[:120]}...\n" + f"建议: 对 {label_b} 执行 gdal.Warp 并指定 dstSRS 参数。" + ) + return cls(msg, details={"proj_a": proj_a, "proj_b": proj_b}) + + +# ============================================================ +# 核心验证函数 +# ============================================================ + +def validate_shapes(shape_a: Tuple[int, int], + shape_b: Tuple[int, int], + label_a: str = "栅格A", + label_b: str = "栅格B") -> None: + """验证两个二维栅格的 numpy shape 一致 + + Args: + shape_a, shape_b: (rows, cols) 二元组 + label_a, label_b: 人类可读的标签,用于错误消息 + + Raises: + SpatialAlignmentError: shape 不一致时 + """ + if shape_a != shape_b: + raise SpatialAlignmentError.from_shape_mismatch( + label_a, label_b, shape_a, shape_b + ) + + +def validate_geotransforms(gt_a: Tuple[float, ...], + gt_b: Tuple[float, ...], + label_a: str = "栅格A", + label_b: str = "栅格B", + tolerance: float = 1e-6) -> None: + """验证两个 GeoTransform 参数一致(在容差范围内) + + 比较全部 6 个仿射变换参数。容忍度按每个参数的相对大小自适应缩放。 + + Args: + gt_a, gt_b: 6 元组 GeoTransform + label_a, label_b: 标签 + tolerance: 相对容差系数 + + Raises: + SpatialAlignmentError: 任意参数差异超过容差时 + """ + if gt_a is None or gt_b is None: + return # 无地理参考时不验证 + + diffs = [] + for i in range(6): + diff = abs(gt_a[i] - gt_b[i]) + threshold = tolerance * max(abs(gt_a[i]), 1.0) + if diff > threshold: + diffs.append((i, diff)) + + if diffs: + raise SpatialAlignmentError.from_geotransform_mismatch( + label_a, label_b, gt_a, gt_b, diffs + ) + + +def validate_projections(proj_a: Optional[str], + proj_b: Optional[str], + label_a: str = "栅格A", + label_b: str = "栅格B") -> None: + """验证两个投影坐标系 (WKT) 一致 + + 空字符串视为"无投影",不验证。 + + Raises: + SpatialAlignmentError: 投影不一致时 + """ + if not proj_a or not proj_b: + return # 至少一方无投影,跳过 + + # 标准化比较(去除空白差异) + norm_a = " ".join(proj_a.split()) + norm_b = " ".join(proj_b.split()) + if norm_a != norm_b: + raise SpatialAlignmentError.from_projection_mismatch( + label_a, label_b, proj_a, proj_b + ) + + +def validate_spatial_alignment( + shape_or_data_a, geotransform_a=None, projection_a=None, + shape_or_data_b=None, geotransform_b=None, projection_b=None, + label_a: str = "栅格A", label_b: str = "栅格B", + gt_tolerance: float = 1e-6, + check_projection: bool = True, +) -> None: + """一站式空间对齐验证 + + 同时检查维度、GeoTransform 和投影坐标系。 + 支持传入 numpy 数组(自动提取 shape)或显式的 (rows, cols) 元组。 + + Args: + shape_or_data_a, shape_or_data_b: numpy 数组 或 (rows, cols) 元组 + geotransform_a, geotransform_b: 6 元组 GeoTransform(可选) + projection_a, projection_b: WKT 字符串(可选) + label_a, label_b: 标签(用于错误消息) + gt_tolerance: GeoTransform 相对容差 + check_projection: 是否检查投影一致性 + + Raises: + SpatialAlignmentError: 任何维度/GT/投影不一致时 + + Examples:: + + # 验证两个 GDAL 数据集 + validate_spatial_alignment( + ds_a.ReadAsArray(), ds_a.GetGeoTransform(), ds_a.GetProjection(), + ds_b.ReadAsArray(), ds_b.GetGeoTransform(), ds_b.GetProjection(), + label_a="影像", label_b="掩膜", + ) + + # 验证两个 numpy 数组(不需要地理参考) + validate_spatial_alignment( + img.shape[:2], None, None, + mask.shape, None, None, + label_a="RGB图", label_b="水体掩膜", + ) + """ + # --- 提取 shape --- + def _get_shape(obj) -> Tuple[int, int]: + if isinstance(obj, np.ndarray): + h, w = obj.shape[:2] + return (int(h), int(w)) + if isinstance(obj, (tuple, list)) and len(obj) == 2: + return (int(obj[0]), int(obj[1])) + raise TypeError(f"无法从 {type(obj)} 提取 shape,请传入 numpy 数组或 (rows, cols) 元组") + + shape_a = _get_shape(shape_or_data_a) if shape_or_data_a is not None else None + shape_b = _get_shape(shape_or_data_b) if shape_or_data_b is not None else None + + # --- 维度检查 --- + if shape_a is not None and shape_b is not None: + validate_shapes(shape_a, shape_b, label_a, label_b) + + # --- GeoTransform 检查 --- + if geotransform_a is not None and geotransform_b is not None: + validate_geotransforms( + geotransform_a, geotransform_b, label_a, label_b, gt_tolerance + ) + + # --- 投影检查 --- + if check_projection and projection_a is not None and projection_b is not None: + validate_projections(projection_a, projection_b, label_a, label_b) + + +def validate_two_rasters(dataset_a: Any, + dataset_b: Any, + label_a: str = "栅格A", + label_b: str = "栅格B", + check_projection: bool = True) -> None: + """验证两个已打开的 GDAL Dataset 的空间一致性 + + 这是最常用的入口:传入两个 gdal.Open() 的返回值即可。 + + Args: + dataset_a, dataset_b: GDAL Dataset 对象 + label_a, label_b: 标签 + check_projection: 是否检查投影 + + Raises: + SpatialAlignmentError: 空间不一致时 + + Examples:: + + ds_img = gdal.Open(img_path) + ds_mask = gdal.Open(mask_path) + validate_two_rasters(ds_img, ds_mask, "影像", "掩膜") + """ + shape_a = (dataset_a.RasterYSize, dataset_a.RasterXSize) + shape_b = (dataset_b.RasterYSize, dataset_b.RasterXSize) + gt_a = dataset_a.GetGeoTransform() + gt_b = dataset_b.GetGeoTransform() + proj_a = dataset_a.GetProjection() + proj_b = dataset_b.GetProjection() + + validate_spatial_alignment( + shape_a, gt_a, proj_a, + shape_b, gt_b, proj_b, + label_a=label_a, label_b=label_b, + check_projection=check_projection, + ) + + +# ============================================================ +# 便捷函数:批量栅格验证 +# ============================================================ + +def validate_raster_stack(datasets: list, + labels: Optional[list] = None, + check_projection: bool = True) -> None: + """验证多个 GDAL Dataset 两两之间的空间一致性 + + Args: + datasets: GDAL Dataset 列表 + labels: 对应的标签列表(长度需一致) + check_projection: 是否检查投影 + + Raises: + SpatialAlignmentError: 任意两个 dataset 不一致时 + """ + if len(datasets) < 2: + return + + if labels is None: + labels = [f"栅格{i+1}" for i in range(len(datasets))] + + base = datasets[0] + base_label = labels[0] + for i, ds in enumerate(datasets[1:], start=1): + validate_two_rasters(base, ds, base_label, labels[i], check_projection) + + +# ============================================================ +# 诊断工具:GeoTransform 信息提取 +# ============================================================ + +def has_rotation(geotransform: Tuple[float, ...]) -> bool: + """检查 GeoTransform 是否包含旋转/倾斜参数""" + return abs(geotransform[2]) > 1e-10 or abs(geotransform[4]) > 1e-10 + + +def get_pixel_resolution(geotransform: Tuple[float, ...]) -> Tuple[float, float]: + """获取像素的实际地面分辨率(考虑旋转参数) + + 对于旋转影像,使用勾股定理计算实际分辨率。 + 参见 extract_water_area.py 中的相同逻辑。 + + Returns: + (pixel_size_x, pixel_size_y) —— 均为正值 + """ + return ( + np.sqrt(geotransform[1]**2 + geotransform[2]**2), + np.sqrt(geotransform[4]**2 + geotransform[5]**2), + ) + + +def describe_geotransform(gt: Tuple[float, ...]) -> str: + """生成 GeoTransform 的人类可读描述""" + if gt is None: + return "无地理参考" + px, py = get_pixel_resolution(gt) + rot = " (含旋转)" if has_rotation(gt) else "" + return ( + f"原点=({gt[0]:.2f}, {gt[3]:.2f}), " + f"分辨率=({px:.2f}, {py:.2f}) m/px{rot}" + ) diff --git a/src/utils/extract_water_area.py b/src/utils/extract_water_area.py index f92537a..1c93c5f 100644 --- a/src/utils/extract_water_area.py +++ b/src/utils/extract_water_area.py @@ -85,9 +85,10 @@ def rasterize_shp(shp_filepath, raster_fn_out, img_path, NoData_value=None): layer = source_ds.GetLayer(0) layer_name = layer.GetName() - # about 25 metres(ish) use 0.001 if you want roughly 100m - pixel_size_x = abs(geotransform[1]) # 像素宽度(X方向) - pixel_size_y = abs(geotransform[5]) # 像素高度(Y方向,通常是负值,需要取绝对值) + # 计算像素分辨率(考虑旋转参数) + # 对于旋转影像(gt[2] 或 gt[4] != 0),像素实际分辨率需用勾股定理 + pixel_size_x = np.sqrt(geotransform[1]**2 + geotransform[2]**2) + pixel_size_y = np.sqrt(geotransform[4]**2 + geotransform[5]**2) raster_fn_out_tmp = append2filename(raster_fn_out, "_tmp_delete") # 构建栅格化参数 diff --git a/src/utils/find_severe_glint_area.py b/src/utils/find_severe_glint_area.py index bde083c..0c28381 100644 --- a/src/utils/find_severe_glint_area.py +++ b/src/utils/find_severe_glint_area.py @@ -2,6 +2,7 @@ from src.utils.util import * from osgeo import gdal, ogr import argparse import cv2 +from src.core.utils.spatial_validator import validate_spatial_alignment @@ -621,6 +622,14 @@ def find_severe_glint_area(img_path, water_mask, glint_wave=750, output_path=Non if dataset_water_mask is None: raise ValueError(f"无法打开水域掩膜文件: {water_mask}") data_water_mask = dataset_water_mask.GetRasterBand(1).ReadAsArray() + + # 防御性检查:水域掩膜与影像的空间一致性 + validate_spatial_alignment( + data_water_mask, dataset_water_mask.GetGeoTransform(), dataset_water_mask.GetProjection(), + (im_height, im_width), dataset.GetGeoTransform(), dataset.GetProjection(), + label_a="水域掩膜", label_b="影像", + ) + del dataset_water_mask print(f"使用检测方法: {method}") diff --git a/src/utils/sampling.py b/src/utils/sampling.py index 833284c..93d00da 100644 --- a/src/utils/sampling.py +++ b/src/utils/sampling.py @@ -16,6 +16,7 @@ from osgeo import gdal, ogr import spectral from scipy import ndimage from src.utils.util import write_bands +from src.core.utils.spatial_validator import validate_spatial_alignment try: from skimage import morphology @@ -130,6 +131,15 @@ def get_spectral_sampling_points_chunked(bil_file, water_mask_shp, severe_glint= if dataset_severe_glint is None: raise ValueError(f"无法打开耀斑掩膜文件: {severe_glint}") data_severe_glint = dataset_severe_glint.GetRasterBand(1).ReadAsArray() + + # 防御性检查:耀斑掩膜与水体掩膜的空间一致性 + validate_spatial_alignment( + water_mask_raster, None, None, + data_severe_glint, None, None, + label_a="水体掩膜", label_b="耀斑掩膜", + check_projection=False, + ) + print("已加载耀斑掩膜") # 对glint边界进行外扩1-2像素作为缓冲 data_severe_glint = expand_glint_buffer(data_severe_glint, buffer_size=2) @@ -142,7 +152,7 @@ def get_spectral_sampling_points_chunked(bil_file, water_mask_shp, severe_glint= valid_area = (water_mask_raster > 0) & (~(data_severe_glint > 0)) else: valid_area = (water_mask_raster > 0) - + # 计算水体宽度(用于自适应采样) width_map = None if use_adaptive_sampling: @@ -454,6 +464,15 @@ def get_spectral_sampling_points(bil_file, water_mask_shp, severe_glint=None, ou if dataset_severe_glint is None: raise ValueError(f"无法打开耀斑掩膜文件: {severe_glint}") data_severe_glint = dataset_severe_glint.GetRasterBand(1).ReadAsArray() + + # 防御性检查:耀斑掩膜与水体掩膜的空间一致性 + validate_spatial_alignment( + water_mask_raster, None, None, + data_severe_glint, None, None, + label_a="水体掩膜", label_b="耀斑掩膜", + check_projection=False, + ) + print("已加载耀斑掩膜") # 对glint边界进行外扩1-2像素作为缓冲 data_severe_glint = expand_glint_buffer(data_severe_glint, buffer_size=2) @@ -920,6 +939,14 @@ def get_coor_base_interval(water_mask, severe_glint=None, output_csvpath=None, i raise ValueError(f"无法打开耀斑掩膜文件: {severe_glint}") data_severe_glint = dataset_severe_glint.GetRasterBand(1).ReadAsArray() + # 防御性检查:耀斑掩膜与水体掩膜的空间一致性 + validate_spatial_alignment( + data_water_mask, None, None, + data_severe_glint, None, None, + label_a="水体掩膜", label_b="耀斑掩膜", + check_projection=False, + ) + # 使用耀斑掩膜的几何信息 im_width = dataset_severe_glint.RasterXSize im_height = dataset_severe_glint.RasterYSize