fix: 防御性编程重构 — 栅格空间对齐 + NoData 处理 + 高危代码加固

**核心修复 (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) 使用勾股定理
  正确计算旋转影像的像素分辨率
This commit is contained in:
duxin
2026-07-07 09:10:40 +08:00
parent c46f78e69d
commit 826f110894
7 changed files with 947 additions and 26 deletions

View File

@ -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

View File

@ -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-basedGDAL 惯例)
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 MaskedArrayNoData 像素被 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 元组 GeoTransformNone 表示无地理参考)
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

View File

@ -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/1NDWI 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

View File

@ -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}"
)

View File

@ -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")
# 构建栅格化参数

View File

@ -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}")

View File

@ -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