diff --git a/src/core/prediction/inference_batch.py b/src/core/prediction/inference_batch.py index 9489a2d..fbf24ea 100644 --- a/src/core/prediction/inference_batch.py +++ b/src/core/prediction/inference_batch.py @@ -361,6 +361,9 @@ class WaterQualityInference: if 'test_rmse' in metadata: print(f" 测试集RMSE: {metadata['test_rmse']:.4f}") + # ★ 加载时打印 MNF 波段选择信息(从已训练的 Pipeline 中提取) + self._print_mnf_info_if_available() + def load_specific_model(self, model_file_path: str): """ 加载指定的模型文件 @@ -381,6 +384,9 @@ class WaterQualityInference: print(f" 模型名称: {self.loaded_model_data['model_name']}") print(f" 模型类型: {type(self.loaded_model_data['model'])}") + # ★ 加载时打印 MNF 波段选择信息 + self._print_mnf_info_if_available() + def _auto_detect_train_wavelengths(self): """自动获取训练波长:优先级 json/txt > 工作目录 CSV > None""" import os as _os, json as _json, glob as _glob @@ -461,14 +467,15 @@ class WaterQualityInference: pass # np.interp 重采样:308/113/任意波段 → 模型训练波长 + # ★ 边缘填充:left/right 使用当前行首尾有效值,杜绝 NaN→0.0 断崖 spec_data = spectra[spec_cols].values.astype(np.float64) src_wl = np.array([float(c) for c in spec_cols], dtype=np.float64) dst_wl = np.array(train_wl, dtype=np.float64) resampled = np.zeros((spec_data.shape[0], len(dst_wl)), dtype=np.float64) for i in range(spec_data.shape[0]): - resampled[i] = np.interp(dst_wl, src_wl, spec_data[i], - left=np.nan, right=np.nan) - resampled = np.nan_to_num(resampled, nan=0.0) + y_vals = spec_data[i] + resampled[i] = np.interp(dst_wl, src_wl, y_vals, + left=y_vals[0], right=y_vals[-1]) result = pd.DataFrame(resampled, columns=[f'{wl:.6f}' for wl in train_wl]) @@ -572,28 +579,12 @@ class WaterQualityInference: resampled = np.zeros((spectral_data.shape[0], len(train_wavelengths)), dtype=np.float64) for i in range(spectral_data.shape[0]): + y_vals = spectral_data[i] resampled[i] = np.interp( - train_wl_arr, target_wl_arr, spectral_data[i], - left=np.nan, right=np.nan + train_wl_arr, target_wl_arr, y_vals, + left=y_vals[0], right=y_vals[-1] ) - # NaN 填充(外推区域):用最近有效值填充 - nan_mask = np.isnan(resampled) - if nan_mask.any(): - print(f"[光谱重采样] {nan_mask.sum()} 个 NaN (外推区域) → 最近邻填充") - for i in range(resampled.shape[0]): - row_nan = np.isnan(resampled[i]) - if row_nan.any(): - valid = ~row_nan - if valid.any(): - valid_idx = np.where(valid)[0] - resampled[i, row_nan] = np.interp( - np.where(row_nan)[0], valid_idx, - resampled[i, valid] - ) - else: - resampled[i] = 0.0 - # 4) 重组为 DataFrame(列名 = 训练波长字符串) wl_col_names = [f"{wl:.6f}" for wl in train_wavelengths] resampled_df = pd.DataFrame(resampled, columns=wl_col_names, @@ -669,24 +660,11 @@ class WaterQualityInference: _resampled = np.zeros((_spec_data.shape[0], expected_features), dtype=np.float64) for i in range(_spec_data.shape[0]): + _y_vals = _spec_data[i] _resampled[i] = np.interp( - fallback_wl, _target_arr, _spec_data[i], - left=np.nan, right=np.nan + fallback_wl, _target_arr, _y_vals, + left=_y_vals[0], right=_y_vals[-1] ) - # NaN 填充 - _nan = np.isnan(_resampled) - if _nan.any(): - for i in range(_resampled.shape[0]): - row_nan = np.isnan(_resampled[i]) - if row_nan.any(): - valid = ~row_nan - if valid.any(): - vi = np.where(valid)[0] - _resampled[i, row_nan] = np.interp( - np.where(row_nan)[0], vi, _resampled[i, valid] - ) - else: - _resampled[i] = 0.0 # 重组 DataFrame _wl_cols = [f"{w:.1f}" for w in fallback_wl] @@ -837,6 +815,72 @@ class WaterQualityInference: f"({(n_clipped_lo + n_clipped_hi) / len(predictions) * 100:.1f}%)") return np.clip(predictions, lo, hi) + def _print_mnf_info_if_available(self): + """从已加载的 Pipeline 模型中提取 MNFTransformer 并打印波段选择信息。""" + if self.loaded_model_data is None: + return + model = self.loaded_model_data.get('model') + if model is None: + return + + from sklearn.pipeline import Pipeline, FeatureUnion + if not isinstance(model, Pipeline): + return + + # 递归查找 MNFTransformer + def _find_mnf(step): + if step.__class__.__name__ == 'MNFTransformer': + return step + if isinstance(step, FeatureUnion): + for _name, _trans in step.transformer_list: + found = _find_mnf(_trans) + if found is not None: + return found + if isinstance(step, Pipeline): + for _sub_name, _sub_step in step.steps: + found = _find_mnf(_sub_step) + if found is not None: + return found + return None + + mnf = _find_mnf(model) + if mnf is not None and hasattr(mnf, 'eigvals_w_'): + target_name = self.loaded_model_data.get('target_column_name', '') + if not target_name and 'metadata' in self.loaded_model_data: + target_name = self.loaded_model_data['metadata'].get('target_column_name', '') + if target_name: + print(f"\n [MNF 波段信息] 目标指数: {target_name}") + mnf._print_band_selection() + + @staticmethod + def _mask_zero_spectra_pixels(spectra: pd.DataFrame, + predictions: np.ndarray) -> np.ndarray: + """将全零光谱行(NDWI 掩膜外的陆地像素)的预测值强制设为 NaN。 + + GUI 制图 / Kriging / IDW 插值天然忽略 NaN,这是最安全的掩膜继承方式。 + + Parameters + ---------- + spectra : pd.DataFrame + 原始光谱 DataFrame(shape: n_pixels × n_bands) + predictions : np.ndarray + 模型预测值数组(shape: n_pixels) + + Returns + ------- + np.ndarray + 处理后的预测值数组 + """ + all_zero = (spectra == 0).all(axis=1) + n_zero = int(all_zero.sum()) + if n_zero == 0: + return predictions + predictions = predictions.astype(np.float64, copy=True) + predictions[all_zero.values] = np.nan + print(f"[全零拦截] 检测到 {n_zero} 个全零光谱像素 " + f"({n_zero / len(predictions) * 100:.1f}%),预测值已设为 NaN") + return predictions + def save_predictions(self, coords: pd.DataFrame, predictions: np.ndarray, output_path: str, prediction_column: str = 'prediction', wqi_columns: Optional[pd.DataFrame] = None): @@ -947,6 +991,9 @@ class WaterQualityInference: print("-" * 40) predictions = self.predict(spectra_processed) + # ★ 全零光谱拦截:NDWI 掩膜外的陆地像素 → NaN + predictions = self._mask_zero_spectra_pixels(spectra, predictions) + # 5. 保存预测结果(透传 WQI 列至最终输出文件) print("\n步骤5: 保存预测结果") print("-" * 40) @@ -1024,6 +1071,7 @@ class WaterQualityInference: coords, spectra, wqi_df = self.load_sampling_data(str(csv_file)) spectra_processed = self.preprocess_spectra(spectra) predictions = self.predict(spectra_processed) + predictions = self._mask_zero_spectra_pixels(spectra, predictions) result_df = self.save_predictions(coords, predictions, str(output_file), prediction_column, wqi_df) @@ -1223,6 +1271,7 @@ class WaterQualityInference: coords, spectra, wqi_df = self.load_sampling_data(str(csv_file)) spectra_processed = self.preprocess_spectra(spectra) predictions = self.predict(spectra_processed) + predictions = self._mask_zero_spectra_pixels(spectra, predictions) result_df = self.save_predictions(coords, predictions, str(output_file), prediction_column, wqi_df)