diff --git a/src/core/glint_removal/Goodman.py b/src/core/glint_removal/Goodman.py index cc154f2..dbb79b3 100644 --- a/src/core/glint_removal/Goodman.py +++ b/src/core/glint_removal/Goodman.py @@ -9,7 +9,7 @@ except ImportError: GDAL_AVAILABLE = False print("警告: GDAL未安装,将使用numpy处理模式") -from src.utils.util import find_band_number +from src.utils.util import find_band_number, create_output_dataset try: from tqdm import tqdm @@ -49,7 +49,8 @@ class Goodman: def __init__(self, im_aligned, nir_lower_wavelength=641.93, nir_upper_wavelength=751.49, 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 耀斑去除算法 — 波长驱动版。 @@ -90,6 +91,13 @@ class Goodman: self.chunk_size = chunk_size self.is_file_path = isinstance(im_aligned, str) 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 元数据动态解析波段索引 ── if self.is_file_path: @@ -123,90 +131,70 @@ class Goodman: self.width = im_aligned.shape[1] self.n_bands = im_aligned.shape[-1] - # 加载水域掩膜(在获取图像尺寸之后) - self.water_mask = self._load_water_mask(water_mask) + # 加载水域掩膜(★ v3:仅打开句柄/保留引用,绝不整幅 ReadAsArray 驻留) + 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。 """ - 加载水域掩膜 - - :param water_mask: 可以是None、numpy数组、文件路径(.dat/.tif)或shapefile路径(.shp) - :return: numpy数组或None,1表示水域,0表示非水域 - """ - if water_mask is None: - return None - - # 如果已经是numpy数组 - 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 self._wm_band is not None or self._wm_np is not None: + return + wm = self.water_mask_path + if wm is None: + return + + if isinstance(wm, np.ndarray): + if wm.shape[:2] != (self.height, self.width): + raise ValueError( + f"掩膜尺寸 {wm.shape[:2]} 与图像尺寸 {(self.height, self.width)} 不匹配") + self._wm_np = wm > 0 + return + + if isinstance(wm, str): if not GDAL_AVAILABLE: raise ValueError("使用文件路径作为掩膜时,必须安装GDAL") - - # 检查是否为shapefile - if water_mask.lower().endswith('.shp'): - # 从shp文件创建掩膜 - if self.is_file_path: - ref_path = self.im_aligned - else: + src = wm + if wm.lower().endswith('.shp'): + if not self.is_file_path: raise ValueError("输入为numpy数组时,无法从shp文件创建掩膜(需要参考栅格)") - - try: - 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() - projection = ref_dataset.GetProjection() - width = ref_dataset.RasterXSize - height = ref_dataset.RasterYSize - - # 创建内存中的栅格数据集 - mem_driver = gdal.GetDriverByName('MEM') - mask_dataset = mem_driver.Create('', width, height, 1, gdal.GDT_Byte) - mask_dataset.SetGeoTransform(geotransform) - mask_dataset.SetProjection(projection) - - 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)}") + 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 + + raise ValueError(f"不支持的掩膜类型: {type(wm)}") + + 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): + """仅供小图/数组模式的旧全幅路径使用;大影像分块模式不会调用。""" + self._open_mask() + if self._wm_band is not None: + return self._wm_band.ReadAsArray() > 0 + return self._wm_np def _get_corrected_bands_numpy(self): """ @@ -224,8 +212,8 @@ class Goodman: diff_640_750 = R_640 - R_750 corrected_bands = [] - # 获取水域掩膜(如果存在) - water_mask_bool = self.water_mask.astype(bool) if self.water_mask is not None else None + # 获取水域掩膜(numpy 数组模式直接引用 _wm_np) + 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): @@ -267,8 +255,8 @@ class Goodman: diff_640_750 = R_640 - R_750 del R_640 # 释放不再需要的 R_640 - # 水域掩膜 + 水体占比检测(控制行段跳跃模式) - water = self.water_mask.astype(bool) if self.water_mask is not None else None + # 水域掩膜 + 水体占比检测(★ 旧全幅路径:仅无 output_path 的传统返回列表模式) + water = self._legacy_full_mask_bool() water_pct = (100.0 * np.count_nonzero(water) / water.size) if water is not None else 100.0 # 行段跳跃: 水域占比 < 50% 时启用 @@ -400,15 +388,8 @@ class Goodman: else: bsq_path = self.output_path - # 使用ENVI驱动(默认就是BSQ格式) - driver = gdal.GetDriverByName('ENVI') - 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}") + # 创建输出数据集(大图自动用 GTiff+BIGTIFF,规避 ENVI int32 偏移溢出) + dataset = create_output_dataset(bsq_path, width, height, n_bands, gdal.GDT_Float32) try: # 设置地理变换和投影 @@ -461,28 +442,23 @@ class Goodman: return self._get_corrected_bands_numpy() def _get_corrected_bands_streaming(self): - """流式处理:逐波段校正并直接写入输出文件,不累积内存 + """★ v3 分块流式:块外循环 × 波段内循环 × 掩膜分块读 × 逐块写盘。 - 适用于大尺度影像(如 6522×13215×150)。 - 内存峰值 ≈ 3 个全波段数组(NIR×2 + 当前波段),而非全部 150 个波段。 + 不再整幅驻留 R640/R750/diff 或整幅掩膜;每块只驻留 block_size² 数据, + 与 Hedley/Kutser/SUGAR 的分块范式完全对齐,内存恒定。 - :return: None(波段已在输出文件中) + :return: None(波段已写入输出文件) """ import os - # ── 创建输出文件 ── base_path, ext = os.path.splitext(self.output_path) bsq_path = base_path + '.bsq' if ext.lower() != '.bsq' else self.output_path output_dir = os.path.dirname(bsq_path) if output_dir and not os.path.exists(output_dir): os.makedirs(output_dir, exist_ok=True) - driver = gdal.GetDriverByName('ENVI') - out_ds = driver.Create( - bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32 - ) - if out_ds is None: - raise ValueError(f"无法创建输出文件: {bsq_path}") + out_ds = create_output_dataset( + bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32) # ── 设置地理参考 ── if self.is_file_path and self.dataset is not None: @@ -490,22 +466,85 @@ class Goodman: out_ds.SetProjection(self.dataset.GetProjection()) try: - # 逐波段处理 + 立即写入(波段在 _get_corrected_bands_gdal 的循环中 - # 由 WriteArray 写入 → FlushCache → del,不会累积) - self._get_corrected_bands_gdal(out_dataset=out_ds) + self._run_block_correction(out_ds) finally: 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' if os.path.exists(hdr_path): - print(f"校正后的图像已保存至: {bsq_path} (BSQ格式, 流式写入)") + print(f"校正后的图像已保存至: {bsq_path} (BSQ格式, 分块流式写入)") else: print(f"校正后的图像已保存至: {bsq_path} (BSQ格式)") print("警告: 未检测到.hdr文件,但GDAL应该已自动创建") 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): """清理资源""" if self.dataset is not None and self.is_file_path: diff --git a/src/core/glint_removal/Hedley.py b/src/core/glint_removal/Hedley.py index 94db60b..8889fcf 100644 --- a/src/core/glint_removal/Hedley.py +++ b/src/core/glint_removal/Hedley.py @@ -7,7 +7,7 @@ try: except ImportError: GDAL_AVAILABLE = False -from src.utils.util import find_band_number +from src.utils.util import find_band_number, create_output_dataset class Hedley: @@ -36,6 +36,10 @@ class Hedley: print(f"[Hedley] 波段解析: NIR={float(nir_wavelength):.1f}nm→band{self.NIR_band}") self.water_mask = None 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.block_size = block_size self.R_min = None @@ -49,32 +53,46 @@ class Hedley: self.height = self.dataset.RasterYSize 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: - return None + return if isinstance(self.water_mask_path, np.ndarray): if self.water_mask_path.shape[:2] != (self.height, self.width): raise ValueError( 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 self.water_mask_path.lower().endswith('.shp'): raise ValueError("请先栅格化shapefile为栅格掩膜文件") - mask_dataset = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly) - if mask_dataset is None: + self._mask_ds = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly) + if self._mask_ds is None: raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}") - mask_array = mask_dataset.GetRasterBand(1).ReadAsArray() - mask_dataset = None - if mask_array.shape != (self.height, self.width): + band = self._mask_ds.GetRasterBand(1) + if band.XSize != self.width or band.YSize != self.height: 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 covariance_NIR(self, NIR, b): @@ -93,7 +111,7 @@ class Hedley: 使用重采样方式扫描,大幅降低内存占用。 """ print(f"[Hedley] 扫描全局统计量(采样步长={sample_step})...") - water_mask = self._load_water_mask() + self._open_mask() nir_samples = [] sample_count = 0 @@ -106,10 +124,9 @@ class Hedley: nir_block = nir_band.ReadAsArray(0, y_off, self.width, block_height) nir_band = None - if water_mask is not None: - mask_block = water_mask[y_off:y_end, :] - mask_bool = mask_block.astype(bool) - else: + # 掩膜同步分块读取(不再整幅驻留) + 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) if mask_bool.any(): @@ -117,7 +134,7 @@ class Hedley: nir_samples.append(nir_sampled) sample_count += nir_sampled.size - del nir_block, mask_block + del nir_block if sample_count == 0: self.R_min = 0.0 @@ -136,7 +153,7 @@ class Hedley: 由于需要相关性计算,需要足够多的样本,取sample_step=5 """ print(f"[Hedley] 计算全局协方差系数列表(采样步长={sample_step})...") - water_mask = self._load_water_mask() + self._open_mask() # 预收集NIR和每个波段的样本数据 nir_samples = [] @@ -151,11 +168,9 @@ class Hedley: nir_block = nir_band.ReadAsArray(0, y_off, self.width, block_height).astype(np.float32) nir_band = None - # 取 NIR 样本(每块只取一次,放在波段循环外) - if water_mask is not None: - mask_block = water_mask[y_off:y_end, :] - mask_bool = mask_block.astype(bool) - else: + # 取 NIR 样本(每块只取一次,放在波段循环外;掩膜同步分块读) + 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) if mask_bool.any(): @@ -208,14 +223,8 @@ class Hedley: # 预计算 NIR - R_min NIR_diff = NIR - self.R_min - # 获取掩膜 - water_mask = self._load_water_mask() - 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 + # 获取掩膜(同步分块读取,不再整幅驻留) + mask_block = self._read_mask_block(x_off, y_off, x_size, y_size) # 逐波段处理 corrected_bands = [] @@ -262,11 +271,8 @@ class Hedley: geotransform = self.dataset.GetGeoTransform() projection = self.dataset.GetProjection() - driver = gdal.GetDriverByName('ENVI') - out_dataset = driver.Create(bsq_path, self.width, self.height, - self.n_bands, gdal.GDT_Float32) - if out_dataset is None: - raise ValueError(f"无法创建输出文件: {bsq_path}") + out_dataset = create_output_dataset( + bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32) out_dataset.SetGeoTransform(geotransform) out_dataset.SetProjection(projection) diff --git a/src/core/glint_removal/Kutser.py b/src/core/glint_removal/Kutser.py index e27c7db..3fc19e2 100644 --- a/src/core/glint_removal/Kutser.py +++ b/src/core/glint_removal/Kutser.py @@ -7,7 +7,7 @@ try: except ImportError: GDAL_AVAILABLE = False -from src.utils.util import find_band_number +from src.utils.util import find_band_number, create_output_dataset class Kutser: @@ -49,6 +49,10 @@ class Kutser: f"NIR={float(nir_wavelength):.1f}nm→band{self.NIR_band}") self.water_mask = None # 延迟加载,在处理前初始化 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.block_size = block_size self.R_min = None # 全局R_min(来自重采样扫描) @@ -80,32 +84,59 @@ class Kutser: except Exception: 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: - return None + return if isinstance(self.water_mask_path, np.ndarray): if self.water_mask_path.shape[:2] != (self.height, self.width): raise ValueError( 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 self.water_mask_path.lower().endswith('.shp'): raise ValueError("请先栅格化shapefile为栅格掩膜文件") - mask_dataset = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly) - if mask_dataset is None: + self._mask_ds = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly) + if self._mask_ds is None: raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}") - mask_array = mask_dataset.GetRasterBand(1).ReadAsArray() - mask_dataset = None - if mask_array.shape != (self.height, self.width): + band = self._mask_ds.GetRasterBand(1) + if band.XSize != self.width or band.YSize != self.height: 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 def _scan_global_stats(self, sample_step=20): @@ -116,7 +147,7 @@ class Kutser: 内存峰值 ≈ 单波段块大小 + 几个掩膜数组 ≈ block_size² × 4~8MB """ print(f"[Kutser] 扫描全局统计量(采样步长={sample_step})...") - water_mask = self._load_water_mask() + self._open_mask() # 预分配采样数组(NIR波段和D值) nir_samples = [] @@ -148,14 +179,10 @@ class Kutser: # 计算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) - # 获取掩膜(整块) - if water_mask is not None: - mask_block = water_mask[y_off:y_end, :] - else: - mask_block = np.ones((block_height, self.width), dtype=np.uint8) - - # 对掩膜区域进行采样 - mask_bool = mask_block.astype(bool) + # 获取掩膜(与波段同步分块读取,不再整幅驻留) + 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) if mask_bool.any(): # 按步长采样 @@ -166,7 +193,7 @@ class Kutser: 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: @@ -189,7 +216,7 @@ class Kutser: 使用全分辨率扫描,但逐波段读取,每波段内存 ≈ block_size² """ print(f"[Kutser] 计算全局G值列表(n_bands={self.n_bands})...") - water_mask = self._load_water_mask() + self._open_mask() # 初始化G_max和G_min为极值 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) 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): band = self.dataset.GetRasterBand(b + 1) block = band.ReadAsArray(0, y_off, self.width, block_height).astype(np.float32) band = None - if water_mask is not None: - mask_block = water_mask[y_off:y_end, :] - mask_bool = mask_block.astype(bool) - if mask_bool.any(): - 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()) + if mask_bool.any(): + 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()) del block @@ -254,14 +280,8 @@ class Kutser: # 释放临时块 del lower_block, upper_block, oxy_block, D - # 获取当前块的水域掩膜 - water_mask = self._load_water_mask() - 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 + # 获取当前块的水域掩膜(同步分块读取,不再整幅驻留) + mask_block = self._read_mask_block(x_off, y_off, x_size, y_size) # 逐波段处理 corrected_bands = [] @@ -296,8 +316,10 @@ class Kutser: if self.output_path is None: raise ValueError("output_path 必须提供,分块处理需要直接写入文件") - # 全波段顺序模式:每个波段完整读取一次,顺序IO - return self.run_fast() + # ★ v3:默认走分块 run()(掩膜同步分块读取,内存恒定)。 + # 旧的 run_fast()(整幅参考波段 + 整幅掩膜)仅保留给明确需要 + # 顺序 IO 加速的小图,不再作为默认入口,避免大影像内存峰值。 + return self.run() def run(self): """原版分块模式(低内存兼容)。默认使用 run_fast() 代替。""" @@ -322,11 +344,8 @@ class Kutser: geotransform = self.dataset.GetGeoTransform() projection = self.dataset.GetProjection() - driver = gdal.GetDriverByName('ENVI') - out_dataset = driver.Create(bsq_path, self.width, self.height, - self.n_bands, gdal.GDT_Float32) - if out_dataset is None: - raise ValueError(f"无法创建输出文件: {bsq_path}") + out_dataset = create_output_dataset( + bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32) out_dataset.SetGeoTransform(geotransform) out_dataset.SetProjection(projection) @@ -384,11 +403,8 @@ class Kutser: if self.output_path is None: raise ValueError("output_path 必须提供") - water_mask = self._load_water_mask() - if water_mask is not None: - mask_bool = water_mask.astype(bool) - else: - mask_bool = None + # 仅 run_fast 需要整幅掩膜;大影像请用默认 run() + mask_bool = self._legacy_full_mask_bool() chunk_mb = self.width * self.height * self.n_bands * 2 / (1024 * 1024) print(f"[Kutser-Fast] 数据立方体 {chunk_mb:.0f}MB, " @@ -446,9 +462,8 @@ class Kutser: base_path, ext = os.path.splitext(self.output_path) bsq_path = base_path + '.bsq' if ext.lower() != '.bsq' else self.output_path - driver = gdal.GetDriverByName('ENVI') - out_dataset = driver.Create(bsq_path, self.width, self.height, - self.n_bands, gdal.GDT_Float32) + out_dataset = create_output_dataset( + bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32) out_dataset.SetGeoTransform(self.dataset.GetGeoTransform()) out_dataset.SetProjection(self.dataset.GetProjection()) diff --git a/src/core/glint_removal/SUGAR.py b/src/core/glint_removal/SUGAR.py index 014308b..49cb11d 100644 --- a/src/core/glint_removal/SUGAR.py +++ b/src/core/glint_removal/SUGAR.py @@ -4,6 +4,8 @@ import os from scipy import ndimage from scipy.optimize import minimize_scalar +from src.utils.util import create_output_dataset + try: from osgeo import gdal GDAL_AVAILABLE = True @@ -98,6 +100,10 @@ class SUGAR: self.glint_mask_method = glint_mask_method self.water_mask = None 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.block_size = block_size @@ -117,32 +123,46 @@ class SUGAR: self.glint_pixel_indices = [] # list of (block_idx, row, col) 索引 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: - return None + return if isinstance(self.water_mask_path, np.ndarray): if self.water_mask_path.shape[:2] != (self.height, self.width): raise ValueError( 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 self.water_mask_path.lower().endswith('.shp'): raise ValueError("请先栅格化shapefile为栅格掩膜文件") - mask_dataset = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly) - if mask_dataset is None: + self._mask_ds = gdal.Open(self.water_mask_path, gdal.GA_ReadOnly) + if self._mask_ds is None: raise ValueError(f"无法打开掩膜文件: {self.water_mask_path}") - mask_array = mask_dataset.GetRasterBand(1).ReadAsArray() - mask_dataset = None - if mask_array.shape != (self.height, self.width): + band = self._mask_ds.GetRasterBand(1) + if band.XSize != self.width or band.YSize != self.height: 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 _compute_threshold(self, im): @@ -163,14 +183,10 @@ class SUGAR: thresh = self.thresholds[self._current_band] glint_mask = (log_im < thresh).astype(np.uint8) - # 应用水域掩膜 - water_mask = self._load_water_mask() - if water_mask is not None: - y_off = self._current_y - 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] + # 应用水域掩膜(同步分块读取,不再整幅驻留) + mask_block = self._read_mask_block( + self._current_x, self._current_y, band_data.shape[1], band_data.shape[0]) + if mask_block is not None: glint_mask = glint_mask * mask_block return log_im, glint_mask @@ -214,7 +230,7 @@ class SUGAR: 内存:仅存储每波段的阈值(float)和 glint 像素位置索引 """ print(f"[SUGAR] 步骤1: 扫描全图收集glint像素...") - water_mask = self._load_water_mask() + self._open_mask() # 初始化阈值列表 self.thresholds = [None] * self.n_bands @@ -231,6 +247,9 @@ class SUGAR: x_size = x_end - x_off n_blocks += 1 + # 掩膜块每块只读一次(同步分块,不再整幅驻留) + _mask_block = self._read_mask_block(x_off, y_off, x_size, y_size) + for b in range(self.n_bands): band = self.dataset.GetRasterBand(b + 1) 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) - # mask_block 在波段循环外初始化,每块只计算一次 - 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 is not None: if _mask_block.any(): log_collections[b].append(log_im[_mask_block]) else: @@ -250,9 +265,6 @@ class SUGAR: del block, log_im - if water_mask is not None: - del _mask_block - # 计算每波段的全局阈值(需要所有LoG值) print(f"[SUGAR] 计算 {self.n_bands} 个波段的全局阈值...") for b in range(self.n_bands): @@ -275,7 +287,7 @@ class SUGAR: 内存:只存储 1D 数组(所有 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_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_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): band = self.dataset.GetRasterBand(b + 1) 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] glint_mask = (log_im < thresh).astype(np.uint8) - if water_mask is not None: - mask_block = water_mask[y_off:y_end, x_off:x_end] - glint_mask = glint_mask * mask_block + if _mask_block is not None: + glint_mask = glint_mask * _mask_block # 背景 if self.estimate_background: @@ -353,7 +367,8 @@ class SUGAR: """ 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): band = self.dataset.GetRasterBand(b + 1) @@ -365,8 +380,7 @@ class SUGAR: thresh = self.thresholds[b] glint_mask = (log_im < thresh).astype(np.uint8) - if water_mask is not None: - mask_block = water_mask[y_off:y_off + y_size, x_off:x_off + x_size] + if mask_block is not None: glint_mask = glint_mask * mask_block glint_bool = glint_mask.astype(bool) @@ -411,11 +425,8 @@ class SUGAR: geotransform = self.dataset.GetGeoTransform() projection = self.dataset.GetProjection() - driver = gdal.GetDriverByName('ENVI') - out_dataset = driver.Create(bsq_path, self.width, self.height, - self.n_bands, gdal.GDT_Float32) - if out_dataset is None: - raise ValueError(f"无法创建输出文件: {bsq_path}") + out_dataset = create_output_dataset( + bsq_path, self.width, self.height, self.n_bands, gdal.GDT_Float32) out_dataset.SetGeoTransform(geotransform) out_dataset.SetProjection(projection)