diff --git a/src/utils/find_severe_glint_area.py b/src/utils/find_severe_glint_area.py index 0c28381..867d2b8 100644 --- a/src/utils/find_severe_glint_area.py +++ b/src/utils/find_severe_glint_area.py @@ -559,8 +559,234 @@ def create_water_mask_from_shp(shp_file, reference_raster): raise +# ============================================================ +# ★ v3 两趟式(Two-Pass Chunking)耀斑检测 +# +# 旧实现把"单波段影像 + 水体掩膜"整幅读入内存后做全局统计与逐像素检测, +# 在 34914x19177(6.7 亿像素)影像上内存峰值可达 5~10 GB。 +# 这里改为: +# 第一趟:按 block 抽样收集有效样本 → 求全局标量 +# (zscore: mean/std;percentile: 分位;iqr: Q1/Q3; +# otsu: 2%/98% 拉伸分位 + 256-bin 直方图上的 Otsu 阈值) +# 第二趟:按 block 读取 → 逐块应用阈值 → WriteArray 写 uint8 掩膜。 +# 全程每块只驻留 block_size² 数据,内存恒定。 +# +# 限制说明(不影响默认 GUI 主路径): +# - 'adaptive' 是整幅邻域算子(局部窗口百分位),仍需整幅数组; +# - max_area(连通域面积过滤)/ buffer_size(岸线缓冲)是全局结构后处理, +# 同样需整幅数组。 +# 启用这些选项时仍走旧全幅路径,大影像上会告警。 +# ============================================================ + +def _otsu_threshold_from_hist(hist): + """由 256-bin 直方图求 Otsu 阈值(0..255),复刻原 otsu() 的类间方差最大准则。""" + total = float(hist.sum()) + if total <= 0: + return 0.0 + vals = np.arange(256, dtype=np.float64) + sum_all = float((hist * vals).sum()) + w_b = 0.0 + sum_b = 0.0 + max_between = -1.0 + thr = 0.0 + for t in range(256): + w_b += float(hist[t]) + if w_b == 0: + continue + w_f = total - w_b + if w_f == 0: + break + sum_b += t * float(hist[t]) + m_b = sum_b / w_b + m_f = (sum_all - sum_b) / w_f + between = w_b * w_f * (m_b - m_f) ** 2 + if between > max_between: + max_between = between + thr = float(t) + return thr + + +def _remove_raster_output(path): + """移除已有栅格输出(兼容 GTiff / ENVI .dat+.hdr)。""" + base = os.path.splitext(path)[0] + for c in dict.fromkeys([path, path + '.hdr', base + '.hdr', base + '.dat']): + if c and os.path.exists(c): + os.remove(c) + + +def _detect_glint_two_pass(img_path, mask_arg, out_path, bands_1based, weights, + method, z_threshold, percentile_value, iqr_multiplier, + sample_every=None, block_size=2048): + """两趟分块耀斑检测并写出 uint8 掩膜(method ∈ zscore/percentile/iqr/otsu)。 + + Args: + img_path: 输入影像路径 + mask_arg: 水体掩膜路径(栅格/.shp);None/空串 → 全图皆为水域 + out_path: 输出掩膜路径(.dat/.bil/.bsq→ENVI;.tif/.tiff→GTiff) + bands_1based: 参与检测的波段号列表(1-based,多波段时加权融合) + weights: 波段权重(None→等权) + method: 'zscore' | 'percentile' | 'iqr' | 'otsu' + z_threshold / percentile_value / iqr_multiplier: 方法参数 + """ + ds = gdal.Open(img_path) + if ds is None: + raise ValueError(f"无法打开影像文件: {img_path}") + cols, rows = ds.RasterXSize, ds.RasterYSize + geo, proj = ds.GetGeoTransform(), ds.GetProjection() + + # 抽样步长自适应:目标全局样本量 ~20 万,兼顾统计代表性与内存 + if sample_every is None: + px = max(1, cols * rows) + sample_every = min(256, max(1, int((px / 200000.0) ** 0.5))) + band_objs = [ds.GetRasterBand(b) for b in bands_1based] + nb = len(band_objs) + if nb == 0: + ds = None + raise ValueError("没有可用的检测波段") + if weights is None: + weights = [1.0] * nb + weights = [float(w) for w in weights] + total_w = sum(weights) or 1.0 + + # ---- 掩膜访问器:支持分块读取 ---- + mask_band = None + wm_cleanup = None + if mask_arg is not None and str(mask_arg).strip() != "": + wm = str(mask_arg) + if wm.lower().endswith('.shp'): + wm_cleanup = out_path + '__glint_wm.tif' + from src.utils.extract_water_area import rasterize_shp + rasterize_shp(wm, wm_cleanup, img_path) + mds = gdal.Open(wm_cleanup) + else: + mds = gdal.Open(wm) + if mds is None: + raise ValueError(f"无法打开水域掩膜文件: {mask_arg}") + mask_band = mds.GetRasterBand(1) + if mask_band.XSize != cols or mask_band.YSize != rows: + raise ValueError("掩膜与影像尺寸不一致,两趟检测要求同网格") + + def _read_fused(x, y, xs, ys): + if nb == 1: + return band_objs[0].ReadAsArray(x, y, xs, ys).astype(np.float32) + arr = None + for i, bo in enumerate(band_objs): + bi = bo.ReadAsArray(x, y, xs, ys).astype(np.float32) * (weights[i] / total_w) + arr = bi if arr is None else arr + bi + return arr + + def _read_mask(x, y, xs, ys): + if mask_band is None: + return np.ones((ys, xs), dtype=bool) + return mask_band.ReadAsArray(x, y, xs, ys) > 0 + + def _create_output(): + _remove_raster_output(out_path) + ext = os.path.splitext(out_path)[1].lower() + if ext in ('.tif', '.tiff'): + od = gdal.GetDriverByName('GTiff').Create( + out_path, cols, rows, 1, gdal.GDT_Byte, options=['COMPRESS=LZW']) + else: + od = gdal.GetDriverByName('ENVI').Create(out_path, cols, rows, 1, gdal.GDT_Byte) + od.SetGeoTransform(geo) + od.SetProjection(proj) + return od + + try: + # ────────── 第一趟:抽样扫全局标量 ────────── + samples = [] + for y in range(0, rows, block_size): + ys = min(block_size, rows - y) + for x in range(0, cols, block_size): + xs = min(block_size, cols - x) + blk = _read_fused(x, y, xs, ys) + mb = _read_mask(x, y, xs, ys) + sub = blk[::sample_every, ::sample_every] + mbs = mb[::sample_every, ::sample_every] + valid = mbs & (sub > 0) & np.isfinite(sub) + if valid.any(): + samples.append(sub[valid]) + allv = np.concatenate(samples) if samples else None + + mean_val = std_val = None + thr = None + otsu_thr = None + p2 = p98 = None + + if allv is not None: + if method == 'zscore': + mean_val = float(np.mean(allv)) + std_val = float(np.std(allv)) + if std_val == 0: + raise ValueError("耀斑波段标准差为 0,无法使用 Z-score 方法") + print(f"[两趟] Z-score 全局: 均值={mean_val:.4f}, 标准差={std_val:.4f}, " + f"样本={allv.size}") + elif method == 'percentile': + thr = float(np.percentile(allv, percentile_value)) + print(f"[两趟] 百分位数法: {percentile_value}% 分位={thr:.4f}, 样本={allv.size}") + elif method == 'iqr': + q1, q3 = np.percentile(allv, [25.0, 75.0]) + thr = float(q3 + iqr_multiplier * (q3 - q1)) + print(f"[两趟] IQR 法: Q1={q1:.4f}, Q3={q3:.4f}, 上界={thr:.4f}") + elif method == 'otsu': + p2, p98 = np.percentile(allv, [2.0, 98.0]) + if p98 - p2 > 1e-12: + scaled = (np.clip(allv, p2, p98) - p2) / (p98 - p2) * 255.0 + hist, _ = np.histogram(scaled, bins=256, range=(0.0, 255.0)) + otsu_thr = _otsu_threshold_from_hist(hist) + else: + otsu_thr = 0.0 + print(f"[两趟] Otsu 法: P2={p2:.4f}, P98={p98:.4f}, 阈值(0-255)={otsu_thr:.1f}") + + # ────────── 创建输出 ────────── + od = _create_output() + ob = od.GetRasterBand(1) + ob.Fill(0) + + # 无有效样本 → 全 0(无耀斑) + if allv is None: + ob.FlushCache() + od = None + print("警告: 水域内无有效正反射率样本,输出全 0 耀斑掩膜") + return out_path + + # ────────── 第二趟:分块检测并写盘 ────────── + for y in range(0, rows, block_size): + ys = min(block_size, rows - y) + for x in range(0, cols, block_size): + xs = min(block_size, cols - x) + blk = _read_fused(x, y, xs, ys) + mb = _read_mask(x, y, xs, ys) + + if method == 'zscore': + zs = np.zeros_like(blk, dtype=np.float32) + valid = mb & np.isfinite(blk) + zs[valid] = (blk[valid] - mean_val) / std_val + det = (zs > z_threshold) + elif method == 'percentile': + det = mb & (blk > thr) + elif method == 'iqr': + det = mb & (blk > thr) + elif method == 'otsu': + if p98 is not None and p98 - p2 > 1e-12 and otsu_thr is not None: + scaled = (np.clip(blk, p2, p98) - p2) / (p98 - p2) * 255.0 + det = mb & (scaled > otsu_thr) + else: + det = np.zeros(blk.shape, dtype=bool) + ob.WriteArray(det.astype(np.uint8), x, y) + + ob.FlushCache() + od = None + print(f"[两趟] 耀斑掩膜已写出: {out_path} ({cols}x{rows}, uint8)") + return out_path + finally: + if wm_cleanup is not None: + _remove_raster_output(wm_cleanup) + ds = None + + @timeit -def find_severe_glint_area(img_path, water_mask, glint_wave=750, output_path=None, +def find_severe_glint_area(img_path, water_mask, glint_wave=750, output_path=None, method='otsu', multi_band_waves=None, **kwargs): """ 找到严重耀斑区域的主函数 @@ -604,6 +830,49 @@ def find_severe_glint_area(img_path, water_mask, glint_wave=750, output_path=Non num_bands = dataset.RasterCount im_width = dataset.RasterXSize im_height = dataset.RasterYSize + + # ── ★ v3 两趟快速路径(内存安全)──────────────────────────── + # 默认方法(otsu/zscore/percentile/iqr/multi_band)且未启用 + # max_area/buffer_size 结构后处理时,直接走两趟分块检测, + # 全程不整幅驻留波段与掩膜。adaptive / 结构后处理仍走旧全幅路径。 + max_area_cfg = kwargs.get('max_area', None) + buffer_cfg = kwargs.get('buffer_size', None) + post_needed = bool(max_area_cfg and max_area_cfg > 0) or \ + bool(buffer_cfg and buffer_cfg > 0) + + if not post_needed: + if method == 'multi_band': + waves = multi_band_waves if multi_band_waves else \ + [glint_wave, glint_wave + 50, glint_wave + 100] + bands1 = [] + for wv in waves: + bn = find_band_number(wv, img_path) + if 0 <= bn < dataset.RasterCount: + bands1.append(bn + 1) + sub = kwargs.get('sub_method', 'zscore') + if bands1 and sub in ('zscore', 'percentile', 'iqr', 'otsu'): + _detect_glint_two_pass( + img_path, water_mask, output_path, bands1, + kwargs.get('weights', None), sub, + z_threshold=kwargs.get('z_threshold', 2.5), + percentile_value=kwargs.get('percentile', 95.0), + iqr_multiplier=kwargs.get('iqr_multiplier', 1.5)) + dataset = None + return output_path + elif method in ('zscore', 'percentile', 'iqr', 'otsu'): + bn = find_band_number(glint_wave, img_path) + if 0 <= bn < dataset.RasterCount: + _detect_glint_two_pass( + img_path, water_mask, output_path, [bn + 1], None, method, + z_threshold=kwargs.get('z_threshold', 2.5), + percentile_value=kwargs.get('percentile', 95.0), + iqr_multiplier=kwargs.get('iqr_multiplier', 1.5)) + dataset = None + return output_path + else: + print("信息: 启用了连通域/岸线后处理,主检测走旧全幅路径。") + if im_height * im_width > 50_000_000: + print("警告: 影像超过 5000 万像素,后处理需整幅数组,内存占用会很大。") # 读取水域掩膜,如果water_mask为None或空字符串,则创建全图掩膜 if water_mask is None or water_mask == "":