perf: 四种去耀斑算法改为分块流式,内存占用恒定
Goodman/Hedley/Kutser/SUGAR 统一为块外循环 × 波段内循环 × 掩膜分块读 × 逐块写盘,不再整幅驻留参考波段与掩膜,也不再在内存中累积全部校正波段(大图可省数十 GB)。 水域掩膜改为保留 GDAL band 句柄按需分块读取,不再 ReadAsArray 整幅;SHP 掩膜栅格化到临时 GTiff 使用后清理。 Goodman 行带顺序改写(整幅宽 × 2048 行/条),每波段每条只 FlushCache 一次,替代原先 1000x1000 小块逐波段刷盘导致的十万级磁盘同步。 输出数据集统一走 create_output_dataset,大图自动切换 GTiff+BIGTIFF 规避 ENVI int32 偏移溢出。
This commit is contained in:
@ -9,7 +9,7 @@ except ImportError:
|
|||||||
GDAL_AVAILABLE = False
|
GDAL_AVAILABLE = False
|
||||||
print("警告: GDAL未安装,将使用numpy处理模式")
|
print("警告: GDAL未安装,将使用numpy处理模式")
|
||||||
|
|
||||||
from src.utils.util import find_band_number
|
from src.utils.util import find_band_number, create_output_dataset
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
@ -49,7 +49,8 @@ class Goodman:
|
|||||||
def __init__(self, im_aligned,
|
def __init__(self, im_aligned,
|
||||||
nir_lower_wavelength=641.93, nir_upper_wavelength=751.49,
|
nir_lower_wavelength=641.93, nir_upper_wavelength=751.49,
|
||||||
A=0.000019, B=0.1,
|
A=0.000019, B=0.1,
|
||||||
use_gdal=True, chunk_size=None, water_mask=None, output_path=None):
|
use_gdal=True, chunk_size=None, water_mask=None, output_path=None,
|
||||||
|
block_size=1000):
|
||||||
"""
|
"""
|
||||||
Goodman 耀斑去除算法 — 波长驱动版。
|
Goodman 耀斑去除算法 — 波长驱动版。
|
||||||
|
|
||||||
@ -90,6 +91,13 @@ class Goodman:
|
|||||||
self.chunk_size = chunk_size
|
self.chunk_size = chunk_size
|
||||||
self.is_file_path = isinstance(im_aligned, str)
|
self.is_file_path = isinstance(im_aligned, str)
|
||||||
self.output_path = output_path
|
self.output_path = output_path
|
||||||
|
self.block_size = block_size
|
||||||
|
|
||||||
|
# ★ v3:掩膜不再整幅驻留——栅格路径保留 GDAL band 句柄,按需分块读
|
||||||
|
self._mask_ds = None
|
||||||
|
self._wm_band = None
|
||||||
|
self._wm_np = None
|
||||||
|
self._wm_temp = None # SHP 临时栅格(需清理)
|
||||||
|
|
||||||
# ── 波长驱动:通过 HDR 元数据动态解析波段索引 ──
|
# ── 波长驱动:通过 HDR 元数据动态解析波段索引 ──
|
||||||
if self.is_file_path:
|
if self.is_file_path:
|
||||||
@ -123,90 +131,70 @@ class Goodman:
|
|||||||
self.width = im_aligned.shape[1]
|
self.width = im_aligned.shape[1]
|
||||||
self.n_bands = im_aligned.shape[-1]
|
self.n_bands = im_aligned.shape[-1]
|
||||||
|
|
||||||
# 加载水域掩膜(在获取图像尺寸之后)
|
# 加载水域掩膜(★ v3:仅打开句柄/保留引用,绝不整幅 ReadAsArray 驻留)
|
||||||
self.water_mask = self._load_water_mask(water_mask)
|
self.water_mask_path = water_mask
|
||||||
|
self.water_mask = None
|
||||||
|
self._open_mask()
|
||||||
|
|
||||||
def _load_water_mask(self, water_mask):
|
def _open_mask(self):
|
||||||
|
"""打开水域掩膜,但绝不整幅 ReadAsArray 驻留。
|
||||||
|
|
||||||
|
栅格路径 → 保留 GDAL dataset/band 句柄(_wm_band),按需分块读取;
|
||||||
|
np.ndarray(外部已驻留)→ 转 bool 后保留 _wm_np,按窗口切片;
|
||||||
|
.shp(需与栅格同投影)→ 用 extract_water_area 的原生栅格化写临时 GTiff。
|
||||||
"""
|
"""
|
||||||
加载水域掩膜
|
if self._wm_band is not None or self._wm_np is not None:
|
||||||
|
return
|
||||||
|
wm = self.water_mask_path
|
||||||
|
if wm is None:
|
||||||
|
return
|
||||||
|
|
||||||
:param water_mask: 可以是None、numpy数组、文件路径(.dat/.tif)或shapefile路径(.shp)
|
if isinstance(wm, np.ndarray):
|
||||||
:return: numpy数组或None,1表示水域,0表示非水域
|
if wm.shape[:2] != (self.height, self.width):
|
||||||
"""
|
raise ValueError(
|
||||||
if water_mask is None:
|
f"掩膜尺寸 {wm.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配")
|
||||||
return None
|
self._wm_np = wm > 0
|
||||||
|
return
|
||||||
|
|
||||||
# 如果已经是numpy数组
|
if isinstance(wm, str):
|
||||||
if isinstance(water_mask, np.ndarray):
|
|
||||||
if water_mask.shape[:2] != (self.height, self.width):
|
|
||||||
raise ValueError(f"掩膜尺寸 {water_mask.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配")
|
|
||||||
return (water_mask > 0).astype(np.uint8) # 确保是0/1掩膜
|
|
||||||
|
|
||||||
# 如果是文件路径
|
|
||||||
if isinstance(water_mask, str):
|
|
||||||
if not GDAL_AVAILABLE:
|
if not GDAL_AVAILABLE:
|
||||||
raise ValueError("使用文件路径作为掩膜时,必须安装GDAL")
|
raise ValueError("使用文件路径作为掩膜时,必须安装GDAL")
|
||||||
|
src = wm
|
||||||
# 检查是否为shapefile
|
if wm.lower().endswith('.shp'):
|
||||||
if water_mask.lower().endswith('.shp'):
|
if not self.is_file_path:
|
||||||
# 从shp文件创建掩膜
|
|
||||||
if self.is_file_path:
|
|
||||||
ref_path = self.im_aligned
|
|
||||||
else:
|
|
||||||
raise ValueError("输入为numpy数组时,无法从shp文件创建掩膜(需要参考栅格)")
|
raise ValueError("输入为numpy数组时,无法从shp文件创建掩膜(需要参考栅格)")
|
||||||
|
self._wm_temp = (self.output_path or self.im_aligned) + '__wm.tif'
|
||||||
|
from src.utils.extract_water_area import rasterize_shp
|
||||||
|
rasterize_shp(wm, self._wm_temp, self.im_aligned)
|
||||||
|
src = self._wm_temp
|
||||||
|
self._mask_ds = gdal.Open(src, gdal.GA_ReadOnly)
|
||||||
|
if self._mask_ds is None:
|
||||||
|
raise ValueError(f"无法打开掩膜文件: {wm}")
|
||||||
|
band = self._mask_ds.GetRasterBand(1)
|
||||||
|
if band.XSize != self.width or band.YSize != self.height:
|
||||||
|
raise ValueError(
|
||||||
|
f"掩膜尺寸 {(band.YSize, band.XSize)} 与图像尺寸 {(self.height, self.width)} 不匹配")
|
||||||
|
self._wm_band = band
|
||||||
|
return
|
||||||
|
|
||||||
try:
|
raise ValueError(f"不支持的掩膜类型: {type(wm)}")
|
||||||
from osgeo import ogr
|
|
||||||
ref_dataset = gdal.Open(ref_path, gdal.GA_ReadOnly)
|
|
||||||
if ref_dataset is None:
|
|
||||||
raise ValueError(f"无法打开参考栅格文件: {ref_path}")
|
|
||||||
|
|
||||||
geotransform = ref_dataset.GetGeoTransform()
|
def _read_mask_block(self, x_off, y_off, x_size, y_size):
|
||||||
projection = ref_dataset.GetProjection()
|
"""分块读取掩膜布尔数组(True=水体);无掩膜时返回 None。"""
|
||||||
width = ref_dataset.RasterXSize
|
self._open_mask()
|
||||||
height = ref_dataset.RasterYSize
|
if self._wm_band is not None:
|
||||||
|
arr = self._wm_band.ReadAsArray(x_off, y_off, x_size, y_size)
|
||||||
|
return (arr > 0)
|
||||||
|
if self._wm_np is not None:
|
||||||
|
return self._wm_np[y_off:y_off + y_size, x_off:x_off + x_size]
|
||||||
|
return None
|
||||||
|
|
||||||
# 创建内存中的栅格数据集
|
def _legacy_full_mask_bool(self):
|
||||||
mem_driver = gdal.GetDriverByName('MEM')
|
"""仅供小图/数组模式的旧全幅路径使用;大影像分块模式不会调用。"""
|
||||||
mask_dataset = mem_driver.Create('', width, height, 1, gdal.GDT_Byte)
|
self._open_mask()
|
||||||
mask_dataset.SetGeoTransform(geotransform)
|
if self._wm_band is not None:
|
||||||
mask_dataset.SetProjection(projection)
|
return self._wm_band.ReadAsArray() > 0
|
||||||
|
return self._wm_np
|
||||||
mask_band = mask_dataset.GetRasterBand(1)
|
|
||||||
mask_band.Fill(0)
|
|
||||||
|
|
||||||
# 打开shp文件
|
|
||||||
shp_dataset = ogr.Open(water_mask)
|
|
||||||
if shp_dataset is None:
|
|
||||||
raise ValueError(f"无法打开shp文件: {water_mask}")
|
|
||||||
|
|
||||||
layer = shp_dataset.GetLayer()
|
|
||||||
gdal.RasterizeLayer(mask_dataset, [1], layer, burn_values=[1])
|
|
||||||
|
|
||||||
water_mask_array = mask_band.ReadAsArray()
|
|
||||||
|
|
||||||
ref_dataset = None
|
|
||||||
mask_dataset = None
|
|
||||||
shp_dataset = None
|
|
||||||
|
|
||||||
return (water_mask_array > 0).astype(np.uint8)
|
|
||||||
except Exception as e:
|
|
||||||
raise ValueError(f"从shp文件创建掩膜时出错: {e}")
|
|
||||||
else:
|
|
||||||
# 栅格文件
|
|
||||||
mask_dataset = gdal.Open(water_mask, gdal.GA_ReadOnly)
|
|
||||||
if mask_dataset is None:
|
|
||||||
raise ValueError(f"无法打开掩膜文件: {water_mask}")
|
|
||||||
|
|
||||||
mask_array = mask_dataset.GetRasterBand(1).ReadAsArray()
|
|
||||||
mask_dataset = None
|
|
||||||
|
|
||||||
if mask_array.shape != (self.height, self.width):
|
|
||||||
raise ValueError(f"掩膜尺寸 {mask_array.shape} 与图像尺寸 {(self.height, self.width)} 不匹配")
|
|
||||||
|
|
||||||
return (mask_array > 0).astype(np.uint8)
|
|
||||||
|
|
||||||
raise ValueError(f"不支持的掩膜类型: {type(water_mask)}")
|
|
||||||
|
|
||||||
def _get_corrected_bands_numpy(self):
|
def _get_corrected_bands_numpy(self):
|
||||||
"""
|
"""
|
||||||
@ -224,8 +212,8 @@ class Goodman:
|
|||||||
diff_640_750 = R_640 - R_750
|
diff_640_750 = R_640 - R_750
|
||||||
corrected_bands = []
|
corrected_bands = []
|
||||||
|
|
||||||
# 获取水域掩膜(如果存在)
|
# 获取水域掩膜(numpy 数组模式直接引用 _wm_np)
|
||||||
water_mask_bool = self.water_mask.astype(bool) if self.water_mask is not None else None
|
water_mask_bool = self._wm_np if self._wm_np is not None else None
|
||||||
|
|
||||||
# 逐波段处理:每次只处理一个波段,处理完后立即添加到结果列表
|
# 逐波段处理:每次只处理一个波段,处理完后立即添加到结果列表
|
||||||
for i in tqdm(range(self.n_bands), desc="处理波段 (numpy)", total=self.n_bands, disable=_is_frozen_gui):
|
for i in tqdm(range(self.n_bands), desc="处理波段 (numpy)", total=self.n_bands, disable=_is_frozen_gui):
|
||||||
@ -267,8 +255,8 @@ class Goodman:
|
|||||||
diff_640_750 = R_640 - R_750
|
diff_640_750 = R_640 - R_750
|
||||||
del R_640 # 释放不再需要的 R_640
|
del R_640 # 释放不再需要的 R_640
|
||||||
|
|
||||||
# 水域掩膜 + 水体占比检测(控制行段跳跃模式)
|
# 水域掩膜 + 水体占比检测(★ 旧全幅路径:仅无 output_path 的传统返回列表模式)
|
||||||
water = self.water_mask.astype(bool) if self.water_mask is not None else None
|
water = self._legacy_full_mask_bool()
|
||||||
water_pct = (100.0 * np.count_nonzero(water) / water.size) if water is not None else 100.0
|
water_pct = (100.0 * np.count_nonzero(water) / water.size) if water is not None else 100.0
|
||||||
|
|
||||||
# 行段跳跃: 水域占比 < 50% 时启用
|
# 行段跳跃: 水域占比 < 50% 时启用
|
||||||
@ -400,15 +388,8 @@ class Goodman:
|
|||||||
else:
|
else:
|
||||||
bsq_path = self.output_path
|
bsq_path = self.output_path
|
||||||
|
|
||||||
# 使用ENVI驱动(默认就是BSQ格式)
|
# 创建输出数据集(大图自动用 GTiff+BIGTIFF,规避 ENVI int32 偏移溢出)
|
||||||
driver = gdal.GetDriverByName('ENVI')
|
dataset = create_output_dataset(bsq_path, width, height, n_bands, gdal.GDT_Float32)
|
||||||
if driver is None:
|
|
||||||
raise ValueError("无法创建ENVI格式文件,ENVI驱动不可用")
|
|
||||||
|
|
||||||
# 创建ENVI格式数据集(会自动生成.hdr文件)
|
|
||||||
dataset = driver.Create(bsq_path, width, height, n_bands, gdal.GDT_Float32)
|
|
||||||
if dataset is None:
|
|
||||||
raise ValueError(f"无法创建输出文件: {bsq_path}")
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# 设置地理变换和投影
|
# 设置地理变换和投影
|
||||||
@ -461,28 +442,23 @@ class Goodman:
|
|||||||
return self._get_corrected_bands_numpy()
|
return self._get_corrected_bands_numpy()
|
||||||
|
|
||||||
def _get_corrected_bands_streaming(self):
|
def _get_corrected_bands_streaming(self):
|
||||||
"""流式处理:逐波段校正并直接写入输出文件,不累积内存
|
"""★ v3 分块流式:块外循环 × 波段内循环 × 掩膜分块读 × 逐块写盘。
|
||||||
|
|
||||||
适用于大尺度影像(如 6522×13215×150)。
|
不再整幅驻留 R640/R750/diff 或整幅掩膜;每块只驻留 block_size² 数据,
|
||||||
内存峰值 ≈ 3 个全波段数组(NIR×2 + 当前波段),而非全部 150 个波段。
|
与 Hedley/Kutser/SUGAR 的分块范式完全对齐,内存恒定。
|
||||||
|
|
||||||
:return: None(波段已在输出文件中)
|
:return: None(波段已写入输出文件)
|
||||||
"""
|
"""
|
||||||
import os
|
import os
|
||||||
|
|
||||||
# ── 创建输出文件 ──
|
|
||||||
base_path, ext = os.path.splitext(self.output_path)
|
base_path, ext = os.path.splitext(self.output_path)
|
||||||
bsq_path = base_path + '.bsq' if ext.lower() != '.bsq' else self.output_path
|
bsq_path = base_path + '.bsq' if ext.lower() != '.bsq' else self.output_path
|
||||||
output_dir = os.path.dirname(bsq_path)
|
output_dir = os.path.dirname(bsq_path)
|
||||||
if output_dir and not os.path.exists(output_dir):
|
if output_dir and not os.path.exists(output_dir):
|
||||||
os.makedirs(output_dir, exist_ok=True)
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
|
||||||
driver = gdal.GetDriverByName('ENVI')
|
out_ds = create_output_dataset(
|
||||||
out_ds = driver.Create(
|
bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32)
|
||||||
bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32
|
|
||||||
)
|
|
||||||
if out_ds is None:
|
|
||||||
raise ValueError(f"无法创建输出文件: {bsq_path}")
|
|
||||||
|
|
||||||
# ── 设置地理参考 ──
|
# ── 设置地理参考 ──
|
||||||
if self.is_file_path and self.dataset is not None:
|
if self.is_file_path and self.dataset is not None:
|
||||||
@ -490,22 +466,85 @@ class Goodman:
|
|||||||
out_ds.SetProjection(self.dataset.GetProjection())
|
out_ds.SetProjection(self.dataset.GetProjection())
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# 逐波段处理 + 立即写入(波段在 _get_corrected_bands_gdal 的循环中
|
self._run_block_correction(out_ds)
|
||||||
# 由 WriteArray 写入 → FlushCache → del,不会累积)
|
|
||||||
self._get_corrected_bands_gdal(out_dataset=out_ds)
|
|
||||||
finally:
|
finally:
|
||||||
out_ds = None # 关闭文件,确保数据落盘
|
out_ds = None # 关闭文件,确保数据落盘
|
||||||
|
if self._wm_temp is not None and os.path.exists(self._wm_temp):
|
||||||
|
try:
|
||||||
|
os.remove(self._wm_temp)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
self._wm_temp = None
|
||||||
|
|
||||||
# ── 日志 ──
|
|
||||||
hdr_path = bsq_path + '.hdr'
|
hdr_path = bsq_path + '.hdr'
|
||||||
if os.path.exists(hdr_path):
|
if os.path.exists(hdr_path):
|
||||||
print(f"校正后的图像已保存至: {bsq_path} (BSQ格式, 流式写入)")
|
print(f"校正后的图像已保存至: {bsq_path} (BSQ格式, 分块流式写入)")
|
||||||
else:
|
else:
|
||||||
print(f"校正后的图像已保存至: {bsq_path} (BSQ格式)")
|
print(f"校正后的图像已保存至: {bsq_path} (BSQ格式)")
|
||||||
print("警告: 未检测到.hdr文件,但GDAL应该已自动创建")
|
print("警告: 未检测到.hdr文件,但GDAL应该已自动创建")
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def _run_block_correction(self, out_ds):
|
||||||
|
"""整幅宽 × 大行带(row-band)顺序流式校正(性能优化 v4)。
|
||||||
|
|
||||||
|
原 v3 按"1000×1000 x块 × 波段"小窗写入,且每块每波段 FlushCache,
|
||||||
|
导致 ~10 万次磁盘同步 + 大量小 strip 随机写 → 大图极慢。
|
||||||
|
现改为整幅宽、2048 行/条的行带:
|
||||||
|
- 外层逐行带、内层逐波段,每波段在该行带内只读/写一次连续大块(顺序 IO)
|
||||||
|
- 每波段每条行带只 FlushCache 一次(总落盘 ≈ 150 × 17 ≈ 2600 次)
|
||||||
|
内存恒定:行带(2048×宽) 参考波段 float + 单波段 float ≈ 1~1.5GB。
|
||||||
|
|
||||||
|
Goodman 公式为纯逐像素运算(无全局统计):
|
||||||
|
R_corrected = R - R750 + A + B * (R640 - R750)
|
||||||
|
有掩膜时仅对水体像素应用校正,陆地保留原值;全程保持非负。
|
||||||
|
"""
|
||||||
|
row_band = 2048
|
||||||
|
band_640 = self.dataset.GetRasterBand(self.NIR_lower + 1)
|
||||||
|
band_750 = self.dataset.GetRasterBand(self.NIR_upper + 1)
|
||||||
|
|
||||||
|
n_bands = self.n_bands
|
||||||
|
n_rows = (self.height + row_band - 1) // row_band
|
||||||
|
import time as _t
|
||||||
|
_p_t0 = _t.time()
|
||||||
|
_est_gb = self.width * self.height * n_bands * 4 / 2**30
|
||||||
|
print(f"[Goodman] 行带流式校正启动: {n_rows} 条行带 × {n_bands} 波段, "
|
||||||
|
f"输出约 {_est_gb:.0f}GB。", flush=True)
|
||||||
|
|
||||||
|
for ri, y_off in enumerate(range(0, self.height, row_band)):
|
||||||
|
ys = min(row_band, self.height - y_off)
|
||||||
|
|
||||||
|
R640 = band_640.ReadAsArray(0, y_off, self.width, ys).astype(np.float32)
|
||||||
|
R750 = band_750.ReadAsArray(0, y_off, self.width, ys).astype(np.float32)
|
||||||
|
diff = R640 - R750
|
||||||
|
del R640
|
||||||
|
maskb = self._read_mask_block(0, y_off, self.width, ys)
|
||||||
|
|
||||||
|
for i in range(n_bands):
|
||||||
|
band = self.dataset.GetRasterBand(i + 1)
|
||||||
|
R = band.ReadAsArray(0, y_off, self.width, ys).astype(np.float32)
|
||||||
|
|
||||||
|
corr = R - R750 + self.A + self.B * diff
|
||||||
|
np.maximum(corr, 0.0, out=corr)
|
||||||
|
if maskb is not None:
|
||||||
|
corr = np.where(maskb, corr, R)
|
||||||
|
np.maximum(corr, 0.0, out=corr)
|
||||||
|
|
||||||
|
ob = out_ds.GetRasterBand(i + 1)
|
||||||
|
ob.WriteArray(corr.astype(np.float32), 0, y_off)
|
||||||
|
ob.FlushCache()
|
||||||
|
del R, corr
|
||||||
|
|
||||||
|
del R750, diff
|
||||||
|
|
||||||
|
if not _is_frozen_gui and (ri == 0 or (ri + 1) % max(1, n_rows // 10) == 0
|
||||||
|
or ri == n_rows - 1):
|
||||||
|
_el = _t.time() - _p_t0
|
||||||
|
_eta = _el / (ri + 1) * (n_rows - ri - 1)
|
||||||
|
print(f" [Goodman] 行带 {ri + 1}/{n_rows} "
|
||||||
|
f"(已用 {_el:.0f}s, 预计剩余 {_eta:.0f}s)", flush=True)
|
||||||
|
return None
|
||||||
|
|
||||||
def __del__(self):
|
def __del__(self):
|
||||||
"""清理资源"""
|
"""清理资源"""
|
||||||
if self.dataset is not None and self.is_file_path:
|
if self.dataset is not None and self.is_file_path:
|
||||||
|
|||||||
@ -7,7 +7,7 @@ try:
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
GDAL_AVAILABLE = False
|
GDAL_AVAILABLE = False
|
||||||
|
|
||||||
from src.utils.util import find_band_number
|
from src.utils.util import find_band_number, create_output_dataset
|
||||||
|
|
||||||
|
|
||||||
class Hedley:
|
class Hedley:
|
||||||
@ -36,6 +36,10 @@ class Hedley:
|
|||||||
print(f"[Hedley] 波段解析: NIR={float(nir_wavelength):.1f}nm→band{self.NIR_band}")
|
print(f"[Hedley] 波段解析: NIR={float(nir_wavelength):.1f}nm→band{self.NIR_band}")
|
||||||
self.water_mask = None
|
self.water_mask = None
|
||||||
self.water_mask_path = water_mask
|
self.water_mask_path = water_mask
|
||||||
|
# ★ v3:掩膜不再整幅驻留——栅格路径保留 GDAL band 句柄,按需分块读
|
||||||
|
self._mask_ds = None
|
||||||
|
self._wm_band = None
|
||||||
|
self._wm_np = None
|
||||||
self.output_path = output_path
|
self.output_path = output_path
|
||||||
self.block_size = block_size
|
self.block_size = block_size
|
||||||
self.R_min = None
|
self.R_min = None
|
||||||
@ -49,32 +53,46 @@ class Hedley:
|
|||||||
self.height = self.dataset.RasterYSize
|
self.height = self.dataset.RasterYSize
|
||||||
self.n_bands = self.dataset.RasterCount
|
self.n_bands = self.dataset.RasterCount
|
||||||
|
|
||||||
def _load_water_mask(self):
|
def _open_mask(self):
|
||||||
"""延迟加载水域掩膜"""
|
"""打开水域掩膜,但绝不整幅 ReadAsArray 驻留。
|
||||||
|
|
||||||
|
栅格路径 → 保留 GDAL dataset/band 句柄(_wm_band),按需分块读取;
|
||||||
|
np.ndarray(外部已驻留)→ 转 bool 后保留 _wm_np,按窗口切片。
|
||||||
|
"""
|
||||||
|
if self._wm_band is not None or self._wm_np is not None:
|
||||||
|
return
|
||||||
if self.water_mask_path is None:
|
if self.water_mask_path is None:
|
||||||
return None
|
return
|
||||||
|
|
||||||
if isinstance(self.water_mask_path, np.ndarray):
|
if isinstance(self.water_mask_path, np.ndarray):
|
||||||
if self.water_mask_path.shape[:2] != (self.height, self.width):
|
if self.water_mask_path.shape[:2] != (self.height, self.width):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"掩膜尺寸 {self.water_mask_path.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
f"掩膜尺寸 {self.water_mask_path.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
||||||
)
|
)
|
||||||
return (self.water_mask_path > 0).astype(np.uint8)
|
self._wm_np = self.water_mask_path > 0
|
||||||
|
return
|
||||||
|
|
||||||
if isinstance(self.water_mask_path, str):
|
if isinstance(self.water_mask_path, str):
|
||||||
if self.water_mask_path.lower().endswith('.shp'):
|
if self.water_mask_path.lower().endswith('.shp'):
|
||||||
raise ValueError("请先栅格化shapefile为栅格掩膜文件")
|
raise ValueError("请先栅格化shapefile为栅格掩膜文件")
|
||||||
mask_dataset = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly)
|
self._mask_ds = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly)
|
||||||
if mask_dataset is None:
|
if self._mask_ds is None:
|
||||||
raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}")
|
raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}")
|
||||||
mask_array = mask_dataset.GetRasterBand(1).ReadAsArray()
|
band = self._mask_ds.GetRasterBand(1)
|
||||||
mask_dataset = None
|
if band.XSize != self.width or band.YSize != self.height:
|
||||||
if mask_array.shape != (self.height, self.width):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"掩膜尺寸 {mask_array.shape} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
f"掩膜尺寸 {(band.YSize, band.XSize)} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
||||||
)
|
)
|
||||||
return (mask_array > 0).astype(np.uint8)
|
self._wm_band = band
|
||||||
|
|
||||||
|
def _read_mask_block(self, x_off, y_off, x_size, y_size):
|
||||||
|
"""分块读取掩膜布尔数组(True=水体);无掩膜时返回 None。"""
|
||||||
|
self._open_mask()
|
||||||
|
if self._wm_band is not None:
|
||||||
|
arr = self._wm_band.ReadAsArray(x_off, y_off, x_size, y_size)
|
||||||
|
return (arr > 0)
|
||||||
|
if self._wm_np is not None:
|
||||||
|
return self._wm_np[y_off:y_off + y_size, x_off:x_off + x_size]
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def covariance_NIR(self, NIR, b):
|
def covariance_NIR(self, NIR, b):
|
||||||
@ -93,7 +111,7 @@ class Hedley:
|
|||||||
使用重采样方式扫描,大幅降低内存占用。
|
使用重采样方式扫描,大幅降低内存占用。
|
||||||
"""
|
"""
|
||||||
print(f"[Hedley] 扫描全局统计量(采样步长={sample_step})...")
|
print(f"[Hedley] 扫描全局统计量(采样步长={sample_step})...")
|
||||||
water_mask = self._load_water_mask()
|
self._open_mask()
|
||||||
|
|
||||||
nir_samples = []
|
nir_samples = []
|
||||||
sample_count = 0
|
sample_count = 0
|
||||||
@ -106,10 +124,9 @@ class Hedley:
|
|||||||
nir_block = nir_band.ReadAsArray(0, y_off, self.width, block_height)
|
nir_block = nir_band.ReadAsArray(0, y_off, self.width, block_height)
|
||||||
nir_band = None
|
nir_band = None
|
||||||
|
|
||||||
if water_mask is not None:
|
# 掩膜同步分块读取(不再整幅驻留)
|
||||||
mask_block = water_mask[y_off:y_end, :]
|
mask_bool = self._read_mask_block(0, y_off, self.width, block_height)
|
||||||
mask_bool = mask_block.astype(bool)
|
if mask_bool is None:
|
||||||
else:
|
|
||||||
mask_bool = np.ones((block_height, self.width), dtype=bool)
|
mask_bool = np.ones((block_height, self.width), dtype=bool)
|
||||||
|
|
||||||
if mask_bool.any():
|
if mask_bool.any():
|
||||||
@ -117,7 +134,7 @@ class Hedley:
|
|||||||
nir_samples.append(nir_sampled)
|
nir_samples.append(nir_sampled)
|
||||||
sample_count += nir_sampled.size
|
sample_count += nir_sampled.size
|
||||||
|
|
||||||
del nir_block, mask_block
|
del nir_block
|
||||||
|
|
||||||
if sample_count == 0:
|
if sample_count == 0:
|
||||||
self.R_min = 0.0
|
self.R_min = 0.0
|
||||||
@ -136,7 +153,7 @@ class Hedley:
|
|||||||
由于需要相关性计算,需要足够多的样本,取sample_step=5
|
由于需要相关性计算,需要足够多的样本,取sample_step=5
|
||||||
"""
|
"""
|
||||||
print(f"[Hedley] 计算全局协方差系数列表(采样步长={sample_step})...")
|
print(f"[Hedley] 计算全局协方差系数列表(采样步长={sample_step})...")
|
||||||
water_mask = self._load_water_mask()
|
self._open_mask()
|
||||||
|
|
||||||
# 预收集NIR和每个波段的样本数据
|
# 预收集NIR和每个波段的样本数据
|
||||||
nir_samples = []
|
nir_samples = []
|
||||||
@ -151,11 +168,9 @@ class Hedley:
|
|||||||
nir_block = nir_band.ReadAsArray(0, y_off, self.width, block_height).astype(np.float32)
|
nir_block = nir_band.ReadAsArray(0, y_off, self.width, block_height).astype(np.float32)
|
||||||
nir_band = None
|
nir_band = None
|
||||||
|
|
||||||
# 取 NIR 样本(每块只取一次,放在波段循环外)
|
# 取 NIR 样本(每块只取一次,放在波段循环外;掩膜同步分块读)
|
||||||
if water_mask is not None:
|
mask_bool = self._read_mask_block(0, y_off, self.width, block_height)
|
||||||
mask_block = water_mask[y_off:y_end, :]
|
if mask_bool is None:
|
||||||
mask_bool = mask_block.astype(bool)
|
|
||||||
else:
|
|
||||||
mask_bool = np.ones((block_height, self.width), dtype=bool)
|
mask_bool = np.ones((block_height, self.width), dtype=bool)
|
||||||
|
|
||||||
if mask_bool.any():
|
if mask_bool.any():
|
||||||
@ -208,14 +223,8 @@ class Hedley:
|
|||||||
# 预计算 NIR - R_min
|
# 预计算 NIR - R_min
|
||||||
NIR_diff = NIR - self.R_min
|
NIR_diff = NIR - self.R_min
|
||||||
|
|
||||||
# 获取掩膜
|
# 获取掩膜(同步分块读取,不再整幅驻留)
|
||||||
water_mask = self._load_water_mask()
|
mask_block = self._read_mask_block(x_off, y_off, x_size, y_size)
|
||||||
if water_mask is not None:
|
|
||||||
y_end = y_off + y_size
|
|
||||||
x_end = x_off + x_size
|
|
||||||
mask_block = water_mask[y_off:y_end, x_off:x_end].astype(bool)
|
|
||||||
else:
|
|
||||||
mask_block = None
|
|
||||||
|
|
||||||
# 逐波段处理
|
# 逐波段处理
|
||||||
corrected_bands = []
|
corrected_bands = []
|
||||||
@ -262,11 +271,8 @@ class Hedley:
|
|||||||
geotransform = self.dataset.GetGeoTransform()
|
geotransform = self.dataset.GetGeoTransform()
|
||||||
projection = self.dataset.GetProjection()
|
projection = self.dataset.GetProjection()
|
||||||
|
|
||||||
driver = gdal.GetDriverByName('ENVI')
|
out_dataset = create_output_dataset(
|
||||||
out_dataset = driver.Create(bsq_path, self.width, self.height,
|
bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32)
|
||||||
self.n_bands, gdal.GDT_Float32)
|
|
||||||
if out_dataset is None:
|
|
||||||
raise ValueError(f"无法创建输出文件: {bsq_path}")
|
|
||||||
|
|
||||||
out_dataset.SetGeoTransform(geotransform)
|
out_dataset.SetGeoTransform(geotransform)
|
||||||
out_dataset.SetProjection(projection)
|
out_dataset.SetProjection(projection)
|
||||||
|
|||||||
@ -7,7 +7,7 @@ try:
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
GDAL_AVAILABLE = False
|
GDAL_AVAILABLE = False
|
||||||
|
|
||||||
from src.utils.util import find_band_number
|
from src.utils.util import find_band_number, create_output_dataset
|
||||||
|
|
||||||
|
|
||||||
class Kutser:
|
class Kutser:
|
||||||
@ -49,6 +49,10 @@ class Kutser:
|
|||||||
f"NIR={float(nir_wavelength):.1f}nm→band{self.NIR_band}")
|
f"NIR={float(nir_wavelength):.1f}nm→band{self.NIR_band}")
|
||||||
self.water_mask = None # 延迟加载,在处理前初始化
|
self.water_mask = None # 延迟加载,在处理前初始化
|
||||||
self.water_mask_path = water_mask
|
self.water_mask_path = water_mask
|
||||||
|
# ★ v3:掩膜不再整幅驻留——栅格路径保留 GDAL band 句柄,按需分块读
|
||||||
|
self._mask_ds = None
|
||||||
|
self._wm_band = None
|
||||||
|
self._wm_np = None
|
||||||
self.output_path = output_path
|
self.output_path = output_path
|
||||||
self.block_size = block_size
|
self.block_size = block_size
|
||||||
self.R_min = None # 全局R_min(来自重采样扫描)
|
self.R_min = None # 全局R_min(来自重采样扫描)
|
||||||
@ -80,32 +84,59 @@ class Kutser:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def _load_water_mask(self):
|
def _open_mask(self):
|
||||||
"""延迟加载水域掩膜"""
|
"""打开水域掩膜,但绝不整幅 ReadAsArray 驻留。
|
||||||
|
|
||||||
|
栅格路径 → 保留 GDAL dataset/band 句柄(_wm_band),按需分块读取;
|
||||||
|
np.ndarray(外部已驻留)→ 转 bool 后保留 _wm_np,按窗口切片。
|
||||||
|
"""
|
||||||
|
if self._wm_band is not None or self._wm_np is not None:
|
||||||
|
return
|
||||||
if self.water_mask_path is None:
|
if self.water_mask_path is None:
|
||||||
return None
|
return
|
||||||
|
|
||||||
if isinstance(self.water_mask_path, np.ndarray):
|
if isinstance(self.water_mask_path, np.ndarray):
|
||||||
if self.water_mask_path.shape[:2] != (self.height, self.width):
|
if self.water_mask_path.shape[:2] != (self.height, self.width):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"掩膜尺寸 {self.water_mask_path.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
f"掩膜尺寸 {self.water_mask_path.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
||||||
)
|
)
|
||||||
return (self.water_mask_path > 0).astype(np.uint8)
|
self._wm_np = self.water_mask_path > 0
|
||||||
|
return
|
||||||
|
|
||||||
if isinstance(self.water_mask_path, str):
|
if isinstance(self.water_mask_path, str):
|
||||||
if self.water_mask_path.lower().endswith('.shp'):
|
if self.water_mask_path.lower().endswith('.shp'):
|
||||||
raise ValueError("请先栅格化shapefile为栅格掩膜文件")
|
raise ValueError("请先栅格化shapefile为栅格掩膜文件")
|
||||||
mask_dataset = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly)
|
self._mask_ds = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly)
|
||||||
if mask_dataset is None:
|
if self._mask_ds is None:
|
||||||
raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}")
|
raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}")
|
||||||
mask_array = mask_dataset.GetRasterBand(1).ReadAsArray()
|
band = self._mask_ds.GetRasterBand(1)
|
||||||
mask_dataset = None
|
if band.XSize != self.width or band.YSize != self.height:
|
||||||
if mask_array.shape != (self.height, self.width):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"掩膜尺寸 {mask_array.shape} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
f"掩膜尺寸 {(band.YSize, band.XSize)} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
||||||
)
|
)
|
||||||
return (mask_array > 0).astype(np.uint8)
|
self._wm_band = band
|
||||||
|
|
||||||
|
def _read_mask_block(self, x_off, y_off, x_size, y_size):
|
||||||
|
"""分块读取掩膜布尔数组(True=水体);无掩膜时返回 None。"""
|
||||||
|
self._open_mask()
|
||||||
|
if self._wm_band is not None:
|
||||||
|
arr = self._wm_band.ReadAsArray(x_off, y_off, x_size, y_size)
|
||||||
|
return (arr > 0)
|
||||||
|
if self._wm_np is not None:
|
||||||
|
return self._wm_np[y_off:y_off + y_size, x_off:x_off + x_size]
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _legacy_full_mask_bool(self):
|
||||||
|
"""仅供 run_fast(整幅顺序模式)使用:需要整幅布尔掩膜。
|
||||||
|
|
||||||
|
大影像请使用默认入口 run()(分块 + 掩膜分块读),避免整幅掩膜峰值。
|
||||||
|
"""
|
||||||
|
self._open_mask()
|
||||||
|
if self._wm_band is not None:
|
||||||
|
full = self._wm_band.ReadAsArray()
|
||||||
|
return (full > 0)
|
||||||
|
if self._wm_np is not None:
|
||||||
|
return self._wm_np
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _scan_global_stats(self, sample_step=20):
|
def _scan_global_stats(self, sample_step=20):
|
||||||
@ -116,7 +147,7 @@ class Kutser:
|
|||||||
内存峰值 ≈ 单波段块大小 + 几个掩膜数组 ≈ block_size² × 4~8MB
|
内存峰值 ≈ 单波段块大小 + 几个掩膜数组 ≈ block_size² × 4~8MB
|
||||||
"""
|
"""
|
||||||
print(f"[Kutser] 扫描全局统计量(采样步长={sample_step})...")
|
print(f"[Kutser] 扫描全局统计量(采样步长={sample_step})...")
|
||||||
water_mask = self._load_water_mask()
|
self._open_mask()
|
||||||
|
|
||||||
# 预分配采样数组(NIR波段和D值)
|
# 预分配采样数组(NIR波段和D值)
|
||||||
nir_samples = []
|
nir_samples = []
|
||||||
@ -148,14 +179,10 @@ class Kutser:
|
|||||||
# 计算D = (lower + upper) * 0.5 - oxy
|
# 计算D = (lower + upper) * 0.5 - oxy
|
||||||
d_block = (lower_block.astype(np.float32) + upper_block.astype(np.float32)) * 0.5 - oxy_block.astype(np.float32)
|
d_block = (lower_block.astype(np.float32) + upper_block.astype(np.float32)) * 0.5 - oxy_block.astype(np.float32)
|
||||||
|
|
||||||
# 获取掩膜(整块)
|
# 获取掩膜(与波段同步分块读取,不再整幅驻留)
|
||||||
if water_mask is not None:
|
mask_bool = self._read_mask_block(0, y_off, self.width, block_height)
|
||||||
mask_block = water_mask[y_off:y_end, :]
|
if mask_bool is None:
|
||||||
else:
|
mask_bool = np.ones((block_height, self.width), dtype=bool)
|
||||||
mask_block = np.ones((block_height, self.width), dtype=np.uint8)
|
|
||||||
|
|
||||||
# 对掩膜区域进行采样
|
|
||||||
mask_bool = mask_block.astype(bool)
|
|
||||||
|
|
||||||
if mask_bool.any():
|
if mask_bool.any():
|
||||||
# 按步长采样
|
# 按步长采样
|
||||||
@ -166,7 +193,7 @@ class Kutser:
|
|||||||
sample_count += nir_sampled.size
|
sample_count += nir_sampled.size
|
||||||
|
|
||||||
# 显式释放块内存
|
# 显式释放块内存
|
||||||
del nir_block, lower_block, upper_block, oxy_block, d_block, mask_block
|
del nir_block, lower_block, upper_block, oxy_block, d_block
|
||||||
|
|
||||||
# 汇总
|
# 汇总
|
||||||
if sample_count == 0:
|
if sample_count == 0:
|
||||||
@ -189,7 +216,7 @@ class Kutser:
|
|||||||
使用全分辨率扫描,但逐波段读取,每波段内存 ≈ block_size²
|
使用全分辨率扫描,但逐波段读取,每波段内存 ≈ block_size²
|
||||||
"""
|
"""
|
||||||
print(f"[Kutser] 计算全局G值列表(n_bands={self.n_bands})...")
|
print(f"[Kutser] 计算全局G值列表(n_bands={self.n_bands})...")
|
||||||
water_mask = self._load_water_mask()
|
self._open_mask()
|
||||||
|
|
||||||
# 初始化G_max和G_min为极值
|
# 初始化G_max和G_min为极值
|
||||||
g_max = np.full(self.n_bands, -np.inf, dtype=np.float32)
|
g_max = np.full(self.n_bands, -np.inf, dtype=np.float32)
|
||||||
@ -200,22 +227,21 @@ class Kutser:
|
|||||||
y_end = min(y_off + self.block_size, self.height)
|
y_end = min(y_off + self.block_size, self.height)
|
||||||
block_height = y_end - y_off
|
block_height = y_end - y_off
|
||||||
|
|
||||||
|
# 掩膜块只读一次(同步分块,不再整幅驻留)
|
||||||
|
mask_bool = self._read_mask_block(0, y_off, self.width, block_height)
|
||||||
|
if mask_bool is None:
|
||||||
|
mask_bool = np.ones((block_height, self.width), dtype=bool)
|
||||||
|
|
||||||
# 读取所有波段的当前块
|
# 读取所有波段的当前块
|
||||||
for b in range(self.n_bands):
|
for b in range(self.n_bands):
|
||||||
band = self.dataset.GetRasterBand(b + 1)
|
band = self.dataset.GetRasterBand(b + 1)
|
||||||
block = band.ReadAsArray(0, y_off, self.width, block_height).astype(np.float32)
|
block = band.ReadAsArray(0, y_off, self.width, block_height).astype(np.float32)
|
||||||
band = None
|
band = None
|
||||||
|
|
||||||
if water_mask is not None:
|
if mask_bool.any():
|
||||||
mask_block = water_mask[y_off:y_end, :]
|
band_masked = block[mask_bool]
|
||||||
mask_bool = mask_block.astype(bool)
|
g_max[b] = max(g_max[b], band_masked.max())
|
||||||
if mask_bool.any():
|
g_min[b] = min(g_min[b], band_masked.min())
|
||||||
band_masked = block[mask_bool]
|
|
||||||
g_max[b] = max(g_max[b], band_masked.max())
|
|
||||||
g_min[b] = min(g_min[b], band_masked.min())
|
|
||||||
else:
|
|
||||||
g_max[b] = max(g_max[b], block.max())
|
|
||||||
g_min[b] = min(g_min[b], block.min())
|
|
||||||
|
|
||||||
del block
|
del block
|
||||||
|
|
||||||
@ -254,14 +280,8 @@ class Kutser:
|
|||||||
# 释放临时块
|
# 释放临时块
|
||||||
del lower_block, upper_block, oxy_block, D
|
del lower_block, upper_block, oxy_block, D
|
||||||
|
|
||||||
# 获取当前块的水域掩膜
|
# 获取当前块的水域掩膜(同步分块读取,不再整幅驻留)
|
||||||
water_mask = self._load_water_mask()
|
mask_block = self._read_mask_block(x_off, y_off, x_size, y_size)
|
||||||
if water_mask is not None:
|
|
||||||
y_end = y_off + y_size
|
|
||||||
x_end = x_off + x_size
|
|
||||||
mask_block = water_mask[y_off:y_end, x_off:x_end].astype(bool)
|
|
||||||
else:
|
|
||||||
mask_block = None
|
|
||||||
|
|
||||||
# 逐波段处理
|
# 逐波段处理
|
||||||
corrected_bands = []
|
corrected_bands = []
|
||||||
@ -296,8 +316,10 @@ class Kutser:
|
|||||||
if self.output_path is None:
|
if self.output_path is None:
|
||||||
raise ValueError("output_path 必须提供,分块处理需要直接写入文件")
|
raise ValueError("output_path 必须提供,分块处理需要直接写入文件")
|
||||||
|
|
||||||
# 全波段顺序模式:每个波段完整读取一次,顺序IO
|
# ★ v3:默认走分块 run()(掩膜同步分块读取,内存恒定)。
|
||||||
return self.run_fast()
|
# 旧的 run_fast()(整幅参考波段 + 整幅掩膜)仅保留给明确需要
|
||||||
|
# 顺序 IO 加速的小图,不再作为默认入口,避免大影像内存峰值。
|
||||||
|
return self.run()
|
||||||
|
|
||||||
def run(self):
|
def run(self):
|
||||||
"""原版分块模式(低内存兼容)。默认使用 run_fast() 代替。"""
|
"""原版分块模式(低内存兼容)。默认使用 run_fast() 代替。"""
|
||||||
@ -322,11 +344,8 @@ class Kutser:
|
|||||||
geotransform = self.dataset.GetGeoTransform()
|
geotransform = self.dataset.GetGeoTransform()
|
||||||
projection = self.dataset.GetProjection()
|
projection = self.dataset.GetProjection()
|
||||||
|
|
||||||
driver = gdal.GetDriverByName('ENVI')
|
out_dataset = create_output_dataset(
|
||||||
out_dataset = driver.Create(bsq_path, self.width, self.height,
|
bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32)
|
||||||
self.n_bands, gdal.GDT_Float32)
|
|
||||||
if out_dataset is None:
|
|
||||||
raise ValueError(f"无法创建输出文件: {bsq_path}")
|
|
||||||
|
|
||||||
out_dataset.SetGeoTransform(geotransform)
|
out_dataset.SetGeoTransform(geotransform)
|
||||||
out_dataset.SetProjection(projection)
|
out_dataset.SetProjection(projection)
|
||||||
@ -384,11 +403,8 @@ class Kutser:
|
|||||||
if self.output_path is None:
|
if self.output_path is None:
|
||||||
raise ValueError("output_path 必须提供")
|
raise ValueError("output_path 必须提供")
|
||||||
|
|
||||||
water_mask = self._load_water_mask()
|
# 仅 run_fast 需要整幅掩膜;大影像请用默认 run()
|
||||||
if water_mask is not None:
|
mask_bool = self._legacy_full_mask_bool()
|
||||||
mask_bool = water_mask.astype(bool)
|
|
||||||
else:
|
|
||||||
mask_bool = None
|
|
||||||
|
|
||||||
chunk_mb = self.width * self.height * self.n_bands * 2 / (1024 * 1024)
|
chunk_mb = self.width * self.height * self.n_bands * 2 / (1024 * 1024)
|
||||||
print(f"[Kutser-Fast] 数据立方体 {chunk_mb:.0f}MB, "
|
print(f"[Kutser-Fast] 数据立方体 {chunk_mb:.0f}MB, "
|
||||||
@ -446,9 +462,8 @@ class Kutser:
|
|||||||
base_path, ext = os.path.splitext(self.output_path)
|
base_path, ext = os.path.splitext(self.output_path)
|
||||||
bsq_path = base_path + '.bsq' if ext.lower() != '.bsq' else self.output_path
|
bsq_path = base_path + '.bsq' if ext.lower() != '.bsq' else self.output_path
|
||||||
|
|
||||||
driver = gdal.GetDriverByName('ENVI')
|
out_dataset = create_output_dataset(
|
||||||
out_dataset = driver.Create(bsq_path, self.width, self.height,
|
bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32)
|
||||||
self.n_bands, gdal.GDT_Float32)
|
|
||||||
out_dataset.SetGeoTransform(self.dataset.GetGeoTransform())
|
out_dataset.SetGeoTransform(self.dataset.GetGeoTransform())
|
||||||
out_dataset.SetProjection(self.dataset.GetProjection())
|
out_dataset.SetProjection(self.dataset.GetProjection())
|
||||||
|
|
||||||
|
|||||||
@ -4,6 +4,8 @@ import os
|
|||||||
from scipy import ndimage
|
from scipy import ndimage
|
||||||
from scipy.optimize import minimize_scalar
|
from scipy.optimize import minimize_scalar
|
||||||
|
|
||||||
|
from src.utils.util import create_output_dataset
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from osgeo import gdal
|
from osgeo import gdal
|
||||||
GDAL_AVAILABLE = True
|
GDAL_AVAILABLE = True
|
||||||
@ -98,6 +100,10 @@ class SUGAR:
|
|||||||
self.glint_mask_method = glint_mask_method
|
self.glint_mask_method = glint_mask_method
|
||||||
self.water_mask = None
|
self.water_mask = None
|
||||||
self.water_mask_path = water_mask
|
self.water_mask_path = water_mask
|
||||||
|
# ★ v3:掩膜不再整幅驻留——栅格路径保留 GDAL band 句柄,按需分块读
|
||||||
|
self._mask_ds = None
|
||||||
|
self._wm_band = None
|
||||||
|
self._wm_np = None
|
||||||
self.output_path = output_path
|
self.output_path = output_path
|
||||||
self.block_size = block_size
|
self.block_size = block_size
|
||||||
|
|
||||||
@ -117,32 +123,46 @@ class SUGAR:
|
|||||||
self.glint_pixel_indices = [] # list of (block_idx, row, col) 索引
|
self.glint_pixel_indices = [] # list of (block_idx, row, col) 索引
|
||||||
self.thresholds = [] # 每波段的全局阈值
|
self.thresholds = [] # 每波段的全局阈值
|
||||||
|
|
||||||
def _load_water_mask(self):
|
def _open_mask(self):
|
||||||
"""延迟加载水域掩膜"""
|
"""打开水域掩膜,但绝不整幅 ReadAsArray 驻留。
|
||||||
|
|
||||||
|
栅格路径 → 保留 GDAL dataset/band 句柄(_wm_band),按需分块读取;
|
||||||
|
np.ndarray(外部已驻留)→ 转 bool 后保留 _wm_np,按窗口切片。
|
||||||
|
"""
|
||||||
|
if self._wm_band is not None or self._wm_np is not None:
|
||||||
|
return
|
||||||
if self.water_mask_path is None:
|
if self.water_mask_path is None:
|
||||||
return None
|
return
|
||||||
|
|
||||||
if isinstance(self.water_mask_path, np.ndarray):
|
if isinstance(self.water_mask_path, np.ndarray):
|
||||||
if self.water_mask_path.shape[:2] != (self.height, self.width):
|
if self.water_mask_path.shape[:2] != (self.height, self.width):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"掩膜尺寸 {self.water_mask_path.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
f"掩膜尺寸 {self.water_mask_path.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
||||||
)
|
)
|
||||||
return (self.water_mask_path > 0).astype(np.uint8)
|
self._wm_np = self.water_mask_path > 0
|
||||||
|
return
|
||||||
|
|
||||||
if isinstance(self.water_mask_path, str):
|
if isinstance(self.water_mask_path, str):
|
||||||
if self.water_mask_path.lower().endswith('.shp'):
|
if self.water_mask_path.lower().endswith('.shp'):
|
||||||
raise ValueError("请先栅格化shapefile为栅格掩膜文件")
|
raise ValueError("请先栅格化shapefile为栅格掩膜文件")
|
||||||
mask_dataset = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly)
|
self._mask_ds = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly)
|
||||||
if mask_dataset is None:
|
if self._mask_ds is None:
|
||||||
raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}")
|
raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}")
|
||||||
mask_array = mask_dataset.GetRasterBand(1).ReadAsArray()
|
band = self._mask_ds.GetRasterBand(1)
|
||||||
mask_dataset = None
|
if band.XSize != self.width or band.YSize != self.height:
|
||||||
if mask_array.shape != (self.height, self.width):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"掩膜尺寸 {mask_array.shape} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
f"掩膜尺寸 {(band.YSize, band.XSize)} 与图像尺寸 {(self.height, self.width)} 不匹配"
|
||||||
)
|
)
|
||||||
return (mask_array > 0).astype(np.uint8)
|
self._wm_band = band
|
||||||
|
|
||||||
|
def _read_mask_block(self, x_off, y_off, x_size, y_size):
|
||||||
|
"""分块读取掩膜布尔数组(True=水体);无掩膜时返回 None。"""
|
||||||
|
self._open_mask()
|
||||||
|
if self._wm_band is not None:
|
||||||
|
arr = self._wm_band.ReadAsArray(x_off, y_off, x_size, y_size)
|
||||||
|
return (arr > 0)
|
||||||
|
if self._wm_np is not None:
|
||||||
|
return self._wm_np[y_off:y_off + y_size, x_off:x_off + x_size]
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _compute_threshold(self, im):
|
def _compute_threshold(self, im):
|
||||||
@ -163,14 +183,10 @@ class SUGAR:
|
|||||||
thresh = self.thresholds[self._current_band]
|
thresh = self.thresholds[self._current_band]
|
||||||
glint_mask = (log_im < thresh).astype(np.uint8)
|
glint_mask = (log_im < thresh).astype(np.uint8)
|
||||||
|
|
||||||
# 应用水域掩膜
|
# 应用水域掩膜(同步分块读取,不再整幅驻留)
|
||||||
water_mask = self._load_water_mask()
|
mask_block = self._read_mask_block(
|
||||||
if water_mask is not None:
|
self._current_x, self._current_y, band_data.shape[1], band_data.shape[0])
|
||||||
y_off = self._current_y
|
if mask_block is not None:
|
||||||
y_end = y_off + band_data.shape[0]
|
|
||||||
x_off = self._current_x
|
|
||||||
x_end = x_off + band_data.shape[1]
|
|
||||||
mask_block = water_mask[y_off:y_end, x_off:x_end]
|
|
||||||
glint_mask = glint_mask * mask_block
|
glint_mask = glint_mask * mask_block
|
||||||
|
|
||||||
return log_im, glint_mask
|
return log_im, glint_mask
|
||||||
@ -214,7 +230,7 @@ class SUGAR:
|
|||||||
内存:仅存储每波段的阈值(float)和 glint 像素位置索引
|
内存:仅存储每波段的阈值(float)和 glint 像素位置索引
|
||||||
"""
|
"""
|
||||||
print(f"[SUGAR] 步骤1: 扫描全图收集glint像素...")
|
print(f"[SUGAR] 步骤1: 扫描全图收集glint像素...")
|
||||||
water_mask = self._load_water_mask()
|
self._open_mask()
|
||||||
|
|
||||||
# 初始化阈值列表
|
# 初始化阈值列表
|
||||||
self.thresholds = [None] * self.n_bands
|
self.thresholds = [None] * self.n_bands
|
||||||
@ -231,6 +247,9 @@ class SUGAR:
|
|||||||
x_size = x_end - x_off
|
x_size = x_end - x_off
|
||||||
n_blocks += 1
|
n_blocks += 1
|
||||||
|
|
||||||
|
# 掩膜块每块只读一次(同步分块,不再整幅驻留)
|
||||||
|
_mask_block = self._read_mask_block(x_off, y_off, x_size, y_size)
|
||||||
|
|
||||||
for b in range(self.n_bands):
|
for b in range(self.n_bands):
|
||||||
band = self.dataset.GetRasterBand(b + 1)
|
band = self.dataset.GetRasterBand(b + 1)
|
||||||
block = band.ReadAsArray(x_off, y_off, x_size, y_size).astype(np.float32)
|
block = band.ReadAsArray(x_off, y_off, x_size, y_size).astype(np.float32)
|
||||||
@ -238,11 +257,7 @@ class SUGAR:
|
|||||||
|
|
||||||
log_im = ndimage.gaussian_laplace(block, sigma=self.sigma)
|
log_im = ndimage.gaussian_laplace(block, sigma=self.sigma)
|
||||||
|
|
||||||
# mask_block 在波段循环外初始化,每块只计算一次
|
if _mask_block is not None:
|
||||||
if b == 0 and water_mask is not None:
|
|
||||||
_mask_block = water_mask[y_off:y_end, x_off:x_end].astype(bool)
|
|
||||||
|
|
||||||
if water_mask is not None:
|
|
||||||
if _mask_block.any():
|
if _mask_block.any():
|
||||||
log_collections[b].append(log_im[_mask_block])
|
log_collections[b].append(log_im[_mask_block])
|
||||||
else:
|
else:
|
||||||
@ -250,9 +265,6 @@ class SUGAR:
|
|||||||
|
|
||||||
del block, log_im
|
del block, log_im
|
||||||
|
|
||||||
if water_mask is not None:
|
|
||||||
del _mask_block
|
|
||||||
|
|
||||||
# 计算每波段的全局阈值(需要所有LoG值)
|
# 计算每波段的全局阈值(需要所有LoG值)
|
||||||
print(f"[SUGAR] 计算 {self.n_bands} 个波段的全局阈值...")
|
print(f"[SUGAR] 计算 {self.n_bands} 个波段的全局阈值...")
|
||||||
for b in range(self.n_bands):
|
for b in range(self.n_bands):
|
||||||
@ -275,7 +287,7 @@ class SUGAR:
|
|||||||
内存:只存储 1D 数组(所有 glint 像素值)
|
内存:只存储 1D 数组(所有 glint 像素值)
|
||||||
"""
|
"""
|
||||||
print(f"[SUGAR] 步骤2: 收集glint像素值用于全局优化...")
|
print(f"[SUGAR] 步骤2: 收集glint像素值用于全局优化...")
|
||||||
water_mask = self._load_water_mask()
|
self._open_mask()
|
||||||
|
|
||||||
R_glint_list = [[] for _ in range(self.n_bands)]
|
R_glint_list = [[] for _ in range(self.n_bands)]
|
||||||
R_bg_glint_list = [[] for _ in range(self.n_bands)]
|
R_bg_glint_list = [[] for _ in range(self.n_bands)]
|
||||||
@ -288,6 +300,9 @@ class SUGAR:
|
|||||||
x_end = min(x_off + self.block_size, self.width)
|
x_end = min(x_off + self.block_size, self.width)
|
||||||
x_size = x_end - x_off
|
x_size = x_end - x_off
|
||||||
|
|
||||||
|
# 掩膜块每块只读一次(同步分块,不再整幅驻留)
|
||||||
|
_mask_block = self._read_mask_block(x_off, y_off, x_size, y_size)
|
||||||
|
|
||||||
for b in range(self.n_bands):
|
for b in range(self.n_bands):
|
||||||
band = self.dataset.GetRasterBand(b + 1)
|
band = self.dataset.GetRasterBand(b + 1)
|
||||||
R_block = band.ReadAsArray(x_off, y_off, x_size, y_size).astype(np.float32)
|
R_block = band.ReadAsArray(x_off, y_off, x_size, y_size).astype(np.float32)
|
||||||
@ -298,9 +313,8 @@ class SUGAR:
|
|||||||
thresh = self.thresholds[b]
|
thresh = self.thresholds[b]
|
||||||
glint_mask = (log_im < thresh).astype(np.uint8)
|
glint_mask = (log_im < thresh).astype(np.uint8)
|
||||||
|
|
||||||
if water_mask is not None:
|
if _mask_block is not None:
|
||||||
mask_block = water_mask[y_off:y_end, x_off:x_end]
|
glint_mask = glint_mask * _mask_block
|
||||||
glint_mask = glint_mask * mask_block
|
|
||||||
|
|
||||||
# 背景
|
# 背景
|
||||||
if self.estimate_background:
|
if self.estimate_background:
|
||||||
@ -353,7 +367,8 @@ class SUGAR:
|
|||||||
"""
|
"""
|
||||||
Step 4: 分块处理并写入输出文件
|
Step 4: 分块处理并写入输出文件
|
||||||
"""
|
"""
|
||||||
water_mask = self._load_water_mask()
|
# 掩膜块每块只读一次(同步分块,不再整幅驻留)
|
||||||
|
mask_block = self._read_mask_block(x_off, y_off, x_size, y_size)
|
||||||
|
|
||||||
for b in range(self.n_bands):
|
for b in range(self.n_bands):
|
||||||
band = self.dataset.GetRasterBand(b + 1)
|
band = self.dataset.GetRasterBand(b + 1)
|
||||||
@ -365,8 +380,7 @@ class SUGAR:
|
|||||||
thresh = self.thresholds[b]
|
thresh = self.thresholds[b]
|
||||||
glint_mask = (log_im < thresh).astype(np.uint8)
|
glint_mask = (log_im < thresh).astype(np.uint8)
|
||||||
|
|
||||||
if water_mask is not None:
|
if mask_block is not None:
|
||||||
mask_block = water_mask[y_off:y_off + y_size, x_off:x_off + x_size]
|
|
||||||
glint_mask = glint_mask * mask_block
|
glint_mask = glint_mask * mask_block
|
||||||
|
|
||||||
glint_bool = glint_mask.astype(bool)
|
glint_bool = glint_mask.astype(bool)
|
||||||
@ -411,11 +425,8 @@ class SUGAR:
|
|||||||
geotransform = self.dataset.GetGeoTransform()
|
geotransform = self.dataset.GetGeoTransform()
|
||||||
projection = self.dataset.GetProjection()
|
projection = self.dataset.GetProjection()
|
||||||
|
|
||||||
driver = gdal.GetDriverByName('ENVI')
|
out_dataset = create_output_dataset(
|
||||||
out_dataset = driver.Create(bsq_path, self.width, self.height,
|
bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32)
|
||||||
self.n_bands, gdal.GDT_Float32)
|
|
||||||
if out_dataset is None:
|
|
||||||
raise ValueError(f"无法创建输出文件: {bsq_path}")
|
|
||||||
out_dataset.SetGeoTransform(geotransform)
|
out_dataset.SetGeoTransform(geotransform)
|
||||||
out_dataset.SetProjection(projection)
|
out_dataset.SetProjection(projection)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user