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:
@ -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
|
||||
|
||||
291
src/core/utils/nodata_handler.py
Normal file
291
src/core/utils/nodata_handler.py
Normal 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-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
|
||||
@ -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
|
||||
|
||||
334
src/core/utils/spatial_validator.py
Normal file
334
src/core/utils/spatial_validator.py
Normal 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}"
|
||||
)
|
||||
@ -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")
|
||||
|
||||
# 构建栅格化参数
|
||||
|
||||
@ -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}")
|
||||
|
||||
@ -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
|
||||
|
||||
Reference in New Issue
Block a user