diff --git a/src/core/utils/preview_generator.py b/src/core/utils/preview_generator.py index eeb93c9..bf63b74 100644 --- a/src/core/utils/preview_generator.py +++ b/src/core/utils/preview_generator.py @@ -22,6 +22,42 @@ plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans', ' plt.rcParams['axes.unicode_minus'] = False +# 预览图最长边像素上限:超过此尺寸的影像一律用 GDAL buf 降采样读取, +# 保证 Python 侧驻留的内存恒定(不随原始分辨率膨胀)。 +MAX_PREVIEW_DIM = 2048 + + +# ============================================================ +# 辅助函数:带最大分辨率限制的底层降维读取 +# ============================================================ + +def _preview_size(x_size: int, y_size: int, max_dim: int = MAX_PREVIEW_DIM): + """按最长边约束计算降采样后的 (buf_w, buf_h)。 + + 两个不同数据集若想叠加(如 RGB 底图 + 掩膜),必须对 XSize/YSize + 相同的源调用同一函数,得到的 buf 尺寸才会严格一致。 + """ + if max(x_size, y_size) <= max_dim: + return x_size, y_size + scale = max_dim / max(x_size, y_size) + return max(1, int(x_size * scale)), max(1, int(y_size * scale)) + + +def _read_downsampled_band(band, max_dim: int = MAX_PREVIEW_DIM): + """利用 GDAL 底层 C++ 直接读取降采样的数组。 + + 通过 band.ReadAsArray(buf_xsize=..., buf_ysize=...) 让 GDAL 在读取时 + 完成抽样——Python 内存里只会出现 buf_y*buf_x 大小的数组,绝不先把 + 整幅(如 34914x19177)读进内存再 resize。重采样用最近邻,预览足够。 + """ + x_size = band.XSize + y_size = band.YSize + if max(x_size, y_size) <= max_dim: + return band.ReadAsArray() + buf_w, buf_h = _preview_size(x_size, y_size, max_dim) + return band.ReadAsArray(buf_xsize=buf_w, buf_ysize=buf_h) + + # ============================================================ # 辅助函数:波段选择 # ============================================================ @@ -169,35 +205,43 @@ def _normalize_mask(mask_data: np.ndarray, def _warp_mask_to_image(mask_path: str, base_dataset: gdal.Dataset, - nodata_output: float = _WARP_NODATA): - """使用 gdal.Warp 将掩膜重采样到与底图完全一致的像素网格 + nodata_output: float = _WARP_NODATA, + max_dim: int = MAX_PREVIEW_DIM): + """使用 gdal.Warp 将掩膜重采样到与底图预览(buf)网格一致 这是解决"矩形色块错位/倾斜"问题的核心函数。当掩膜与 RGB 底图 具有不同的 GeoTransform(旋转参数)或投影时,numpy 级别的 crop/pad 在地理空间上是错误的——只有通过真正的重投影/重采样, 掩膜才能在每个像素位置上与底图精确对齐。 - 内部使用 GDAL MEM 驱动在内存中完成 warp,不产生临时文件。 + ★ v3:支持降采样。Warp 的目标尺寸直接按 max_dim 约束的 buf 尺寸 + (width/height 参数)输出,因此不会构建全分辨率 MEM 数据集, + 内存不随原始分辨率膨胀。返回数组与底图 _read_downsampled_band + 的尺寸严格一致,可直接用于叠加。 Args: mask_path: 掩膜文件路径 base_dataset: 已打开的 RGB 底图 GDAL Dataset nodata_output: warp 后输出栅格的 NoData 填充值 + max_dim: 预览最长边约束(与底图读取用同一值,保证 buf 尺寸一致) Returns: (warped_mask_2d, nodata_used) - - warped_mask_2d: shape=(base_h, base_w) 的 float32 数组, + - warped_mask_2d: shape=(buf_h, buf_w) 的 float32 数组, 水体像素保留原值,背景/越界像素 = nodata_output - nodata_used: 实际使用的 NoData 值(= nodata_output) Raises: - RuntimeError: 若 warp 后尺寸与底图不一致(理论上不应发生) + RuntimeError: 若 warp 后尺寸与 buf 目标差异过大(理论上不应发生) """ base_w = base_dataset.RasterXSize base_h = base_dataset.RasterYSize base_gt = base_dataset.GetGeoTransform() base_proj = base_dataset.GetProjection() + # 目标预览尺寸:掩膜 warp 到与底图 buf 完全一致的网格 + buf_w, buf_h = _preview_size(base_w, base_h, max_dim) + # ---- 计算底图的地理范围(支持旋转影像) ---- # 四个角点的地理坐标 corners_x = [ @@ -232,9 +276,11 @@ def _warp_mask_to_image(mask_path: str, if (mask_w == base_w and mask_h == base_h and abs(mask_gt[1] - base_gt[1]) < 1e-9 and abs(mask_gt[5] - base_gt[5]) < 1e-9): - data = mask_ds.GetRasterBand(1).ReadAsArray().astype(np.float32) + # 同分辨率同范围:与底图同网格,直接 buf 降采样读取, + # 因 XSize/YSize 相同 → 降采样尺寸与底图 buf 严格一致。 + data = _read_downsampled_band(mask_ds.GetRasterBand(1), max_dim).astype(np.float32) mask_ds = None - print("[预览] 掩膜与底图分辨率一致,跳过 Warp 直接读取") + print("[预览] 掩膜与底图分辨率一致,跳过 Warp 直接读取(buf 降采样)") return _normalize_mask(data, nodata_value=src_nodata), nodata_output mask_proj = mask_ds.GetProjection() @@ -242,13 +288,19 @@ def _warp_mask_to_image(mask_path: str, warp_kwargs = { 'format': 'MEM', 'outputBounds': [minx, miny, maxx, maxy], - 'xRes': base_gt[1], # 保留符号 - 'yRes': base_gt[5], # 保留符号(通常为负,表示北在上) + # ★ 不传 xRes/yRes,改用 width/height 直接指定 buf 尺寸输出, + # 避免构建全分辨率 MEM 数据集导致内存爆炸。 + 'width': buf_w, + 'height': buf_h, 'dstSRS': base_proj, 'srcSRS': mask_proj if mask_proj else base_proj, 'dstNodata': nodata_output, 'resampleAlg': gdal.GRA_NearestNeighbour, - 'targetAlignedPixels': True, + # ★ 输出统一 Float32:nodata_output(-9999) 在 Byte 目标上会被 clamp 到 0, + # 使陆地(0)被误当作 nodata 而烧成 1(全图水体)。Float32 则能如实保留 -9999。 + 'outputType': gdal.GDT_Float32, + # 注意:使用 width/height 精确控制输出尺寸时不能带 targetAlignedPixels + # (-tap 必须搭配 -tr/xRes)。固定 outputBounds + width/height 已精确对齐。 'warpOptions': ['NUM_THREADS=ALL_CPUS'], } if src_nodata is not None: @@ -258,10 +310,9 @@ def _warp_mask_to_image(mask_path: str, 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}" + f"dst=({buf_w}x{buf_h} buf)" ) - print(f"[GDAL Warp] 正在将掩膜重采样到与底图一致的像素网格: {warp_kwargs_str}") + print(f"[GDAL Warp] 正在将掩膜重采样到与底图一致的预览网格: {warp_kwargs_str}") try: warped_ds = gdal.Warp('', mask_ds, **warp_kwargs) @@ -277,30 +328,30 @@ def _warp_mask_to_image(mask_path: str, mask_ds = None # 源掩膜已不再需要 # ---- 验证输出尺寸(容忍 ±WARP_SIZE_TOLERANCE 像素的浮点舍入) ---- - # 坐标系转换时的浮点舍入可能导致 warp 输出与底图相差 1-2 像素, + # 坐标系转换时的浮点舍入可能导致 warp 输出与 buf 目标相差 1-2 像素, # 此时地理位置已经对齐,不应回退到直接读取模式。 _WARP_SIZE_TOLERANCE = 2 warped_w = warped_ds.RasterXSize warped_h = warped_ds.RasterYSize - dw = warped_w - base_w - dh = warped_h - base_h + dw = warped_w - buf_w + dh = warped_h - buf_h if abs(dw) > _WARP_SIZE_TOLERANCE or abs(dh) > _WARP_SIZE_TOLERANCE: warped_ds = None raise RuntimeError( - f"gdal.Warp 输出尺寸 ({warped_w}x{warped_h}) 与底图 " - f"({base_w}x{base_h}) 差异过大 (dw={dw}, dh={dh})!" + f"gdal.Warp 输出尺寸 ({warped_w}x{warped_h}) 与 buf 目标 " + f"({buf_w}x{buf_h}) 差异过大 (dw={dw}, dh={dh})!" f"请检查投影/分辨率参数。" ) - # ---- 读取 warped 数据 ---- + # ---- 读取 warped 数据(MEM 数据集本身已是 buf 尺寸,内存很小) ---- warped_data = warped_ds.GetRasterBand(1).ReadAsArray().astype(np.float32) warped_ds = None - # ---- 微调尺寸到 target_shape(slice 多余边缘 或 pad 缺失边缘) ---- + # ---- 微调尺寸到 buf 目标(slice 多余边缘 或 pad 缺失边缘) ---- if dw != 0 or dh != 0: warped_data = _snap_array_to_shape( - warped_data, base_h, base_w, fill_value=nodata_output + warped_data, buf_h, buf_w, fill_value=nodata_output ) return warped_data, nodata_output @@ -402,10 +453,15 @@ def generate_image_preview(img_path: str, else: bands = [0, 0, 0] - # 读取波段 - r_data = dataset.GetRasterBand(bands[0] + 1).ReadAsArray().astype(np.float32) - g_data = r_data if band_count == 1 else dataset.GetRasterBand(bands[1] + 1).ReadAsArray().astype(np.float32) - b_data = r_data if band_count <= 2 else dataset.GetRasterBand(bands[2] + 1).ReadAsArray().astype(np.float32) + # 读取波段(★ buf 降采样:无论影像多大,Python 内存只驻留 <=2048 边长的数组) + def _read_rgb_band(idx): + return _read_downsampled_band( + dataset.GetRasterBand(bands[idx] + 1), MAX_PREVIEW_DIM + ).astype(np.float32) + + r_data = _read_rgb_band(0) + g_data = r_data if band_count == 1 else _read_rgb_band(1) + b_data = r_data if band_count <= 2 else _read_rgb_band(2) r_data[r_data <= 0] = np.nan if band_count > 1: @@ -501,9 +557,18 @@ def generate_water_mask_overlay(img_path: str, else: bands = [0, 0, 0] - r_data = dataset.GetRasterBand(bands[0] + 1).ReadAsArray().astype(np.float32) - g_data = r_data if band_count == 1 else dataset.GetRasterBand(bands[1] + 1).ReadAsArray().astype(np.float32) - b_data = r_data if band_count <= 2 else dataset.GetRasterBand(bands[2] + 1).ReadAsArray().astype(np.float32) + # 读取波段(★ buf 降采样:与掩膜用同一 MAX_PREVIEW_DIM → buf 尺寸严格一致) + def _read_rgb_band(idx): + return _read_downsampled_band( + dataset.GetRasterBand(bands[idx] + 1), MAX_PREVIEW_DIM + ).astype(np.float32) + + r_data = _read_rgb_band(0) + g_data = r_data if band_count == 1 else _read_rgb_band(1) + b_data = r_data if band_count <= 2 else _read_rgb_band(2) + + # 预览 buf 尺寸(面积统计时按比例放大回原始分辨率像元数) + buf_w, buf_h = _preview_size(width, height, MAX_PREVIEW_DIM) r_data[r_data <= 0] = np.nan if band_count > 1: @@ -541,9 +606,10 @@ def generate_water_mask_overlay(img_path: str, print(f"警告: 掩膜文件不存在: {mask_path}") else: try: - # Step A: 地理空间 warp(核心) + # Step A: 地理空间 warp(核心,buf 降采样输出) mask_data_raw, mask_nodata = _warp_mask_to_image( - mask_path, dataset, nodata_output=_WARP_NODATA + mask_path, dataset, nodata_output=_WARP_NODATA, + max_dim=MAX_PREVIEW_DIM, ) # Step B: 值域归一化(显式传入 warp 的 nodata,确保背景透明) @@ -561,7 +627,11 @@ def generate_water_mask_overlay(img_path: str, # 回退:直接读取掩膜(假设与底图已对齐) 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) + # 回退路径假设掩膜与底图已对齐(同网格): + # XSize/YSize 相同 → 与底图用同一 MAX_PREVIEW_DIM 降采样后 buf 尺寸一致 + mask_data_raw = _read_downsampled_band( + mask_ds.GetRasterBand(1), MAX_PREVIEW_DIM + ).astype(np.float32) # 读取源掩膜的 nodata(如果存在) try: src_nd = mask_ds.GetRasterBand(1).GetNoDataValue() @@ -571,16 +641,16 @@ def generate_water_mask_overlay(img_path: str, 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) + # 若 buf 尺寸不一致(极端情形),做简单 crop/pad 到 buf + if mask_data_raw.shape[0] != buf_h or mask_data_raw.shape[1] != buf_w: + aligned = np.zeros((buf_h, buf_w), dtype=np.float32) + copy_h = min(mask_data_raw.shape[0], buf_h) + copy_w = min(mask_data_raw.shape[1], buf_w) 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)。" + f"[回退对齐] 掩膜 → buf ({buf_w}x{buf_h})," + f"已裁剪/填充(非 warp)。" ) mask_data = _normalize_mask(mask_data_raw, nodata_value=mask_nodata) @@ -622,7 +692,11 @@ def generate_water_mask_overlay(img_path: str, 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 + # 统计在 buf 数组上进行:buf 中 1 像元 ≈ scale_x*scale_y 个原始像元。 + # 乘回该比例,面积/占比仍以原始分辨率(m²/px)为基准,保持可信。 + scale_x = (width / buf_w) if buf_w else 1.0 + scale_y = (height / buf_h) if buf_h else 1.0 + pixel_area = pixel_size_x * pixel_size_y * scale_x * scale_y if mask_data is not None: water_pixels = int(np.sum(mask_data > 0))