fix: 采样光谱源头量级收敛(全图探测+均值缩放),避免整块重铸OOM

- 新增 _detect_global_reflectance_scale():采样前全图散布窗口粗扫,判定 0-10000 放大反射率
- chunked/non-chunked 在采样均值层面按统一因子收敛到 0~1(均值线性,零大数组开销)
- 修复整分块 astype/除法导致的 10.9GB 二次分配 OOM
(注:该文件另含先前本地未提交改动,一并入库)
This commit is contained in:
duxin
2026-09-08 17:31:55 +08:00
parent 8f416c728b
commit 61a0a9dde4

View File

@ -15,7 +15,7 @@ import pandas as pd
from osgeo import gdal, ogr
import spectral
from scipy import ndimage
from src.utils.util import write_bands
from src.utils.util import write_bands, atomic_filepath
from src.core.utils.spatial_validator import validate_spatial_alignment
try:
@ -38,6 +38,8 @@ def get_wavelengths_from_bil_header(bil_file):
list - 波长列表如果无法获取则返回None
"""
try:
import glob # 同目录 *ref*.hdr 回退所需
# 获取头文件路径(多命名规范兼容 .bsq/.bil/.bip/.dat
hdr_candidates = [
os.path.splitext(bil_file)[0] + ".hdr", # 3ref.hdr
@ -45,15 +47,30 @@ def get_wavelengths_from_bil_header(bil_file):
os.path.splitext(bil_file)[0] + ".HDR", # 3ref.HDR
bil_file + ".HDR", # 3ref.bip.HDR
]
header_file = None
for candidate in hdr_candidates:
if os.path.exists(candidate):
header_file = candidate
break
# ==========================================
# ★ 新增:如果当前文件没头(如GTiff),去同目录找原始的 ref.bip.hdr
# ==========================================
if header_file is None:
print(f"警告: 找不到头文件,已尝试: {hdr_candidates}")
return None
dir_name = os.path.dirname(bil_file)
# 搜索包含 ref 的原始头文件
fallback_hdrs = glob.glob(os.path.join(dir_name, "*ref*.hdr")) + \
glob.glob(os.path.join(dir_name, "*ref*.HDR"))
if fallback_hdrs:
header_file = fallback_hdrs[0]
print(f"[*] 提示: 找不到当前影像头文件,已自动回退读取原始影像头文件: "
f"{os.path.basename(header_file)}")
else:
print(f"警告: 找不到任何可用的头文件,已尝试当前文件及同目录下 *ref*.hdr")
return None
# ==========================================
# 使用spectral库读取头文件
import spectral.io.envi as envi
@ -84,6 +101,45 @@ def get_wavelengths_from_bil_header(bil_file):
return None
def _detect_global_reflectance_scale(dataset_bil, threshold=10.0, divisor=10000.0,
n_probes=300, patch=16, seed=7):
"""
探测整幅影像的光谱量级,判定是否属 0-10000 放大反射率格式。
水体多为暗像素,单个采样窗口的最大值可能远小于阈值,因此在采样前对全图
散布随机窗口粗扫:只要任一窗口最大值 > threshold即判定整幅为放大格式并
返回 divisor默认 10000否则返回 1.0(影像已是 0~1 物理反射率)。
返回:
float: 1.0(无需缩放)或 divisor下游将每个窗口统一除以该值
"""
if dataset_bil is None:
return 1.0
im_w = dataset_bil.RasterXSize
im_h = dataset_bil.RasterYSize
if im_w < patch or im_h < patch:
return 1.0
rng = np.random.RandomState(seed)
xs = rng.randint(0, im_w - patch, size=n_probes)
ys = rng.randint(0, im_h - patch, size=n_probes)
for x, y in zip(xs, ys):
try:
block = dataset_bil.ReadAsArray(int(x), int(y), patch, patch)
except Exception:
continue
if block is None or block.size == 0:
continue
if not np.issubdtype(block.dtype, np.floating):
block = block.astype(np.float32)
m = float(np.nanmax(block))
if m > threshold:
print(f"[量级探测] 影像疑似 0-{int(divisor):d} 放大反射率 "
f"(探测窗口最大值 {m:.2f}),后续光谱统一 /{int(divisor):d}")
return divisor
print("[量级探测] 影像反射率量级正常 (<=10),采样光谱保持原值")
return 1.0
def get_spectral_sampling_points_chunked(bil_file, water_mask_shp, severe_glint=None, output_csvpath=None,
interval=100, sample_radius=1, chunk_size=1000,
use_adaptive_sampling=True, min_interval=10, max_interval=200,
@ -212,6 +268,10 @@ def get_spectral_sampling_points_chunked(bil_file, water_mask_shp, severe_glint=
sample_count = 0
sampled_pixels = set() # 用于记录已采样的像素,避免重复
# ★ 量级探测:整幅若为 0-10000 放大反射率,先求统一缩放因子;
# 均值在 add_sample_point_chunked 内收敛,避免整块重铸/除法造成二次大分配 OOM。
_global_scale = _detect_global_reflectance_scale(dataset_bil)
# 辅助函数:添加采样点(分块版本)
def add_sample_point_chunked(x, y, local_y, spectral_chunk, valid_chunk, sample_radius,
geotransform_input, num_bands, f, x_out, y_out,
@ -251,6 +311,11 @@ def get_spectral_sampling_points_chunked(bil_file, water_mask_shp, severe_glint=
mean_value = np.nan
spectral_sample.append(mean_value)
# ★ 量级收敛:整幅影像为 0-10000 放大反射率时,把每个采样均值收敛到 0~1。
# (均值线性 ⇒ mean(X/scale)=mean(X)/scale故在均值后缩放零大数组开销
if _global_scale > 1:
spectral_sample = [v / _global_scale for v in spectral_sample]
# 转换为地理坐标
geo_x, geo_y = gdal.ApplyGeoTransform(
geotransform_input,
@ -291,6 +356,8 @@ def get_spectral_sampling_points_chunked(bil_file, water_mask_shp, severe_glint=
0, read_start, im_width, read_end - read_start
) # shape: (bands, chunk_height, width)
# 分块按原始 dtype 读取(不做整块重铸/除法,避免 10GB+ 级二次分配 OOM
# 量级收敛统一在 add_sample_point_chunked 内按 _global_scale 对均值缩放。
# 获取对应的有效区域掩膜和宽度图
valid_chunk = valid_area[read_start:read_end, :]
water_chunk = water_mask_raster[read_start:read_end, :]
@ -457,9 +524,10 @@ def get_spectral_sampling_points(bil_file, water_mask_shp, severe_glint=None, ou
print(f"bil文件信息: 宽度={im_width}, 高度={im_height}, 波段数={num_bands}")
# 读取光谱数据(所有波段)
print("正在读取光谱数据...")
spectral_data_full = dataset_bil.ReadAsArray() # shape: (bands, height, width)
# ★ v3不再整幅读取全波段——spectral_data_full 会把
# 6.7 亿像素 × 波段数的影像一次吃进十几 GB 内存。
# dataset_bil 保留为打开的 GDAL 句柄,待采样点确定后
# 按需用窗口 ReadAsArray(x_off, y_off, ws, ws) 即时读取。
# 创建水体掩膜栅格
print("正在处理水体掩膜...")
@ -525,9 +593,17 @@ def get_spectral_sampling_points(bil_file, water_mask_shp, severe_glint=None, ou
y_out = []
spectral_out = []
# 如果没有提供输出路径,则不保存文件
# 如果没有提供输出路径,则不保存文件(★ 原子:先写 .__wip成功才替换
_csv_ok = False
_wip_csv = None
if output_csvpath:
f = open(output_csvpath, "w")
_wip_csv = output_csvpath + ".__wip"
if os.path.exists(_wip_csv):
try:
os.remove(_wip_csv)
except OSError:
pass
f = open(_wip_csv, "w")
# 写入CSV头部
header = "x_coord,y_coord,pixel_x,pixel_y"
@ -546,43 +622,71 @@ def get_spectral_sampling_points(bil_file, water_mask_shp, severe_glint=None, ou
else:
f = None
# ★ 量级探测:整幅影像若为 0-10000 放大反射率,先求统一缩放因子,
# 使每个采样窗口(即使水体暗像素单窗 <10都被同一因子收敛到 0~1。
_global_scale = _detect_global_reflectance_scale(dataset_bil)
try:
print("正在生成采样点...")
sample_count = 0
sampled_pixels = set() # 用于记录已采样的像素,避免重复
# 辅助函数:添加采样点
def add_sample_point(x, y, spectral_data_full, valid_area, sample_radius,
geotransform_input, num_bands, f, x_out, y_out,
spectral_out, sampled_pixels):
# 辅助函数:添加采样点(★ v3光谱改为 GDAL 窗口按需读取)
def add_sample_point(x, y, dataset_bil, valid_area, sample_radius,
geotransform_input, num_bands, im_width, im_height,
f, x_out, y_out, spectral_out, sampled_pixels):
"""添加单个采样点"""
# 检查是否已采样
if (x, y) in sampled_pixels:
return False
# 检查边界
if (x < sample_radius or x >= im_width - sample_radius or
# 检查边界(防御;扫描循环范围已保证 x_off/y_off 不越界)
if (x < sample_radius or x >= im_width - sample_radius or
y < sample_radius or y >= im_height - sample_radius):
return False
# ★ v3窗口按需读取的几何参数
r = sample_radius
ws = 2 * r + 1
x_off = x - r
y_off = y - r
# 检查采样点周围区域水体占比
sample_area = valid_area[y - sample_radius:y + sample_radius + 1,
x - sample_radius:x + sample_radius + 1]
sample_area = valid_area[y_off:y_off + ws,
x_off:x_off + ws]
# ★ v2: 允许窗口内部分非水体像素窄水体友好默认≥60%水体即通过)
water_ratio = np.mean(sample_area.astype(np.float32))
if np.isnan(water_ratio):
water_ratio = 0.0
if water_ratio >= water_ratio_threshold:
# 提取光谱数据(采样区域内的平均值)
# ★ v3仅当确定采这个点才从磁盘即时读取该小窗口的所有波段。
# GDAL 读取 shape = (num_bands, ws, ws),不驻留全幅光谱。
try:
window_spectral_data = dataset_bil.ReadAsArray(
x_off, y_off, ws, ws).astype(np.float32)
except Exception:
return False
if window_spectral_data.ndim != 3 or \
window_spectral_data.shape != (num_bands, ws, ws):
# GDAL 越界会静默裁剪而非抛异常 → shape 不匹配时放弃该点
return False
# ★ 量级统管:整幅判定为放大反射率时,把本窗口统一收敛到 0~1
# 避免暗像素单窗 <10 造成“该缩不缩”的量级撕裂。
if _global_scale > 1:
window_spectral_data = window_spectral_data / _global_scale
# 提取光谱数据(窗口内水体像素的波段平均)
spectral_sample = []
for band_idx in range(num_bands):
band_data = spectral_data_full[band_idx,
y - sample_radius:y + sample_radius + 1,
x - sample_radius:x + sample_radius + 1]
# 计算平均值,忽略无效值
valid_pixels = band_data[sample_area]
if len(valid_pixels) > 0:
band_data = window_spectral_data[band_idx] # (ws, ws)
# ★ 绝对净水器NaN/Inf/负反射率一律挡在均值之外(宁缺毋滥)
raw_pixels = band_data[sample_area]
clean = np.isfinite(raw_pixels) & (raw_pixels >= 0)
valid_pixels = raw_pixels[clean]
# 干净像素占比须过半才采纳,否则判 NaN 交下游清洗
if len(valid_pixels) > (len(raw_pixels) * 0.5):
mean_value = np.mean(valid_pixels)
else:
mean_value = np.nan
@ -643,8 +747,9 @@ def get_spectral_sampling_points(bil_file, water_mask_shp, severe_glint=None, ou
adaptive_interval = base_interval
# 尝试添加采样点
if add_sample_point(x, y, spectral_data_full, valid_area, sample_radius,
geotransform_input, num_bands, f, x_out, y_out,
if add_sample_point(x, y, dataset_bil, valid_area, sample_radius,
geotransform_input, num_bands, im_width, im_height,
f, x_out, y_out,
spectral_out, sampled_pixels):
sample_count += 1
@ -666,16 +771,30 @@ def get_spectral_sampling_points(bil_file, water_mask_shp, severe_glint=None, ou
print(f"使用固定间隔采样(间隔: {interval}...")
for y in range(sample_radius, im_height - sample_radius, interval):
for x in range(sample_radius, im_width - sample_radius, interval):
if add_sample_point(x, y, spectral_data_full, valid_area, sample_radius,
geotransform_input, num_bands, f, x_out, y_out,
if add_sample_point(x, y, dataset_bil, valid_area, sample_radius,
geotransform_input, num_bands, im_width, im_height,
f, x_out, y_out,
spectral_out, sampled_pixels):
sample_count += 1
print(f"成功生成 {sample_count} 个采样点")
_csv_ok = True
finally:
if f:
f.close()
# ★ 原子提交:成功才替换为最终路径并打 .done失败则丢弃 .__wip
if _wip_csv is not None:
if _csv_ok and os.path.exists(_wip_csv):
from src.utils.util import mark_file_complete
os.replace(_wip_csv, output_csvpath)
mark_file_complete(output_csvpath)
else:
try:
if os.path.exists(_wip_csv):
os.remove(_wip_csv)
except OSError:
pass
return x_out, y_out, np.array(spectral_out)